diff --git a/internal/codegen/golang/opts/options.go b/internal/codegen/golang/opts/options.go index 646bf1e066..74f9cf254d 100644 --- a/internal/codegen/golang/opts/options.go +++ b/internal/codegen/golang/opts/options.go @@ -158,6 +158,23 @@ func parseGlobalOpts(req *plugin.GenerateRequest) (*GlobalOptions, error) { return nil, err } } + + // A global override that names an engine only applies to that engine. The + // config requires `engine` on global overrides whenever more than one + // engine is in use, so without this filter every generated package would + // pick up whichever rule happened to be listed first and the last struct + // tag, regardless of the engine being generated. + engine := req.Settings.GetEngine() + if engine != "" { + filtered := make([]Override, 0, len(options.Overrides)) + for _, override := range options.Overrides { + if override.Engine != "" && override.Engine != engine { + continue + } + filtered = append(filtered, override) + } + options.Overrides = filtered + } return &options, nil } diff --git a/internal/endtoend/testdata/overrides_global_engine/mysql/db.go b/internal/endtoend/testdata/overrides_global_engine/mysql/db.go new file mode 100644 index 0000000000..8d745411c0 --- /dev/null +++ b/internal/endtoend/testdata/overrides_global_engine/mysql/db.go @@ -0,0 +1,31 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 + +package mysql + +import ( + "context" + "database/sql" +) + +type DBTX interface { + ExecContext(context.Context, string, ...any) (sql.Result, error) + PrepareContext(context.Context, string) (*sql.Stmt, error) + QueryContext(context.Context, string, ...any) (*sql.Rows, error) + QueryRowContext(context.Context, string, ...any) *sql.Row +} + +func New(db DBTX) *Queries { + return &Queries{db: db} +} + +type Queries struct { + db DBTX +} + +func (q *Queries) WithTx(tx *sql.Tx) *Queries { + return &Queries{ + db: tx, + } +} diff --git a/internal/endtoend/testdata/overrides_global_engine/mysql/models.go b/internal/endtoend/testdata/overrides_global_engine/mysql/models.go new file mode 100644 index 0000000000..eb70e855b8 --- /dev/null +++ b/internal/endtoend/testdata/overrides_global_engine/mysql/models.go @@ -0,0 +1,15 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 + +package mysql + +import ( + "database/sql" +) + +type Entry struct { + ID int32 + Value sql.NullString `backend:"mysql"` + Code []byte `backend:"mysql"` +} diff --git a/internal/endtoend/testdata/overrides_global_engine/mysql/query.sql.go b/internal/endtoend/testdata/overrides_global_engine/mysql/query.sql.go new file mode 100644 index 0000000000..8325e7d1fe --- /dev/null +++ b/internal/endtoend/testdata/overrides_global_engine/mysql/query.sql.go @@ -0,0 +1,28 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 +// source: query.sql + +package mysql + +import ( + "context" + "database/sql" +) + +const findEntry = `-- name: FindEntry :one +SELECT id, value, code FROM entries +WHERE value = ? AND code = ? +` + +type FindEntryParams struct { + Value sql.NullString `backend:"mysql"` + Code []byte `backend:"mysql"` +} + +func (q *Queries) FindEntry(ctx context.Context, arg FindEntryParams) (Entry, error) { + row := q.db.QueryRowContext(ctx, findEntry, arg.Value, arg.Code) + var i Entry + err := row.Scan(&i.ID, &i.Value, &i.Code) + return i, err +} diff --git a/internal/endtoend/testdata/overrides_global_engine/postgresql/db.go b/internal/endtoend/testdata/overrides_global_engine/postgresql/db.go new file mode 100644 index 0000000000..6614d155e5 --- /dev/null +++ b/internal/endtoend/testdata/overrides_global_engine/postgresql/db.go @@ -0,0 +1,31 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 + +package pg + +import ( + "context" + "database/sql" +) + +type DBTX interface { + ExecContext(context.Context, string, ...any) (sql.Result, error) + PrepareContext(context.Context, string) (*sql.Stmt, error) + QueryContext(context.Context, string, ...any) (*sql.Rows, error) + QueryRowContext(context.Context, string, ...any) *sql.Row +} + +func New(db DBTX) *Queries { + return &Queries{db: db} +} + +type Queries struct { + db DBTX +} + +func (q *Queries) WithTx(tx *sql.Tx) *Queries { + return &Queries{ + db: tx, + } +} diff --git a/internal/endtoend/testdata/overrides_global_engine/postgresql/models.go b/internal/endtoend/testdata/overrides_global_engine/postgresql/models.go new file mode 100644 index 0000000000..d543e260e0 --- /dev/null +++ b/internal/endtoend/testdata/overrides_global_engine/postgresql/models.go @@ -0,0 +1,15 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 + +package pg + +import ( + "encoding/json" +) + +type Entry struct { + ID int32 + Value string `backend:"postgresql"` + Code json.RawMessage `backend:"postgresql"` +} diff --git a/internal/endtoend/testdata/overrides_global_engine/postgresql/query.sql.go b/internal/endtoend/testdata/overrides_global_engine/postgresql/query.sql.go new file mode 100644 index 0000000000..b186bff3e3 --- /dev/null +++ b/internal/endtoend/testdata/overrides_global_engine/postgresql/query.sql.go @@ -0,0 +1,28 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 +// source: query.sql + +package pg + +import ( + "context" + "encoding/json" +) + +const findEntry = `-- name: FindEntry :one +SELECT id, value, code FROM entries +WHERE value = $1 AND code = $2 +` + +type FindEntryParams struct { + Value string `backend:"postgresql"` + Code json.RawMessage `backend:"postgresql"` +} + +func (q *Queries) FindEntry(ctx context.Context, arg FindEntryParams) (Entry, error) { + row := q.db.QueryRowContext(ctx, findEntry, arg.Value, arg.Code) + var i Entry + err := row.Scan(&i.ID, &i.Value, &i.Code) + return i, err +} diff --git a/internal/endtoend/testdata/overrides_global_engine/query.sql b/internal/endtoend/testdata/overrides_global_engine/query.sql new file mode 100644 index 0000000000..09e6a55fc5 --- /dev/null +++ b/internal/endtoend/testdata/overrides_global_engine/query.sql @@ -0,0 +1,3 @@ +-- name: FindEntry :one +SELECT id, value, code FROM entries +WHERE value = sqlc.arg(value) AND code = sqlc.arg(code); diff --git a/internal/endtoend/testdata/overrides_global_engine/schema.sql b/internal/endtoend/testdata/overrides_global_engine/schema.sql new file mode 100644 index 0000000000..3d66f2ec20 --- /dev/null +++ b/internal/endtoend/testdata/overrides_global_engine/schema.sql @@ -0,0 +1,5 @@ +CREATE TABLE entries ( + id integer NOT NULL, + value text NOT NULL, + code text NOT NULL +); diff --git a/internal/endtoend/testdata/overrides_global_engine/sqlc.yaml b/internal/endtoend/testdata/overrides_global_engine/sqlc.yaml new file mode 100644 index 0000000000..e74f50d3a2 --- /dev/null +++ b/internal/endtoend/testdata/overrides_global_engine/sqlc.yaml @@ -0,0 +1,53 @@ +version: "2" +overrides: + go: + overrides: + - db_type: text + engine: postgresql + go_type: string + go_struct_tag: 'backend:"postgresql"' + - db_type: text + engine: mysql + go_type: "database/sql.NullString" + go_struct_tag: 'backend:"mysql"' + - db_type: text + engine: sqlite + go_type: + type: byte + slice: true + go_struct_tag: 'backend:"sqlite"' + - column: entries.code + engine: postgresql + go_type: + import: encoding/json + type: RawMessage + - column: entries.code + engine: mysql + go_type: + type: byte + slice: true + - column: entries.code + engine: sqlite + go_type: string +sql: + - engine: postgresql + schema: schema.sql + queries: query.sql + gen: + go: + package: pg + out: postgresql + - engine: mysql + schema: schema.sql + queries: query.sql + gen: + go: + package: mysql + out: mysql + - engine: sqlite + schema: schema.sql + queries: query.sql + gen: + go: + package: sqlite + out: sqlite diff --git a/internal/endtoend/testdata/overrides_global_engine/sqlite/db.go b/internal/endtoend/testdata/overrides_global_engine/sqlite/db.go new file mode 100644 index 0000000000..be3742276e --- /dev/null +++ b/internal/endtoend/testdata/overrides_global_engine/sqlite/db.go @@ -0,0 +1,31 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 + +package sqlite + +import ( + "context" + "database/sql" +) + +type DBTX interface { + ExecContext(context.Context, string, ...any) (sql.Result, error) + PrepareContext(context.Context, string) (*sql.Stmt, error) + QueryContext(context.Context, string, ...any) (*sql.Rows, error) + QueryRowContext(context.Context, string, ...any) *sql.Row +} + +func New(db DBTX) *Queries { + return &Queries{db: db} +} + +type Queries struct { + db DBTX +} + +func (q *Queries) WithTx(tx *sql.Tx) *Queries { + return &Queries{ + db: tx, + } +} diff --git a/internal/endtoend/testdata/overrides_global_engine/sqlite/models.go b/internal/endtoend/testdata/overrides_global_engine/sqlite/models.go new file mode 100644 index 0000000000..031862d729 --- /dev/null +++ b/internal/endtoend/testdata/overrides_global_engine/sqlite/models.go @@ -0,0 +1,11 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 + +package sqlite + +type Entry struct { + ID int64 + Value []byte `backend:"sqlite"` + Code string `backend:"sqlite"` +} diff --git a/internal/endtoend/testdata/overrides_global_engine/sqlite/query.sql.go b/internal/endtoend/testdata/overrides_global_engine/sqlite/query.sql.go new file mode 100644 index 0000000000..fb70006c88 --- /dev/null +++ b/internal/endtoend/testdata/overrides_global_engine/sqlite/query.sql.go @@ -0,0 +1,27 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 +// source: query.sql + +package sqlite + +import ( + "context" +) + +const findEntry = `-- name: FindEntry :one +SELECT id, value, code FROM entries +WHERE value = ?1 AND code = ?2 +` + +type FindEntryParams struct { + Value []byte `backend:"sqlite"` + Code string `backend:"sqlite"` +} + +func (q *Queries) FindEntry(ctx context.Context, arg FindEntryParams) (Entry, error) { + row := q.db.QueryRowContext(ctx, findEntry, arg.Value, arg.Code) + var i Entry + err := row.Scan(&i.ID, &i.Value, &i.Code) + return i, err +}