OSDN Git Service

feat: init cross_tx keepers (#146)
[bytom/vapor.git] / vendor / github.com / jinzhu / gorm / dialect_mysql.go
1 package gorm
2
3 import (
4         "crypto/sha1"
5         "fmt"
6         "reflect"
7         "regexp"
8         "strconv"
9         "strings"
10         "time"
11         "unicode/utf8"
12 )
13
14 type mysql struct {
15         commonDialect
16 }
17
18 func init() {
19         RegisterDialect("mysql", &mysql{})
20 }
21
22 func (mysql) GetName() string {
23         return "mysql"
24 }
25
26 func (mysql) Quote(key string) string {
27         return fmt.Sprintf("`%s`", key)
28 }
29
30 // Get Data Type for MySQL Dialect
31 func (s *mysql) DataTypeOf(field *StructField) string {
32         var dataValue, sqlType, size, additionalType = ParseFieldStructForDialect(field, s)
33
34         // MySQL allows only one auto increment column per table, and it must
35         // be a KEY column.
36         if _, ok := field.TagSettingsGet("AUTO_INCREMENT"); ok {
37                 if _, ok = field.TagSettingsGet("INDEX"); !ok && !field.IsPrimaryKey {
38                         field.TagSettingsDelete("AUTO_INCREMENT")
39                 }
40         }
41
42         if sqlType == "" {
43                 switch dataValue.Kind() {
44                 case reflect.Bool:
45                         sqlType = "boolean"
46                 case reflect.Int8:
47                         if s.fieldCanAutoIncrement(field) {
48                                 field.TagSettingsSet("AUTO_INCREMENT", "AUTO_INCREMENT")
49                                 sqlType = "tinyint AUTO_INCREMENT"
50                         } else {
51                                 sqlType = "tinyint"
52                         }
53                 case reflect.Int, reflect.Int16, reflect.Int32:
54                         if s.fieldCanAutoIncrement(field) {
55                                 field.TagSettingsSet("AUTO_INCREMENT", "AUTO_INCREMENT")
56                                 sqlType = "int AUTO_INCREMENT"
57                         } else {
58                                 sqlType = "int"
59                         }
60                 case reflect.Uint8:
61                         if s.fieldCanAutoIncrement(field) {
62                                 field.TagSettingsSet("AUTO_INCREMENT", "AUTO_INCREMENT")
63                                 sqlType = "tinyint unsigned AUTO_INCREMENT"
64                         } else {
65                                 sqlType = "tinyint unsigned"
66                         }
67                 case reflect.Uint, reflect.Uint16, reflect.Uint32, reflect.Uintptr:
68                         if s.fieldCanAutoIncrement(field) {
69                                 field.TagSettingsSet("AUTO_INCREMENT", "AUTO_INCREMENT")
70                                 sqlType = "int unsigned AUTO_INCREMENT"
71                         } else {
72                                 sqlType = "int unsigned"
73                         }
74                 case reflect.Int64:
75                         if s.fieldCanAutoIncrement(field) {
76                                 field.TagSettingsSet("AUTO_INCREMENT", "AUTO_INCREMENT")
77                                 sqlType = "bigint AUTO_INCREMENT"
78                         } else {
79                                 sqlType = "bigint"
80                         }
81                 case reflect.Uint64:
82                         if s.fieldCanAutoIncrement(field) {
83                                 field.TagSettingsSet("AUTO_INCREMENT", "AUTO_INCREMENT")
84                                 sqlType = "bigint unsigned AUTO_INCREMENT"
85                         } else {
86                                 sqlType = "bigint unsigned"
87                         }
88                 case reflect.Float32, reflect.Float64:
89                         sqlType = "double"
90                 case reflect.String:
91                         if size > 0 && size < 65532 {
92                                 sqlType = fmt.Sprintf("varchar(%d)", size)
93                         } else {
94                                 sqlType = "longtext"
95                         }
96                 case reflect.Struct:
97                         if _, ok := dataValue.Interface().(time.Time); ok {
98                                 precision := ""
99                                 if p, ok := field.TagSettingsGet("PRECISION"); ok {
100                                         precision = fmt.Sprintf("(%s)", p)
101                                 }
102
103                                 if _, ok := field.TagSettingsGet("NOT NULL"); ok {
104                                         sqlType = fmt.Sprintf("timestamp%v", precision)
105                                 } else {
106                                         sqlType = fmt.Sprintf("timestamp%v NULL", precision)
107                                 }
108                         }
109                 default:
110                         if IsByteArrayOrSlice(dataValue) {
111                                 if size > 0 && size < 65532 {
112                                         sqlType = fmt.Sprintf("varbinary(%d)", size)
113                                 } else {
114                                         sqlType = "longblob"
115                                 }
116                         }
117                 }
118         }
119
120         if sqlType == "" {
121                 panic(fmt.Sprintf("invalid sql type %s (%s) for mysql", dataValue.Type().Name(), dataValue.Kind().String()))
122         }
123
124         if strings.TrimSpace(additionalType) == "" {
125                 return sqlType
126         }
127         return fmt.Sprintf("%v %v", sqlType, additionalType)
128 }
129
130 func (s mysql) RemoveIndex(tableName string, indexName string) error {
131         _, err := s.db.Exec(fmt.Sprintf("DROP INDEX %v ON %v", indexName, s.Quote(tableName)))
132         return err
133 }
134
135 func (s mysql) ModifyColumn(tableName string, columnName string, typ string) error {
136         _, err := s.db.Exec(fmt.Sprintf("ALTER TABLE %v MODIFY COLUMN %v %v", tableName, columnName, typ))
137         return err
138 }
139
140 func (s mysql) LimitAndOffsetSQL(limit, offset interface{}) (sql string) {
141         if limit != nil {
142                 if parsedLimit, err := strconv.ParseInt(fmt.Sprint(limit), 0, 0); err == nil && parsedLimit >= 0 {
143                         sql += fmt.Sprintf(" LIMIT %d", parsedLimit)
144
145                         if offset != nil {
146                                 if parsedOffset, err := strconv.ParseInt(fmt.Sprint(offset), 0, 0); err == nil && parsedOffset >= 0 {
147                                         sql += fmt.Sprintf(" OFFSET %d", parsedOffset)
148                                 }
149                         }
150                 }
151         }
152         return
153 }
154
155 func (s mysql) HasForeignKey(tableName string, foreignKeyName string) bool {
156         var count int
157         currentDatabase, tableName := currentDatabaseAndTable(&s, tableName)
158         s.db.QueryRow("SELECT count(*) FROM INFORMATION_SCHEMA.TABLE_CONSTRAINTS WHERE CONSTRAINT_SCHEMA=? AND TABLE_NAME=? AND CONSTRAINT_NAME=? AND CONSTRAINT_TYPE='FOREIGN KEY'", currentDatabase, tableName, foreignKeyName).Scan(&count)
159         return count > 0
160 }
161
162 func (s mysql) CurrentDatabase() (name string) {
163         s.db.QueryRow("SELECT DATABASE()").Scan(&name)
164         return
165 }
166
167 func (mysql) SelectFromDummyTable() string {
168         return "FROM DUAL"
169 }
170
171 func (s mysql) BuildKeyName(kind, tableName string, fields ...string) string {
172         keyName := s.commonDialect.BuildKeyName(kind, tableName, fields...)
173         if utf8.RuneCountInString(keyName) <= 64 {
174                 return keyName
175         }
176         h := sha1.New()
177         h.Write([]byte(keyName))
178         bs := h.Sum(nil)
179
180         // sha1 is 40 characters, keep first 24 characters of destination
181         destRunes := []rune(regexp.MustCompile("[^a-zA-Z0-9]+").ReplaceAllString(fields[0], "_"))
182         if len(destRunes) > 24 {
183                 destRunes = destRunes[:24]
184         }
185
186         return fmt.Sprintf("%s%x", string(destRunes), bs)
187 }
188
189 func (mysql) DefaultValueStr() string {
190         return "VALUES()"
191 }