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
120 changes: 103 additions & 17 deletions generator/templates/builders_create.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -297,6 +297,7 @@ func loadRelation[P any, C any](
scan func(*sql.Rows, *C) error,
childKey func(*C) (string, bool),
assign func(*P, []*C),
params QueryParams,
) ([]*C, error) {
var parentKeys []any
for _, p := range parents {
Expand All @@ -311,25 +312,26 @@ func loadRelation[P any, C any](
return nil, nil
}

var sb strings.Builder
sb.Grow(128 + len(returningCols)*15 + len(table) + len(fkCol) + len(parentKeys)*3)
sb.WriteString("SELECT ")
for i, col := range returningCols {
if i > 0 {
sb.WriteString(", ")
}
sb.WriteString(q.dialect.Quote(col))
// Prepend parent ID checks to filters using StandardPredicate
allPreds := append([]Predicate{
StandardPredicate{
Data: PredicateData{
Column: fkCol,
Operator: "IN",
Value: parentKeys,
IsLogical: false,
},
},
}, params.Where...)

whereClause, vals := CompilePredicates(q.dialect, allPreds)
if whereClause != "" {
whereClause = " WHERE " + whereClause
}
sb.WriteString(" FROM ")
sb.WriteString(q.dialect.Quote(table))
sb.WriteString(" WHERE ")
sb.WriteString(q.dialect.Quote(fkCol))
sb.WriteString(" IN (")
sb.WriteString(q.bindVars(len(parentKeys)))
sb.WriteString(")")
query := sb.String()

rows, err := q.query(ctx, query, parentKeys...)
query := compileRelationSQL(q.dialect, table, fkCol, returningCols, whereClause, params)

rows, err := q.query(ctx, query, vals...)
if err != nil {
return nil, err
}
Expand Down Expand Up @@ -363,3 +365,87 @@ func loadRelation[P any, C any](

return allChildren, nil
}

func compileRelationSQL(dialect Dialect, table, fkCol string, cols []string, where string, params QueryParams) string {
if params.Take != nil || params.Skip != nil {
return compilePartitionedRelationSQL(dialect, table, fkCol, cols, where, params)
}
return compileSimpleRelationSQL(dialect, table, cols, where, params)
}

func compilePartitionedRelationSQL(dialect Dialect, table, fkCol string, cols []string, where string, params QueryParams) string {
var innerSb strings.Builder
innerSb.WriteString("SELECT ")
for i, col := range cols {
if i > 0 {
innerSb.WriteString(", ")
}
innerSb.WriteString(dialect.Quote(col))
}
innerSb.WriteString(", ROW_NUMBER() OVER (PARTITION BY ")
innerSb.WriteString(dialect.Quote(fkCol))
innerSb.WriteString(" ORDER BY ")
if len(params.OrderBy) > 0 {
for i, ord := range params.OrderBy {
if i > 0 {
innerSb.WriteString(", ")
}
innerSb.WriteString(dialect.Quote(ord.Field))
innerSb.WriteString(" ")
innerSb.WriteString(string(ord.Direction))
}
} else {
innerSb.WriteString(dialect.Quote("id"))
innerSb.WriteString(" ASC")
}
innerSb.WriteString(") as row_num FROM ")
innerSb.WriteString(dialect.Quote(table))
innerSb.WriteString(where)

var outerSb strings.Builder
outerSb.WriteString("SELECT ")
for i, col := range cols {
if i > 0 {
outerSb.WriteString(", ")
}
outerSb.WriteString(dialect.Quote(col))
}
outerSb.WriteString(" FROM (")
outerSb.WriteString(innerSb.String())
outerSb.WriteString(") t WHERE ")

if params.Take != nil && params.Skip != nil {
outerSb.WriteString(fmt.Sprintf("row_num > %d AND row_num <= %d", *params.Skip, *params.Skip+*params.Take))
} else if params.Take != nil {
outerSb.WriteString(fmt.Sprintf("row_num <= %d", *params.Take))
} else if params.Skip != nil {
outerSb.WriteString(fmt.Sprintf("row_num > %d", *params.Skip))
}
return outerSb.String()
}

func compileSimpleRelationSQL(dialect Dialect, table string, cols []string, where string, params QueryParams) string {
var sb strings.Builder
sb.WriteString("SELECT ")
for i, col := range cols {
if i > 0 {
sb.WriteString(", ")
}
sb.WriteString(dialect.Quote(col))
}
sb.WriteString(" FROM ")
sb.WriteString(dialect.Quote(table))
sb.WriteString(where)
if len(params.OrderBy) > 0 {
sb.WriteString(" ORDER BY ")
for i, ord := range params.OrderBy {
if i > 0 {
sb.WriteString(", ")
}
sb.WriteString(dialect.Quote(ord.Field))
sb.WriteString(" ")
sb.WriteString(string(ord.Direction))
}
}
return sb.String()
}
32 changes: 32 additions & 0 deletions generator/templates/client.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -485,6 +485,14 @@ func (f Field[T]) IsNotNull() Predicate {
}
}

func (f Field[T]) Asc() OrderBy {
return OrderBy{Field: f.Column, Direction: Asc}
}

func (f Field[T]) Desc() OrderBy {
return OrderBy{Field: f.Column, Direction: Desc}
}

type UniqueField[T any] struct {
Column string
}
Expand Down Expand Up @@ -596,6 +604,14 @@ func (f UniqueField[T]) IsNotNull() Predicate {
}
}

func (f UniqueField[T]) Asc() OrderBy {
return OrderBy{Field: f.Column, Direction: Asc}
}

func (f UniqueField[T]) Desc() OrderBy {
return OrderBy{Field: f.Column, Direction: Desc}
}

type StringField struct {
Column string
}
Expand Down Expand Up @@ -712,6 +728,14 @@ func (f StringField) IsNotNull() Predicate {
}
}

func (f StringField) Asc() OrderBy {
return OrderBy{Field: f.Column, Direction: Asc}
}

func (f StringField) Desc() OrderBy {
return OrderBy{Field: f.Column, Direction: Desc}
}

type StringUniqueField struct {
Column string
}
Expand Down Expand Up @@ -830,6 +854,14 @@ func (f StringUniqueField) IsNotNull() Predicate {
}
}

func (f StringUniqueField) Asc() OrderBy {
return OrderBy{Field: f.Column, Direction: Asc}
}

func (f StringUniqueField) Desc() OrderBy {
return OrderBy{Field: f.Column, Direction: Desc}
}

func CompilePredicates(dialect Dialect, preds []Predicate) (string, []any) {
if len(preds) == 0 {
return "", nil
Expand Down
6 changes: 6 additions & 0 deletions generator/templates/model_predicate.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,12 @@ func (p UniquePredicate) Validate() error {

type Select = {{ .ParentPackageName }}.{{ .Model.Name }}Select
type Omit = {{ .ParentPackageName }}.{{ .Model.Name }}Omit
type QueryBuilder = {{ .ParentPackageName }}.{{ .Model.Name }}QueryBuilder

func Query() *QueryBuilder {
return &QueryBuilder{}
}

func Record(assignments ...{{ .ParentPackageName }}.FieldAssignment) {{ .ParentPackageName }}.RecordInput {
return {{ .ParentPackageName }}.RecordInput{Assignments: assignments}
}
Expand Down
6 changes: 4 additions & 2 deletions generator/templates/model_relations.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,8 @@ func (q *Queries) load{{ .Model.Name }}Relations(ctx context.Context, records []
{{- else }}
{{- $forceCol = (index $relation.Inverse.FKFields 0).EffectiveColName }}
{{- end }}
returningCols := q.select{{ $relation.TargetModelName }}Cols(selects.{{ capitalize $relation.Name }}, nil, "{{ $forceCol }}")
relationSelects, relationOmits, relationParams := selects.{{ capitalize $relation.Name }}.GetRelationParams()
returningCols := q.select{{ $relation.TargetModelName }}Cols(relationSelects, relationOmits, "{{ $forceCol }}")

{{- if gt (len $relation.FKFields) 0 }}
// Current model holds the FK: {{ $.Model.Name }}.{{ (index $relation.FKFields 0).Name }}
Expand Down Expand Up @@ -51,11 +52,12 @@ func (q *Queries) load{{ .Model.Name }}Relations(ctx context.Context, records []
{{- else }}
setOne(func(p *{{ $.Model.Name }}, c *{{ $relation.TargetModelName }}) { p.{{ capitalize $relation.Name }} = c }),
{{- end }}
relationParams,
)
if err != nil {
return fmt.Errorf("loading {{ $relation.Name }}: %w", err)
}
if err := q.load{{ $relation.TargetModelName }}Relations(ctx, allChildren, selects.{{ capitalize $relation.Name }}); err != nil {
if err := q.load{{ $relation.TargetModelName }}Relations(ctx, allChildren, relationSelects); err != nil {
return err
}
}
Expand Down
65 changes: 61 additions & 4 deletions generator/templates/model_structs.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ type {{ .Model.Name }}Select struct {
{{ capitalize $field.Name }} bool `json:"{{ $field.Name }}"`
{{- end }}
{{- range $relation := .Model.RelationFields }}
{{ capitalize $relation.Name }} *{{ $relation.TargetModelName }}Select `json:"{{ $relation.Name }},omitempty"`
{{ capitalize $relation.Name }} {{ $relation.TargetModelName }}SelectQuery `json:"{{ $relation.Name }},omitempty"`
{{- end }}
}

Expand All @@ -30,9 +30,66 @@ type {{ .Model.Name }}Omit struct {
{{- range $field := .Model.ScalarFields }}
{{ capitalize $field.Name }} bool `json:"{{ $field.Name }}"`
{{- end }}
{{- range $relation := .Model.RelationFields }}
{{ capitalize $relation.Name }} *{{ $relation.TargetModelName }}Omit `json:"{{ $relation.Name }},omitempty"`
{{- end }}
}

type {{ .Model.Name }}SelectQuery interface {
GetRelationParams() (*{{ .Model.Name }}Select, *{{ .Model.Name }}Omit, QueryParams)
}

func (s *{{ .Model.Name }}Select) GetRelationParams() (*{{ .Model.Name }}Select, *{{ .Model.Name }}Omit, QueryParams) {
return s, nil, QueryParams{}
}

// {{ .Model.Name }}QueryBuilder builds a query for the relation {{ .Model.Name }}
type {{ .Model.Name }}QueryBuilder struct {
selects *{{ .Model.Name }}Select
omits *{{ .Model.Name }}Omit
where []Predicate
take *int
skip *int
orderBy []OrderBy
}

func (b *{{ .Model.Name }}QueryBuilder) Where(preds ...Predicate) *{{ .Model.Name }}QueryBuilder {
b.where = append(b.where, preds...)
return b
}

func (b *{{ .Model.Name }}QueryBuilder) Take(limit int) *{{ .Model.Name }}QueryBuilder {
b.take = &limit
return b
}

func (b *{{ .Model.Name }}QueryBuilder) Skip(offset int) *{{ .Model.Name }}QueryBuilder {
b.skip = &offset
return b
}

func (b *{{ .Model.Name }}QueryBuilder) OrderBy(orders ...OrderBy) *{{ .Model.Name }}QueryBuilder {
b.orderBy = append(b.orderBy, orders...)
return b
}

func (b *{{ .Model.Name }}QueryBuilder) Select(s {{ .Model.Name }}Select) *{{ .Model.Name }}QueryBuilder {
b.selects = &s
return b
}

func (b *{{ .Model.Name }}QueryBuilder) Omit(o {{ .Model.Name }}Omit) *{{ .Model.Name }}QueryBuilder {
b.omits = &o
return b
}

func (b *{{ .Model.Name }}QueryBuilder) GetRelationParams() (*{{ .Model.Name }}Select, *{{ .Model.Name }}Omit, QueryParams) {
if b == nil {
return nil, nil, QueryParams{}
}
return b.selects, b.omits, QueryParams{
Where: b.where,
Take: b.take,
Skip: b.skip,
OrderBy: b.orderBy,
}
}

type {{ .Model.Name }}Delegate struct {
Expand Down
19 changes: 16 additions & 3 deletions generator/templates/runtime.gotpl
Original file line number Diff line number Diff line change
Expand Up @@ -180,8 +180,21 @@ func ValidateInt(errs *ValidationError, fieldName string, val int, rule string)
}
}

type OrderDirection string

const (
Asc OrderDirection = "ASC"
Desc OrderDirection = "DESC"
)

type OrderBy struct {
Field string
Direction OrderDirection
}

type QueryParams struct {
Where []Predicate
Take *int
Skip *int
Where []Predicate
Take *int
Skip *int
OrderBy []OrderBy
}
Loading
Loading