From 8b3ab646c119adb16ff3499de5bfbbaf8c8209fc Mon Sep 17 00:00:00 2001 From: Clancy Date: Fri, 31 Jul 2026 19:57:31 +0300 Subject: [PATCH 1/5] feat(migration): add scalar list (array) and JSON check constraint support in migration dialects 1- Add support for scalar lists (arrays) in postgresDialect and sqliteDialect to generate accurate SQL types (text[], TEXT) and array default values ('{}', '[]') during migrations 2- Add buildJsonArrayCheckConstraints() to convertToAtlasSchema to generate json_valid() check constraints for SQLite JSON array columns 3- Update prepareSchema.js to convert scalar array fields to Json with /// @elementType doc comments for SQLite --- integration/prepareSchema.js | 10 ++++++++++ migration/convertToAtlasSchema.go | 23 ++++++++++++++++++++++- migration/postgresDialect.go | 24 +++++++++++++++++++----- migration/sqliteDialect.go | 11 +++++++++++ 4 files changed, 62 insertions(+), 6 deletions(-) diff --git a/integration/prepareSchema.js b/integration/prepareSchema.js index 1821451..fa26a9b 100644 --- a/integration/prepareSchema.js +++ b/integration/prepareSchema.js @@ -45,6 +45,16 @@ function prepare(mode) { // Strip any @db.something attributes currentLine = currentLine.replace(/@db\.[A-Za-z0-9_]+(?:\([^)]*\))?/g, ''); + // Convert scalar array types to Json for SQLite with doc comment preserving element type + const arrayMatch = currentLine.match(/^(\s*)([a-zA-Z_]\w*)\s+(String|Int|BigInt|Float|Decimal|Boolean|DateTime|Bytes)\[\]/); + if (arrayMatch) { + const indent = arrayMatch[1]; + const elemType = arrayMatch[3]; + out.push(indent + '/// @elementType ' + elemType); + currentLine = currentLine.replace(/\b(String|Int|BigInt|Float|Decimal|Boolean|DateTime|Bytes)\[\]/g, 'Json'); + currentLine = currentLine.replace(/@default\(\[\]\)/g, '@default("[]")'); + } + out.push(currentLine); } diff --git a/migration/convertToAtlasSchema.go b/migration/convertToAtlasSchema.go index 792411b..e69f392 100644 --- a/migration/convertToAtlasSchema.go +++ b/migration/convertToAtlasSchema.go @@ -95,6 +95,7 @@ func (b *atlasSchemaBuilder) buildTablesMap() error { b.buildUniqueConstraints(model, table) b.buildIndexes(model, table) b.buildEnumCheckConstraints(model, table) + b.buildJsonArrayCheckConstraints(model, table) b.targetSchema.Tables = append(b.targetSchema.Tables, table) b.tablesMap[tableName] = table @@ -171,7 +172,11 @@ func (b *atlasSchemaBuilder) buildColumns(model *vs.Model, table *schema.Table) defaultVal = getSQLDefault(sf.Default, sf.Type, b.provider) } if defaultVal != "" { - column.Default = &schema.RawExpr{X: defaultVal} + if b.provider == providers.Sqlite && strings.HasPrefix(defaultVal, "'") { + column.Default = &schema.Literal{V: defaultVal} + } else { + column.Default = &schema.RawExpr{X: defaultVal} + } } } @@ -297,6 +302,22 @@ func (b *atlasSchemaBuilder) buildEnumCheckConstraints(model *vs.Model, table *s } } +func (b *atlasSchemaBuilder) buildJsonArrayCheckConstraints(model *vs.Model, table *schema.Table) { + if b.provider != providers.Sqlite || b.dialect == nil { + return + } + for _, sf := range model.ScalarFields { + if sf.IsArray { + colName := sf.EffectiveColName() + checkConstraint := &schema.Check{ + Name: table.Name + "_" + colName + "_check", + Expr: "json_valid(" + b.dialect.QuoteIdent(colName) + ")", + } + table.Attrs = append(table.Attrs, checkConstraint) + } + } +} + func (b *atlasSchemaBuilder) buildForeignKeys() error { for _, model := range b.schemaDef.Models { tableName := model.EffectiveTableName() diff --git a/migration/postgresDialect.go b/migration/postgresDialect.go index 2c3d85d..427a21c 100644 --- a/migration/postgresDialect.go +++ b/migration/postgresDialect.go @@ -17,22 +17,36 @@ func (d PostgresDialect) QuoteIdent(name string) string { } func (d PostgresDialect) GetSQLType(sf *schema.ScalarField) string { - // If its a custom PG enum, return the quoted enum name + var baseType string if sf.EnumRef != nil { enumName := sf.EnumRef.Name if sf.EnumRef.TableMapName != "" { enumName = sf.EnumRef.TableMapName } - return d.QuoteIdent(enumName) + baseType = d.QuoteIdent(enumName) + } else if sf.NativeType != nil && len(sf.NativeType.Args) > 0 { + baseType = fmt.Sprintf("%s(%s)", strings.ToLower(sf.SQLType), strings.Join(sf.NativeType.Args, ", ")) + } else { + baseType = strings.ToLower(sf.SQLType) } - if sf.NativeType != nil && len(sf.NativeType.Args) > 0 { - return fmt.Sprintf("%s(%s)", strings.ToLower(sf.SQLType), strings.Join(sf.NativeType.Args, ", ")) + if sf.IsArray { + return baseType + "[]" } - return strings.ToLower(sf.SQLType) + return baseType } func (d PostgresDialect) GetSQLDefault(dv *schema.DefaultValue, pslType string) string { switch dv.Kind { + case schema.DefaultArray: + if len(dv.ArrayValues) == 0 { + return "'{}'" + } + var escaped []string + for _, val := range dv.ArrayValues { + escaped = append(escaped, fmt.Sprintf("%q", val)) + } + return fmt.Sprintf("'{%s}'", strings.Join(escaped, ",")) + case schema.DefaultLiteral: if pslType == schema.TypeBoolean { return strings.ToUpper(dv.Literal) diff --git a/migration/sqliteDialect.go b/migration/sqliteDialect.go index 95368eb..af15d1f 100644 --- a/migration/sqliteDialect.go +++ b/migration/sqliteDialect.go @@ -2,6 +2,7 @@ package migration import ( "database/sql" + "encoding/json" "fmt" "strings" @@ -18,6 +19,9 @@ func (SqliteDialect) QuoteIdent(name string) string { } func (SqliteDialect) GetSQLType(sf *schema.ScalarField) string { + if sf.IsArray { + return "TEXT" + } sqlType := strings.ToUpper(sf.SQLType) switch sqlType { case "VARCHAR", "TEXT", "UUID": @@ -37,6 +41,13 @@ func (SqliteDialect) GetSQLType(sf *schema.ScalarField) string { func (SqliteDialect) GetSQLDefault(dv *schema.DefaultValue, pslType string) string { switch dv.Kind { + case schema.DefaultArray: + if len(dv.ArrayValues) == 0 { + return "'[]'" + } + bytes, _ := json.Marshal(dv.ArrayValues) + return fmt.Sprintf("'%s'", string(bytes)) + case schema.DefaultLiteral: if pslType == schema.TypeBoolean { return strings.ToUpper(dv.Literal) From 61ac0c80264101eab6f45da2f17e9ced1ebef4ac Mon Sep 17 00:00:00 2001 From: Clancy Date: Fri, 31 Jul 2026 20:00:44 +0300 Subject: [PATCH 2/5] refactor(predicates): scope nullability operators to optional fields and enforce compile-time safety 1- Remove IsNull() and IsNotNull() from non-optional base predicate types (Field, UniqueField, StringField, StringUniqueField, ArrayField) so required fields cannot be queried for nullability at compile time 2- Add OptionalField, OptionalUniqueField, OptionalStringField, and OptionalStringUniqueField struct wrappers that embed base predicate types and exclusively expose IsNull() and IsNotNull() for optional schema fields --- generator/templates/client.gotpl | 306 +++++++++++++++++++++- generator/templates/model_predicate.gotpl | 33 ++- 2 files changed, 323 insertions(+), 16 deletions(-) diff --git a/generator/templates/client.gotpl b/generator/templates/client.gotpl index 23c41a5..5ae496d 100644 --- a/generator/templates/client.gotpl +++ b/generator/templates/client.gotpl @@ -9,11 +9,17 @@ func newDialect() Dialect { SupportsLimitMinusOne: false, SupportsBulkInsert: true, SupportsDefaultKeyword: true, + RequiresConflictTarget: true, ConflictKeyword: "ON CONFLICT", ConflictIgnore: "DO NOTHING", ConflictUpdate: "DO UPDATE SET", ConflictExcluded: "EXCLUDED.", ConflictExcludedEnd: "", + ILikeFmt: "{col} ILIKE {val}", + ArrayHasFmt: "{val} = ANY({col})", + ArrayHasEveryFmt: "{col} @> ARRAY[{vals}]", + ArrayHasSomeFmt: "{col} && ARRAY[{vals}]", + ArrayIsEmptyFmt: "(cardinality({col}) = 0 OR {col} IS NULL)", } } {{- else if or (eq .Schema.Datasource.Provider "sqlite") (eq .Schema.Datasource.Provider "sqlite3") }} @@ -27,11 +33,17 @@ func newDialect() Dialect { SupportsLimitMinusOne: true, SupportsBulkInsert: true, SupportsDefaultKeyword: false, + RequiresConflictTarget: true, ConflictKeyword: "ON CONFLICT", ConflictIgnore: "DO NOTHING", ConflictUpdate: "DO UPDATE SET", ConflictExcluded: "EXCLUDED.", ConflictExcludedEnd: "", + ILikeFmt: "LOWER({col}) LIKE LOWER({val})", + ArrayHasFmt: "{val} IN (SELECT value FROM json_each({col}))", + ArrayHasEveryFmt: "", + ArrayHasSomeFmt: "", + ArrayIsEmptyFmt: "({col} IS NULL OR json_array_length({col}) = 0)", } } {{- else }} @@ -45,11 +57,17 @@ func newDialect() Dialect { SupportsLimitMinusOne: false, SupportsBulkInsert: false, SupportsDefaultKeyword: false, + RequiresConflictTarget: false, ConflictKeyword: "", ConflictIgnore: "", ConflictUpdate: "", ConflictExcluded: "", ConflictExcludedEnd: "", + ILikeFmt: "LOWER({col}) LIKE LOWER({val})", + ArrayHasFmt: "JSON_CONTAINS({col}, {val})", + ArrayHasEveryFmt: "", + ArrayHasSomeFmt: "", + ArrayIsEmptyFmt: "({col} IS NULL OR JSON_LENGTH({col}) = 0)", } } {{- end }} @@ -563,7 +581,31 @@ func (f Field[M, T]) In(vals []T) Predicate[M] { } } -func (f Field[M, T]) IsNull() Predicate[M] { +func (f Field[M, T]) NotIn(vals []T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "NOT IN", + Value: vals, + }, + } +} + +func (f Field[M, T]) Between(min T, max T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "BETWEEN", + Value: []any{min, max}, + }, + } +} + +type OptionalField[M any, T any] struct { + Field[M, T] +} + +func (f OptionalField[M, T]) IsNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -572,7 +614,7 @@ func (f Field[M, T]) IsNull() Predicate[M] { } } -func (f Field[M, T]) IsNotNull() Predicate[M] { +func (f OptionalField[M, T]) IsNotNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -671,7 +713,31 @@ func (f UniqueField[M, T]) In(vals []T) Predicate[M] { } } -func (f UniqueField[M, T]) IsNull() Predicate[M] { +func (f UniqueField[M, T]) NotIn(vals []T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "NOT IN", + Value: vals, + }, + } +} + +func (f UniqueField[M, T]) Between(min T, max T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "BETWEEN", + Value: []any{min, max}, + }, + } +} + +type OptionalUniqueField[M any, T any] struct { + UniqueField[M, T] +} + +func (f OptionalUniqueField[M, T]) IsNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -680,7 +746,7 @@ func (f UniqueField[M, T]) IsNull() Predicate[M] { } } -func (f UniqueField[M, T]) IsNotNull() Predicate[M] { +func (f OptionalUniqueField[M, T]) IsNotNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -775,6 +841,26 @@ func (f StringField[M]) In(vals []string) Predicate[M] { } } +func (f StringField[M]) NotIn(vals []string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "NOT IN", + Value: vals, + }, + } +} + +func (f StringField[M]) Between(min string, max string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "BETWEEN", + Value: []any{min, max}, + }, + } +} + func (f StringField[M]) Like(val string) Predicate[M] { return Predicate[M]{ Data: PredicateData{ @@ -785,6 +871,36 @@ func (f StringField[M]) Like(val string) Predicate[M] { } } +func (f StringField[M]) ILike(val string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "ILIKE", + Value: val, + }, + } +} + +func (f StringField[M]) HasPrefix(prefix string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "LIKE", + Value: prefix + "%", + }, + } +} + +func (f StringField[M]) HasSuffix(suffix string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "LIKE", + Value: "%" + suffix, + }, + } +} + func (f StringField[M]) Contains(val string) Predicate[M] { return Predicate[M]{ Data: PredicateData{ @@ -795,7 +911,11 @@ func (f StringField[M]) Contains(val string) Predicate[M] { } } -func (f StringField[M]) IsNull() Predicate[M] { +type OptionalStringField[M any] struct { + StringField[M] +} + +func (f OptionalStringField[M]) IsNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -804,7 +924,7 @@ func (f StringField[M]) IsNull() Predicate[M] { } } -func (f StringField[M]) IsNotNull() Predicate[M] { +func (f OptionalStringField[M]) IsNotNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -903,6 +1023,26 @@ func (f StringUniqueField[M]) In(vals []string) Predicate[M] { } } +func (f StringUniqueField[M]) NotIn(vals []string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "NOT IN", + Value: vals, + }, + } +} + +func (f StringUniqueField[M]) Between(min string, max string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "BETWEEN", + Value: []any{min, max}, + }, + } +} + func (f StringUniqueField[M]) Like(val string) Predicate[M] { return Predicate[M]{ Data: PredicateData{ @@ -913,6 +1053,36 @@ func (f StringUniqueField[M]) Like(val string) Predicate[M] { } } +func (f StringUniqueField[M]) ILike(val string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "ILIKE", + Value: val, + }, + } +} + +func (f StringUniqueField[M]) HasPrefix(prefix string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "LIKE", + Value: prefix + "%", + }, + } +} + +func (f StringUniqueField[M]) HasSuffix(suffix string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "LIKE", + Value: "%" + suffix, + }, + } +} + func (f StringUniqueField[M]) Contains(val string) Predicate[M] { return Predicate[M]{ Data: PredicateData{ @@ -923,7 +1093,11 @@ func (f StringUniqueField[M]) Contains(val string) Predicate[M] { } } -func (f StringUniqueField[M]) IsNull() Predicate[M] { +type OptionalStringUniqueField[M any] struct { + StringUniqueField[M] +} + +func (f OptionalStringUniqueField[M]) IsNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -932,7 +1106,7 @@ func (f StringUniqueField[M]) IsNull() Predicate[M] { } } -func (f StringUniqueField[M]) IsNotNull() Predicate[M] { +func (f OptionalStringUniqueField[M]) IsNotNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -949,6 +1123,55 @@ func (f StringUniqueField[M]) Desc() OrderBy[M] { return OrderBy[M]{Field: f.Column, Direction: Desc} } +type ArrayField[M any, T any] struct { + Column string +} + +func (f ArrayField[M, T]) Set(vals []T) FieldAssignmentOf[M] { + return FieldAssignmentOf[M]{Col: f.Column, Val: vals} +} + +func (f ArrayField[M, T]) Has(val T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "ARRAY_HAS", + Value: val, + }, + } +} + +func (f ArrayField[M, T]) HasEvery(vals []T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "ARRAY_HAS_EVERY", + Value: vals, + }, + } +} + +func (f ArrayField[M, T]) HasSome(vals []T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "ARRAY_HAS_SOME", + Value: vals, + }, + } +} + +func (f ArrayField[M, T]) IsEmpty() Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "ARRAY_IS_EMPTY", + }, + } +} + + + func CompilePredicates[M any](dialect Dialect, preds []PredicateOf[M], startBindIdx ...int) (string, []any, int) { bindIdx := 1 if len(startBindIdx) > 0 && startBindIdx[0] > 0 { @@ -1023,6 +1246,71 @@ func CompilePredicateData(dialect Dialect, data []PredicateData, startBindIdx .. args = append(args, val) } return fmt.Sprintf("%s IN (%s)", dialect.Quote(p.Column), strings.Join(placeHolders, ", ")) + case "NOT IN": + valSlice := unpackSlice(p.Value) + if len(valSlice) == 0 { + return "1=1" + } + var placeHolders []string + for range valSlice { + placeHolders = append(placeHolders, dialect.BindVar(bindIdx)) + bindIdx++ + } + for _, val := range valSlice { + args = append(args, val) + } + return fmt.Sprintf("%s NOT IN (%s)", dialect.Quote(p.Column), strings.Join(placeHolders, ", ")) + case "BETWEEN": + valSlice := unpackSlice(p.Value) + if len(valSlice) < 2 { + return "" + } + p1 := dialect.BindVar(bindIdx) + bindIdx++ + p2 := dialect.BindVar(bindIdx) + bindIdx++ + args = append(args, valSlice[0], valSlice[1]) + return fmt.Sprintf("%s BETWEEN %s AND %s", dialect.Quote(p.Column), p1, p2) + case "ILIKE": + placeholder := dialect.BindVar(bindIdx) + bindIdx++ + args = append(args, p.Value) + return dialect.FormatILike(p.Column, placeholder) + case "ARRAY_HAS": + placeholder := dialect.BindVar(bindIdx) + bindIdx++ + args = append(args, p.Value) + return dialect.FormatArrayHas(p.Column, placeholder) + case "ARRAY_HAS_EVERY": + valSlice := unpackSlice(p.Value) + if len(valSlice) == 0 { + return "1=1" + } + var placeHolders []string + for range valSlice { + placeHolders = append(placeHolders, dialect.BindVar(bindIdx)) + bindIdx++ + } + for _, val := range valSlice { + args = append(args, val) + } + return dialect.FormatArrayHasEvery(p.Column, placeHolders) + case "ARRAY_HAS_SOME": + valSlice := unpackSlice(p.Value) + if len(valSlice) == 0 { + return "1=0" + } + var placeHolders []string + for range valSlice { + placeHolders = append(placeHolders, dialect.BindVar(bindIdx)) + bindIdx++ + } + for _, val := range valSlice { + args = append(args, val) + } + return dialect.FormatArrayHasSome(p.Column, placeHolders) + case "ARRAY_IS_EMPTY": + return dialect.FormatArrayIsEmpty(p.Column) default: placeholder := dialect.BindVar(bindIdx) bindIdx++ @@ -1091,7 +1379,7 @@ func unpackSlice(val any) []any { res[i] = x } return res - {{- if hasTimeAnywhere .Schema }} + {{- if hasType .Schema "DateTime" }} case []time.Time: res := make([]any, len(v)) for i, x := range v { diff --git a/generator/templates/model_predicate.gotpl b/generator/templates/model_predicate.gotpl index c1e9d00..1d55c2b 100644 --- a/generator/templates/model_predicate.gotpl +++ b/generator/templates/model_predicate.gotpl @@ -2,10 +2,10 @@ package {{ .PackageName }} import ( "context" - {{- if hasJsonField .Model }} + {{- if hasModelType .Model "Json" }} "encoding/json" {{- end }} - {{- if hasTimeField .Model }} + {{- if hasModelType .Model "DateTime" }} "time" {{- end }} "{{ .ParentImportPath }}" @@ -46,17 +46,36 @@ func Not(pred {{ .ParentPackageName }}.PredicateOf[{{ .ParentPackageName }}.{{ . {{- $isUnique := or $field.IsID $field.IsUnique -}} {{- $fieldType := fieldPredType $field $.ParentPackageName -}} {{- $col := $field.EffectiveColName -}} -{{- if eq $field.Type "String" }} - {{- if $isUnique }} -var {{ capitalize $field.Name }} = {{ $.ParentPackageName }}.StringUniqueField[{{ $.ParentPackageName }}.{{ $.Model.Name }}]{Column: "{{ $col }}"} +{{- if $field.IsArray }} + {{- $elemType := trimPrefix $fieldType "[]" }} +var {{ capitalize $field.Name }} = {{ $.ParentPackageName }}.ArrayField[{{ $.ParentPackageName }}.{{ $.Model.Name }}, {{ $elemType }}]{Column: "{{ $col }}"} +{{- else if eq $field.Type "String" }} + {{- if $field.Optional }} + {{- if $isUnique }} +var {{ capitalize $field.Name }} = {{ $.ParentPackageName }}.OptionalStringUniqueField[{{ $.ParentPackageName }}.{{ $.Model.Name }}]{StringUniqueField: {{ $.ParentPackageName }}.StringUniqueField[{{ $.ParentPackageName }}.{{ $.Model.Name }}]{Column: "{{ $col }}"}} + {{- else }} +var {{ capitalize $field.Name }} = {{ $.ParentPackageName }}.OptionalStringField[{{ $.ParentPackageName }}.{{ $.Model.Name }}]{StringField: {{ $.ParentPackageName }}.StringField[{{ $.ParentPackageName }}.{{ $.Model.Name }}]{Column: "{{ $col }}"}} + {{- end }} {{- else }} + {{- if $isUnique }} +var {{ capitalize $field.Name }} = {{ $.ParentPackageName }}.StringUniqueField[{{ $.ParentPackageName }}.{{ $.Model.Name }}]{Column: "{{ $col }}"} + {{- else }} var {{ capitalize $field.Name }} = {{ $.ParentPackageName }}.StringField[{{ $.ParentPackageName }}.{{ $.Model.Name }}]{Column: "{{ $col }}"} + {{- end }} {{- end }} {{- else }} - {{- if $isUnique }} -var {{ capitalize $field.Name }} = {{ $.ParentPackageName }}.UniqueField[{{ $.ParentPackageName }}.{{ $.Model.Name }}, {{ $fieldType }}]{Column: "{{ $col }}"} + {{- if $field.Optional }} + {{- if $isUnique }} +var {{ capitalize $field.Name }} = {{ $.ParentPackageName }}.OptionalUniqueField[{{ $.ParentPackageName }}.{{ $.Model.Name }}, {{ $fieldType }}]{UniqueField: {{ $.ParentPackageName }}.UniqueField[{{ $.ParentPackageName }}.{{ $.Model.Name }}, {{ $fieldType }}]{Column: "{{ $col }}"}} + {{- else }} +var {{ capitalize $field.Name }} = {{ $.ParentPackageName }}.OptionalField[{{ $.ParentPackageName }}.{{ $.Model.Name }}, {{ $fieldType }}]{Field: {{ $.ParentPackageName }}.Field[{{ $.ParentPackageName }}.{{ $.Model.Name }}, {{ $fieldType }}]{Column: "{{ $col }}"}} + {{- end }} {{- else }} + {{- if $isUnique }} +var {{ capitalize $field.Name }} = {{ $.ParentPackageName }}.UniqueField[{{ $.ParentPackageName }}.{{ $.Model.Name }}, {{ $fieldType }}]{Column: "{{ $col }}"} + {{- else }} var {{ capitalize $field.Name }} = {{ $.ParentPackageName }}.Field[{{ $.ParentPackageName }}.{{ $.Model.Name }}, {{ $fieldType }}]{Column: "{{ $col }}"} + {{- end }} {{- end }} {{- end }} {{ end }} From 7c016d86c0e896abc664261ebd4eede58312636a Mon Sep 17 00:00:00 2001 From: Clancy Date: Fri, 31 Jul 2026 20:17:56 +0300 Subject: [PATCH 3/5] refactor(generator & runtime): optimize array serialization, dialect abstraction, and @updatedAt handling 1- Make lib/pq and database/sql/driver imports conditional in header.gotpl and runtime.gotpl, ensuring lib/pq is only imported when provider is postgres AND array fields exist 2- Update ArrayVal and ArrayScan to use pure Go json.Marshal/json.Unmarshal for SQLite and other providers while preserving pq.Array for PostgreSQL 3- Add RequiresConflictTarget boolean configuration to Dialect struct and refactor BuildConflictClause to remove string-based dialect comparisons 4- Consolidate repetitive type checking helpers in generator/helpers.go into generic hasType and hasModelType helpers, and replace imperative Need... default loops in generator.go with hasDefaultFunc 5- Fix pointer vs value assignments for @updatedAt timestamp fields during model creation and updates --- generator/generator.go | 54 ++-------- generator/helpers.go | 114 +++++++++++---------- generator/templates/header.gotpl | 47 +++++---- generator/templates/model_create.gotpl | 41 ++++++-- generator/templates/model_header.gotpl | 11 +- generator/templates/model_structs.gotpl | 27 +++-- generator/templates/runtime.gotpl | 130 ++++++++++++++++++++++-- 7 files changed, 281 insertions(+), 143 deletions(-) diff --git a/generator/generator.go b/generator/generator.go index e7e0ac8..62f44f7 100644 --- a/generator/generator.go +++ b/generator/generator.go @@ -23,12 +23,6 @@ type templateData struct { DefaultDiskPath string Schema schema.Schema DefaultLogs []string - NeedCUID bool - NeedCUID2 bool - NeedUUID bool - NeedUUID7 bool - NeedULID bool - NeedNanoID bool } type modelTemplateData struct { @@ -36,6 +30,7 @@ type modelTemplateData struct { Model *schema.Model ParentImportPath string ParentPackageName string + Schema schema.Schema } func ResolveImportPath(clientDir string) (string, error) { @@ -111,18 +106,14 @@ func GenerateClient(sch schema.Schema, pkgName string, parentImportPath string, "fieldPredType": fieldPredType, "hasLog": hasLog, "hasAnyLog": hasAnyLog, - "hasJsonField": hasJsonField, - "hasJsonAnywhere": hasJsonAnywhere, - "hasTimeField": hasTimeField, - "hasTimeAnywhere": hasTimeAnywhere, + "hasType": hasType, + "hasModelType": hasModelType, "trimPrefix": strings.TrimPrefix, "isKnownDefaultFunc": isKnownDefaultFunc, "defaultFuncCall": defaultFuncCall, - "hasHstoreAnywhere": hasHstoreAnywhere, - "hasNetAnywhere": hasNetAnywhere, - "hasUuidAnywhere": hasUuidAnywhere, - "hasFloatAnywhere": hasFloatAnywhere, - "hasDecimalAnywhere": hasDecimalAnywhere, + "isPostgresProvider": isPostgresProvider, + "needsPQImport": needsPQImport, + "hasDefaultFunc": hasDefaultFunc, "hstoreExpr": hstoreExpr, }) tmpl, err := tmpl.ParseFS(templatesFS, "templates/*.gotpl") @@ -130,31 +121,6 @@ func GenerateClient(sch schema.Schema, pkgName string, parentImportPath string, return nil, err } - var needCUID, needUUID, needUUID7, needCUID2, needULID, needNanoID bool - for _, m := range sch.Models { - for _, sf := range m.ScalarFields { - if sf.Default != nil && sf.Default.Kind == schema.DefaultFunc { - switch sf.Default.FuncName { - case "cuid", "cuid(1)": - needCUID = true - case "cuid(2)": - needCUID2 = true - case "uuid", "uuid(4)": - needUUID = true - case "uuid(7)": - needUUID7 = true - case "ulid": - needULID = true - case "nanoid": - needNanoID = true - } - } - if sf.IsID && sf.GoType == "string" && sf.Default == nil { - needCUID = true - } - } - } - var embedDir string if embedPath != "" { embedDir = filepath.ToSlash(filepath.Dir(embedPath)) @@ -167,12 +133,6 @@ func GenerateClient(sch schema.Schema, pkgName string, parentImportPath string, DefaultDiskPath: defaultDiskPath, Schema: sch, DefaultLogs: defaultLogs, - NeedCUID: needCUID, - NeedCUID2: needCUID2, - NeedUUID: needUUID, - NeedUUID7: needUUID7, - NeedULID: needULID, - NeedNanoID: needNanoID, } outputs := make(map[string]string) @@ -210,6 +170,7 @@ func GenerateClient(sch schema.Schema, pkgName string, parentImportPath string, Model: m, ParentImportPath: parentImportPath, ParentPackageName: pkgName, + Schema: sch, } if err := tmpl.ExecuteTemplate(&mBuf, "model_header.gotpl", mData); err != nil { @@ -244,6 +205,7 @@ func GenerateClient(sch schema.Schema, pkgName string, parentImportPath string, Model: m, ParentImportPath: parentImportPath, ParentPackageName: pkgName, + Schema: sch, } if err := tmpl.ExecuteTemplate(&pBuf, "model_predicate.gotpl", pData); err != nil { return nil, err diff --git a/generator/helpers.go b/generator/helpers.go index 8beffb8..1b0321e 100644 --- a/generator/helpers.go +++ b/generator/helpers.go @@ -4,6 +4,7 @@ import ( "slices" "strings" + providers "github.com/voidclancy/valk/dbProviders" "github.com/voidclancy/valk/schema" ) @@ -66,26 +67,45 @@ func fieldPredType(f *schema.ScalarField, parentPkg string) string { return t } -func hasFieldWhere(m *schema.Model, pred func(*schema.ScalarField) bool) bool { - return slices.ContainsFunc(m.ScalarFields, pred) +func hasModelType(m *schema.Model, targetTypes ...string) bool { + for _, sf := range m.ScalarFields { + for _, t := range targetTypes { + switch t { + case "Array": + if sf.IsArray { + return true + } + case "Hstore": + if strings.Contains(sf.GoType, "map[string]*string") { + return true + } + case "DateTime", "Time": + if sf.Type == "DateTime" || strings.Contains(sf.GoType, "time.Time") { + return true + } + case "Json": + if (sf.Type == "Json" || strings.Contains(sf.GoType, "json.RawMessage")) && !sf.IsArray { + return true + } + default: + if sf.Type == t || (sf.NativeType != nil && sf.NativeType.Name == t) { + return true + } + } + } + } + return false } -func hasJsonField(m *schema.Model) bool { - return hasFieldWhere(m, func(sf *schema.ScalarField) bool { - return sf.Type == "Json" || strings.Contains(sf.GoType, "json.RawMessage") - }) -} -func hasJsonAnywhere(sch schema.Schema) bool { - return slices.ContainsFunc(sch.Models, hasJsonField) -} -func hasTimeField(m *schema.Model) bool { - return hasFieldWhere(m, func(sf *schema.ScalarField) bool { - return sf.Type == "DateTime" || strings.Contains(sf.GoType, "time.Time") - }) -} -func hasTimeAnywhere(sch schema.Schema) bool { - return slices.ContainsFunc(sch.Models, hasTimeField) +func hasType(sch schema.Schema, targetTypes ...string) bool { + for _, m := range sch.Models { + if hasModelType(m, targetTypes...) { + return true + } + } + return false } + func isKnownDefaultFunc(funcName string) bool { val, ok := DEFAULT_FUNCS[funcName] return ok && val != "" @@ -94,45 +114,29 @@ func isKnownDefaultFunc(funcName string) bool { func defaultFuncCall(funcName string) string { return DEFAULT_FUNCS[funcName] } -func hasUuidField(m *schema.Model) bool { - return hasFieldWhere(m, func(sf *schema.ScalarField) bool { - return sf.NativeType != nil && sf.NativeType.Name == "Uuid" - }) -} -func hasUuidAnywhere(sch schema.Schema) bool { - return slices.ContainsFunc(sch.Models, hasUuidField) -} -func hasFloatField(m *schema.Model) bool { - return hasFieldWhere(m, func(sf *schema.ScalarField) bool { - return sf.Type == "Float" - }) -} -func hasFloatAnywhere(sch schema.Schema) bool { - return slices.ContainsFunc(sch.Models, hasFloatField) -} -func hasDecimalField(m *schema.Model) bool { - return hasFieldWhere(m, func(sf *schema.ScalarField) bool { - return sf.Type == "Decimal" - }) -} -func hasDecimalAnywhere(sch schema.Schema) bool { - return slices.ContainsFunc(sch.Models, hasDecimalField) -} -func hasNetField(m *schema.Model) bool { - return hasFieldWhere(m, func(sf *schema.ScalarField) bool { - return sf.NativeType != nil && sf.NativeType.Name == "Inet" - }) -} -func hasNetAnywhere(sch schema.Schema) bool { - return slices.ContainsFunc(sch.Models, hasNetField) -} -func hasHstoreField(m *schema.Model) bool { - return hasFieldWhere(m, func(sf *schema.ScalarField) bool { - return strings.TrimPrefix(sf.GoType, "*") == "map[string]*string" - }) + +func isPostgresProvider(sch schema.Schema) bool { + p := sch.Datasource.Provider + return p == providers.Postgres || p == providers.Postgresql } -func hasHstoreAnywhere(sch schema.Schema) bool { - return slices.ContainsFunc(sch.Models, hasHstoreField) + +func needsPQImport(sch schema.Schema) bool { + return isPostgresProvider(sch) && hasType(sch, "Array") +} +func hasDefaultFunc(sch schema.Schema, names ...string) bool { + for _, m := range sch.Models { + for _, sf := range m.ScalarFields { + if sf.Default != nil && sf.Default.Kind == schema.DefaultFunc { + if slices.Contains(names, sf.Default.FuncName) { + return true + } + } + if sf.IsID && sf.GoType == "string" && sf.Default == nil && slices.Contains(names, "cuid") { + return true + } + } + } + return false } func hstoreExpr(goType string, expr string) string { if strings.TrimPrefix(goType, "*") == "map[string]*string" { diff --git a/generator/templates/header.gotpl b/generator/templates/header.gotpl index b0dce29..01f27f2 100644 --- a/generator/templates/header.gotpl +++ b/generator/templates/header.gotpl @@ -2,13 +2,16 @@ package {{ .PackageName }} import ( "context" - {{- if hasFloatAnywhere .Schema }} + {{- if hasType .Schema "Float" }} "math" {{- end }} - {{- if or .NeedCUID .NeedCUID2 .NeedULID .NeedNanoID }} + {{- if hasDefaultFunc .Schema "cuid" "cuid(1)" "cuid(2)" "ulid" "nanoid" }} "crypto/rand" {{- end }} "database/sql" + {{- if hasType .Schema "Array" }} + "database/sql/driver" + {{- end }} "encoding/json" {{- if .EmbedPath }} "embed" @@ -17,50 +20,56 @@ import ( {{- if hasAnyLog }} "log" {{- end }} - {{- if hasNetAnywhere .Schema }} + {{- if hasType .Schema "Inet" }} "net" {{- end }} - {{- if or (hasUuidAnywhere .Schema) (hasDecimalAnywhere .Schema) }} + {{- if hasType .Schema "Uuid" "Decimal" }} "regexp" {{- end }} "strconv" "slices" "strings" "sync" - {{- if or (hasTimeAnywhere .Schema) .NeedCUID .NeedCUID2 .NeedULID }} + {{- if or (hasType .Schema "DateTime") (hasDefaultFunc .Schema "cuid" "cuid(1)" "cuid(2)" "ulid") }} "time" {{- end }} "unicode/utf8" - {{- if hasHstoreAnywhere .Schema }} + {{- if hasType .Schema "Hstore" }} "github.com/lib/pq/hstore" {{- end }} + {{- if needsPQImport .Schema }} + "github.com/lib/pq" + {{- end }} "github.com/pressly/goose/v3" - {{- if or .NeedUUID .NeedUUID7 }} + {{- if hasDefaultFunc .Schema "uuid" "uuid(4)" "uuid(7)" }} "github.com/google/uuid" {{- end }} ) -{{- if or (hasTimeAnywhere .Schema) .NeedCUID .NeedCUID2 .NeedULID }} +{{- if or (hasType .Schema "DateTime") (hasDefaultFunc .Schema "cuid" "cuid(1)" "cuid(2)" "ulid") }} var _ = time.Time{} {{- end }} -{{- if hasHstoreAnywhere .Schema }} +{{- if hasType .Schema "Hstore" }} var _ = hstore.Hstore{} {{- end }} -{{- if hasNetAnywhere .Schema }} +{{- if needsPQImport .Schema }} +var _ = pq.Array +{{- end }} +{{- if hasType .Schema "Inet" }} var _ = net.ParseIP {{- end }} var _ = json.RawMessage{} var _ = strings.Join var _ = slices.Clone[[]any] -{{- if or .NeedUUID .NeedUUID7 }} +{{- if hasDefaultFunc .Schema "uuid" "uuid(4)" "uuid(7)" }} var _ = uuid.New {{- end }} -{{- if or .NeedCUID .NeedCUID2 .NeedULID .NeedNanoID }} +{{- if hasDefaultFunc .Schema "cuid" "cuid(1)" "cuid(2)" "ulid" "nanoid" }} var _ = rand.Read {{- end }} -{{- if or .NeedCUID .NeedCUID2 .NeedULID }} +{{- if hasDefaultFunc .Schema "cuid" "cuid(1)" "cuid(2)" "ulid" }} var _ = strconv.AppendUint {{- end }} @@ -69,7 +78,7 @@ var _ = strconv.AppendUint var migrationsFS embed.FS {{- end }} -{{- if .NeedCUID }} +{{- if hasDefaultFunc .Schema "cuid" "cuid(1)" }} func generateCUID() string { now := uint64(time.Now().UnixMilli()) b := make([]byte, 8) @@ -86,13 +95,13 @@ func generateCUID() string { } {{- end }} -{{- if .NeedUUID }} +{{- if hasDefaultFunc .Schema "uuid" "uuid(4)" }} func generateUUID() string { return uuid.New().String() } {{- end }} -{{- if .NeedUUID7 }} +{{- if hasDefaultFunc .Schema "uuid(7)" }} func generateUUID7() string { id, err := uuid.NewV7() if err != nil { @@ -102,7 +111,7 @@ func generateUUID7() string { } {{- end }} -{{- if .NeedCUID2 }} +{{- if hasDefaultFunc .Schema "cuid(2)" }} func generateCUID2() string { now := uint64(time.Now().UnixMilli()) b := make([]byte, 12) @@ -118,7 +127,7 @@ func generateCUID2() string { } {{- end }} -{{- if .NeedULID }} +{{- if hasDefaultFunc .Schema "ulid" }} func generateULID() string { now := uint64(time.Now().UnixMilli()) b := make([]byte, 10) @@ -148,7 +157,7 @@ func generateULID() string { } {{- end }} -{{- if .NeedNanoID }} +{{- if hasDefaultFunc .Schema "nanoid" }} func generateNanoID() string { const alphabet = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZ_abcdefghijklmnopqrstuvwxyz-" b := make([]byte, 21) diff --git a/generator/templates/model_create.gotpl b/generator/templates/model_create.gotpl index dbba952..bffa2a8 100644 --- a/generator/templates/model_create.gotpl +++ b/generator/templates/model_create.gotpl @@ -80,7 +80,7 @@ func assignmentsTo{{ .Model.Name }}Create(assignments []FieldAssignment) ({{ .Mo {{- range $field := .Model.ScalarFields }} {{- $col := $field.EffectiveColName }} {{- $coreType := trimPrefix $field.GoType "*" }} - {{- $isPointer := or (ne $field.Default nil) $field.Optional (ne $coreType $field.GoType) }} + {{- $isPointer := and (or (ne $field.Default nil) $field.Optional (ne $coreType $field.GoType)) (not $field.IsArray) }} case "{{ $col }}": provided |= provided{{ $.Model.Name }}{{ capitalize $field.Name }} {{- if $field.EnumRef }} @@ -208,7 +208,9 @@ func assignmentsTo{{ .Model.Name }}Create(assignments []FieldAssignment) ({{ .Mo {{- range $field := .Model.ScalarFields }} {{- $fieldName := capitalize $field.Name }} - {{- if and (eq $field.Default nil) (not $field.Optional) (not $field.IsArray) }} + {{- $coreType := trimPrefix $field.GoType "*" }} + {{- $isPointer := and (or (ne $field.Default nil) $field.Optional (ne $coreType $field.GoType)) (not $field.IsArray) }} + {{- if and (eq $field.Default nil) (not $field.Optional) (not $field.IsArray) (not $field.IsUpdatedAt) }} if provided&provided{{ $.Model.Name }}{{ $fieldName }} == 0 { {{- if or $field.EnumRef (ne $field.GoType "string") }} errs.Add("{{ $field.EffectiveColName }}", nil, "required", "field {{ $fieldName }} is required") @@ -217,6 +219,16 @@ func assignmentsTo{{ .Model.Name }}Create(assignments []FieldAssignment) ({{ .Mo {{- end }} } {{- end }} + {{- if $field.IsUpdatedAt }} + if provided&provided{{ $.Model.Name }}{{ $fieldName }} == 0 { + {{- if $isPointer }} + now := time.Now().Truncate(time.Microsecond) + input.{{ $fieldName }} = &now + {{- else }} + input.{{ $fieldName }} = time.Now().Truncate(time.Microsecond) + {{- end }} + } + {{- end }} {{- end }} if errs.HasErrors() { @@ -232,12 +244,12 @@ func (s *{{ .Model.Name }}Create) ToColsVals() (cols []string, vals []any) { {{- $col := $field.EffectiveColName }} {{- $fieldName := capitalize $field.Name }} {{- $coreType := trimPrefix $field.GoType "*" }} - {{- $isPointer := or (ne $field.Default nil) $field.Optional (ne $coreType $field.GoType) }} + {{- $isPointer := and (or (ne $field.Default nil) $field.Optional (ne $coreType $field.GoType)) (not $field.IsArray) }} {{- if $field.EnumRef }} {{- if $field.IsArray }} if s.{{ $fieldName }} != nil { cols = append(cols, "{{ $col }}") - vals = append(vals, s.{{ $fieldName }}) + vals = append(vals, ArrayVal(s.{{ $fieldName }})) } {{- else }} if s.{{ $fieldName }} != nil { @@ -248,7 +260,7 @@ func (s *{{ .Model.Name }}Create) ToColsVals() (cols []string, vals []any) { {{- else if $field.IsArray }} if s.{{ $fieldName }} != nil { cols = append(cols, "{{ $col }}") - vals = append(vals, s.{{ $fieldName }}) + vals = append(vals, ArrayVal(s.{{ $fieldName }})) } {{- else }} {{- if and $field.Default (eq $field.Default.Kind.String "Func") }} @@ -274,6 +286,13 @@ func (s *{{ .Model.Name }}Create) ToColsVals() (cols []string, vals []any) { if s.{{ $fieldName }} != nil { cols = append(cols, "{{ $col }}") vals = append(vals, {{ hstoreExpr $field.GoType (printf "*s.%s" $fieldName) }}) + } + {{- else if $field.IsUpdatedAt }} + cols = append(cols, "{{ $col }}") + if !s.{{ $fieldName }}.IsZero() { + vals = append(vals, {{ hstoreExpr $field.GoType (printf "s.%s" $fieldName) }}) + } else { + vals = append(vals, time.Now().Truncate(time.Microsecond)) } {{- else if $field.Optional }} if s.{{ $fieldName }} != nil { @@ -720,7 +739,7 @@ func (d *{{ .Model.Name }}Delegate) buildBulkInsertSQL(q *Queries, batch []*{{ . switch col { {{- range $field := .Model.ScalarFields }} {{- $coreType := trimPrefix $field.GoType "*" }} - {{- $isPointer := or (ne $field.Default nil) $field.Optional (ne $coreType $field.GoType) }} + {{- $isPointer := and (or (ne $field.Default nil) $field.Optional (ne $coreType $field.GoType)) (not $field.IsArray) }} case "{{ $field.EffectiveColName }}": {{- if and $field.IsID (eq $field.GoType "string") }} {{- if $isPointer }} @@ -742,10 +761,16 @@ func (d *{{ .Model.Name }}Delegate) buildBulkInsertSQL(q *Queries, batch []*{{ . } else { vals = append(vals, {{ defaultFuncCall $field.Default.FuncName }}) } + {{- else if $field.IsUpdatedAt }} + if !input.{{ capitalize $field.Name }}.IsZero() { + vals = append(vals, {{ hstoreExpr $field.GoType (printf "input.%s" (capitalize $field.Name)) }}) + } else { + vals = append(vals, time.Now().Truncate(time.Microsecond)) + } {{- else if $field.EnumRef }} {{- if $field.IsArray }} if input.{{ capitalize $field.Name }} != nil { - vals = append(vals, input.{{ capitalize $field.Name }}) + vals = append(vals, ArrayVal(input.{{ capitalize $field.Name }})) } else { writeDefault = true } @@ -760,7 +785,7 @@ func (d *{{ .Model.Name }}Delegate) buildBulkInsertSQL(q *Queries, batch []*{{ . {{- end }} {{- else if $field.IsArray }} if input.{{ capitalize $field.Name }} != nil { - vals = append(vals, input.{{ capitalize $field.Name }}) + vals = append(vals, ArrayVal(input.{{ capitalize $field.Name }})) } else { writeDefault = true } diff --git a/generator/templates/model_header.gotpl b/generator/templates/model_header.gotpl index 30b1e36..bfd95c9 100644 --- a/generator/templates/model_header.gotpl +++ b/generator/templates/model_header.gotpl @@ -3,17 +3,24 @@ package {{ .PackageName }} import ( "context" "database/sql" - {{- if hasJsonField .Model }} + {{- if hasModelType .Model "Json" }} "encoding/json" {{- end }} "fmt" "slices" "strings" - {{- if hasTimeField .Model }} + {{- if hasModelType .Model "DateTime" }} "time" {{- end }} + {{- if and (hasModelType .Model "Array") (isPostgresProvider .Schema) }} + "github.com/lib/pq" + {{- end }} {{- if ne .PackageName .ParentPackageName }} "{{ .ParentImportPath }}" {{- end }} ) +{{- if and (hasModelType .Model "Array") (isPostgresProvider .Schema) }} +var _ = pq.Array +{{- end }} + diff --git a/generator/templates/model_structs.gotpl b/generator/templates/model_structs.gotpl index f436e52..def7a58 100644 --- a/generator/templates/model_structs.gotpl +++ b/generator/templates/model_structs.gotpl @@ -25,7 +25,7 @@ type {{ .Model.Name }} struct { type {{ .Model.Name }}Create struct { {{- range $field := .Model.ScalarFields }} {{- $coreType := trimPrefix $field.GoType "*" }} - {{- $isPointer := or (ne $field.Default nil) $field.Optional (ne $coreType $field.GoType) }} + {{- $isPointer := and (or (ne $field.Default nil) $field.Optional (ne $coreType $field.GoType)) (not $field.IsArray) }} {{ capitalize $field.Name }} {{ if $field.EnumRef }}{{ if $field.IsArray }}[]{{ $field.EnumRef.Name }}Type{{ else }}{{ if $isPointer }}*{{ end }}{{ $field.EnumRef.Name }}Type{{ end }}{{ else }}{{ if $field.IsArray }}{{ $field.GoType }}{{ else }}{{ if $isPointer }}*{{ end }}{{ $coreType }}{{ end }}{{ end }} `json:"{{ $field.Name }}"` {{- end }} } @@ -34,7 +34,7 @@ func (s *{{ .Model.Name }}Create) colMask() uint64 { var mask uint64 {{- range $i, $field := .Model.ScalarFields }} {{- $coreType := trimPrefix $field.GoType "*" }} - {{- $isClientDefault := or (and $field.IsID (eq $field.GoType "string")) (and $field.Default (eq $field.Default.Kind.String "Func") ($field.Default.FuncName | isKnownDefaultFunc)) }} + {{- $isClientDefault := or (and $field.IsID (eq $field.GoType "string")) (and $field.Default (eq $field.Default.Kind.String "Func") ($field.Default.FuncName | isKnownDefaultFunc)) $field.IsUpdatedAt }} {{- $isPointer := and (or (ne $field.Default nil) $field.Optional (ne $coreType $field.GoType)) (not $isClientDefault) }} {{- if $isClientDefault }} mask |= 1 << {{ $i }} @@ -196,12 +196,16 @@ func assignmentsTo{{ .Model.Name }}Update(assignments []FieldAssignment) ({{ .Mo } {{- else }} if v, ok := a.Val.({{ $coreType }}); ok { - input.{{ capitalize $field.Name }} = &v - } else if v, ok := a.Val.(*{{ $coreType }}); ok { + {{- if $field.IsArray }} input.{{ capitalize $field.Name }} = v - } else if v, ok := a.Val.({{ $field.GoType }}); ok { - {{- if eq (trimPrefix $field.GoType "*") $field.GoType }} + {{- else }} input.{{ capitalize $field.Name }} = &v + {{- end }} + } else if v, ok := a.Val.(*{{ $coreType }}); ok { + {{- if $field.IsArray }} + if v != nil { + input.{{ capitalize $field.Name }} = *v + } {{- else }} input.{{ capitalize $field.Name }} = v {{- end }} @@ -213,6 +217,15 @@ func assignmentsTo{{ .Model.Name }}Update(assignments []FieldAssignment) ({{ .Mo } } + {{- range $field := .Model.ScalarFields }} + {{- if $field.IsUpdatedAt }} + if input.{{ capitalize $field.Name }} == nil { + now := time.Now().Truncate(time.Microsecond) + input.{{ capitalize $field.Name }} = &now + } + {{- end }} + {{- end }} + if errs.HasErrors() { return input, errs } @@ -668,6 +681,8 @@ func (m *{{ .Model.Name }}) ScanFields(cols []string) []any { case "{{ $field.EffectiveColName }}": {{- if eq (trimPrefix $field.GoType "*") "map[string]*string" }} targets[i] = HstoreScan{P: &m.{{ capitalize $field.Name }}} + {{- else if $field.IsArray }} + targets[i] = ArrayScan(&m.{{ capitalize $field.Name }}) {{- else }} targets[i] = &m.{{ capitalize $field.Name }} {{- end }} diff --git a/generator/templates/runtime.gotpl b/generator/templates/runtime.gotpl index 9c7f937..20a5b64 100644 --- a/generator/templates/runtime.gotpl +++ b/generator/templates/runtime.gotpl @@ -104,11 +104,17 @@ type Dialect struct { SupportsLimitMinusOne bool SupportsBulkInsert bool SupportsDefaultKeyword bool + RequiresConflictTarget bool ConflictKeyword string // "ON CONFLICT" (Pg/SQLite) or "ON DUPLICATE KEY" (MySQL) ConflictIgnore string // "DO NOTHING" or "" ConflictUpdate string // "DO UPDATE SET" or "UPDATE" ConflictExcluded string // "EXCLUDED." or "VALUES(" ConflictExcludedEnd string // "" or ")" + ILikeFmt string // "{col} ILIKE {val}" or "LOWER({col}) LIKE LOWER({val})" + ArrayHasFmt string // "{val} = ANY({col})" or "JSON_CONTAINS({col}, {val})" + ArrayHasEveryFmt string // "{col} @> ARRAY[{vals}]" (empty = compose from ArrayHasFmt) + ArrayHasSomeFmt string // "{col} && ARRAY[{vals}]" (empty = compose from ArrayHasFmt) + ArrayIsEmptyFmt string // "(cardinality({col}) = 0 OR {col} IS NULL)" or "({col} IS NULL OR json_array_length({col}) = 0)" } func (d Dialect) WriteQuotedIdent(sb *strings.Builder, ident string) { @@ -160,6 +166,116 @@ func (d Dialect) FormatLimitOffset(take *int, skip *int) string { return "" } +func (d Dialect) FormatILike(col string, placeholder string) string { + r := strings.NewReplacer("{col}", d.Quote(col), "{val}", placeholder) + return r.Replace(d.ILikeFmt) +} + +func (d Dialect) FormatArrayHas(col string, placeholder string) string { + r := strings.NewReplacer("{col}", d.Quote(col), "{val}", placeholder) + return r.Replace(d.ArrayHasFmt) +} + +func (d Dialect) FormatArrayHasEvery(col string, placeholders []string) string { + if d.ArrayHasEveryFmt != "" { + r := strings.NewReplacer("{col}", d.Quote(col), "{vals}", strings.Join(placeholders, ", ")) + return r.Replace(d.ArrayHasEveryFmt) + } + var parts []string + for _, p := range placeholders { + parts = append(parts, d.FormatArrayHas(col, p)) + } + if len(parts) == 0 { + return "1=1" + } + return "(" + strings.Join(parts, " AND ") + ")" +} + +func (d Dialect) FormatArrayHasSome(col string, placeholders []string) string { + if d.ArrayHasSomeFmt != "" { + r := strings.NewReplacer("{col}", d.Quote(col), "{vals}", strings.Join(placeholders, ", ")) + return r.Replace(d.ArrayHasSomeFmt) + } + var parts []string + for _, p := range placeholders { + parts = append(parts, d.FormatArrayHas(col, p)) + } + if len(parts) == 0 { + return "1=0" + } + return "(" + strings.Join(parts, " OR ") + ")" +} + +func (d Dialect) FormatArrayIsEmpty(col string) string { + r := strings.NewReplacer("{col}", d.Quote(col)) + return r.Replace(d.ArrayIsEmptyFmt) +} + +{{- if hasType .Schema "Array" }} +func ArrayVal[T any](v []T) any { + return ArrayValWrapper[T]{V: v} +} + +type ArrayValWrapper[T any] struct { + V []T +} + +func (a ArrayValWrapper[T]) Value() (driver.Value, error) { + if a.V == nil { + return nil, nil + } + {{- if isPostgresProvider .Schema }} + return pq.Array(a.V).Value() + {{- else }} + b, err := json.Marshal(a.V) + if err != nil { + return nil, err + } + return string(b), nil + {{- end }} +} + +func ArrayScan[T any](p *[]T) any { + return ArrayScanWrapper[T]{P: p} +} + +type ArrayScanWrapper[T any] struct { + P *[]T +} + +func (a ArrayScanWrapper[T]) Scan(src any) error { + if src == nil { + *a.P = nil + return nil + } + {{- if isPostgresProvider .Schema }} + switch v := src.(type) { + case []byte: + if len(v) > 0 && v[0] == '{' { + return pq.Array(a.P).Scan(v) + } + return json.Unmarshal(v, a.P) + case string: + if len(v) > 0 && v[0] == '{' { + return pq.Array(a.P).Scan(v) + } + return json.Unmarshal([]byte(v), a.P) + default: + return pq.Array(a.P).Scan(src) + } + {{- else }} + switch v := src.(type) { + case []byte: + return json.Unmarshal(v, a.P) + case string: + return json.Unmarshal([]byte(v), a.P) + default: + return fmt.Errorf("cannot scan %T into array", src) + } + {{- end }} +} +{{- end }} + type rawDefault struct{} type DBTX interface { @@ -221,7 +337,7 @@ type RecordInput struct { Assignments []FieldAssignment } -{{- if hasHstoreAnywhere .Schema }} +{{- if hasType .Schema "Hstore" }} func ToHstore(m map[string]*string) hstore.Hstore { result := hstore.Hstore{Map: make(map[string]sql.NullString, len(m))} for k, v := range m { @@ -279,7 +395,7 @@ func (errs *ValidationError) ValidateString(fieldName string, val string, isRequ errs.Add(fieldName, val, "format", "bit string must contain only '0' and '1'") } } - {{- if hasNetAnywhere .Schema }} + {{- if hasType .Schema "Inet" }} if isInet { if net.ParseIP(val) == nil { if _, _, err := net.ParseCIDR(val); err != nil { @@ -341,7 +457,7 @@ func (errs *ValidationError) ValidateInt(fieldName string, val int, rule string) } } -{{- if hasUuidAnywhere .Schema }} +{{- if hasType .Schema "Uuid" }} var uuidRegex = regexp.MustCompile(`^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$`) func (errs *ValidationError) ValidateUUID(fieldName string, val string) { @@ -351,14 +467,14 @@ func (errs *ValidationError) ValidateUUID(fieldName string, val string) { } {{- end }} -{{- if hasFloatAnywhere .Schema }} +{{- if hasType .Schema "Float" }} func (errs *ValidationError) ValidateFloat(fieldName string, val float64) { if math.IsInf(val, 0) || math.IsNaN(val) { errs.Add(fieldName, val, "range", "field must be a finite number") } } {{- end }} -{{- if hasDecimalAnywhere .Schema }} +{{- if hasType .Schema "Decimal" }} var decimalRegex = regexp.MustCompile(`^[-+]?(\d+(\.\d*)?|\.\d+)$`) func (errs *ValidationError) ValidateDecimal(fieldName string, val string, scale int) { @@ -377,7 +493,7 @@ func (errs *ValidationError) ValidateDecimal(fieldName string, val string, scale } } {{- end }} -{{- if hasJsonAnywhere .Schema }} +{{- if hasType .Schema "Json" }} func (errs *ValidationError) ValidateJson(fieldName string, val json.RawMessage) { if val == nil { return @@ -430,7 +546,7 @@ func (d Dialect) BuildConflictClause(conflictCols []string, action *ConflictActi return "", nil } var colsStr string - if len(conflictCols) > 0 && d.ConflictKeyword == "ON CONFLICT" { + if len(conflictCols) > 0 && d.RequiresConflictTarget { var quoted []string for _, col := range conflictCols { quoted = append(quoted, d.Quote(col)) From 3da807e116c6287eff6f6b5666cd5215f0e8e13a Mon Sep 17 00:00:00 2001 From: Clancy Date: Fri, 31 Jul 2026 20:20:47 +0300 Subject: [PATCH 4/5] test(integration): add predicate operators test suite and update schema for @updatedAt and array fields 1- Add integration test suite in integration/predicate_operators_test.go verifying IsNull/IsNotNull predicate scoping and Has/HasEvery/HasSome/IsEmpty array operators 2- Update integration/schema.prisma and migrations with updatedAt, updatedAtNoDefualt, updatedAtnoDecorator, updatedAtOptional, and tags fields to verify default, required, optional, and @updatedAt decorator behavior --- integration/main.go | 22 +- integration/predicate_operators_test.go | 232 +++++++++++++++++++++ integration/schema.prisma | 33 +-- integration/valk/allFieldsSoFar.go | 26 ++- integration/valk/migrations/00001_init.sql | 6 + 5 files changed, 289 insertions(+), 30 deletions(-) create mode 100644 integration/predicate_operators_test.go diff --git a/integration/main.go b/integration/main.go index 93e3d04..73efef9 100644 --- a/integration/main.go +++ b/integration/main.go @@ -55,9 +55,10 @@ func main() { if err != nil { log.Fatal("Error loading .env file") } - db := openConn() - // db := openPGConn() - // defer dbReset(db) + // db := openConn() + db := openPGConn() + dbReset(db) + defer dbReset(db) defer db.Close() rawDB := db.Raw() rawDB.SetMaxOpenConns(10) @@ -69,9 +70,20 @@ func main() { // runBlockBasedTransaction(db, ctx) // runPaginationExamples(db, ctx) // runExtensionExamples(db, ctx) - runCTP(db, ctx) - db.User.FindMany().Exec(ctx) + // runCTP(db, ctx) + // + + usr, err := db.User.Create().SetEmail("xx@yy.com").SetPhoneNum("11").Exec(ctx) + if err != nil { + panic(err) + } + db.Post.Create().SetTitle("Post 1").SetAuthorId(usr.Id).SetTags([]string{"aa", "bb"}).Exec(ctx) + posts, err := db.Post.FindMany(post.Tags.HasEvery([]string{"aa", "bb"})).Exec(ctx) + if err != nil { + panic(err) + } + printJSON(posts) } // ============================================================================= diff --git a/integration/predicate_operators_test.go b/integration/predicate_operators_test.go new file mode 100644 index 0000000..10d2109 --- /dev/null +++ b/integration/predicate_operators_test.go @@ -0,0 +1,232 @@ +package main + +import ( + "context" + + "testing" + + "integration/valk" + "integration/valk/post" + "integration/valk/user" + + _ "github.com/lib/pq" + _ "github.com/mattn/go-sqlite3" +) + +func seedUsers(t *testing.T, ctx context.Context, db *valk.DB) ([]*valk.User, error) { + t.Helper() + + u1, err := db.User.Create(). + SetEmail("alpha@example.com"). + SetPhoneNum("1111111111"). + SetLoginCount(10). + Exec(ctx) + if err != nil { + return nil, err + } + + u2, err := db.User.Create(). + SetEmail("beta@domain.org"). + SetPhoneNum("2222222222"). + SetLoginCount(25). + Exec(ctx) + if err != nil { + return nil, err + } + + u3, err := db.User.Create(). + SetEmail("gamma@example.com"). + SetPhoneNum("3333333333"). + SetLoginCount(50). + Exec(ctx) + if err != nil { + return nil, err + } + + return []*valk.User{u1, u2, u3}, nil +} +func TestPredicateOperators(t *testing.T) { + db, cleanup := setupTestDB(t) + defer cleanup() + ctx := context.Background() + + users, err := seedUsers(t, ctx, db) + if err != nil { + t.Fatalf("failed to seed users: %v", err) + } + + u1, u2, u3 := users[0], users[1], users[2] + t.Run("NotIn", func(t *testing.T) { + res, err := db.User.FindMany(user.Email.NotIn([]string{u1.Email, u2.Email})).Exec(ctx) + if err != nil { + t.Fatalf("NotIn failed: %v", err) + } + if len(res) != 1 || res[0].Id != u3.Id { + t.Errorf("expected u3, got %v", res) + } + }) + + t.Run("Between", func(t *testing.T) { + res, err := db.User.FindMany(user.LoginCount.Between(20, 60)).Exec(ctx) + if err != nil { + t.Fatalf("Between failed: %v", err) + } + if len(res) != 2 { + t.Errorf("expected 2 records, got %d", len(res)) + } + }) + + t.Run("HasPrefix", func(t *testing.T) { + res, err := db.User.FindMany(user.Email.HasPrefix("alp")).Exec(ctx) + if err != nil { + t.Fatalf("HasPrefix failed: %v", err) + } + if len(res) != 1 || res[0].Id != u1.Id { + t.Errorf("expected u1, got %v", res) + } + }) + + t.Run("HasSuffix", func(t *testing.T) { + res, err := db.User.FindMany(user.Email.HasSuffix("@example.com")).Exec(ctx) + if err != nil { + t.Fatalf("HasSuffix failed: %v", err) + } + if len(res) != 2 { + t.Errorf("expected 2 records, got %d", len(res)) + } + }) + + t.Run("ILike", func(t *testing.T) { + res, err := db.User.FindMany(user.Email.ILike("ALPHA@%")).Exec(ctx) + if err != nil { + t.Fatalf("ILike failed: %v", err) + } + if len(res) != 1 || res[0].Id != u1.Id { + t.Errorf("expected u1, got %v", res) + } + }) + + t.Run("Logical Composition And/Or/Not", func(t *testing.T) { + // (Email ends with @example.com AND loginCount = 50) OR (phoneNum = 2222222222) + comp := valk.Or( + valk.And( + user.Email.HasSuffix("@example.com"), + user.LoginCount.EQ(50), + ), + user.PhoneNum.EQ("2222222222"), + ) + res, err := db.User.FindMany(comp).Exec(ctx) + if err != nil { + t.Fatalf("Logical And/Or failed: %v", err) + } + if len(res) != 2 { + t.Errorf("expected 2 records, got %d", len(res)) + } + + // NOT (Email ends with @example.com) + notRes, err := db.User.FindMany(valk.Not(user.Email.HasSuffix("@example.com"))).Exec(ctx) + if err != nil { + t.Fatalf("Logical Not failed: %v", err) + } + if len(notRes) != 1 || notRes[0].Id != u2.Id { + t.Errorf("expected u2, got %v", notRes) + } + }) + + t.Run("IsNull and IsNotNull on optional field", func(t *testing.T) { + nullRes, err := db.User.FindMany(user.ReferredById.IsNull()).Exec(ctx) + if err != nil { + t.Fatalf("IsNull failed: %v", err) + } + if len(nullRes) != 3 { + t.Errorf("expected 3 users with null ReferredById, got %d", len(nullRes)) + } + + notNullRes, err := db.User.FindMany(user.ReferredById.IsNotNull()).Exec(ctx) + if err != nil { + t.Fatalf("IsNotNull failed: %v", err) + } + if len(notNullRes) != 0 { + t.Errorf("expected 0 users with non-null ReferredById, got %d", len(notNullRes)) + } + }) +} + +func TestArrayOperators(t *testing.T) { + db, cleanup := setupTestDB(t) + defer cleanup() + ctx := context.Background() + + users, err := seedUsers(t, ctx, db) + if err != nil { + t.Fatalf("seed failed: %v", err) + } + + p1, err := db.Post.Create(). + SetTitle("Go ORM"). + SetAuthorId(users[0].Id). + SetTags([]string{"golang", "orm", "database"}). + Exec(ctx) + if err != nil { + t.Fatalf("failed creating p1: %v", err) + } + + p2, err := db.Post.Create(). + SetTitle("Python Web"). + SetAuthorId(users[1].Id). + SetTags([]string{"python", "web"}). + Exec(ctx) + if err != nil { + t.Fatalf("failed creating p2: %v", err) + } + _ = p2 + + p3, err := db.Post.Create(). + SetTitle("Empty Post"). + SetAuthorId(users[2].Id). + SetTags([]string{}). + Exec(ctx) + if err != nil { + t.Fatalf("failed creating p3: %v", err) + } + + t.Run("Has", func(t *testing.T) { + res, err := db.Post.FindMany(post.Tags.Has("golang")).Exec(ctx) + if err != nil { + t.Fatalf("Has failed: %v", err) + } + if len(res) != 1 || res[0].Id != p1.Id { + t.Errorf("expected p1, got %v", res) + } + }) + + t.Run("HasEvery", func(t *testing.T) { + res, err := db.Post.FindMany(post.Tags.HasEvery([]string{"golang", "orm"})).Exec(ctx) + if err != nil { + t.Fatalf("HasEvery failed: %v", err) + } + if len(res) != 1 || res[0].Id != p1.Id { + t.Errorf("expected p1, got %v", res) + } + }) + + t.Run("HasSome", func(t *testing.T) { + res, err := db.Post.FindMany(post.Tags.HasSome([]string{"python", "golang"})).Exec(ctx) + if err != nil { + t.Fatalf("HasSome failed: %v", err) + } + if len(res) != 2 { + t.Errorf("expected 2 records, got %d", len(res)) + } + }) + + t.Run("IsEmpty", func(t *testing.T) { + res, err := db.Post.FindMany(post.Tags.IsEmpty()).Exec(ctx) + if err != nil { + t.Fatalf("IsEmpty failed: %v", err) + } + if len(res) != 1 || res[0].Id != p3.Id { + t.Errorf("expected p3, got %v", res) + } + }) +} diff --git a/integration/schema.prisma b/integration/schema.prisma index 917dfbf..6050d64 100644 --- a/integration/schema.prisma +++ b/integration/schema.prisma @@ -11,20 +11,24 @@ enum UserRole { } model User { - id String @id @default(cuid()) - email String @unique - phoneNum String @unique - password String? - role UserRole @default(STUDENT) - roleOptional UserRole? - profile Profile? - posts Post[] - comments Comment[] - loginCount Int @default(0) - - referredById String? - referredBy User? @relation("UserReferrals", fields: [referredById], references: [id]) - referrals User[] @relation("UserReferrals") + id String @id @default(cuid()) + email String @unique + phoneNum String @unique + password String? + role UserRole @default(STUDENT) + roleOptional UserRole? + profile Profile? + posts Post[] + comments Comment[] + loginCount Int @default(0) + referredById String? + referredBy User? @relation("UserReferrals", fields: [referredById], references: [id]) + referrals User[] @relation("UserReferrals") + createdAt DateTime @default(now()) + updatedAt DateTime @default(now()) @updatedAt + updatedAtNoDefualt DateTime @default(now()) + updatedAtnoDecorator DateTime @default(now()) + updatedAtOptional DateTime? @@unique([email, phoneNum], name: "emailPhone") } @@ -46,6 +50,7 @@ model Post { author User @relation(fields: [authorId], references: [id]) comments Comment[] categories CategoryToPost[] + tags String[] @default([]) } model Comment { diff --git a/integration/valk/allFieldsSoFar.go b/integration/valk/allFieldsSoFar.go index 995bb36..5c790be 100644 --- a/integration/valk/allFieldsSoFar.go +++ b/integration/valk/allFieldsSoFar.go @@ -961,8 +961,6 @@ func assignmentsToAllFieldsSoFarUpdate(assignments []FieldAssignment) (AllFields input.JsonReq = &v } else if v, ok := a.Val.(*json.RawMessage); ok { input.JsonReq = v - } else if v, ok := a.Val.(json.RawMessage); ok { - input.JsonReq = &v } else { errs.Add("jsonReq", a.Val, "type", "field jsonReq must be of type json.RawMessage") } @@ -971,8 +969,6 @@ func assignmentsToAllFieldsSoFarUpdate(assignments []FieldAssignment) (AllFields input.JsonOpt = &v } else if v, ok := a.Val.(*json.RawMessage); ok { input.JsonOpt = v - } else if v, ok := a.Val.(*json.RawMessage); ok { - input.JsonOpt = v } else { errs.Add("jsonOpt", a.Val, "type", "field jsonOpt must be of type *json.RawMessage") } @@ -981,8 +977,6 @@ func assignmentsToAllFieldsSoFarUpdate(assignments []FieldAssignment) (AllFields input.JsonVal = &v } else if v, ok := a.Val.(*json.RawMessage); ok { input.JsonVal = v - } else if v, ok := a.Val.(json.RawMessage); ok { - input.JsonVal = &v } else { errs.Add("jsonVal", a.Val, "type", "field jsonVal must be of type json.RawMessage") } @@ -1007,8 +1001,6 @@ func assignmentsToAllFieldsSoFarUpdate(assignments []FieldAssignment) (AllFields input.HstoreField = &v } else if v, ok := a.Val.(*map[string]*string); ok { input.HstoreField = v - } else if v, ok := a.Val.(*map[string]*string); ok { - input.HstoreField = v } else { errs.Add("hstoreField", a.Val, "type", "field hstoreField must be of type *map[string]*string") } @@ -1032,6 +1024,10 @@ func assignmentsToAllFieldsSoFarUpdate(assignments []FieldAssignment) (AllFields } } } + if input.UpdatedAt == nil { + now := time.Now().Truncate(time.Microsecond) + input.UpdatedAt = &now + } if errs.HasErrors() { return input, errs @@ -2884,7 +2880,7 @@ func assignmentsToAllFieldsSoFarCreate(assignments []FieldAssignment) (AllFields errs.Add("dateTimeReq", nil, "required", "field DateTimeReq is required") } if provided&providedAllFieldsSoFarUpdatedAt == 0 { - errs.Add("updatedAt", nil, "required", "field UpdatedAt is required") + input.UpdatedAt = time.Now().Truncate(time.Microsecond) } if provided&providedAllFieldsSoFarDateTimeTz == 0 { errs.Add("dateTimeTz", nil, "required", "field DateTimeTz is required") @@ -3058,7 +3054,11 @@ func (s *AllFieldsSoFarCreate) ToColsVals() (cols []string, vals []any) { vals = append(vals, time.Now()) } cols = append(cols, "updatedAt") - vals = append(vals, s.UpdatedAt) + if !s.UpdatedAt.IsZero() { + vals = append(vals, s.UpdatedAt) + } else { + vals = append(vals, time.Now().Truncate(time.Microsecond)) + } cols = append(cols, "dateTimeTz") vals = append(vals, s.DateTimeTz) cols = append(cols, "timestampVal") @@ -3674,7 +3674,11 @@ func (d *AllFieldsSoFarDelegate) buildBulkInsertSQL(q *Queries, batch []*AllFiel vals = append(vals, time.Now()) } case "updatedAt": - vals = append(vals, input.UpdatedAt) + if !input.UpdatedAt.IsZero() { + vals = append(vals, input.UpdatedAt) + } else { + vals = append(vals, time.Now().Truncate(time.Microsecond)) + } case "dateTimeTz": vals = append(vals, input.DateTimeTz) case "timestampVal": diff --git a/integration/valk/migrations/00001_init.sql b/integration/valk/migrations/00001_init.sql index abf1f59..d366533 100644 --- a/integration/valk/migrations/00001_init.sql +++ b/integration/valk/migrations/00001_init.sql @@ -13,6 +13,11 @@ CREATE TABLE "public"."User" ( "roleOptional" "public"."user_roles" NULL, "loginCount" integer NOT NULL DEFAULT 0, "referredById" text NULL, + "createdAt" timestamp NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updatedAt" timestamp NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updatedAtNoDefualt" timestamp NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updatedAtnoDecorator" timestamp NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updatedAtOptional" timestamp NULL, PRIMARY KEY ("id"), CONSTRAINT "User_referredById_fkey" FOREIGN KEY ("referredById") REFERENCES "public"."User" ("id") ON UPDATE NO ACTION ON DELETE NO ACTION ); @@ -91,6 +96,7 @@ CREATE TABLE "public"."Post" ( "content" text NULL, "published" boolean NOT NULL DEFAULT FALSE, "authorId" text NOT NULL, + "tags" text[] NOT NULL DEFAULT '{}', PRIMARY KEY ("id"), CONSTRAINT "Post_authorId_fkey" FOREIGN KEY ("authorId") REFERENCES "public"."User" ("id") ON UPDATE NO ACTION ON DELETE NO ACTION ); From 556a817410dd18f0812f0c36aede855db89ac514 Mon Sep 17 00:00:00 2001 From: Clancy Date: Fri, 31 Jul 2026 20:21:58 +0300 Subject: [PATCH 5/5] updated generated client after optimizing array serialization, dialect abstraction, and @updatedAt handling, adding scalar list and JSON check constraint support in migration dialects, and scoping nullability operators --- .../valk/allFieldsSoFar/allFieldsSoFar.go | 24 +- integration/valk/client.go | 433 +++++++++++- integration/valk/comment.go | 2 - integration/valk/comment/comment.go | 2 +- integration/valk/post.go | 143 +++- integration/valk/post/post.go | 4 +- integration/valk/profile/profile.go | 2 +- integration/valk/user.go | 619 ++++++++++++++---- integration/valk/user/user.go | 17 +- 9 files changed, 1022 insertions(+), 224 deletions(-) diff --git a/integration/valk/allFieldsSoFar/allFieldsSoFar.go b/integration/valk/allFieldsSoFar/allFieldsSoFar.go index 073a05f..7cf77f8 100644 --- a/integration/valk/allFieldsSoFar/allFieldsSoFar.go +++ b/integration/valk/allFieldsSoFar/allFieldsSoFar.go @@ -42,7 +42,7 @@ var Id = valk.UniqueField[valk.AllFieldsSoFar, int32]{Column: "id"} var StringReq = valk.StringField[valk.AllFieldsSoFar]{Column: "stringReq"} -var StringOpt = valk.StringField[valk.AllFieldsSoFar]{Column: "stringOpt"} +var StringOpt = valk.OptionalStringField[valk.AllFieldsSoFar]{StringField: valk.StringField[valk.AllFieldsSoFar]{Column: "stringOpt"}} var StringDefault = valk.StringField[valk.AllFieldsSoFar]{Column: "stringDefault"} @@ -78,7 +78,7 @@ var UuidDb = valk.StringField[valk.AllFieldsSoFar]{Column: "uuidDb"} var IntReq = valk.Field[valk.AllFieldsSoFar, int32]{Column: "intReq"} -var IntOpt = valk.Field[valk.AllFieldsSoFar, int32]{Column: "intOpt"} +var IntOpt = valk.OptionalField[valk.AllFieldsSoFar, int32]{Field: valk.Field[valk.AllFieldsSoFar, int32]{Column: "intOpt"}} var IntDefault = valk.Field[valk.AllFieldsSoFar, int32]{Column: "intDefault"} @@ -92,17 +92,17 @@ var OidVal = valk.Field[valk.AllFieldsSoFar, int32]{Column: "oidVal"} var BigIntReq = valk.Field[valk.AllFieldsSoFar, int64]{Column: "bigIntReq"} -var BigIntOpt = valk.Field[valk.AllFieldsSoFar, int64]{Column: "bigIntOpt"} +var BigIntOpt = valk.OptionalField[valk.AllFieldsSoFar, int64]{Field: valk.Field[valk.AllFieldsSoFar, int64]{Column: "bigIntOpt"}} var FloatReq = valk.Field[valk.AllFieldsSoFar, float64]{Column: "floatReq"} -var FloatOpt = valk.Field[valk.AllFieldsSoFar, float64]{Column: "floatOpt"} +var FloatOpt = valk.OptionalField[valk.AllFieldsSoFar, float64]{Field: valk.Field[valk.AllFieldsSoFar, float64]{Column: "floatOpt"}} var RealVal = valk.Field[valk.AllFieldsSoFar, float64]{Column: "realVal"} var DecimalReq = valk.Field[valk.AllFieldsSoFar, string]{Column: "decimalReq"} -var DecimalOpt = valk.Field[valk.AllFieldsSoFar, string]{Column: "decimalOpt"} +var DecimalOpt = valk.OptionalField[valk.AllFieldsSoFar, string]{Field: valk.Field[valk.AllFieldsSoFar, string]{Column: "decimalOpt"}} var DecimalPrecise = valk.Field[valk.AllFieldsSoFar, string]{Column: "decimalPrecise"} @@ -110,13 +110,13 @@ var MoneyVal = valk.Field[valk.AllFieldsSoFar, string]{Column: "moneyVal"} var BoolReq = valk.Field[valk.AllFieldsSoFar, bool]{Column: "boolReq"} -var BoolOpt = valk.Field[valk.AllFieldsSoFar, bool]{Column: "boolOpt"} +var BoolOpt = valk.OptionalField[valk.AllFieldsSoFar, bool]{Field: valk.Field[valk.AllFieldsSoFar, bool]{Column: "boolOpt"}} var BoolDefault = valk.Field[valk.AllFieldsSoFar, bool]{Column: "boolDefault"} var DateTimeReq = valk.Field[valk.AllFieldsSoFar, time.Time]{Column: "dateTimeReq"} -var DateTimeOpt = valk.Field[valk.AllFieldsSoFar, time.Time]{Column: "dateTimeOpt"} +var DateTimeOpt = valk.OptionalField[valk.AllFieldsSoFar, time.Time]{Field: valk.Field[valk.AllFieldsSoFar, time.Time]{Column: "dateTimeOpt"}} var DateTimeDefault = valk.Field[valk.AllFieldsSoFar, time.Time]{Column: "dateTimeDefault"} @@ -132,19 +132,19 @@ var TimetzVal = valk.Field[valk.AllFieldsSoFar, time.Time]{Column: "timetzVal"} var JsonReq = valk.Field[valk.AllFieldsSoFar, json.RawMessage]{Column: "jsonReq"} -var JsonOpt = valk.Field[valk.AllFieldsSoFar, json.RawMessage]{Column: "jsonOpt"} +var JsonOpt = valk.OptionalField[valk.AllFieldsSoFar, json.RawMessage]{Field: valk.Field[valk.AllFieldsSoFar, json.RawMessage]{Column: "jsonOpt"}} var JsonVal = valk.Field[valk.AllFieldsSoFar, json.RawMessage]{Column: "jsonVal"} var BytesReq = valk.Field[valk.AllFieldsSoFar, []byte]{Column: "bytesReq"} -var BytesOpt = valk.Field[valk.AllFieldsSoFar, []byte]{Column: "bytesOpt"} +var BytesOpt = valk.OptionalField[valk.AllFieldsSoFar, []byte]{Field: valk.Field[valk.AllFieldsSoFar, []byte]{Column: "bytesOpt"}} -var HstoreField = valk.Field[valk.AllFieldsSoFar, map[string]*string]{Column: "hstoreField"} +var HstoreField = valk.OptionalField[valk.AllFieldsSoFar, map[string]*string]{Field: valk.Field[valk.AllFieldsSoFar, map[string]*string]{Column: "hstoreField"}} -var LtreeField = valk.Field[valk.AllFieldsSoFar, string]{Column: "ltreeField"} +var LtreeField = valk.OptionalField[valk.AllFieldsSoFar, string]{Field: valk.Field[valk.AllFieldsSoFar, string]{Column: "ltreeField"}} -var CitextField = valk.Field[valk.AllFieldsSoFar, string]{Column: "citextField"} +var CitextField = valk.OptionalField[valk.AllFieldsSoFar, string]{Field: valk.Field[valk.AllFieldsSoFar, string]{Column: "citextField"}} type CreateInput = valk.AllFieldsSoFarCreate type Create = valk.AllFieldsSoFarCreate diff --git a/integration/valk/client.go b/integration/valk/client.go index a76a243..4a0f959 100644 --- a/integration/valk/client.go +++ b/integration/valk/client.go @@ -4,9 +4,14 @@ import ( "context" "crypto/rand" "database/sql" + "database/sql/driver" "embed" "encoding/json" "fmt" + "github.com/google/uuid" + "github.com/lib/pq" + "github.com/lib/pq/hstore" + "github.com/pressly/goose/v3" "math" "net" "regexp" @@ -16,14 +21,11 @@ import ( "sync" "time" "unicode/utf8" - - "github.com/google/uuid" - "github.com/lib/pq/hstore" - "github.com/pressly/goose/v3" ) var _ = time.Time{} var _ = hstore.Hstore{} +var _ = pq.Array var _ = net.ParseIP var _ = json.RawMessage{} var _ = strings.Join @@ -113,11 +115,11 @@ type UserRoleType string const ( // Admin maps to "ADMIN" - UserRoleTypeAdmin UserRoleType = "ADMIN" + userRoleTypeAdmin UserRoleType = "ADMIN" // Student maps to "student" - UserRoleTypeStudent UserRoleType = "student" + userRoleTypeStudent UserRoleType = "student" // Teacher maps to "TEACHER" - UserRoleTypeTeacher UserRoleType = "TEACHER" + userRoleTypeTeacher UserRoleType = "TEACHER" ) type userRoleNamespace struct { @@ -135,14 +137,14 @@ type userRoleNamespace struct { // STUDENT student // TEACHER TEACHER var UserRole = userRoleNamespace{ - Admin: UserRoleTypeAdmin, - Student: UserRoleTypeStudent, - Teacher: UserRoleTypeTeacher, + Admin: userRoleTypeAdmin, + Student: userRoleTypeStudent, + Teacher: userRoleTypeTeacher, } func (e UserRoleType) IsValid() bool { switch e { - case UserRoleTypeAdmin, UserRoleTypeStudent, UserRoleTypeTeacher: + case userRoleTypeAdmin, userRoleTypeStudent, userRoleTypeTeacher: return true } return false @@ -254,11 +256,17 @@ type Dialect struct { SupportsLimitMinusOne bool SupportsBulkInsert bool SupportsDefaultKeyword bool + RequiresConflictTarget bool ConflictKeyword string // "ON CONFLICT" (Pg/SQLite) or "ON DUPLICATE KEY" (MySQL) ConflictIgnore string // "DO NOTHING" or "" ConflictUpdate string // "DO UPDATE SET" or "UPDATE" ConflictExcluded string // "EXCLUDED." or "VALUES(" ConflictExcludedEnd string // "" or ")" + ILikeFmt string // "{col} ILIKE {val}" or "LOWER({col}) LIKE LOWER({val})" + ArrayHasFmt string // "{val} = ANY({col})" or "JSON_CONTAINS({col}, {val})" + ArrayHasEveryFmt string // "{col} @> ARRAY[{vals}]" (empty = compose from ArrayHasFmt) + ArrayHasSomeFmt string // "{col} && ARRAY[{vals}]" (empty = compose from ArrayHasFmt) + ArrayIsEmptyFmt string // "(cardinality({col}) = 0 OR {col} IS NULL)" or "({col} IS NULL OR json_array_length({col}) = 0)" } func (d Dialect) WriteQuotedIdent(sb *strings.Builder, ident string) { @@ -310,6 +318,94 @@ func (d Dialect) FormatLimitOffset(take *int, skip *int) string { return "" } +func (d Dialect) FormatILike(col string, placeholder string) string { + r := strings.NewReplacer("{col}", d.Quote(col), "{val}", placeholder) + return r.Replace(d.ILikeFmt) +} + +func (d Dialect) FormatArrayHas(col string, placeholder string) string { + r := strings.NewReplacer("{col}", d.Quote(col), "{val}", placeholder) + return r.Replace(d.ArrayHasFmt) +} + +func (d Dialect) FormatArrayHasEvery(col string, placeholders []string) string { + if d.ArrayHasEveryFmt != "" { + r := strings.NewReplacer("{col}", d.Quote(col), "{vals}", strings.Join(placeholders, ", ")) + return r.Replace(d.ArrayHasEveryFmt) + } + var parts []string + for _, p := range placeholders { + parts = append(parts, d.FormatArrayHas(col, p)) + } + if len(parts) == 0 { + return "1=1" + } + return "(" + strings.Join(parts, " AND ") + ")" +} + +func (d Dialect) FormatArrayHasSome(col string, placeholders []string) string { + if d.ArrayHasSomeFmt != "" { + r := strings.NewReplacer("{col}", d.Quote(col), "{vals}", strings.Join(placeholders, ", ")) + return r.Replace(d.ArrayHasSomeFmt) + } + var parts []string + for _, p := range placeholders { + parts = append(parts, d.FormatArrayHas(col, p)) + } + if len(parts) == 0 { + return "1=0" + } + return "(" + strings.Join(parts, " OR ") + ")" +} + +func (d Dialect) FormatArrayIsEmpty(col string) string { + r := strings.NewReplacer("{col}", d.Quote(col)) + return r.Replace(d.ArrayIsEmptyFmt) +} +func ArrayVal[T any](v []T) any { + return ArrayValWrapper[T]{V: v} +} + +type ArrayValWrapper[T any] struct { + V []T +} + +func (a ArrayValWrapper[T]) Value() (driver.Value, error) { + if a.V == nil { + return nil, nil + } + return pq.Array(a.V).Value() +} + +func ArrayScan[T any](p *[]T) any { + return ArrayScanWrapper[T]{P: p} +} + +type ArrayScanWrapper[T any] struct { + P *[]T +} + +func (a ArrayScanWrapper[T]) Scan(src any) error { + if src == nil { + *a.P = nil + return nil + } + switch v := src.(type) { + case []byte: + if len(v) > 0 && v[0] == '{' { + return pq.Array(a.P).Scan(v) + } + return json.Unmarshal(v, a.P) + case string: + if len(v) > 0 && v[0] == '{' { + return pq.Array(a.P).Scan(v) + } + return json.Unmarshal([]byte(v), a.P) + default: + return pq.Array(a.P).Scan(src) + } +} + type rawDefault struct{} type DBTX interface { @@ -568,7 +664,7 @@ func (d Dialect) BuildConflictClause(conflictCols []string, action *ConflictActi return "", nil } var colsStr string - if len(conflictCols) > 0 && d.ConflictKeyword == "ON CONFLICT" { + if len(conflictCols) > 0 && d.RequiresConflictTarget { var quoted []string for _, col := range conflictCols { quoted = append(quoted, d.Quote(col)) @@ -634,11 +730,17 @@ func newDialect() Dialect { SupportsLimitMinusOne: false, SupportsBulkInsert: true, SupportsDefaultKeyword: true, + RequiresConflictTarget: true, ConflictKeyword: "ON CONFLICT", ConflictIgnore: "DO NOTHING", ConflictUpdate: "DO UPDATE SET", ConflictExcluded: "EXCLUDED.", ConflictExcludedEnd: "", + ILikeFmt: "{col} ILIKE {val}", + ArrayHasFmt: "{val} = ANY({col})", + ArrayHasEveryFmt: "{col} @> ARRAY[{vals}]", + ArrayHasSomeFmt: "{col} && ARRAY[{vals}]", + ArrayIsEmptyFmt: "(cardinality({col}) = 0 OR {col} IS NULL)", } } @@ -651,14 +753,19 @@ type Queries struct { mu *sync.RWMutex // User provides CRUD operations for User. // - // id string default: cuid() - // email string required - // phoneNum string required - // password string optional - // role UserRole default: STUDENT - // roleOptional UserRole optional - // loginCount int32 default: 0 - // referredById string optional + // id string default: cuid() + // email string required + // phoneNum string required + // password string optional + // role UserRole default: STUDENT + // roleOptional UserRole optional + // loginCount int32 default: 0 + // referredById string optional + // createdAt time.Time default: now() + // updatedAt time.Time default: now() + // updatedAtNoDefualt time.Time default: now() + // updatedAtnoDecorator time.Time default: now() + // updatedAtOptional time.Time optional User *UserDelegate // Profile provides CRUD operations for Profile. // @@ -669,11 +776,12 @@ type Queries struct { Profile *ProfileDelegate // Post provides CRUD operations for Post. // - // id string default: cuid() - // title string required - // content string optional - // published bool default: false - // authorId string required + // id string default: cuid() + // title string required + // content string optional + // published bool default: false + // authorId string required + // tags []string default: [] Post *PostDelegate // Comment provides CRUD operations for Comment. // @@ -1197,7 +1305,31 @@ func (f Field[M, T]) In(vals []T) Predicate[M] { } } -func (f Field[M, T]) IsNull() Predicate[M] { +func (f Field[M, T]) NotIn(vals []T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "NOT IN", + Value: vals, + }, + } +} + +func (f Field[M, T]) Between(min T, max T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "BETWEEN", + Value: []any{min, max}, + }, + } +} + +type OptionalField[M any, T any] struct { + Field[M, T] +} + +func (f OptionalField[M, T]) IsNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -1206,7 +1338,7 @@ func (f Field[M, T]) IsNull() Predicate[M] { } } -func (f Field[M, T]) IsNotNull() Predicate[M] { +func (f OptionalField[M, T]) IsNotNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -1305,7 +1437,31 @@ func (f UniqueField[M, T]) In(vals []T) Predicate[M] { } } -func (f UniqueField[M, T]) IsNull() Predicate[M] { +func (f UniqueField[M, T]) NotIn(vals []T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "NOT IN", + Value: vals, + }, + } +} + +func (f UniqueField[M, T]) Between(min T, max T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "BETWEEN", + Value: []any{min, max}, + }, + } +} + +type OptionalUniqueField[M any, T any] struct { + UniqueField[M, T] +} + +func (f OptionalUniqueField[M, T]) IsNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -1314,7 +1470,7 @@ func (f UniqueField[M, T]) IsNull() Predicate[M] { } } -func (f UniqueField[M, T]) IsNotNull() Predicate[M] { +func (f OptionalUniqueField[M, T]) IsNotNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -1409,6 +1565,26 @@ func (f StringField[M]) In(vals []string) Predicate[M] { } } +func (f StringField[M]) NotIn(vals []string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "NOT IN", + Value: vals, + }, + } +} + +func (f StringField[M]) Between(min string, max string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "BETWEEN", + Value: []any{min, max}, + }, + } +} + func (f StringField[M]) Like(val string) Predicate[M] { return Predicate[M]{ Data: PredicateData{ @@ -1419,6 +1595,36 @@ func (f StringField[M]) Like(val string) Predicate[M] { } } +func (f StringField[M]) ILike(val string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "ILIKE", + Value: val, + }, + } +} + +func (f StringField[M]) HasPrefix(prefix string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "LIKE", + Value: prefix + "%", + }, + } +} + +func (f StringField[M]) HasSuffix(suffix string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "LIKE", + Value: "%" + suffix, + }, + } +} + func (f StringField[M]) Contains(val string) Predicate[M] { return Predicate[M]{ Data: PredicateData{ @@ -1429,7 +1635,11 @@ func (f StringField[M]) Contains(val string) Predicate[M] { } } -func (f StringField[M]) IsNull() Predicate[M] { +type OptionalStringField[M any] struct { + StringField[M] +} + +func (f OptionalStringField[M]) IsNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -1438,7 +1648,7 @@ func (f StringField[M]) IsNull() Predicate[M] { } } -func (f StringField[M]) IsNotNull() Predicate[M] { +func (f OptionalStringField[M]) IsNotNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -1538,7 +1748,6 @@ func (f StringUniqueField[M]) In(vals []string) Predicate[M] { } func (f StringUniqueField[M]) NotIn(vals []string) Predicate[M] { - return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -1548,6 +1757,16 @@ func (f StringUniqueField[M]) NotIn(vals []string) Predicate[M] { } } +func (f StringUniqueField[M]) Between(min string, max string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "BETWEEN", + Value: []any{min, max}, + }, + } +} + func (f StringUniqueField[M]) Like(val string) Predicate[M] { return Predicate[M]{ Data: PredicateData{ @@ -1558,6 +1777,36 @@ func (f StringUniqueField[M]) Like(val string) Predicate[M] { } } +func (f StringUniqueField[M]) ILike(val string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "ILIKE", + Value: val, + }, + } +} + +func (f StringUniqueField[M]) HasPrefix(prefix string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "LIKE", + Value: prefix + "%", + }, + } +} + +func (f StringUniqueField[M]) HasSuffix(suffix string) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "LIKE", + Value: "%" + suffix, + }, + } +} + func (f StringUniqueField[M]) Contains(val string) Predicate[M] { return Predicate[M]{ Data: PredicateData{ @@ -1568,7 +1817,11 @@ func (f StringUniqueField[M]) Contains(val string) Predicate[M] { } } -func (f StringUniqueField[M]) IsNull() Predicate[M] { +type OptionalStringUniqueField[M any] struct { + StringUniqueField[M] +} + +func (f OptionalStringUniqueField[M]) IsNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -1577,7 +1830,7 @@ func (f StringUniqueField[M]) IsNull() Predicate[M] { } } -func (f StringUniqueField[M]) IsNotNull() Predicate[M] { +func (f OptionalStringUniqueField[M]) IsNotNull() Predicate[M] { return Predicate[M]{ Data: PredicateData{ Column: f.Column, @@ -1594,6 +1847,53 @@ func (f StringUniqueField[M]) Desc() OrderBy[M] { return OrderBy[M]{Field: f.Column, Direction: Desc} } +type ArrayField[M any, T any] struct { + Column string +} + +func (f ArrayField[M, T]) Set(vals []T) FieldAssignmentOf[M] { + return FieldAssignmentOf[M]{Col: f.Column, Val: vals} +} + +func (f ArrayField[M, T]) Has(val T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "ARRAY_HAS", + Value: val, + }, + } +} + +func (f ArrayField[M, T]) HasEvery(vals []T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "ARRAY_HAS_EVERY", + Value: vals, + }, + } +} + +func (f ArrayField[M, T]) HasSome(vals []T) Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "ARRAY_HAS_SOME", + Value: vals, + }, + } +} + +func (f ArrayField[M, T]) IsEmpty() Predicate[M] { + return Predicate[M]{ + Data: PredicateData{ + Column: f.Column, + Operator: "ARRAY_IS_EMPTY", + }, + } +} + func CompilePredicates[M any](dialect Dialect, preds []PredicateOf[M], startBindIdx ...int) (string, []any, int) { bindIdx := 1 if len(startBindIdx) > 0 && startBindIdx[0] > 0 { @@ -1668,6 +1968,71 @@ func CompilePredicateData(dialect Dialect, data []PredicateData, startBindIdx .. args = append(args, val) } return fmt.Sprintf("%s IN (%s)", dialect.Quote(p.Column), strings.Join(placeHolders, ", ")) + case "NOT IN": + valSlice := unpackSlice(p.Value) + if len(valSlice) == 0 { + return "1=1" + } + var placeHolders []string + for range valSlice { + placeHolders = append(placeHolders, dialect.BindVar(bindIdx)) + bindIdx++ + } + for _, val := range valSlice { + args = append(args, val) + } + return fmt.Sprintf("%s NOT IN (%s)", dialect.Quote(p.Column), strings.Join(placeHolders, ", ")) + case "BETWEEN": + valSlice := unpackSlice(p.Value) + if len(valSlice) < 2 { + return "" + } + p1 := dialect.BindVar(bindIdx) + bindIdx++ + p2 := dialect.BindVar(bindIdx) + bindIdx++ + args = append(args, valSlice[0], valSlice[1]) + return fmt.Sprintf("%s BETWEEN %s AND %s", dialect.Quote(p.Column), p1, p2) + case "ILIKE": + placeholder := dialect.BindVar(bindIdx) + bindIdx++ + args = append(args, p.Value) + return dialect.FormatILike(p.Column, placeholder) + case "ARRAY_HAS": + placeholder := dialect.BindVar(bindIdx) + bindIdx++ + args = append(args, p.Value) + return dialect.FormatArrayHas(p.Column, placeholder) + case "ARRAY_HAS_EVERY": + valSlice := unpackSlice(p.Value) + if len(valSlice) == 0 { + return "1=1" + } + var placeHolders []string + for range valSlice { + placeHolders = append(placeHolders, dialect.BindVar(bindIdx)) + bindIdx++ + } + for _, val := range valSlice { + args = append(args, val) + } + return dialect.FormatArrayHasEvery(p.Column, placeHolders) + case "ARRAY_HAS_SOME": + valSlice := unpackSlice(p.Value) + if len(valSlice) == 0 { + return "1=0" + } + var placeHolders []string + for range valSlice { + placeHolders = append(placeHolders, dialect.BindVar(bindIdx)) + bindIdx++ + } + for _, val := range valSlice { + args = append(args, val) + } + return dialect.FormatArrayHasSome(p.Column, placeHolders) + case "ARRAY_IS_EMPTY": + return dialect.FormatArrayIsEmpty(p.Column) default: placeholder := dialect.BindVar(bindIdx) bindIdx++ diff --git a/integration/valk/comment.go b/integration/valk/comment.go index 52dc7bc..2db58d7 100644 --- a/integration/valk/comment.go +++ b/integration/valk/comment.go @@ -184,8 +184,6 @@ func assignmentsToCommentUpdate(assignments []FieldAssignment) (CommentUpdate, e input.Meta = &v } else if v, ok := a.Val.(*json.RawMessage); ok { input.Meta = v - } else if v, ok := a.Val.(*json.RawMessage); ok { - input.Meta = v } else { errs.Add("meta", a.Val, "type", "field meta must be of type *json.RawMessage") } diff --git a/integration/valk/comment/comment.go b/integration/valk/comment/comment.go index cc2f5ac..ce672e4 100644 --- a/integration/valk/comment/comment.go +++ b/integration/valk/comment/comment.go @@ -51,7 +51,7 @@ var PostId = valk.StringField[valk.Comment]{Column: "postId"} var AuthorId = valk.StringField[valk.Comment]{Column: "authorId"} -var Meta = valk.Field[valk.Comment, json.RawMessage]{Column: "meta"} +var Meta = valk.OptionalField[valk.Comment, json.RawMessage]{Field: valk.Field[valk.Comment, json.RawMessage]{Column: "meta"}} type CreateInput = valk.CommentCreate type Create = valk.CommentCreate diff --git a/integration/valk/post.go b/integration/valk/post.go index 206eac4..84acd79 100644 --- a/integration/valk/post.go +++ b/integration/valk/post.go @@ -4,10 +4,13 @@ import ( "context" "database/sql" "fmt" + "github.com/lib/pq" "slices" "strings" ) +var _ = pq.Array + // Post represents the database model type Post struct { Id string `db:"id" json:"id"` @@ -15,6 +18,7 @@ type Post struct { Content *string `db:"content" json:"content,omitempty"` Published bool `db:"published" json:"published"` AuthorId string `db:"authorId" json:"authorId"` + Tags []string `db:"tags" json:"tags"` Author *User `json:"author,omitempty"` Comments []*Comment `json:"comments,omitempty"` Categories []*CategoryToPost `json:"categories,omitempty"` @@ -24,17 +28,19 @@ type Post struct { // // Fields for Post: // -// id string default: cuid() -// title string required -// content string optional -// published bool default: false -// authorId string required +// id string default: cuid() +// title string required +// content string optional +// published bool default: false +// authorId string required +// tags []string default: [] type PostCreate struct { - Id *string `json:"id"` - Title string `json:"title"` - Content *string `json:"content"` - Published *bool `json:"published"` - AuthorId string `json:"authorId"` + Id *string `json:"id"` + Title string `json:"title"` + Content *string `json:"content"` + Published *bool `json:"published"` + AuthorId string `json:"authorId"` + Tags []string `json:"tags"` } // colMask returns a bit mask of columns that are set @@ -49,16 +55,20 @@ func (s *PostCreate) colMask() uint64 { mask |= 1 << 3 } mask |= 1 << 4 + if s.Tags != nil { + mask |= 1 << 5 + } return mask } // PostUpdate contains model input fields for Post update operations. type PostUpdate struct { - Id *string `json:"id"` - Title *string `json:"title"` - Content *string `json:"content"` - Published *bool `json:"published"` - AuthorId *string `json:"authorId"` + Id *string `json:"id"` + Title *string `json:"title"` + Content *string `json:"content"` + Published *bool `json:"published"` + AuthorId *string `json:"authorId"` + Tags []string `json:"tags"` } func (u *PostUpdate) ToColsVals() ([]string, []any) { @@ -84,6 +94,10 @@ func (u *PostUpdate) ToColsVals() ([]string, []any) { cols = append(cols, "authorId") vals = append(vals, u.AuthorId) } + if u.Tags != nil { + cols = append(cols, "tags") + vals = append(vals, u.Tags) + } return cols, vals } @@ -137,6 +151,16 @@ func assignmentsToPostUpdate(assignments []FieldAssignment) (PostUpdate, error) } else { errs.Add("authorId", a.Val, "type", "field authorId must be of type string") } + case "tags": + if v, ok := a.Val.([]string); ok { + input.Tags = v + } else if v, ok := a.Val.(*[]string); ok { + if v != nil { + input.Tags = *v + } + } else { + errs.Add("tags", a.Val, "type", "field tags must be of type []string") + } } } @@ -156,6 +180,7 @@ func assignmentsToPostUpdate(assignments []FieldAssignment) (PostUpdate, error) // content (bool) // published (bool) // authorId (bool) +// tags (bool) // -- Relations -- // author (User) // comments ([]Comment) @@ -166,6 +191,7 @@ type PostSelect struct { Content bool `json:"content"` Published bool `json:"published"` AuthorId bool `json:"authorId"` + Tags bool `json:"tags"` Author *UserSelect `json:"author,omitempty"` Comments CommentSelectQuery `json:"comments,omitempty"` Categories CategoryToPostSelectQuery `json:"categories,omitempty"` @@ -177,6 +203,7 @@ var fullPostSelectVal = &PostSelect{ Content: true, Published: true, AuthorId: true, + Tags: true, } func fullPostSelect() *PostSelect { @@ -187,7 +214,7 @@ func (s *PostSelect) hasAnyScalar() bool { if s == nil { return false } - return s.Id || s.Title || s.Content || s.Published || s.AuthorId + return s.Id || s.Title || s.Content || s.Published || s.AuthorId || s.Tags } func (s *PostSelect) hasAnySelected() bool { @@ -203,6 +230,7 @@ type PostOmit struct { Content bool `json:"content"` Published bool `json:"published"` AuthorId bool `json:"authorId"` + Tags bool `json:"tags"` } type PostSelectQuery interface { @@ -275,11 +303,12 @@ func (b *PostQueryBuilder) GetRelationParams() (*PostSelect, *PostOmit, QueryPar // // Fields for Post: // -// id string default: cuid() -// title string required -// content string optional -// published bool default: false -// authorId string required +// id string default: cuid() +// title string required +// content string optional +// published bool default: false +// authorId string required +// tags []string default: [] // // Relations for Post: // @@ -301,11 +330,12 @@ type PostCreateArgs struct { // // Fields for Post: // -// id string default: cuid() -// title string required -// content string optional -// published bool default: false -// authorId string required +// id string default: cuid() +// title string required +// content string optional +// published bool default: false +// authorId string required +// tags []string default: [] type PostCreateManyArgs struct { // Data is the slice of model inputs to bulk insert. Data []*PostCreate @@ -330,11 +360,12 @@ func (a *PostCreateManyArgs) AppendData(builders ...*PostCreateBuilder) *PostCre // // Fields for Post: // -// id string default: cuid() -// title string required -// content string optional -// published bool default: false -// authorId string required +// id string default: cuid() +// title string required +// content string optional +// published bool default: false +// authorId string required +// tags []string default: [] // // Relations for Post: // @@ -607,6 +638,8 @@ func (m *Post) ScanFields(cols []string) []any { targets[i] = &m.Published case "authorId": targets[i] = &m.AuthorId + case "tags": + targets[i] = ArrayScan(&m.Tags) } } return targets @@ -618,6 +651,7 @@ var postDefaultCols = []string{ "content", "published", "authorId", + "tags", } var postPKCols = []string{ @@ -633,7 +667,7 @@ func selectPostCols(selects *PostSelect, omits *PostOmit, forceCols ...string) [ return postDefaultCols } - anySelected := selects != nil && (selects.Id || selects.Title || selects.Content || selects.Published || selects.AuthorId || selects.Author != nil || selects.Comments != nil || selects.Categories != nil) + anySelected := selects != nil && (selects.Id || selects.Title || selects.Content || selects.Published || selects.AuthorId || selects.Tags || selects.Author != nil || selects.Comments != nil || selects.Categories != nil) specs := []colSpec{ {"id", selects != nil && selects.Id, omits != nil && omits.Id, selects != nil && selects.hasAnyRelation()}, @@ -641,6 +675,7 @@ func selectPostCols(selects *PostSelect, omits *PostOmit, forceCols ...string) [ {"content", selects != nil && selects.Content, omits != nil && omits.Content, false}, {"published", selects != nil && selects.Published, omits != nil && omits.Published, false}, {"authorId", selects != nil && selects.AuthorId, omits != nil && omits.AuthorId, selects != nil && selects.Author != nil}, + {"tags", selects != nil && selects.Tags, omits != nil && omits.Tags, false}, } cols := computeCols(specs, selects != nil, anySelected) @@ -706,6 +741,10 @@ func (b *PostCreateBuilder) SetAuthorId(v string) *PostCreateBuilder { b.assignments = append(b.assignments, FieldAssignment{Col: "authorId", Val: v}) return b } +func (b *PostCreateBuilder) SetTags(v []string) *PostCreateBuilder { + b.assignments = append(b.assignments, FieldAssignment{Col: "tags", Val: v}) + return b +} func (b *PostCreateBuilder) Assignments(assignments ...FieldAssignmentOf[Post]) *PostCreateBuilder { for _, a := range assignments { @@ -728,6 +767,7 @@ const ( providedPostContent uint64 = 1 << 2 providedPostPublished uint64 = 1 << 3 providedPostAuthorId uint64 = 1 << 4 + providedPostTags uint64 = 1 << 5 ) func assignmentsToPostCreate(assignments []FieldAssignment) (PostCreate, error) { @@ -776,6 +816,13 @@ func assignmentsToPostCreate(assignments []FieldAssignment) (PostCreate, error) } else { errs.Add("authorId", a.Val, "type", "field authorId must be of type string") } + case "tags": + provided |= providedPostTags + if v, ok := a.Val.([]string); ok { + input.Tags = v + } else { + errs.Add("tags", a.Val, "type", "field tags must be of type []string") + } } } if provided&providedPostTitle == 0 { @@ -792,8 +839,8 @@ func assignmentsToPostCreate(assignments []FieldAssignment) (PostCreate, error) } func (s *PostCreate) ToColsVals() (cols []string, vals []any) { - cols = make([]string, 0, 5) - vals = make([]any, 0, 5) + cols = make([]string, 0, 6) + vals = make([]any, 0, 6) cols = append(cols, "id") if s.Id != nil { vals = append(vals, *s.Id) @@ -812,6 +859,10 @@ func (s *PostCreate) ToColsVals() (cols []string, vals []any) { } cols = append(cols, "authorId") vals = append(vals, s.AuthorId) + if s.Tags != nil { + cols = append(cols, "tags") + vals = append(vals, ArrayVal(s.Tags)) + } return } @@ -1197,7 +1248,7 @@ func (d *PostDelegate) buildBulkInsertSQL(q *Queries, batch []*PostCreate, param colMask |= input.colMask() } - cols = make([]string, 0, 5) + cols = make([]string, 0, 6) for i, c := range postDefaultCols { if colMask&(1<