Files
gorm/gorm.go
T

545 lines
13 KiB
Go
Raw Normal View History

2020-01-28 23:01:35 +08:00
package gorm
import (
2020-01-29 19:22:44 +08:00
"context"
2020-06-05 10:08:22 +08:00
"database/sql"
2020-03-09 20:37:01 +08:00
"fmt"
"reflect"
2021-04-09 11:43:24 +08:00
"sort"
2020-02-02 14:40:44 +08:00
"sync"
2020-01-28 23:01:35 +08:00
"time"
2020-06-02 09:16:07 +08:00
"gorm.io/gorm/clause"
"gorm.io/gorm/logger"
"gorm.io/gorm/schema"
2020-01-28 23:01:35 +08:00
)
// for Config.cacheStore store PreparedStmtDB key
const preparedStmtDBKey = "preparedStmt"
2020-01-28 23:01:35 +08:00
// Config GORM config
type Config struct {
// GORM perform single create, update, delete operations in transactions by default to ensure database data integrity
2020-02-02 14:40:44 +08:00
// You can disable it by setting `SkipDefaultTransaction` to true
2025-05-25 15:40:40 +08:00
SkipDefaultTransaction bool
DefaultTransactionTimeout time.Duration
2020-01-31 14:17:02 +08:00
// NamingStrategy tables, columns naming strategy
NamingStrategy schema.Namer
// FullSaveAssociations full save associations
FullSaveAssociations bool
2020-01-28 23:01:35 +08:00
// Logger
Logger logger.Interface
// NowFunc the function to be used when creating a new timestamp
NowFunc func() time.Time
2020-06-01 21:26:23 +08:00
// DryRun generate sql without execute
DryRun bool
2020-06-05 10:08:22 +08:00
// PrepareStmt executes the given query in cached statement
PrepareStmt bool
// PrepareStmt cache support LRU expired,
// default maxsize=int64 Max value and ttl=1h
PrepareStmtMaxSize int
PrepareStmtTTL time.Duration
2020-06-05 21:23:20 +08:00
// DisableAutomaticPing
DisableAutomaticPing bool
// DisableForeignKeyConstraintWhenMigrating
DisableForeignKeyConstraintWhenMigrating bool
// IgnoreRelationshipsWhenMigrating
IgnoreRelationshipsWhenMigrating bool
2020-12-16 19:33:35 +08:00
// DisableNestedTransaction disable nested transaction
DisableNestedTransaction bool
2020-08-23 20:08:23 +08:00
// AllowGlobalUpdate allow global update
AllowGlobalUpdate bool
// QueryFields executes the SQL query with all fields of the table
QueryFields bool
2020-12-02 14:59:50 +08:00
// CreateBatchSize default create batch size
CreateBatchSize int
// TranslateError enabling error translation
TranslateError bool
// PropagateUnscoped propagate Unscoped to every other nested statement
PropagateUnscoped bool
2020-06-05 10:08:22 +08:00
2020-03-09 13:10:48 +08:00
// ClauseBuilders clause builder
ClauseBuilders map[string]clause.ClauseBuilder
// ConnPool db conn pool
ConnPool ConnPool
// Dialector database dialector
Dialector
2020-06-23 09:38:51 +08:00
// Plugins registered plugins
Plugins map[string]Plugin
2020-03-09 13:10:48 +08:00
2020-05-31 11:34:59 +08:00
callbacks *callbacks
cacheStore *sync.Map
2020-02-05 11:14:58 +08:00
}
2022-02-09 17:39:01 +08:00
// Apply update config to new config
2021-03-04 17:14:08 +08:00
func (c *Config) Apply(config *Config) error {
2021-03-04 20:37:39 +08:00
if config != c {
*config = *c
}
2021-03-04 17:14:08 +08:00
return nil
}
2022-02-09 17:39:01 +08:00
// AfterInitialize initialize plugins after db connected
2021-03-04 17:14:08 +08:00
func (c *Config) AfterInitialize(db *DB) error {
if db != nil {
for _, plugin := range c.Plugins {
if err := plugin.Initialize(db); err != nil {
return err
}
}
}
return nil
}
2022-02-09 17:39:01 +08:00
// Option gorm option interface
2021-03-04 17:14:08 +08:00
type Option interface {
Apply(*Config) error
AfterInitialize(*DB) error
}
2020-01-28 23:01:35 +08:00
// DB GORM DB definition
type DB struct {
*Config
2020-03-09 13:10:48 +08:00
Error error
RowsAffected int64
Statement *Statement
2020-05-31 23:55:56 +08:00
clone int
2020-01-30 15:14:48 +08:00
}
2020-02-02 14:40:44 +08:00
// Session session config when create session with Session() method
2020-01-30 15:14:48 +08:00
type Session struct {
2020-12-16 19:33:35 +08:00
DryRun bool
PrepareStmt bool
NewDB bool
2022-01-28 19:26:10 +08:00
Initialized bool
2020-12-16 19:33:35 +08:00
SkipHooks bool
SkipDefaultTransaction bool
DisableNestedTransaction bool
AllowGlobalUpdate bool
FullSaveAssociations bool
PropagateUnscoped bool
2020-12-16 19:33:35 +08:00
QueryFields bool
Context context.Context
Logger logger.Interface
NowFunc func() time.Time
CreateBatchSize int
2020-01-30 15:14:48 +08:00
}
// Open initialize db session based on dialector
2021-03-04 17:14:08 +08:00
func Open(dialector Dialector, opts ...Option) (db *DB, err error) {
config := &Config{}
2021-04-09 11:43:24 +08:00
sort.Slice(opts, func(i, j int) bool {
_, isConfig := opts[i].(*Config)
_, isConfig2 := opts[j].(*Config)
return isConfig && !isConfig2
})
if len(opts) > 0 {
if c, ok := opts[0].(*Config); ok {
config = c
} else {
opts = append([]Option{config}, opts...)
}
}
2025-05-22 10:53:47 +08:00
var skipAfterInitialize bool
2021-03-04 17:14:08 +08:00
for _, opt := range opts {
if opt != nil {
2022-03-31 20:57:20 +08:00
if applyErr := opt.Apply(config); applyErr != nil {
return nil, applyErr
2021-03-04 17:14:08 +08:00
}
2021-03-26 14:20:42 +08:00
defer func(opt Option) {
2025-05-22 10:53:47 +08:00
if skipAfterInitialize {
return
}
2021-03-29 18:36:01 +08:00
if errr := opt.AfterInitialize(db); errr != nil {
err = errr
}
2021-03-26 14:20:42 +08:00
}(opt)
2021-03-04 17:14:08 +08:00
}
2020-02-02 08:35:01 +08:00
}
2021-03-11 10:29:52 +08:00
if d, ok := dialector.(interface{ Apply(*Config) error }); ok {
if err = d.Apply(config); err != nil {
return
}
}
2020-01-31 14:17:02 +08:00
if config.NamingStrategy == nil {
2023-05-29 19:00:48 -07:00
config.NamingStrategy = schema.NamingStrategy{IdentifierMaxLength: 64} // Default Identifier length is 64
2020-01-31 14:17:02 +08:00
}
2020-02-02 08:35:01 +08:00
if config.Logger == nil {
config.Logger = logger.Default
}
if config.NowFunc == nil {
config.NowFunc = func() time.Time { return time.Now().Local() }
2020-02-02 08:35:01 +08:00
}
2020-03-09 13:10:48 +08:00
if dialector != nil {
config.Dialector = dialector
}
2020-06-23 10:36:45 +08:00
if config.Plugins == nil {
config.Plugins = map[string]Plugin{}
}
2020-03-09 13:10:48 +08:00
if config.cacheStore == nil {
config.cacheStore = &sync.Map{}
}
2020-05-31 23:55:56 +08:00
db = &DB{Config: config, clone: 1}
2020-02-02 08:35:01 +08:00
2020-03-09 20:37:01 +08:00
db.callbacks = initializeCallbacks(db)
2020-02-02 19:32:27 +08:00
2020-05-29 22:34:35 +08:00
if config.ClauseBuilders == nil {
config.ClauseBuilders = map[string]clause.ClauseBuilder{}
}
2020-06-05 21:23:20 +08:00
if config.Dialector != nil {
err = config.Dialector.Initialize(db)
if err != nil {
if db, _ := db.DB(); db != nil {
_ = db.Close()
}
2025-05-22 10:53:47 +08:00
// DB is not initialized, so we skip AfterInitialize
skipAfterInitialize = true
return
}
if config.TranslateError {
if _, ok := db.Dialector.(ErrorTranslator); !ok {
config.Logger.Warn(context.Background(), "The TranslateError option is enabled, but the Dialector %s does not implement ErrorTranslator.", db.Dialector.Name())
}
}
2020-02-02 08:35:01 +08:00
}
2020-06-02 11:09:17 +08:00
2020-06-05 10:08:22 +08:00
if config.PrepareStmt {
preparedStmt := NewPreparedStmtDB(db.ConnPool, config.PrepareStmtMaxSize, config.PrepareStmtTTL)
db.cacheStore.Store(preparedStmtDBKey, preparedStmt)
2020-07-28 14:26:09 +08:00
db.ConnPool = preparedStmt
2020-06-05 10:08:22 +08:00
}
2020-06-05 21:23:20 +08:00
db.Statement = &Statement{
DB: db,
ConnPool: db.ConnPool,
Context: context.Background(),
Clauses: map[string]clause.Clause{},
2020-06-05 10:08:22 +08:00
}
2020-06-05 21:23:20 +08:00
if err == nil && !config.DisableAutomaticPing {
2020-06-02 11:09:17 +08:00
if pinger, ok := db.ConnPool.(interface{ Ping() error }); ok {
err = pinger.Ping()
}
}
if err != nil {
config.Logger.Error(context.Background(), "failed to initialize database, got error %v", err)
}
2020-02-02 08:35:01 +08:00
return
2020-01-30 15:14:48 +08:00
}
// Session create new db session
2020-03-09 20:37:01 +08:00
func (db *DB) Session(config *Session) *DB {
2020-01-30 15:14:48 +08:00
var (
2020-05-31 23:55:56 +08:00
txConfig = *db.Config
tx = &DB{
Config: &txConfig,
Statement: db.Statement,
2021-01-07 11:45:40 +08:00
Error: db.Error,
2020-05-31 23:55:56 +08:00
clone: 1,
}
2020-01-30 15:14:48 +08:00
)
2020-12-02 14:59:50 +08:00
if config.CreateBatchSize > 0 {
tx.Config.CreateBatchSize = config.CreateBatchSize
}
if config.SkipDefaultTransaction {
tx.Config.SkipDefaultTransaction = true
}
2020-08-23 20:08:23 +08:00
if config.AllowGlobalUpdate {
txConfig.AllowGlobalUpdate = true
}
if config.FullSaveAssociations {
txConfig.FullSaveAssociations = true
}
if config.PropagateUnscoped {
txConfig.PropagateUnscoped = true
}
2020-11-17 15:19:58 +08:00
if config.Context != nil || config.PrepareStmt || config.SkipHooks {
2020-06-05 21:23:20 +08:00
tx.Statement = tx.Statement.clone()
tx.Statement.DB = tx
2020-11-17 15:19:58 +08:00
}
if config.Context != nil {
2020-05-31 23:55:56 +08:00
tx.Statement.Context = config.Context
}
2020-06-05 10:08:22 +08:00
if config.PrepareStmt {
var preparedStmt *PreparedStmtDB
if v, ok := db.cacheStore.Load(preparedStmtDBKey); ok {
preparedStmt = v.(*PreparedStmtDB)
} else {
preparedStmt = NewPreparedStmtDB(db.ConnPool, db.PrepareStmtMaxSize, db.PrepareStmtTTL)
db.cacheStore.Store(preparedStmtDBKey, preparedStmt)
2020-06-05 10:08:22 +08:00
}
switch t := tx.Statement.ConnPool.(type) {
case Tx:
tx.Statement.ConnPool = &PreparedStmtTX{
Tx: t,
PreparedStmtDB: preparedStmt,
}
default:
tx.Statement.ConnPool = &PreparedStmtDB{
ConnPool: db.Config.ConnPool,
Mux: preparedStmt.Mux,
Stmts: preparedStmt.Stmts,
}
}
txConfig.ConnPool = tx.Statement.ConnPool
txConfig.PrepareStmt = true
2020-06-05 10:08:22 +08:00
}
2020-11-17 15:19:58 +08:00
if config.SkipHooks {
2020-11-17 17:49:43 +08:00
tx.Statement.SkipHooks = true
2020-11-17 15:19:58 +08:00
}
2020-12-16 19:33:35 +08:00
if config.DisableNestedTransaction {
txConfig.DisableNestedTransaction = true
}
if !config.NewDB {
2020-06-05 21:23:20 +08:00
tx.clone = 2
2020-01-30 15:14:48 +08:00
}
2020-06-01 21:26:23 +08:00
if config.DryRun {
tx.Config.DryRun = true
}
if config.QueryFields {
tx.Config.QueryFields = true
}
2020-01-30 15:14:48 +08:00
if config.Logger != nil {
2020-05-31 23:55:56 +08:00
tx.Config.Logger = config.Logger
2020-01-30 15:14:48 +08:00
}
if config.NowFunc != nil {
2020-05-31 23:55:56 +08:00
tx.Config.NowFunc = config.NowFunc
2020-01-30 15:14:48 +08:00
}
2022-01-28 19:26:10 +08:00
if config.Initialized {
tx = tx.getInstance()
}
2020-05-31 23:55:56 +08:00
return tx
2020-01-29 19:22:44 +08:00
}
// WithContext change current instance db's context to ctx
2020-03-09 20:37:01 +08:00
func (db *DB) WithContext(ctx context.Context) *DB {
return db.Session(&Session{Context: ctx})
2020-01-30 15:14:48 +08:00
}
// Debug start debug mode
2020-03-09 20:37:01 +08:00
func (db *DB) Debug() (tx *DB) {
2022-07-18 18:06:45 +08:00
tx = db.getInstance()
return tx.Session(&Session{
Logger: db.Logger.LogMode(logger.Info),
2020-05-31 23:55:56 +08:00
})
2020-01-30 15:14:48 +08:00
}
2020-01-29 19:22:44 +08:00
// Set store value with key into current db instance's context
2020-03-09 20:37:01 +08:00
func (db *DB) Set(key string, value interface{}) *DB {
2020-01-29 19:22:44 +08:00
tx := db.getInstance()
tx.Statement.Settings.Store(key, value)
return tx
}
// Get get value with key from current db instance's context
2020-03-09 20:37:01 +08:00
func (db *DB) Get(key string) (interface{}, bool) {
2020-06-05 21:23:20 +08:00
return db.Statement.Settings.Load(key)
2020-01-29 19:22:44 +08:00
}
2020-05-31 23:55:56 +08:00
// InstanceSet store value with key into current db instance's context
func (db *DB) InstanceSet(key string, value interface{}) *DB {
tx := db.getInstance()
tx.Statement.Settings.Store(fmt.Sprintf("%p", tx.Statement)+key, value)
return tx
}
// InstanceGet get value with key from current db instance's context
func (db *DB) InstanceGet(key string) (interface{}, bool) {
2020-06-05 21:23:20 +08:00
return db.Statement.Settings.Load(fmt.Sprintf("%p", db.Statement) + key)
2020-05-31 23:55:56 +08:00
}
2020-02-02 08:35:01 +08:00
// Callback returns callback manager
2020-03-09 20:37:01 +08:00
func (db *DB) Callback() *callbacks {
2020-02-02 08:35:01 +08:00
return db.callbacks
}
2020-03-09 13:10:48 +08:00
// AddError add error to db
2020-04-19 23:11:56 +08:00
func (db *DB) AddError(err error) error {
if err != nil {
if db.Config.TranslateError {
if errTranslator, ok := db.Dialector.(ErrorTranslator); ok {
err = errTranslator.Translate(err)
}
}
if db.Error == nil {
db.Error = err
} else {
db.Error = fmt.Errorf("%v; %w", db.Error, err)
}
2020-03-09 20:37:01 +08:00
}
2020-04-19 23:11:56 +08:00
return db.Error
2020-03-09 13:10:48 +08:00
}
2020-06-17 19:56:03 +08:00
// DB returns `*sql.DB`
func (db *DB) DB() (*sql.DB, error) {
connPool := db.ConnPool
if db.Statement != nil && db.Statement.ConnPool != nil {
connPool = db.Statement.ConnPool
}
if tx, ok := connPool.(*sql.Tx); ok && tx != nil {
return (*sql.DB)(reflect.ValueOf(tx).Elem().FieldByName("db").UnsafePointer()), nil
}
2021-03-19 15:54:32 +08:00
if dbConnector, ok := connPool.(GetDBConnector); ok && dbConnector != nil {
if sqldb, err := dbConnector.GetDBConn(); sqldb != nil || err != nil {
return sqldb, err
}
2020-06-17 19:56:03 +08:00
}
if sqldb, ok := connPool.(*sql.DB); ok && sqldb != nil {
2020-06-17 19:56:03 +08:00
return sqldb, nil
}
2021-04-19 06:03:39 -07:00
return nil, ErrInvalidDB
2020-06-17 19:56:03 +08:00
}
2020-03-09 20:37:01 +08:00
func (db *DB) getInstance() *DB {
2020-05-31 23:55:56 +08:00
if db.clone > 0 {
2021-04-09 11:07:14 +08:00
tx := &DB{Config: db.Config, Error: db.Error}
2020-05-31 11:34:59 +08:00
2020-06-05 21:23:20 +08:00
if db.clone == 1 {
// clone with new statement
2020-06-05 10:08:22 +08:00
tx.Statement = &Statement{
DB: tx,
ConnPool: db.Statement.ConnPool,
Context: db.Statement.Context,
Clauses: map[string]clause.Clause{},
Vars: make([]interface{}, 0, 8),
SkipHooks: db.Statement.SkipHooks,
2020-06-05 10:08:22 +08:00
}
if db.Config.PropagateUnscoped {
tx.Statement.Unscoped = db.Statement.Unscoped
}
2020-06-05 21:23:20 +08:00
} else {
// with clone statement
tx.Statement = db.Statement.clone()
tx.Statement.DB = tx
2020-01-30 15:14:48 +08:00
}
2020-05-31 23:55:56 +08:00
return tx
2020-01-29 19:22:44 +08:00
}
return db
}
2020-05-30 16:47:16 +08:00
2022-02-09 17:39:01 +08:00
// Expr returns clause.Expr, which can be used to pass SQL expression as params
2020-05-30 16:47:16 +08:00
func Expr(expr string, args ...interface{}) clause.Expr {
return clause.Expr{SQL: expr, Vars: args}
}
2020-06-08 13:45:41 +08:00
2022-02-09 17:39:01 +08:00
// SetupJoinTable setup join table schema
2020-06-08 13:45:41 +08:00
func (db *DB) SetupJoinTable(model interface{}, field string, joinTable interface{}) error {
var (
tx = db.getInstance()
stmt = tx.Statement
modelSchema, joinSchema *schema.Schema
)
err := stmt.Parse(model)
if err != nil {
2020-06-08 13:45:41 +08:00
return err
}
modelSchema = stmt.Schema
2020-06-08 13:45:41 +08:00
err = stmt.Parse(joinTable)
if err != nil {
2020-06-08 13:45:41 +08:00
return err
}
joinSchema = stmt.Schema
2020-06-08 13:45:41 +08:00
relation, ok := modelSchema.Relationships.Relations[field]
isRelation := ok && relation.JoinTable != nil
if !isRelation {
2022-08-12 16:46:18 +03:00
return fmt.Errorf("failed to find relation: %s", field)
2020-06-08 13:45:41 +08:00
}
for _, ref := range relation.References {
f := joinSchema.LookUpField(ref.ForeignKey.DBName)
if f == nil {
return fmt.Errorf("missing field %s for join table", ref.ForeignKey.DBName)
}
f.DataType = ref.ForeignKey.DataType
f.GORMDataType = ref.ForeignKey.GORMDataType
if f.Size == 0 {
f.Size = ref.ForeignKey.Size
}
ref.ForeignKey = f
}
for name, rel := range relation.JoinTable.Relationships.Relations {
if _, ok := joinSchema.Relationships.Relations[name]; !ok {
rel.Schema = joinSchema
joinSchema.Relationships.Relations[name] = rel
}
}
relation.JoinTable = joinSchema
2020-06-08 13:45:41 +08:00
return nil
}
2020-06-23 09:38:51 +08:00
2022-02-09 17:39:01 +08:00
// Use use plugin
func (db *DB) Use(plugin Plugin) error {
2020-06-23 09:38:51 +08:00
name := plugin.Name()
if _, ok := db.Plugins[name]; ok {
2020-06-23 09:38:51 +08:00
return ErrRegistered
}
if err := plugin.Initialize(db); err != nil {
return err
}
db.Plugins[name] = plugin
return nil
2020-06-23 09:38:51 +08:00
}
// ToSQL for generate SQL string.
//
// db.ToSQL(func(tx *gorm.DB) *gorm.DB {
// return tx.Model(&User{}).Where(&User{Name: "foo", Age: 20})
// .Limit(10).Offset(5)
// .Order("name ASC")
// .First(&User{})
// })
func (db *DB) ToSQL(queryFn func(tx *DB) *DB) string {
2025-05-25 15:40:40 +08:00
tx := queryFn(db.Session(&Session{DryRun: true, SkipDefaultTransaction: true}).getInstance())
stmt := tx.Statement
return db.Dialector.Explain(stmt.SQL.String(), stmt.Vars...)
}