aboutsummaryrefslogtreecommitdiff
path: root/plan.go
diff options
context:
space:
mode:
authorStefan Majewsky <majewsky@gmx.net>2026-09-17 16:51:17 +0200
committerStefan Majewsky <majewsky@gmx.net>2026-09-17 16:51:17 +0200
commit1825c3040f5ee71ad26185a7a08f2657ede94df5 (patch)
treebf65a15764fae5a947ef1a53d9d6c54a8ccdae11 /plan.go
parent1b8935ec8cc8ebe9ad74ecd12141e8ad115cebfb (diff)
downloadgo-oblast-1825c3040f5ee71ad26185a7a08f2657ede94df5.tar.gz
rebase onto gg@v1.16.0/oblast
Diffstat (limited to 'plan.go')
-rw-r--r--plan.go488
1 files changed, 0 insertions, 488 deletions
diff --git a/plan.go b/plan.go
deleted file mode 100644
index 6b5f002..0000000
--- a/plan.go
+++ /dev/null
@@ -1,488 +0,0 @@
-// SPDX-FileCopyrightText: 2026 Stefan Majewsky <majewsky@gmx.net>
-// SPDX-License-Identifier: Apache-2.0
-
-package oblast
-
-import (
- "cmp"
- "errors"
- "fmt"
- "reflect"
- "slices"
- "strings"
- "sync"
-)
-
-// planOpts holds additional arguments to buildPlan().
-type planOpts struct {
- ReadOnly bool
- StructTagKey string // defaults to "db"
- TableName string
- PrimaryKeyColumnNames []string
-}
-
-func collectPlanOptions(popts []PlanOption) planOpts {
- opts := planOpts{
- StructTagKey: "db",
- }
- for _, popt := range popts {
- popt(&opts)
- }
- return opts
-}
-
-type planCacheKey struct {
- Type reflect.Type
- Dialect string
- ReadOnly bool
- StructTagKey string
- TableName string
- PrimaryKeyColumnNames string
-}
-
-var (
- generatedPlans = make(map[planCacheKey]plan)
- generatedPlansMutex sync.RWMutex
- generatedTuplePlans = make(map[reflect.Type]plan)
- generatedTuplePlansMutex sync.Mutex
-)
-
-// getOrBuildPlan is like [buildPlan], but caches generated plans and tries to reuse cached plans.
-func getOrBuildPlan(t reflect.Type, dialect Dialect, opts planOpts) (plan, error) {
- key := planCacheKey{
- Type: t,
- Dialect: dialect.String(),
- ReadOnly: opts.ReadOnly,
- StructTagKey: opts.StructTagKey,
- TableName: opts.TableName,
- PrimaryKeyColumnNames: strings.Join(opts.PrimaryKeyColumnNames, "\000"),
- }
-
- generatedPlansMutex.RLock()
- p, ok := generatedPlans[key]
- generatedPlansMutex.RUnlock()
- if ok {
- return p, nil
- }
-
- p, err := buildPlan(t, dialect, opts)
- if err != nil {
- return plan{}, err
- }
- generatedPlansMutex.Lock()
- generatedPlans[key] = p
- generatedPlansMutex.Unlock()
- return p, nil
-}
-
-// getOrBuildTuplePlan is like [getOrBuildPlan], but generates the plans used by TupleSelect() et al
-func getOrBuildTuplePlan(t reflect.Type) plan {
- generatedTuplePlansMutex.Lock()
- defer generatedTuplePlansMutex.Unlock()
-
- p, ok := generatedTuplePlans[t]
- if !ok {
- indexes := make([][]int, t.NumField())
- for idx := range indexes {
- indexes[idx] = []int{idx}
- }
- p = plan{
- TypeName: t.Name(),
- StaticIndexes: indexes,
- }
- generatedTuplePlans[t] = p
- }
- return p
-}
-
-// plan holds all information that we can derive from reflecting on a given type.
-// The queries held within are only valid within the context of a given SQL dialect.
-type plan struct {
- TypeName string // for use in error messages
- TableName string // from info.TableNameIs marker (if any)
- AllColumnNames []string // in order of struct fields (not set for TupleSelect() plans)
- PrimaryKeyColumnNames []string // from info.PrimaryKeyIs marker (if any)
- AutoColumnNames []string // subset of AllColumnNames where field has `,auto` marker
-
- // Field index (i.e. argument for reflect.Value.FieldByIndex()) for each column name. Not set for TupleSelect() plans.
- IndexByColumnName map[string][]int
- // Select indexes for TupleSelect() plans. Always of the form [[0], [1], [2], ..., [N]]. Not set for regular plans.
- StaticIndexes [][]int
- // Pointer-typed fields that need to be initialized before scanning into this type.
- TransparentPointerStructFields []fieldInfo
-
- // Whether the INSERT query uses QueryRow or Exec.
- // - When no auto-generated values are collected, or when a single value can be collected through LastInsertId(),
- // this will be false because Exec() is more memory-efficient than QueryRow(); it does not have to allocate an *sql.Rows instance.
- // - Otherwise, i.e. when auto-generated values are collected with a RETURNING clause,
- // this will be true because Exec() does not support scanning result values.
- InsertUsesQueryRow bool
- // If InsertUsesQueryRow = false and a primary key is collected from LastInsertId(),
- // this decides whether we write it with reflect.Value.SetInt() or reflect.Value.SetUint().
- LastInsertIdIsUnsigned bool
-
- // Planned queries.
- Select plannedQuery // only `SELECT ... FROM ... WHERE `; user supplies the rest during Select{,One}Where()
- Insert plannedQuery
- Upsert plannedQuery
- Update plannedQuery
- Delete plannedQuery
-
- // Whether Insert/Upsert/Update/Delete query planning was inhibited by the ReadOnly option.
- // This information is preserved in order to render more useful error messages.
- ReadOnly bool
-}
-
-// fieldInfo appears in type plan.
-type fieldInfo struct {
- Name string
- Index []int
- ContainsPrimaryKey bool
-}
-
-// plannedQuery appears in type plan.
-type plannedQuery struct {
- // Empty if the respective query type is not supported by this plan for lack of the required marker types.
- Query string
- // Arguments for reflect.Value.FieldByIndex() in the correct order for the query arguments of the above query.
- ArgumentIndexes [][]int
- // Arguments for reflect.Value.FieldByIndex() in the correct order for the Scan() arguments of the above query.
- ScanIndexes [][]int
-}
-
-// buildPlan creates a new plan for the given struct type.
-func buildPlan(t reflect.Type, dialect Dialect, opts planOpts) (plan, error) {
- if t.Kind() != reflect.Struct {
- return plan{}, fmt.Errorf("expected struct type, but got kind %q", t.Kind().String())
- }
-
- var p = plan{
- TypeName: t.Name(),
- TableName: opts.TableName,
- PrimaryKeyColumnNames: opts.PrimaryKeyColumnNames,
- IndexByColumnName: make(map[string][]int),
- ReadOnly: opts.ReadOnly,
- }
-
- var (
- indexesOfOpaqueStructs [][]int
- indexesOfUnusedTransparentStructs [][]int
- )
- isWithin := func(fieldIndex, structIndex []int) bool {
- // returns whether `structIndex` is a prefix of `fieldIndex` (i.e. whether the field is contained within the struct)
- return len(fieldIndex) > len(structIndex) && slices.Equal(fieldIndex[0:len(structIndex)], structIndex)
- }
-
- // discover addressable fields in this type, collect information from markers and tags
- for _, field := range assignableFields(t) {
- // recurse into struct fields (i.e. ignore the struct itself and consider its members instead)
- // unless the field itself has a `db:"..."` tag
- if field.Type.Kind() == reflect.Struct || (field.Type.Kind() == reflect.Pointer && field.Type.Elem().Kind() == reflect.Struct) {
- if field.Tag.Get(opts.StructTagKey) == "" {
- indexesOfUnusedTransparentStructs = append(indexesOfUnusedTransparentStructs, field.Index)
- if field.Type.Kind() == reflect.Pointer {
- // remember that, when scanning into a record of type `t`, we need to write a non-nil zeroed struct into this field
- // to enable taking an address of its mapped member fields
- p.TransparentPointerStructFields = append(p.TransparentPointerStructFields, fieldInfo{
- Name: field.Name,
- Index: field.Index,
- ContainsPrimaryKey: false, // might be set later
- })
- }
- continue
- }
- indexesOfOpaqueStructs = append(indexesOfOpaqueStructs, field.Index)
- }
-
- // ignore fields that are within a struct type that is mapped as a whole
- if slices.ContainsFunc(indexesOfOpaqueStructs, func(index []int) bool {
- return isWithin(field.Index, index)
- }) {
- continue
- }
-
- // check `db:"..."` tag, ignore fields that are declared with column name "-"
- tags := strings.Split(strings.TrimSpace(field.Tag.Get(opts.StructTagKey)), ",")
- columnName, extraTags := cmp.Or(tags[0], field.Name), tags[1:]
- if columnName == "-" {
- continue
- }
-
- if otherIndex := p.IndexByColumnName[columnName]; otherIndex != nil {
- return plan{}, fmt.Errorf(
- "duplicate tag `%s:%q` on field index %v, but also on field index %v",
- opts.StructTagKey, columnName, otherIndex, field.Index,
- )
- }
- p.IndexByColumnName[columnName] = field.Index
- p.AllColumnNames = append(p.AllColumnNames, columnName)
-
- // track whether transparent structs contain fields that are mapped
- restartIteration:
- for idx, index := range indexesOfUnusedTransparentStructs {
- if isWithin(field.Index, index) {
- indexesOfUnusedTransparentStructs = slices.Delete(indexesOfUnusedTransparentStructs, idx, idx+1)
- goto restartIteration
- }
- }
-
- // track which transparent pointer structs contain PK fields
- if slices.Contains(p.PrimaryKeyColumnNames, columnName) {
- for idx, tpsField := range p.TransparentPointerStructFields {
- if isWithin(field.Index, tpsField.Index) {
- p.TransparentPointerStructFields[idx].ContainsPrimaryKey = true
- }
- }
- }
-
- for _, tag := range extraTags {
- switch tag {
- case "auto":
- p.AutoColumnNames = append(p.AutoColumnNames, columnName)
- default:
- return plan{}, fmt.Errorf("unknown option `%s:%q` on field %q", opts.StructTagKey, ","+tag, field.Name)
- }
- }
- }
-
- // validation: transparent structs need to have at least one of their members mapped
- // (this property is most often violated when a user of a library-defined type is not aware that this type is a struct under the hood,
- // e.g. a field like "CreatedAt time.Time" needs to have a tag like `db:"created_at"`,
- // otherwise nothing will be mapped because time.Time does not have any exported fields)
- for _, index := range indexesOfUnusedTransparentStructs {
- field := t.FieldByIndex(index)
- return plan{}, fmt.Errorf(
- "field %q of type %s does not contain any mapped fields (to map this whole field to a DB column, add an explicit `%s:\"...\"` tag)",
- field.Name, field.Type.String(), opts.StructTagKey,
- )
- }
-
- // validation: defining a primary key only makes sense for records that map onto a single table
- if len(p.PrimaryKeyColumnNames) > 0 && p.TableName == "" {
- return plan{}, errors.New("cannot declare a primary key without also providing the TableNameIs option")
- }
-
- // validation: oblast.PrimaryKeyInfo must refer to columns that exist
- for _, columnName := range p.PrimaryKeyColumnNames {
- _, ok := p.IndexByColumnName[columnName]
- if !ok {
- return plan{}, fmt.Errorf("no field has tag `%s:%q`, but a field of this name was declared in the primary key", opts.StructTagKey, columnName)
- }
- }
-
- // pick strategy for INSERT
- if p.TableName != "" {
- switch len(p.AutoColumnNames) {
- case 0:
- p.InsertUsesQueryRow = false
- case 1:
- if dialect.CanUseLastInsertId() {
- columnName := p.AutoColumnNames[0]
- field := t.FieldByIndex(p.IndexByColumnName[columnName])
- switch field.Type.Kind() { //nolint:exhaustive // false positive
- case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
- p.InsertUsesQueryRow = false
- p.LastInsertIdIsUnsigned = false
- case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
- p.InsertUsesQueryRow = false
- p.LastInsertIdIsUnsigned = true
- default:
- p.InsertUsesQueryRow = true
- }
- } else {
- p.InsertUsesQueryRow = true
- }
- default:
- p.InsertUsesQueryRow = true
- }
- }
-
- // prepare query strings
- p.Select = p.buildSelectQueryIfPossible(dialect)
- if !opts.ReadOnly {
- p.Insert = p.buildInsertQueryIfPossible(dialect, false)
- p.Upsert = p.buildInsertQueryIfPossible(dialect, true)
- p.Update = p.buildUpdateQueryIfPossible(dialect)
- p.Delete = p.buildDeleteQueryIfPossible(dialect)
- }
-
- return p, nil
-}
-
-// Like reflect.VisibleFields(), but considers all fields within the type that
-// are assignable (i.e. `v.FieldByIndex(...).Set(...)` does not panic).
-func assignableFields(t reflect.Type) (result []reflect.StructField) {
- for field := range t.Fields() {
- // assignment is allowed for exported or embedded fields only
- if field.IsExported() || field.Anonymous {
- result = append(result, field)
-
- // recurse into struct fields
- ft := field.Type
- if ft.Kind() == reflect.Pointer {
- ft = ft.Elem()
- }
- if ft.Kind() == reflect.Struct {
- for _, subfield := range assignableFields(ft) {
- subfield.Index = append(slices.Clone(field.Index), subfield.Index...)
- result = append(result, subfield)
- }
- }
- }
- }
-
- return result
-}
-
-func (p plan) getNonAutoColumnNames() []string {
- result := make([]string, 0, len(p.AllColumnNames)-len(p.AutoColumnNames))
- for _, columnName := range p.AllColumnNames {
- if !slices.Contains(p.AutoColumnNames, columnName) {
- result = append(result, columnName)
- }
- }
- return result
-}
-
-func (p plan) getNonPrimaryKeyColumnNames() []string {
- result := make([]string, 0, len(p.AllColumnNames)-len(p.PrimaryKeyColumnNames))
- for _, columnName := range p.AllColumnNames {
- if !slices.Contains(p.PrimaryKeyColumnNames, columnName) {
- result = append(result, columnName)
- }
- }
- return result
-}
-
-func (p plan) buildSelectQueryIfPossible(dialect Dialect) plannedQuery {
- if p.TableName == "" {
- return plannedQuery{Query: ""}
- }
-
- var (
- scanIndexes = make([][]int, len(p.AllColumnNames))
- quotedColumnNames = make([]string, len(p.AllColumnNames))
- )
- for idx, columnName := range p.AllColumnNames {
- scanIndexes[idx] = p.IndexByColumnName[columnName]
- quotedColumnNames[idx] = dialect.QuoteIdentifier(columnName)
- }
-
- query := fmt.Sprintf(
- `SELECT %s FROM %s WHERE `,
- strings.Join(quotedColumnNames, ", "),
- dialect.QuoteIdentifier(p.TableName),
- )
- return plannedQuery{query, nil, scanIndexes}
-}
-
-func (p plan) buildInsertQueryIfPossible(dialect Dialect, isUpsert bool) plannedQuery {
- if p.TableName == "" || len(p.AllColumnNames) == 0 {
- return plannedQuery{Query: ""}
- }
- nonAutoColumnNames := p.getNonAutoColumnNames()
- if len(nonAutoColumnNames) == 0 {
- return plannedQuery{Query: ""}
- }
-
- // UPSERT queries specifically are only generated if we have non-auto primary keys:
- // - cannot hit a key conflict if there are no keys
- // - cannot hit a key conflict on insert if all keys are autogenerated (and thus we never supply them during INSERT)
- if isUpsert && !slices.ContainsFunc(p.PrimaryKeyColumnNames, func(n string) bool { return !slices.Contains(p.AutoColumnNames, n) }) {
- return plannedQuery{Query: ""}
- }
-
- var (
- argumentIndexes = make([][]int, len(nonAutoColumnNames))
- scanIndexes [][]int
- quotedColumnNames = make([]string, len(nonAutoColumnNames))
- quotedPlaceholders = make([]string, len(nonAutoColumnNames))
- )
- for idx, columnName := range nonAutoColumnNames {
- argumentIndexes[idx] = p.IndexByColumnName[columnName]
- quotedColumnNames[idx] = dialect.QuoteIdentifier(columnName)
- quotedPlaceholders[idx] = dialect.Placeholder(idx)
- }
- if len(p.AutoColumnNames) > 0 {
- scanIndexes = make([][]int, len(p.AutoColumnNames))
- for idx, columnName := range p.AutoColumnNames {
- scanIndexes[idx] = p.IndexByColumnName[columnName]
- }
- }
-
- query := fmt.Sprintf(
- `INSERT INTO %s (%s) VALUES (%s)`,
- dialect.QuoteIdentifier(p.TableName),
- strings.Join(quotedColumnNames, ", "),
- strings.Join(quotedPlaceholders, ", "),
- )
- if isUpsert {
- query += dialect.UpsertClause(p.PrimaryKeyColumnNames, p.getNonPrimaryKeyColumnNames())
- }
- if len(p.AutoColumnNames) > 0 && p.InsertUsesQueryRow {
- quotedAutoColumns := make([]string, len(p.AutoColumnNames))
- for idx, name := range p.AutoColumnNames {
- quotedAutoColumns[idx] = dialect.QuoteIdentifier(name)
- }
- query += ` RETURNING ` + strings.Join(quotedAutoColumns, ", ")
- }
- return plannedQuery{query, argumentIndexes, scanIndexes}
-}
-
-func (p plan) buildUpdateQueryIfPossible(dialect Dialect) plannedQuery {
- if p.TableName == "" || len(p.PrimaryKeyColumnNames) == 0 {
- return plannedQuery{Query: ""}
- }
- nonPrimaryKeyColumnNames := p.getNonPrimaryKeyColumnNames()
- if len(nonPrimaryKeyColumnNames) == 0 {
- return plannedQuery{Query: ""}
- }
-
- var (
- setArgumentIndexes = make([][]int, len(nonPrimaryKeyColumnNames))
- setClauses = make([]string, len(nonPrimaryKeyColumnNames))
- )
- for idx, columnName := range nonPrimaryKeyColumnNames {
- setArgumentIndexes[idx] = p.IndexByColumnName[columnName]
- setClauses[idx] = fmt.Sprintf("%s = %s", dialect.QuoteIdentifier(columnName), dialect.Placeholder(idx))
- }
-
- var (
- whereArgumentIndexes = make([][]int, len(p.PrimaryKeyColumnNames))
- whereClauses = make([]string, len(p.PrimaryKeyColumnNames))
- )
- for idx, columnName := range p.PrimaryKeyColumnNames {
- whereArgumentIndexes[idx] = p.IndexByColumnName[columnName]
- whereClauses[idx] = fmt.Sprintf("%s = %s", dialect.QuoteIdentifier(columnName), dialect.Placeholder(idx+len(setClauses)))
- }
-
- query := fmt.Sprintf(
- `UPDATE %s SET %s WHERE %s`,
- dialect.QuoteIdentifier(p.TableName),
- strings.Join(setClauses, ", "),
- strings.Join(whereClauses, " AND "),
- )
- return plannedQuery{query, slices.Concat(setArgumentIndexes, whereArgumentIndexes), nil}
-}
-
-func (p plan) buildDeleteQueryIfPossible(dialect Dialect) plannedQuery {
- if p.TableName == "" || len(p.PrimaryKeyColumnNames) == 0 {
- return plannedQuery{Query: ""}
- }
-
- var (
- argumentIndexes = make([][]int, len(p.PrimaryKeyColumnNames))
- clauses = make([]string, len(p.PrimaryKeyColumnNames))
- )
- for idx, columnName := range p.PrimaryKeyColumnNames {
- argumentIndexes[idx] = p.IndexByColumnName[columnName]
- clauses[idx] = fmt.Sprintf("%s = %s", dialect.QuoteIdentifier(columnName), dialect.Placeholder(idx))
- }
-
- query := fmt.Sprintf(
- `DELETE FROM %s WHERE %s`,
- dialect.QuoteIdentifier(p.TableName),
- strings.Join(clauses, " AND "),
- )
- return plannedQuery{query, argumentIndexes, nil}
-}