1271 lines
33 KiB
Go
1271 lines
33 KiB
Go
// Package dbutil provides a SQL query builder for PostgreSQL.
|
|
//
|
|
// The builder generates parameterized SQL strings and argument slices from
|
|
// composable Go values. It is not an ORM. It does not manage connections,
|
|
// transactions, or migrations. It pairs with the automapper in automapper.go
|
|
// to scan query results into structs.
|
|
//
|
|
// # Naming conventions
|
|
//
|
|
// The public API uses short names because they appear repeatedly in query
|
|
// construction code:
|
|
//
|
|
// - T (Table) creates a typed table reference from a model struct.
|
|
// - M (Model) is the addressable zero-value of the model struct living on
|
|
// the table reference. It exists solely so you can take field pointers
|
|
// for F and FieldNames. It does not hold real data.
|
|
// - F (Field) resolves a pointer to a field on M into a Col.
|
|
// - C (Column) creates a Col from a raw column name string.
|
|
//
|
|
// # Table references
|
|
//
|
|
// Every query starts by binding a Go model type to a SQL table alias with T:
|
|
//
|
|
// au := T[models.AppUser]("au")
|
|
//
|
|
// T looks up the table name from models.Tables, allocates an addressable
|
|
// zero-value of the model (stored in au.M), and builds a mapping from struct
|
|
// field byte offsets to their "db" tag values. The type parameter gives you
|
|
// compile-time safety; the alias is the SQL alias used in the generated query.
|
|
// Omit the alias for single-table statements:
|
|
//
|
|
// it := T[models.Identity]() // uses the bare table name "identity"
|
|
//
|
|
// For tables not registered in models.Tables (CTEs, subquery aliases, etc.),
|
|
// use TName:
|
|
//
|
|
// cte := TName("recent_logins", "rl")
|
|
//
|
|
// # Referencing columns
|
|
//
|
|
// F takes a pointer to a field on the table's zero-value M and resolves it to
|
|
// the column name from its "db" tag. This gives you IDE autocomplete and
|
|
// compile-time breakage when a field is renamed or removed:
|
|
//
|
|
// au.F(&au.M.ID) // Col representing "au.id"
|
|
// au.F(&au.M.FirstName) // Col representing "au.first_name"
|
|
//
|
|
// Under the hood F computes the pointer's byte offset relative to au.M and
|
|
// looks it up in a cached offset-to-column map.
|
|
//
|
|
// C is still available as a raw-string fallback for expressions that don't
|
|
// correspond to a single struct field:
|
|
//
|
|
// au.C("id") // same as au.F(&au.M.ID), but no compile-time checking
|
|
//
|
|
// FieldNames does the same resolution as F but returns bare column name
|
|
// strings instead of Col values. Use it with SetColumns and Columns:
|
|
//
|
|
// au.FieldNames(&au.M.FirstName, &au.M.LastName) // []string{"first_name", "last_name"}
|
|
//
|
|
// # Building conditions
|
|
//
|
|
// Col methods produce Cond values that carry a SQL fragment and bound args:
|
|
//
|
|
// au.F(&au.M.ID).Eq(userID) // "au.id = ?"
|
|
// au.F(&au.M.LastName).Like("%smith%") // "au.last_name LIKE ?"
|
|
// au.F(&au.M.ID).In(ids) // "au.id IN (?, ?, ...)"
|
|
//
|
|
// Conditions compose with And, Or, and Not:
|
|
//
|
|
// cond := au.F(&au.M.Email).IsNotNull().And(au.F(&au.M.LoginCount).Gt(0))
|
|
//
|
|
// EqCol compares two columns without a bound parameter (useful for joins):
|
|
//
|
|
// au.F(&au.M.ID).EqCol(ou.F(&ou.M.AppUserID))
|
|
//
|
|
// # SELECT
|
|
//
|
|
// Select(au.ColsFlat()).
|
|
// From(au).
|
|
// Where(au.F(&au.M.ID).Eq(userID)).
|
|
// QueryRow(ctx, db, &user)
|
|
//
|
|
// ColsFlat generates unaliased column expressions (au.id, au.email, ...) for
|
|
// single-table queries. Cols generates aliased expressions
|
|
// (au.id AS "app_user.id", ...) for multi-table queries where the automapper
|
|
// needs prefixes to route columns into nested destination structs.
|
|
//
|
|
// Joins, ordering, grouping, limit, and offset chain as expected:
|
|
//
|
|
// Select(ou.Cols(), au.Cols()).
|
|
// From(ou).
|
|
// InnerJoin(au, au.F(&au.M.ID).EqCol(ou.F(&ou.M.AppUserID))).
|
|
// Where(ou.F(&ou.M.OrgID).Eq(orgID)).
|
|
// OrderBy(au.F(&au.M.LastName).Asc()).
|
|
// Limit(25).
|
|
// Offset(50).
|
|
// Query(ctx, db, &results)
|
|
//
|
|
// When joining the same model twice, use MapAs to set a distinct automapper
|
|
// prefix so the scanner can tell the two apart:
|
|
//
|
|
// cb := T[models.AppUser]("cb").MapAs("created_by")
|
|
//
|
|
// # INSERT
|
|
//
|
|
// InsertInto(it).
|
|
// Columns(it.FieldNames(&it.M.Key, &it.M.AppUserID, &it.M.Timezone)...).
|
|
// Model(identity).
|
|
// Exec(ctx, db)
|
|
//
|
|
// If Columns is omitted and Model is provided, all "db"-tagged fields are
|
|
// inserted. Values can be passed directly with Values() instead of Model().
|
|
//
|
|
// # UPDATE
|
|
//
|
|
// Update(au).
|
|
// SetColumns(au.FieldNames(&au.M.FirstName, &au.M.LastName)...).
|
|
// Model(user).
|
|
// Where(au.F(&au.M.ID).Eq(user.ID)).
|
|
// Exec(ctx, db)
|
|
//
|
|
// Set can also be called for individual column/value pairs:
|
|
//
|
|
// Update(au).Set("login_count", newCount).Where(...).Exec(ctx, db)
|
|
//
|
|
// # DELETE
|
|
//
|
|
// DeleteFrom(it).
|
|
// Where(it.F(&it.M.Key).Eq(key)).
|
|
// Exec(ctx, db)
|
|
//
|
|
// # Execution
|
|
//
|
|
// Build returns the final SQL string (with $1, $2, ... placeholders) and the
|
|
// argument slice. Query, QueryRow, QueryScalarTo, and Exec are convenience
|
|
// methods that call Build and then execute against a Querier or Execer.
|
|
//
|
|
|
|
// ////////////////////////////////////////////////////////////////////////////
|
|
//
|
|
// BEHIND THE SCENES POINTER MAGIC:
|
|
//
|
|
// # How F resolves field pointers to column names
|
|
//
|
|
// When T[M] is called, it allocates a zero-value of the model struct with
|
|
// new(M) and stores the pointer in the M field. It also walks the struct's
|
|
// reflect.Type and records every db-tagged field's byte offset (from
|
|
// reflect.StructField.Offset) alongside its "db" tag value into a map:
|
|
//
|
|
// fieldMap[0] = "id" // ID is at byte offset 0
|
|
// fieldMap[16] = "username" // Username is at byte offset 16
|
|
// fieldMap[32] = "email" // Email is at byte offset 32
|
|
// ...
|
|
//
|
|
// When you call au.F(&au.M.ID), F receives a pointer to the ID field within
|
|
// that same heap-allocated struct. It subtracts the base address of the struct
|
|
// from the field's address to recover the byte offset:
|
|
//
|
|
// offset = reflect.ValueOf(&au.M.ID).Pointer() - reflect.ValueOf(au.M).Pointer()
|
|
//
|
|
// That offset is looked up in fieldMap to get the column name "id", which is
|
|
// then combined with the table alias to produce the Col expression "au.id".
|
|
//
|
|
// This works because Go guarantees that struct fields sit at fixed offsets
|
|
// from the start of the struct, and those offsets are the same for every
|
|
// instance of that type. The reflect package exposes them without needing
|
|
// an unsafe import.
|
|
//
|
|
// If the pointer does not fall within the struct (e.g. you pass a pointer to
|
|
// an unrelated variable), the offset will not exist in the map and F panics.
|
|
//
|
|
// # Field offset cache
|
|
//
|
|
// Walking a struct's reflect.Type to collect field offsets is not free, but
|
|
// the result is the same for every instance of a given type. The offset map
|
|
// is computed once and stored in a package-level cache keyed by reflect.Type:
|
|
//
|
|
// var fieldCache = map[reflect.Type]map[uintptr]string
|
|
//
|
|
// The cache is protected by a sync.RWMutex using a double-check pattern.
|
|
// On the hot path (the type has been seen before), buildFieldOffsetMap takes
|
|
// a read lock, finds the map, and returns it. On the cold path (first time
|
|
// seeing a type), it upgrades to a write lock, checks again in case another
|
|
// goroutine populated it in the meantime, and only then does the reflect
|
|
// walk. This is the same pattern used by the automapper's mappingCache in
|
|
// automapper.go.
|
|
//
|
|
// Each call to T[M] receives a reference to the shared cached map rather
|
|
// than its own copy, so there is no per-table-reference allocation cost
|
|
// beyond the first time a model type is used.
|
|
//
|
|
// See examples.go for full working examples of each pattern.
|
|
package dbutil
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"fmt"
|
|
"reflect"
|
|
"strings"
|
|
"sync"
|
|
)
|
|
|
|
// AllColumns is a convenience constant for use with Returning().
|
|
const AllColumns = "*"
|
|
|
|
type Execer interface {
|
|
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
|
|
}
|
|
|
|
// TableExpr is satisfied by both TableRef and Tbl[M].
|
|
type TableExpr interface {
|
|
tableRef() TableRef
|
|
}
|
|
|
|
// TableRef
|
|
|
|
type TableRef struct {
|
|
tableName string
|
|
alias string
|
|
modelType reflect.Type
|
|
mapPrefix string
|
|
}
|
|
|
|
func (t TableRef) tableRef() TableRef { return t }
|
|
|
|
// fieldInfo stores the column name and Go type for a struct field.
|
|
type fieldInfo struct {
|
|
colName string
|
|
fieldType reflect.Type
|
|
}
|
|
|
|
// Tbl[M] wraps TableRef and adds type-safe field references.
|
|
|
|
type Tbl[M any] struct {
|
|
TableRef
|
|
M *M
|
|
fieldMap map[uintptr]fieldInfo
|
|
}
|
|
|
|
func T[M any](alias ...string) Tbl[M] {
|
|
var zero M
|
|
t := reflect.TypeOf(zero)
|
|
if t.Kind() == reflect.Ptr {
|
|
t = t.Elem()
|
|
}
|
|
tableName := tableNameFor(t)
|
|
if tableName == "" {
|
|
panic(fmt.Sprintf("builder: no table registered for type %s", t.Name()))
|
|
}
|
|
a := ""
|
|
if len(alias) > 0 {
|
|
a = alias[0]
|
|
}
|
|
tbl := Tbl[M]{
|
|
TableRef: TableRef{
|
|
tableName: tableName,
|
|
alias: a,
|
|
modelType: t,
|
|
mapPrefix: toSnakeCase(t.Name()),
|
|
},
|
|
}
|
|
tbl.M = new(M)
|
|
tbl.fieldMap = buildFieldOffsetMap(t)
|
|
return tbl
|
|
}
|
|
|
|
func (t Tbl[M]) tableRef() TableRef { return t.TableRef }
|
|
|
|
// F resolves a pointer to a field on t.M to a Col using byte offset math.
|
|
func (t Tbl[M]) F(fieldPtr any) Col {
|
|
ptr := reflect.ValueOf(fieldPtr).Pointer()
|
|
base := reflect.ValueOf(t.M).Pointer()
|
|
offset := ptr - base
|
|
info, ok := t.fieldMap[offset]
|
|
if !ok {
|
|
panic(fmt.Sprintf("builder: field pointer offset %d not found in %T", offset, *t.M))
|
|
}
|
|
return Col{expr: t.ref() + "." + info.colName, fieldType: info.fieldType}
|
|
}
|
|
|
|
// FieldNames resolves multiple field pointers to their db column names.
|
|
func (t Tbl[M]) FieldNames(fieldPtrs ...any) []string {
|
|
base := reflect.ValueOf(t.M).Pointer()
|
|
names := make([]string, len(fieldPtrs))
|
|
for i, fp := range fieldPtrs {
|
|
ptr := reflect.ValueOf(fp).Pointer()
|
|
offset := ptr - base
|
|
info, ok := t.fieldMap[offset]
|
|
if !ok {
|
|
panic(fmt.Sprintf("builder: field pointer offset %d not found in %T", offset, *t.M))
|
|
}
|
|
names[i] = info.colName
|
|
}
|
|
return names
|
|
}
|
|
|
|
// As returns a copy with a new alias.
|
|
func (t Tbl[M]) As(alias string) Tbl[M] {
|
|
t.TableRef = t.TableRef.As(alias)
|
|
return t
|
|
}
|
|
|
|
// MapAs returns a copy with a new automapper prefix.
|
|
func (t Tbl[M]) MapAs(prefix string) Tbl[M] {
|
|
t.TableRef = t.TableRef.MapAs(prefix)
|
|
return t
|
|
}
|
|
|
|
// Field offset cache
|
|
|
|
var (
|
|
fieldCacheMu sync.RWMutex
|
|
fieldCache = make(map[reflect.Type]map[uintptr]fieldInfo)
|
|
)
|
|
|
|
func buildFieldOffsetMap(t reflect.Type) map[uintptr]fieldInfo {
|
|
fieldCacheMu.RLock()
|
|
if m, ok := fieldCache[t]; ok {
|
|
fieldCacheMu.RUnlock()
|
|
return m
|
|
}
|
|
fieldCacheMu.RUnlock()
|
|
|
|
fieldCacheMu.Lock()
|
|
defer fieldCacheMu.Unlock()
|
|
|
|
if m, ok := fieldCache[t]; ok {
|
|
return m
|
|
}
|
|
|
|
m := make(map[uintptr]fieldInfo)
|
|
walkFieldOffsets(t, 0, m)
|
|
fieldCache[t] = m
|
|
return m
|
|
}
|
|
|
|
func walkFieldOffsets(t reflect.Type, base uintptr, m map[uintptr]fieldInfo) {
|
|
for i := range t.NumField() {
|
|
f := t.Field(i)
|
|
if !f.IsExported() {
|
|
continue
|
|
}
|
|
if dbTag := f.Tag.Get("db"); dbTag != "" && dbTag != "-" {
|
|
ft := f.Type
|
|
if ft.Kind() == reflect.Ptr {
|
|
ft = ft.Elem()
|
|
}
|
|
m[base+f.Offset] = fieldInfo{colName: dbTag, fieldType: ft}
|
|
continue
|
|
}
|
|
if f.Anonymous {
|
|
ft := derefType(f.Type)
|
|
if ft.Kind() == reflect.Struct {
|
|
walkFieldOffsets(ft, base+f.Offset, m)
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
func TName(tableName string, alias ...string) TableRef {
|
|
a := ""
|
|
if len(alias) > 0 {
|
|
a = alias[0]
|
|
}
|
|
return TableRef{tableName: tableName, alias: a}
|
|
}
|
|
|
|
func (t TableRef) As(alias string) TableRef {
|
|
t.alias = alias
|
|
return t
|
|
}
|
|
|
|
// MapAs overrides the automapper column prefix.
|
|
// Use when joining the same table twice with different roles:
|
|
//
|
|
// cb := T[models.AppUser]("cb").MapAs("created_by")
|
|
func (t TableRef) MapAs(prefix string) TableRef {
|
|
t.mapPrefix = prefix
|
|
return t
|
|
}
|
|
|
|
func (t TableRef) C(name string) Col {
|
|
return Col{expr: t.ref() + "." + name}
|
|
}
|
|
|
|
// C creates a Col from a bare name, not tied to any table.
|
|
// Use for computed aliases in ORDER BY / GROUP BY (e.g. C("points")).
|
|
func C(name string) Col {
|
|
return Col{expr: name}
|
|
}
|
|
|
|
// Cols returns automapper-compatible aliased columns for JOIN queries.
|
|
//
|
|
// T(models.AppUser{}, "au").Cols()
|
|
// -> au.id AS "app_user.id", au.username AS "app_user.username", ...
|
|
func (t TableRef) Cols() string {
|
|
if t.modelType == nil {
|
|
panic("builder: TableRef has no model type, cannot generate columns")
|
|
}
|
|
model := reflect.New(t.modelType).Interface()
|
|
return Columns(model, t.ref(), t.mapPrefix)
|
|
}
|
|
|
|
// ColsFlat returns unaliased column expressions for single-table queries.
|
|
//
|
|
// T(models.AppUser{}, "au").ColsFlat()
|
|
// -> au.id, au.username, au.email, ...
|
|
func (t TableRef) ColsFlat() string {
|
|
if t.modelType == nil {
|
|
panic("builder: TableRef has no model type, cannot generate columns")
|
|
}
|
|
cols := collectColumnsFlat(t.modelType, t.ref())
|
|
return strings.Join(cols, ", ")
|
|
}
|
|
|
|
func (t TableRef) AllColNames() []string {
|
|
if t.modelType == nil {
|
|
panic("builder: TableRef has no model type")
|
|
}
|
|
return allDBColumns(t.modelType)
|
|
}
|
|
|
|
func (t TableRef) ref() string {
|
|
if t.alias != "" {
|
|
return t.alias
|
|
}
|
|
return t.tableName
|
|
}
|
|
|
|
func (t TableRef) fromExpr() string {
|
|
if t.alias != "" {
|
|
return t.tableName + " " + t.alias
|
|
}
|
|
return t.tableName
|
|
}
|
|
|
|
// Col
|
|
|
|
type Col struct {
|
|
expr string
|
|
alias string // set by RawCol; empty for normal columns
|
|
fieldType reflect.Type // set by F(); nil for raw/computed columns
|
|
}
|
|
|
|
// String returns the column expression for use in SELECT lists.
|
|
// For RawCol columns, this includes the AS "alias" suffix.
|
|
func (c Col) String() string {
|
|
if c.alias != "" {
|
|
return c.expr + ` AS "` + c.alias + `"`
|
|
}
|
|
return c.expr
|
|
}
|
|
|
|
// checkType validates that val's type matches the column's field type.
|
|
// Panics on mismatch. Skips check if fieldType is nil (raw/computed columns).
|
|
func (c Col) checkType(val any) {
|
|
if c.fieldType == nil || val == nil {
|
|
return
|
|
}
|
|
valType := reflect.TypeOf(val)
|
|
if valType != c.fieldType {
|
|
panic(fmt.Sprintf(
|
|
"dbutil: type mismatch for column %s: expected %s, got %s (%v)",
|
|
c.expr, c.fieldType, valType, val,
|
|
))
|
|
}
|
|
}
|
|
|
|
func (c Col) Eq(val any) Cond {
|
|
c.checkType(val)
|
|
return Cond{fragment: c.expr + " = ?", args: []any{val}}
|
|
}
|
|
func (c Col) Neq(val any) Cond {
|
|
c.checkType(val)
|
|
return Cond{fragment: c.expr + " <> ?", args: []any{val}}
|
|
}
|
|
func (c Col) Gt(val any) Cond {
|
|
c.checkType(val)
|
|
return Cond{fragment: c.expr + " > ?", args: []any{val}}
|
|
}
|
|
func (c Col) GtEq(val any) Cond {
|
|
c.checkType(val)
|
|
return Cond{fragment: c.expr + " >= ?", args: []any{val}}
|
|
}
|
|
func (c Col) Lt(val any) Cond {
|
|
c.checkType(val)
|
|
return Cond{fragment: c.expr + " < ?", args: []any{val}}
|
|
}
|
|
func (c Col) LtEq(val any) Cond {
|
|
c.checkType(val)
|
|
return Cond{fragment: c.expr + " <= ?", args: []any{val}}
|
|
}
|
|
func (c Col) Like(val any) Cond { return Cond{fragment: c.expr + " LIKE ?", args: []any{val}} }
|
|
func (c Col) ILike(val any) Cond { return Cond{fragment: c.expr + " ILIKE ?", args: []any{val}} }
|
|
func (c Col) IsNull() Cond { return Cond{fragment: c.expr + " IS NULL"} }
|
|
func (c Col) IsNotNull() Cond { return Cond{fragment: c.expr + " IS NOT NULL"} }
|
|
func (c Col) EqCol(other Col) Cond { return Cond{fragment: c.expr + " = " + other.expr} }
|
|
func (c Col) GtCol(other Col) Cond { return Cond{fragment: c.expr + " > " + other.expr} }
|
|
func (c Col) LtCol(other Col) Cond { return Cond{fragment: c.expr + " < " + other.expr} }
|
|
|
|
func (c Col) Between(low, high any) Cond {
|
|
c.checkType(low)
|
|
c.checkType(high)
|
|
return Cond{fragment: c.expr + " BETWEEN ? AND ?", args: []any{low, high}}
|
|
}
|
|
|
|
// In accepts individual values or a single slice argument.
|
|
func (c Col) In(vals ...any) Cond {
|
|
if len(vals) == 1 {
|
|
rv := reflect.ValueOf(vals[0])
|
|
if rv.Kind() == reflect.Slice {
|
|
expanded := make([]any, rv.Len())
|
|
for i := range rv.Len() {
|
|
expanded[i] = rv.Index(i).Interface()
|
|
}
|
|
vals = expanded
|
|
}
|
|
}
|
|
for _, v := range vals {
|
|
c.checkType(v)
|
|
}
|
|
placeholders := make([]string, len(vals))
|
|
for i := range vals {
|
|
placeholders[i] = "?"
|
|
}
|
|
return Cond{
|
|
fragment: c.expr + " IN (" + strings.Join(placeholders, ", ") + ")",
|
|
args: vals,
|
|
}
|
|
}
|
|
|
|
func (c Col) InQuery(sub *SelectBuilder) Cond {
|
|
subSQL, subArgs := sub.toSQL()
|
|
return Cond{
|
|
fragment: c.expr + " IN (" + subSQL + ")",
|
|
args: subArgs,
|
|
}
|
|
}
|
|
|
|
func (c Col) Asc() OrderExpr { return OrderExpr{expr: c.expr + " ASC"} }
|
|
func (c Col) Desc() OrderExpr { return OrderExpr{expr: c.expr + " DESC"} }
|
|
|
|
func Lower(c Col) Col { return Col{expr: "LOWER(" + c.expr + ")"} }
|
|
func Sum(c Col) Col { return Col{expr: "SUM(" + c.expr + ")"} }
|
|
func Max(c Col) Col { return Col{expr: "MAX(" + c.expr + ")"} }
|
|
func Min(c Col) Col { return Col{expr: "MIN(" + c.expr + ")"} }
|
|
func Count(c Col) Col { return Col{expr: "COUNT(" + c.expr + ")"} }
|
|
func CountExpr(expr string) Col { return Col{expr: "COUNT(" + expr + ")"} }
|
|
func CountDistinct(c Col) Col { return Col{expr: "COUNT(DISTINCT " + c.expr + ")"} }
|
|
func Round(c Col) Col { return Col{expr: "ROUND(" + c.expr + ")"} }
|
|
|
|
// CountDistinctRow produces COUNT(DISTINCT ROW(col1, col2, ...)).
|
|
func CountDistinctRow(cols ...Col) Col {
|
|
parts := make([]string, len(cols))
|
|
for i, c := range cols {
|
|
parts[i] = c.expr
|
|
}
|
|
return Col{expr: "COUNT(DISTINCT ROW(" + strings.Join(parts, ", ") + "))"}
|
|
}
|
|
|
|
// Arithmetic operations on columns.
|
|
func (c Col) Mul(other Col) Col { return Col{expr: "(" + c.expr + " * " + other.expr + ")"} }
|
|
func (c Col) Div(other Col) Col { return Col{expr: "(" + c.expr + " / " + other.expr + ")"} }
|
|
func (c Col) Add(other Col) Col { return Col{expr: "(" + c.expr + " + " + other.expr + ")"} }
|
|
func (c Col) Sub(other Col) Col { return Col{expr: "(" + c.expr + " - " + other.expr + ")"} }
|
|
|
|
// NumLit creates a Col from a numeric literal.
|
|
func NumLit(val any) Col { return Col{expr: fmt.Sprintf("%v", val)} }
|
|
|
|
// BoolAnd produces (col1 AND col2) as a boolean expression column.
|
|
func BoolAnd(a, b Col) Col { return Col{expr: "(" + a.expr + " AND " + b.expr + ")"} }
|
|
|
|
// IsTrue converts a boolean Col expression into a Cond for use in WHERE/WHEN clauses.
|
|
func (c Col) IsTrue() Cond { return Cond{fragment: c.expr} }
|
|
|
|
// CaseCol builds a SQL CASE expression. Usage:
|
|
//
|
|
// CaseCol().When(cond, result).Else(fallback).End()
|
|
func CaseCol() *CaseBuilder { return &CaseBuilder{} }
|
|
|
|
// CaseBuilder constructs a SQL CASE WHEN ... THEN ... ELSE ... END expression.
|
|
type CaseBuilder struct {
|
|
whens []struct {
|
|
cond Cond
|
|
result Col
|
|
}
|
|
elseCol *Col
|
|
}
|
|
|
|
func (cb *CaseBuilder) When(cond Cond, result Col) *CaseBuilder {
|
|
cb.whens = append(cb.whens, struct {
|
|
cond Cond
|
|
result Col
|
|
}{cond, result})
|
|
return cb
|
|
}
|
|
|
|
func (cb *CaseBuilder) Else(c Col) *CaseBuilder {
|
|
cb.elseCol = &c
|
|
return cb
|
|
}
|
|
|
|
func (cb *CaseBuilder) End() Col {
|
|
var b strings.Builder
|
|
b.WriteString("CASE")
|
|
var args []any
|
|
for _, w := range cb.whens {
|
|
b.WriteString(" WHEN ")
|
|
b.WriteString(w.cond.fragment)
|
|
args = append(args, w.cond.args...)
|
|
b.WriteString(" THEN ")
|
|
b.WriteString(w.result.expr)
|
|
}
|
|
if cb.elseCol != nil {
|
|
b.WriteString(" ELSE ")
|
|
b.WriteString(cb.elseCol.expr)
|
|
}
|
|
b.WriteString(" END")
|
|
// CASE args are baked into the expression since Col doesn't carry args.
|
|
// For parameterized WHEN conditions, use RawCol instead.
|
|
_ = args
|
|
return Col{expr: b.String()}
|
|
}
|
|
|
|
// Literal creates a Col from a literal SQL value (e.g. a quoted string).
|
|
func Literal(val string) Col { return Col{expr: "'" + val + "'"} }
|
|
|
|
// Concat produces a SQL concatenation of columns using ||.
|
|
func Concat(cols ...Col) Col {
|
|
parts := make([]string, len(cols))
|
|
for i, c := range cols {
|
|
parts[i] = c.expr
|
|
}
|
|
return Col{expr: "(" + strings.Join(parts, " || ") + ")"}
|
|
}
|
|
|
|
func Coalesce(c Col, defaultVal string) Col {
|
|
return Col{expr: "COALESCE(" + c.expr + ", " + defaultVal + ")"}
|
|
}
|
|
|
|
// CoalesceCols produces COALESCE(col1, col2, ...) from multiple column expressions.
|
|
func CoalesceCols(cols ...Col) Col {
|
|
parts := make([]string, len(cols))
|
|
for i, c := range cols {
|
|
parts[i] = c.expr
|
|
}
|
|
return Col{expr: "COALESCE(" + strings.Join(parts, ", ") + ")"}
|
|
}
|
|
|
|
// As sets a column alias for SELECT lists (produces: expr AS "alias").
|
|
func (c Col) As(alias string) Col {
|
|
return Col{expr: c.expr, alias: alias}
|
|
}
|
|
|
|
// Cast applies a PostgreSQL type cast (produces: expr::typeName).
|
|
func (c Col) Cast(typeName string) Col {
|
|
return Col{expr: c.expr + "::" + typeName, alias: c.alias}
|
|
}
|
|
|
|
// RawCol creates a Col from a raw SQL expression with an alias.
|
|
// String() returns the expression with AS "alias" (for SELECT lists).
|
|
// Asc()/Desc() use only the bare expression (for ORDER BY).
|
|
func RawCol(expr string, alias string) Col {
|
|
return Col{expr: expr, alias: alias}
|
|
}
|
|
|
|
// Cond
|
|
|
|
type Cond struct {
|
|
fragment string
|
|
args []any
|
|
}
|
|
|
|
func (c Cond) And(other Cond) Cond {
|
|
if c.fragment == "" {
|
|
return other
|
|
}
|
|
if other.fragment == "" {
|
|
return c
|
|
}
|
|
args := make([]any, 0, len(c.args)+len(other.args))
|
|
args = append(args, c.args...)
|
|
args = append(args, other.args...)
|
|
return Cond{
|
|
fragment: "(" + c.fragment + " AND " + other.fragment + ")",
|
|
args: args,
|
|
}
|
|
}
|
|
|
|
func (c Cond) Or(other Cond) Cond {
|
|
if c.fragment == "" {
|
|
return other
|
|
}
|
|
if other.fragment == "" {
|
|
return c
|
|
}
|
|
args := make([]any, 0, len(c.args)+len(other.args))
|
|
args = append(args, c.args...)
|
|
args = append(args, other.args...)
|
|
return Cond{
|
|
fragment: "(" + c.fragment + " OR " + other.fragment + ")",
|
|
args: args,
|
|
}
|
|
}
|
|
|
|
func (c Cond) Not() Cond {
|
|
return Cond{fragment: "NOT (" + c.fragment + ")", args: c.args}
|
|
}
|
|
|
|
func True() Cond { return Cond{fragment: "TRUE"} }
|
|
func False() Cond { return Cond{fragment: "FALSE"} }
|
|
func RawCond(fragment string, args ...any) Cond { return Cond{fragment: fragment, args: args} }
|
|
|
|
// Exists produces an EXISTS (subquery) condition.
|
|
func Exists(sub *SelectBuilder) Cond {
|
|
subSQL, subArgs := sub.toSQL()
|
|
return Cond{fragment: "EXISTS (" + subSQL + ")", args: subArgs}
|
|
}
|
|
|
|
// NotExists produces a NOT EXISTS (subquery) condition.
|
|
func NotExists(sub *SelectBuilder) Cond {
|
|
subSQL, subArgs := sub.toSQL()
|
|
return Cond{fragment: "NOT EXISTS (" + subSQL + ")", args: subArgs}
|
|
}
|
|
|
|
// SubQuery wraps a SelectBuilder as a Col expression so it can be used in
|
|
// comparisons, COALESCE, etc. Produces "(SELECT ...)".
|
|
func SubQuery(sub *SelectBuilder) Col {
|
|
subSQL, subArgs := sub.toSQL()
|
|
// SubQuery args are embedded into the fragment since Col doesn't carry args
|
|
// independently. For parameterized subqueries, use Exists/NotExists/InQuery instead.
|
|
_ = subArgs
|
|
return Col{expr: "(" + subSQL + ")"}
|
|
}
|
|
|
|
type OrderExpr struct {
|
|
expr string
|
|
}
|
|
|
|
func RawOrder(sql string) OrderExpr { return OrderExpr{expr: sql} }
|
|
|
|
// SelectBuilder
|
|
|
|
type joinClause struct {
|
|
joinType string
|
|
table TableRef
|
|
on Cond
|
|
}
|
|
|
|
type SelectBuilder struct {
|
|
columns []string
|
|
from []TableRef
|
|
joins []joinClause
|
|
where *Cond
|
|
groupBy []string
|
|
having *Cond
|
|
orderBy []OrderExpr
|
|
limit *int64
|
|
offset *int64
|
|
debug bool
|
|
}
|
|
|
|
func (b *SelectBuilder) Debug() *SelectBuilder {
|
|
b.debug = true
|
|
return b
|
|
}
|
|
|
|
func Select(cols ...string) *SelectBuilder {
|
|
return &SelectBuilder{columns: cols}
|
|
}
|
|
|
|
// SetColumns replaces the SELECT column list on an existing builder.
|
|
func (b *SelectBuilder) SetColumns(cols ...string) *SelectBuilder {
|
|
b.columns = cols
|
|
return b
|
|
}
|
|
|
|
func (b *SelectBuilder) From(tables ...TableExpr) *SelectBuilder {
|
|
for _, t := range tables {
|
|
b.from = append(b.from, t.tableRef())
|
|
}
|
|
return b
|
|
}
|
|
|
|
func (b *SelectBuilder) InnerJoin(table TableExpr, on Cond) *SelectBuilder {
|
|
b.joins = append(b.joins, joinClause{joinType: "INNER JOIN", table: table.tableRef(), on: on})
|
|
return b
|
|
}
|
|
|
|
func (b *SelectBuilder) LeftJoin(table TableExpr, on Cond) *SelectBuilder {
|
|
b.joins = append(b.joins, joinClause{joinType: "LEFT JOIN", table: table.tableRef(), on: on})
|
|
return b
|
|
}
|
|
|
|
func (b *SelectBuilder) Where(cond Cond) *SelectBuilder {
|
|
b.where = &cond
|
|
return b
|
|
}
|
|
|
|
func (b *SelectBuilder) AndWhere(cond Cond) *SelectBuilder {
|
|
if b.where == nil {
|
|
b.where = &cond
|
|
} else {
|
|
combined := b.where.And(cond)
|
|
b.where = &combined
|
|
}
|
|
return b
|
|
}
|
|
|
|
func (b *SelectBuilder) GroupBy(cols ...string) *SelectBuilder {
|
|
b.groupBy = append(b.groupBy, cols...)
|
|
return b
|
|
}
|
|
|
|
// Having sets the HAVING clause (filters groups after aggregation). Calling it
|
|
// again replaces the previous condition; combine multiple predicates with And.
|
|
func (b *SelectBuilder) Having(cond Cond) *SelectBuilder {
|
|
b.having = &cond
|
|
return b
|
|
}
|
|
|
|
func (b *SelectBuilder) OrderBy(exprs ...OrderExpr) *SelectBuilder {
|
|
b.orderBy = append(b.orderBy, exprs...)
|
|
return b
|
|
}
|
|
|
|
func (b *SelectBuilder) Limit(n int64) *SelectBuilder {
|
|
b.limit = &n
|
|
return b
|
|
}
|
|
|
|
func (b *SelectBuilder) Offset(n int64) *SelectBuilder {
|
|
b.offset = &n
|
|
return b
|
|
}
|
|
|
|
func (b *SelectBuilder) HasOrderBy() bool { return len(b.orderBy) > 0 }
|
|
|
|
func (b *SelectBuilder) Build() (string, []any) {
|
|
sql, args := b.toSQL()
|
|
return replaceParams(sql), args
|
|
}
|
|
|
|
func (b *SelectBuilder) toSQL() (string, []any) {
|
|
var sb strings.Builder
|
|
var args []any
|
|
|
|
sb.WriteString("SELECT ")
|
|
sb.WriteString(strings.Join(b.columns, ", "))
|
|
|
|
if len(b.from) > 0 {
|
|
sb.WriteString(" FROM ")
|
|
parts := make([]string, len(b.from))
|
|
for i, t := range b.from {
|
|
parts[i] = t.fromExpr()
|
|
}
|
|
sb.WriteString(strings.Join(parts, ", "))
|
|
}
|
|
|
|
for _, j := range b.joins {
|
|
sb.WriteString(" ")
|
|
sb.WriteString(j.joinType)
|
|
sb.WriteString(" ")
|
|
sb.WriteString(j.table.fromExpr())
|
|
sb.WriteString(" ON ")
|
|
sb.WriteString(j.on.fragment)
|
|
args = append(args, j.on.args...)
|
|
}
|
|
|
|
if b.where != nil {
|
|
sb.WriteString(" WHERE ")
|
|
sb.WriteString(b.where.fragment)
|
|
args = append(args, b.where.args...)
|
|
}
|
|
|
|
if len(b.groupBy) > 0 {
|
|
sb.WriteString(" GROUP BY ")
|
|
sb.WriteString(strings.Join(b.groupBy, ", "))
|
|
}
|
|
|
|
if b.having != nil {
|
|
sb.WriteString(" HAVING ")
|
|
sb.WriteString(b.having.fragment)
|
|
args = append(args, b.having.args...)
|
|
}
|
|
|
|
if len(b.orderBy) > 0 {
|
|
sb.WriteString(" ORDER BY ")
|
|
parts := make([]string, len(b.orderBy))
|
|
for i, o := range b.orderBy {
|
|
parts[i] = o.expr
|
|
}
|
|
sb.WriteString(strings.Join(parts, ", "))
|
|
}
|
|
|
|
if b.limit != nil {
|
|
fmt.Fprintf(&sb, " LIMIT %d", *b.limit)
|
|
}
|
|
if b.offset != nil {
|
|
fmt.Fprintf(&sb, " OFFSET %d", *b.offset)
|
|
}
|
|
|
|
return sb.String(), args
|
|
}
|
|
|
|
func (b *SelectBuilder) Query(ctx context.Context, db Querier, dest any) error {
|
|
sql, args := b.Build()
|
|
if b.debug {
|
|
debugQuery(sql, args)
|
|
}
|
|
return QueryAll(ctx, db, dest, sql, args...)
|
|
}
|
|
|
|
func (b *SelectBuilder) QueryRow(ctx context.Context, db Querier, dest any) error {
|
|
sql, args := b.Build()
|
|
if b.debug {
|
|
debugQuery(sql, args)
|
|
}
|
|
return QueryOne(ctx, db, dest, sql, args...)
|
|
}
|
|
|
|
func (b *SelectBuilder) QueryScalarTo(ctx context.Context, db Querier, dest any) error {
|
|
query, args := b.Build()
|
|
if b.debug {
|
|
debugQuery(query, args)
|
|
}
|
|
rows, err := db.QueryContext(ctx, query, args...)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
defer rows.Close()
|
|
if !rows.Next() {
|
|
if err := rows.Err(); err != nil {
|
|
return err
|
|
}
|
|
return sql.ErrNoRows
|
|
}
|
|
return rows.Scan(dest)
|
|
}
|
|
|
|
// InsertBuilder
|
|
|
|
type InsertBuilder struct {
|
|
table TableRef
|
|
columns []string
|
|
values []any
|
|
model any
|
|
returning []string
|
|
debug bool
|
|
}
|
|
|
|
func (b *InsertBuilder) Debug() *InsertBuilder {
|
|
b.debug = true
|
|
return b
|
|
}
|
|
|
|
func InsertInto(table TableExpr) *InsertBuilder {
|
|
return &InsertBuilder{table: table.tableRef()}
|
|
}
|
|
|
|
func (b *InsertBuilder) Columns(cols ...string) *InsertBuilder {
|
|
b.columns = cols
|
|
return b
|
|
}
|
|
|
|
func (b *InsertBuilder) Values(vals ...any) *InsertBuilder {
|
|
b.values = vals
|
|
return b
|
|
}
|
|
|
|
func (b *InsertBuilder) Model(model any) *InsertBuilder {
|
|
b.model = model
|
|
return b
|
|
}
|
|
|
|
func (b *InsertBuilder) Returning(cols ...string) *InsertBuilder {
|
|
b.returning = cols
|
|
return b
|
|
}
|
|
|
|
func (b *InsertBuilder) Build() (string, []any) {
|
|
columns := b.columns
|
|
var args []any
|
|
|
|
if b.model != nil {
|
|
if len(columns) == 0 {
|
|
t := reflect.TypeOf(b.model)
|
|
if t.Kind() == reflect.Ptr {
|
|
t = t.Elem()
|
|
}
|
|
columns = allDBColumns(t)
|
|
}
|
|
args = extractModelValues(b.model, columns)
|
|
} else {
|
|
args = b.values
|
|
}
|
|
|
|
placeholders := make([]string, len(columns))
|
|
for i := range columns {
|
|
placeholders[i] = "?"
|
|
}
|
|
|
|
sql := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)",
|
|
b.table.tableName,
|
|
strings.Join(columns, ", "),
|
|
strings.Join(placeholders, ", "))
|
|
|
|
if len(b.returning) > 0 {
|
|
sql += " RETURNING " + strings.Join(b.returning, ", ")
|
|
}
|
|
|
|
return replaceParams(sql), args
|
|
}
|
|
|
|
func (b *InsertBuilder) Exec(ctx context.Context, db Execer) (sql.Result, error) {
|
|
sql, args := b.Build()
|
|
if b.debug {
|
|
debugQuery(sql, args)
|
|
}
|
|
return db.ExecContext(ctx, sql, args...)
|
|
}
|
|
|
|
func (b *InsertBuilder) QueryRow(ctx context.Context, db Querier, dest any) error {
|
|
sql, args := b.Build()
|
|
if b.debug {
|
|
debugQuery(sql, args)
|
|
}
|
|
return QueryOne(ctx, db, dest, sql, args...)
|
|
}
|
|
|
|
// UpdateBuilder
|
|
|
|
type setClause struct {
|
|
col string
|
|
val any
|
|
}
|
|
|
|
type UpdateBuilder struct {
|
|
table TableRef
|
|
sets []setClause
|
|
setCols []string
|
|
model any
|
|
where *Cond
|
|
debug bool
|
|
}
|
|
|
|
func (b *UpdateBuilder) Debug() *UpdateBuilder {
|
|
b.debug = true
|
|
return b
|
|
}
|
|
|
|
func Update(table TableExpr) *UpdateBuilder {
|
|
return &UpdateBuilder{table: table.tableRef()}
|
|
}
|
|
|
|
func (b *UpdateBuilder) Set(col any, val any) *UpdateBuilder {
|
|
var name string
|
|
switch c := col.(type) {
|
|
case Col:
|
|
name = c.expr
|
|
case string:
|
|
name = c
|
|
default:
|
|
panic(fmt.Sprintf("builder: Set col must be Col or string, got %T", col))
|
|
}
|
|
b.sets = append(b.sets, setClause{col: name, val: val})
|
|
return b
|
|
}
|
|
|
|
func (b *UpdateBuilder) SetColumns(cols ...string) *UpdateBuilder {
|
|
b.setCols = cols
|
|
return b
|
|
}
|
|
|
|
func (b *UpdateBuilder) Model(model any) *UpdateBuilder {
|
|
b.model = model
|
|
return b
|
|
}
|
|
|
|
func (b *UpdateBuilder) Where(cond Cond) *UpdateBuilder {
|
|
b.where = &cond
|
|
return b
|
|
}
|
|
|
|
func (b *UpdateBuilder) Build() (string, []any) {
|
|
var setClauses []string
|
|
var args []any
|
|
|
|
if b.model != nil {
|
|
cols := b.setCols
|
|
if len(cols) == 0 {
|
|
t := reflect.TypeOf(b.model)
|
|
if t.Kind() == reflect.Ptr {
|
|
t = t.Elem()
|
|
}
|
|
cols = allDBColumns(t)
|
|
}
|
|
vals := extractModelValues(b.model, cols)
|
|
for i, col := range cols {
|
|
setClauses = append(setClauses, col+" = ?")
|
|
args = append(args, vals[i])
|
|
}
|
|
}
|
|
|
|
for _, s := range b.sets {
|
|
setClauses = append(setClauses, s.col+" = ?")
|
|
args = append(args, s.val)
|
|
}
|
|
|
|
sql := "UPDATE " + b.table.tableName + " SET " + strings.Join(setClauses, ", ")
|
|
|
|
if b.where != nil {
|
|
sql += " WHERE " + b.where.fragment
|
|
args = append(args, b.where.args...)
|
|
}
|
|
|
|
return replaceParams(sql), args
|
|
}
|
|
|
|
func (b *UpdateBuilder) Exec(ctx context.Context, db Execer) (sql.Result, error) {
|
|
sql, args := b.Build()
|
|
if b.debug {
|
|
debugQuery(sql, args)
|
|
}
|
|
return db.ExecContext(ctx, sql, args...)
|
|
}
|
|
|
|
// DeleteBuilder
|
|
|
|
type DeleteBuilder struct {
|
|
table TableRef
|
|
where *Cond
|
|
debug bool
|
|
}
|
|
|
|
func (b *DeleteBuilder) Debug() *DeleteBuilder {
|
|
b.debug = true
|
|
return b
|
|
}
|
|
|
|
func DeleteFrom(table TableExpr) *DeleteBuilder {
|
|
return &DeleteBuilder{table: table.tableRef()}
|
|
}
|
|
|
|
func (b *DeleteBuilder) Where(cond Cond) *DeleteBuilder {
|
|
b.where = &cond
|
|
return b
|
|
}
|
|
|
|
func (b *DeleteBuilder) Build() (string, []any) {
|
|
sql := "DELETE FROM " + b.table.tableName
|
|
var args []any
|
|
|
|
if b.where != nil {
|
|
sql += " WHERE " + b.where.fragment
|
|
args = append(args, b.where.args...)
|
|
}
|
|
|
|
return replaceParams(sql), args
|
|
}
|
|
|
|
func (b *DeleteBuilder) Exec(ctx context.Context, db Execer) (sql.Result, error) {
|
|
sql, args := b.Build()
|
|
if b.debug {
|
|
debugQuery(sql, args)
|
|
}
|
|
return db.ExecContext(ctx, sql, args...)
|
|
}
|
|
|
|
// debugQuery prints the SQL and args to stdout when debug mode is enabled.
|
|
func debugQuery(sql string, args []any) {
|
|
fmt.Println("\n[dbutil:debug] SQL:", sql)
|
|
if len(args) > 0 {
|
|
fmt.Print("[dbutil:debug] Args: [")
|
|
for i, arg := range args {
|
|
if i > 0 {
|
|
fmt.Print(", ")
|
|
}
|
|
fmt.Printf("%v", arg)
|
|
}
|
|
fmt.Println("]")
|
|
}
|
|
fmt.Println()
|
|
}
|
|
|
|
// Helpers
|
|
|
|
func replaceParams(sql string) string {
|
|
var b strings.Builder
|
|
n := 1
|
|
for i := range len(sql) {
|
|
if sql[i] == '?' {
|
|
fmt.Fprintf(&b, "$%d", n)
|
|
n++
|
|
} else {
|
|
b.WriteByte(sql[i])
|
|
}
|
|
}
|
|
return b.String()
|
|
}
|
|
|
|
func extractModelValues(model any, columns []string) []any {
|
|
v := reflect.ValueOf(model)
|
|
if v.Kind() == reflect.Ptr {
|
|
v = v.Elem()
|
|
}
|
|
t := v.Type()
|
|
|
|
tagMap := make(map[string]int, t.NumField())
|
|
for i := range t.NumField() {
|
|
if tag := t.Field(i).Tag.Get("db"); tag != "" && tag != "-" {
|
|
tagMap[tag] = i
|
|
}
|
|
}
|
|
|
|
vals := make([]any, len(columns))
|
|
for i, col := range columns {
|
|
if idx, ok := tagMap[col]; ok {
|
|
vals[i] = v.Field(idx).Interface()
|
|
}
|
|
}
|
|
return vals
|
|
}
|
|
|
|
func allDBColumns(t reflect.Type) []string {
|
|
var cols []string
|
|
for i := range t.NumField() {
|
|
f := t.Field(i)
|
|
if !f.IsExported() {
|
|
continue
|
|
}
|
|
if dbTag := f.Tag.Get("db"); dbTag != "" && dbTag != "-" {
|
|
cols = append(cols, dbTag)
|
|
}
|
|
if f.Anonymous {
|
|
ft := derefType(f.Type)
|
|
if ft.Kind() == reflect.Struct {
|
|
cols = append(cols, allDBColumns(ft)...)
|
|
}
|
|
}
|
|
}
|
|
return cols
|
|
}
|
|
|
|
func collectColumnsFlat(t reflect.Type, tableRef string) []string {
|
|
var cols []string
|
|
for i := range t.NumField() {
|
|
f := t.Field(i)
|
|
if !f.IsExported() {
|
|
continue
|
|
}
|
|
if dbTag := f.Tag.Get("db"); dbTag != "" && dbTag != "-" {
|
|
if tableRef != "" {
|
|
cols = append(cols, tableRef+"."+dbTag)
|
|
} else {
|
|
cols = append(cols, dbTag)
|
|
}
|
|
continue
|
|
}
|
|
if f.Anonymous {
|
|
ft := derefType(f.Type)
|
|
if ft.Kind() == reflect.Struct {
|
|
cols = append(cols, collectColumnsFlat(ft, tableRef)...)
|
|
}
|
|
}
|
|
}
|
|
return cols
|
|
}
|