gorm/dialects/mysql/mysql.go

166 lines
3.8 KiB
Go
Raw Normal View History

2020-02-02 03:35:01 +03:00
package mysql
import (
2020-02-22 12:53:57 +03:00
"database/sql"
"fmt"
"math"
2020-02-02 03:35:01 +03:00
_ "github.com/go-sql-driver/mysql"
"github.com/jinzhu/gorm"
"github.com/jinzhu/gorm/callbacks"
2020-03-09 12:07:00 +03:00
"github.com/jinzhu/gorm/clause"
2020-02-23 07:39:26 +03:00
"github.com/jinzhu/gorm/logger"
2020-02-22 12:53:57 +03:00
"github.com/jinzhu/gorm/migrator"
"github.com/jinzhu/gorm/schema"
2020-02-02 03:35:01 +03:00
)
type Dialector struct {
2020-02-22 12:53:57 +03:00
DSN string
2020-02-02 03:35:01 +03:00
}
func Open(dsn string) gorm.Dialector {
2020-02-22 12:53:57 +03:00
return &Dialector{DSN: dsn}
2020-02-02 03:35:01 +03:00
}
2020-02-22 12:53:57 +03:00
func (dialector Dialector) Initialize(db *gorm.DB) (err error) {
2020-02-02 03:35:01 +03:00
// register callbacks
2020-03-12 08:05:22 +03:00
callbacks.RegisterDefaultCallbacks(db, &callbacks.Config{})
2020-03-09 08:10:48 +03:00
db.ConnPool, err = sql.Open("mysql", dialector.DSN)
2020-05-29 17:34:35 +03:00
for k, v := range dialector.ClauseBuilders() {
db.ClauseBuilders[k] = v
}
2020-02-22 18:08:20 +03:00
return
2020-02-02 03:35:01 +03:00
}
2020-05-29 17:34:35 +03:00
func (dialector Dialector) ClauseBuilders() map[string]clause.ClauseBuilder {
return map[string]clause.ClauseBuilder{
"ON CONFLICT": func(c clause.Clause, builder clause.Builder) {
if onConflict, ok := c.Expression.(clause.OnConflict); ok {
builder.WriteString("ON DUPLICATE KEY UPDATE ")
if len(onConflict.DoUpdates) == 0 {
if s := builder.(*gorm.Statement).Schema; s != nil {
var column clause.Column
onConflict.DoNothing = false
if s.PrioritizedPrimaryField != nil {
column = clause.Column{Name: s.PrioritizedPrimaryField.DBName}
} else {
for _, field := range s.FieldsByDBName {
column = clause.Column{Name: field.DBName}
break
}
}
onConflict.DoUpdates = []clause.Assignment{{Column: column, Value: column}}
}
}
onConflict.DoUpdates.Build(builder)
} else {
c.Build(builder)
}
},
2020-05-31 06:19:45 +03:00
"VALUES": func(c clause.Clause, builder clause.Builder) {
if values, ok := c.Expression.(clause.Values); ok && len(values.Columns) == 0 {
builder.WriteString("VALUES()")
return
}
c.Build(builder)
},
2020-05-29 17:34:35 +03:00
}
}
2020-02-22 12:53:57 +03:00
func (dialector Dialector) Migrator(db *gorm.DB) gorm.Migrator {
2020-02-22 14:41:01 +03:00
return Migrator{migrator.Migrator{Config: migrator.Config{
DB: db,
Dialector: dialector,
}}}
2020-02-02 03:35:01 +03:00
}
2020-03-09 12:59:54 +03:00
func (dialector Dialector) BindVarTo(writer clause.Writer, stmt *gorm.Statement, v interface{}) {
writer.WriteByte('?')
2020-02-02 03:35:01 +03:00
}
2020-02-05 06:14:58 +03:00
2020-03-09 12:07:00 +03:00
func (dialector Dialector) QuoteTo(writer clause.Writer, str string) {
writer.WriteByte('`')
writer.WriteString(str)
writer.WriteByte('`')
2020-02-05 06:14:58 +03:00
}
2020-02-22 12:53:57 +03:00
2020-02-23 07:39:26 +03:00
func (dialector Dialector) Explain(sql string, vars ...interface{}) string {
return logger.ExplainSQL(sql, nil, `"`, vars...)
}
2020-02-22 12:53:57 +03:00
func (dialector Dialector) DataTypeOf(field *schema.Field) string {
switch field.DataType {
case schema.Bool:
return "boolean"
case schema.Int, schema.Uint:
sqlType := "int"
switch {
case field.Size <= 8:
sqlType = "tinyint"
case field.Size <= 16:
sqlType = "smallint"
case field.Size <= 32:
sqlType = "int"
default:
sqlType = "bigint"
}
if field.DataType == schema.Uint {
sqlType += " unsigned"
}
2020-03-12 03:39:42 +03:00
if field.AutoIncrement || field == field.Schema.PrioritizedPrimaryField {
2020-02-22 12:53:57 +03:00
sqlType += " AUTO_INCREMENT"
}
return sqlType
case schema.Float:
if field.Size <= 32 {
return "float"
}
return "double"
case schema.String:
size := field.Size
2020-05-30 19:42:52 +03:00
if size == 0 {
if field.PrimaryKey || field.HasDefaultValue {
size = 256
}
2020-02-22 18:08:20 +03:00
}
2020-02-22 12:53:57 +03:00
if size >= 65536 && size <= int(math.Pow(2, 24)) {
return "mediumtext"
2020-02-22 18:08:20 +03:00
} else if size > int(math.Pow(2, 24)) || size <= 0 {
2020-02-22 12:53:57 +03:00
return "longtext"
}
return fmt.Sprintf("varchar(%d)", size)
case schema.Time:
precision := ""
2020-03-12 03:39:42 +03:00
if field.Precision == 0 {
field.Precision = 3
}
2020-02-22 12:53:57 +03:00
if field.Precision > 0 {
precision = fmt.Sprintf("(%d)", field.Precision)
}
if field.NotNull || field.PrimaryKey {
return "datetime" + precision
}
return "datetime" + precision + " NULL"
case schema.Bytes:
if field.Size > 0 && field.Size < 65536 {
return fmt.Sprintf("varbinary(%d)", field.Size)
}
if field.Size >= 65536 && field.Size <= int(math.Pow(2, 24)) {
return "mediumblob"
}
return "longblob"
}
return ""
}