gorm/gorm.go

312 lines
6.8 KiB
Go
Raw Normal View History

2020-01-28 18:01:35 +03:00
package gorm
import (
2020-01-29 14:22:44 +03:00
"context"
2020-06-05 05:08:22 +03:00
"database/sql"
2020-03-09 15:37:01 +03:00
"fmt"
2020-02-02 09:40:44 +03:00
"sync"
2020-01-28 18:01:35 +03:00
"time"
2020-06-02 04:16:07 +03:00
"gorm.io/gorm/clause"
"gorm.io/gorm/logger"
"gorm.io/gorm/schema"
2020-01-28 18:01:35 +03: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 09:40:44 +03:00
// You can disable it by setting `SkipDefaultTransaction` to true
SkipDefaultTransaction bool
2020-01-31 09:17:02 +03:00
// NamingStrategy tables, columns naming strategy
NamingStrategy schema.Namer
2020-01-28 18:01:35 +03:00
// Logger
Logger logger.Interface
// NowFunc the function to be used when creating a new timestamp
NowFunc func() time.Time
2020-06-01 16:26:23 +03:00
// DryRun generate sql without execute
DryRun bool
2020-01-28 18:01:35 +03:00
2020-06-05 05:08:22 +03:00
// PrepareStmt executes the given query in cached statement
PrepareStmt bool
2020-03-09 08:10:48 +03:00
// ClauseBuilders clause builder
ClauseBuilders map[string]clause.ClauseBuilder
// ConnPool db conn pool
ConnPool ConnPool
// Dialector database dialector
Dialector
2020-05-31 06:34:59 +03:00
callbacks *callbacks
cacheStore *sync.Map
2020-02-05 06:14:58 +03:00
}
2020-01-28 18:01:35 +03:00
// DB GORM DB definition
type DB struct {
*Config
2020-03-09 08:10:48 +03:00
Error error
RowsAffected int64
Statement *Statement
2020-05-31 18:55:56 +03:00
clone int
2020-01-30 10:14:48 +03:00
}
2020-02-02 09:40:44 +03:00
// Session session config when create session with Session() method
2020-01-30 10:14:48 +03:00
type Session struct {
2020-06-01 16:26:23 +03:00
DryRun bool
2020-06-05 05:08:22 +03:00
PrepareStmt bool
2020-05-31 18:55:56 +03:00
WithConditions bool
Context context.Context
Logger logger.Interface
NowFunc func() time.Time
2020-01-30 10:14:48 +03:00
}
// Open initialize db session based on dialector
2020-03-09 15:37:01 +03:00
func Open(dialector Dialector, config *Config) (db *DB, err error) {
2020-02-02 03:35:01 +03:00
if config == nil {
config = &Config{}
}
2020-01-31 09:17:02 +03:00
if config.NamingStrategy == nil {
config.NamingStrategy = schema.NamingStrategy{}
}
2020-02-02 03:35:01 +03: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 03:35:01 +03:00
}
2020-03-09 08:10:48 +03:00
if dialector != nil {
config.Dialector = dialector
}
if config.cacheStore == nil {
config.cacheStore = &sync.Map{}
}
2020-05-31 18:55:56 +03:00
db = &DB{Config: config, clone: 1}
2020-02-02 03:35:01 +03:00
2020-03-09 15:37:01 +03:00
db.callbacks = initializeCallbacks(db)
2020-02-02 14:32:27 +03:00
2020-05-29 17:34:35 +03:00
if config.ClauseBuilders == nil {
config.ClauseBuilders = map[string]clause.ClauseBuilder{}
}
2020-02-02 03:35:01 +03:00
if dialector != nil {
2020-03-09 15:37:01 +03:00
err = dialector.Initialize(db)
2020-02-02 03:35:01 +03:00
}
2020-06-02 06:09:17 +03:00
2020-06-05 05:08:22 +03:00
if config.PrepareStmt {
db.ConnPool = &PreparedStmtDB{
ConnPool: db.ConnPool,
stmts: map[string]*sql.Stmt{},
}
}
if db.Statement == nil {
db.Statement = &Statement{
DB: db,
ConnPool: db.ConnPool,
Context: context.Background(),
Clauses: map[string]clause.Clause{},
}
}
2020-06-02 06:09:17 +03:00
if err == nil {
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 03:35:01 +03:00
return
2020-01-30 10:14:48 +03:00
}
// Session create new db session
2020-03-09 15:37:01 +03:00
func (db *DB) Session(config *Session) *DB {
2020-01-30 10:14:48 +03:00
var (
2020-05-31 18:55:56 +03:00
txConfig = *db.Config
tx = &DB{
Config: &txConfig,
Statement: db.Statement,
clone: 1,
}
2020-01-30 10:14:48 +03:00
)
if config.Context != nil {
2020-05-31 18:55:56 +03:00
if tx.Statement != nil {
tx.Statement = tx.Statement.clone()
2020-06-01 17:31:50 +03:00
tx.Statement.DB = tx
2020-05-31 18:55:56 +03:00
} else {
tx.Statement = &Statement{
DB: tx,
Clauses: map[string]clause.Clause{},
ConnPool: tx.ConnPool,
}
}
tx.Statement.Context = config.Context
}
2020-06-05 05:08:22 +03:00
if config.PrepareStmt {
tx.Statement.ConnPool = &PreparedStmtDB{
ConnPool: db.Config.ConnPool,
stmts: map[string]*sql.Stmt{},
}
}
2020-05-31 18:55:56 +03:00
if config.WithConditions {
tx.clone = 3
2020-01-30 10:14:48 +03:00
}
2020-06-01 16:26:23 +03:00
if config.DryRun {
tx.Config.DryRun = true
}
2020-01-30 10:14:48 +03:00
if config.Logger != nil {
2020-05-31 18:55:56 +03:00
tx.Config.Logger = config.Logger
2020-01-30 10:14:48 +03:00
}
if config.NowFunc != nil {
2020-05-31 18:55:56 +03:00
tx.Config.NowFunc = config.NowFunc
2020-01-30 10:14:48 +03:00
}
2020-05-31 18:55:56 +03:00
return tx
2020-01-29 14:22:44 +03:00
}
// WithContext change current instance db's context to ctx
2020-03-09 15:37:01 +03:00
func (db *DB) WithContext(ctx context.Context) *DB {
2020-05-31 18:55:56 +03:00
return db.Session(&Session{WithConditions: true, Context: ctx})
2020-01-30 10:14:48 +03:00
}
// Debug start debug mode
2020-03-09 15:37:01 +03:00
func (db *DB) Debug() (tx *DB) {
2020-05-31 18:55:56 +03:00
return db.Session(&Session{
WithConditions: true,
Logger: db.Logger.LogMode(logger.Info),
})
2020-01-30 10:14:48 +03:00
}
2020-01-29 14:22:44 +03:00
// Set store value with key into current db instance's context
2020-03-09 15:37:01 +03:00
func (db *DB) Set(key string, value interface{}) *DB {
2020-01-29 14:22:44 +03: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 15:37:01 +03:00
func (db *DB) Get(key string) (interface{}, bool) {
2020-01-29 14:22:44 +03:00
if db.Statement != nil {
return db.Statement.Settings.Load(key)
}
return nil, false
}
2020-05-31 18:55:56 +03: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) {
if db.Statement != nil {
return db.Statement.Settings.Load(fmt.Sprintf("%p", db.Statement) + key)
}
return nil, false
}
2020-06-01 17:31:50 +03:00
func (db *DB) SetupJoinTable(model interface{}, field string, joinTable interface{}) error {
var (
tx = db.getInstance()
stmt = tx.Statement
modelSchema, joinSchema *schema.Schema
)
if err := stmt.Parse(model); err == nil {
modelSchema = stmt.Schema
} else {
return err
}
if err := stmt.Parse(joinTable); err == nil {
joinSchema = stmt.Schema
} else {
return err
}
if relation, ok := modelSchema.Relationships.Relations[field]; ok && relation.JoinTable != nil {
for _, ref := range relation.References {
if f := joinSchema.LookUpField(ref.ForeignKey.DBName); f != nil {
2020-06-02 02:28:29 +03:00
f.DataType = ref.ForeignKey.DataType
2020-06-01 17:31:50 +03:00
ref.ForeignKey = f
} else {
return fmt.Errorf("missing field %v for join table", ref.ForeignKey.DBName)
}
}
relation.JoinTable = joinSchema
} else {
return fmt.Errorf("failed to found relation: %v", field)
}
return nil
}
2020-02-02 03:35:01 +03:00
// Callback returns callback manager
2020-03-09 15:37:01 +03:00
func (db *DB) Callback() *callbacks {
2020-02-02 03:35:01 +03:00
return db.callbacks
}
2020-02-22 14:41:01 +03:00
// AutoMigrate run auto migration for given models
2020-03-09 15:37:01 +03:00
func (db *DB) AutoMigrate(dst ...interface{}) error {
2020-02-22 14:41:01 +03:00
return db.Migrator().AutoMigrate(dst...)
}
2020-03-09 08:10:48 +03:00
// AddError add error to db
2020-04-19 18:11:56 +03:00
func (db *DB) AddError(err error) error {
2020-03-09 15:37:01 +03:00
if db.Error == nil {
db.Error = err
} else if err != nil {
db.Error = fmt.Errorf("%v; %w", db.Error, err)
}
2020-04-19 18:11:56 +03:00
return db.Error
2020-03-09 08:10:48 +03:00
}
2020-03-09 15:37:01 +03:00
func (db *DB) getInstance() *DB {
2020-05-31 18:55:56 +03:00
if db.clone > 0 {
tx := &DB{Config: db.Config}
switch db.clone {
case 1: // clone with new statement
2020-06-05 05:08:22 +03:00
tx.Statement = &Statement{
DB: tx,
ConnPool: db.Statement.ConnPool,
Context: db.Statement.Context,
Clauses: map[string]clause.Clause{},
}
2020-05-31 18:55:56 +03:00
case 2: // with old statement, generate new statement for future call, used to pass to callbacks
db.clone = 1
tx.Statement = db.Statement
case 3: // with clone statement
if db.Statement != nil {
tx.Statement = db.Statement.clone()
tx.Statement.DB = tx
}
}
return tx
2020-01-29 14:22:44 +03:00
}
return db
}
2020-05-30 11:47:16 +03:00
func Expr(expr string, args ...interface{}) clause.Expr {
return clause.Expr{SQL: expr, Vars: args}
}