package schema import ( "strings" "github.com/samber/lo" ) type Blueprint struct { table string // the table the blueprint describes. columns []*ColumnDefinition // columns that should be added to the table commands []*Command // Temporary bool // Whether to make the table temporary. Charset string // The default character set that should be used for the table. Collation string // The collation that should be used for the table. Engine string // The engine that should be used for the table. Comment string } type Command struct { Type string CommandOptions } type CommandOptions struct { Index string // 索引名称 Columns []string // 索引字段 Algorithm string // 索引类型如: USING BTREE To string // 新名词 From string // 旧名词 } func NewBlueprint(table string) *Blueprint { return &Blueprint{table: table} } // 字符串 func (this *Blueprint) Char(column string, length int) *ColumnDefinition { if length == 0 { length = 255 } return this.addColumn("char", column, &ColumnOptions{Length: length}) } // 可变长度字符串 func (this *Blueprint) String(column string, length int) *ColumnDefinition { if length == 0 { length = 255 } return this.addColumn("string", column, &ColumnOptions{Length: length}) } // 文本 func (this *Blueprint) Text(column string) *ColumnDefinition { return this.addColumn("text", column, nil) } // 整型 func (this *Blueprint) Integer(column string, params ...bool) *ColumnDefinition { autoIncrement := false unsigned := false if len(params) > 0 { autoIncrement = params[0] } if len(params) > 1 { unsigned = params[1] } return this.addColumn("integer", column, &ColumnOptions{autoIncrement: autoIncrement, unsigned: unsigned}) } // 迷你整型 1 byte func (this *Blueprint) TinyInteger(column string, params ...bool) *ColumnDefinition { autoIncrement := false unsigned := false if len(params) > 0 { autoIncrement = params[0] } if len(params) > 1 { unsigned = params[1] } return this.addColumn("tinyInteger", column, &ColumnOptions{autoIncrement: autoIncrement, unsigned: unsigned}) } // 小整型 2 byte func (this *Blueprint) SmallInteger(column string, params ...bool) *ColumnDefinition { autoIncrement := false unsigned := false if len(params) > 0 { autoIncrement = params[0] } if len(params) > 1 { unsigned = params[1] } return this.addColumn("smallInteger", column, &ColumnOptions{autoIncrement: autoIncrement, unsigned: unsigned}) } // 大整型 2 byte func (this *Blueprint) BigInteger(column string, params ...bool) *ColumnDefinition { autoIncrement := false unsigned := false if len(params) > 0 { autoIncrement = params[0] } if len(params) > 1 { unsigned = params[1] } return this.addColumn("bigInteger", column, &ColumnOptions{autoIncrement: autoIncrement, unsigned: unsigned}) } // 无符号整型 func (this *Blueprint) UnsignedInteger(column string, params ...bool) *ColumnDefinition { autoIncrement := false if len(params) > 0 { autoIncrement = params[0] } return this.Integer(column, autoIncrement, true) } // 无符号迷你整型 1 byte func (this *Blueprint) UnsignedTinyInteger(column string, params ...bool) *ColumnDefinition { autoIncrement := false if len(params) > 0 { autoIncrement = params[0] } return this.TinyInteger(column, autoIncrement, true) } // 无符号小整型 2 byte func (this *Blueprint) UnsignedSmallInteger(column string, params ...bool) *ColumnDefinition { autoIncrement := false if len(params) > 0 { autoIncrement = params[0] } return this.SmallInteger(column, autoIncrement, true) } // 无符号大整型 func (this *Blueprint) UnsignedBigInteger(column string, params ...bool) *ColumnDefinition { autoIncrement := false if len(params) > 0 { autoIncrement = params[0] } return this.BigInteger(column, autoIncrement, true) } // 精确小数 func (this *Blueprint) Decimal(column string, total, places int) *ColumnDefinition { if total == 0 { total = 8 } if places == 0 { places = 2 } return this.addColumn("decimal", column, &ColumnOptions{Total: total, Places: places}) } // 无符号精确销售 func (this *Blueprint) UnsignedDecimal(column string, total, places int) *ColumnDefinition { if total == 0 { total = 8 } if places == 0 { places = 2 } return this.addColumn("decimal", column, &ColumnOptions{Total: total, Places: places, unsigned: true}) } // 布尔值 func (this *Blueprint) Boolean(column string) *ColumnDefinition { return this.addColumn("boolean", column, nil) } // 枚举类型 func (this *Blueprint) Enum(column string, allowed []string) *ColumnDefinition { return this.addColumn("enum", column, &ColumnOptions{Allowed: allowed}) } // JSON func (this *Blueprint) Json(column string) *ColumnDefinition { return this.addColumn("json", column, nil) } // 日期类型 func (this *Blueprint) Date(column string) *ColumnDefinition { return this.addColumn("date", column, nil) } // 日期时间类型 func (this *Blueprint) DateTime(column string, precision ...int) *ColumnDefinition { if len(precision) > 0 { return this.addColumn("datetime", column, &ColumnOptions{Precision: precision[0]}) } return this.addColumn("datetime", column, nil) } // 时间类型 func (this *Blueprint) Time(column string, precision ...int) *ColumnDefinition { if len(precision) > 0 { return this.addColumn("time", column, &ColumnOptions{Precision: precision[0]}) } return this.addColumn("time", column, nil) } // 时间戳 func (this *Blueprint) Timestamp(column string, precision ...int) *ColumnDefinition { if len(precision) > 0 { return this.addColumn("timestamp", column, &ColumnOptions{Precision: precision[0]}) } return this.addColumn("timestamp", column, nil) } // 年 func (this *Blueprint) Year(column string) *ColumnDefinition { return this.addColumn("year", column, nil) } // 二进制数据 func (this *Blueprint) Binary(column string) *ColumnDefinition { return this.addColumn("binary", column, nil) } // UUID func (this *Blueprint) Uuid(column string) *ColumnDefinition { return this.addColumn("uuid", column, nil) } // 自增字段 func (this *Blueprint) Increments(column string) *ColumnDefinition { return this.UnsignedInteger(column, true) } // 自增Big字段 func (this *Blueprint) BigIncrements(column string) *ColumnDefinition { return this.UnsignedBigInteger(column, true) } // 添加主键 func (this *Blueprint) Primary(columns ...string) *Command { return this.addCommand("primary", CommandOptions{Index: this.generateIndexName("pk", columns), Columns: columns}) } // 唯一键 func (this *Blueprint) Unique(columns ...string) *Command { return this.addCommand("unique", CommandOptions{Index: this.generateIndexName("unique", columns), Columns: columns}) } // 普通索引 func (this *Blueprint) Index(columns ...string) *Command { return this.addCommand("index", CommandOptions{Index: this.generateIndexName("index", columns), Columns: columns}) } // 空间索引 func (this *Blueprint) SpatialIndex(columns ...string) *Command { return this.addCommand("spatialIndex", CommandOptions{Index: this.generateIndexName("spatial_index", columns), Columns: columns}) } // 删除列 func (this *Blueprint) DropColumn(columns ...string) *Command { return this.addCommand("dropColumn", CommandOptions{Columns: columns}) } // 创建表 func (this *Blueprint) Create() *Command { return this.addCommand("create", CommandOptions{}) } // 修改表名 func (this *Blueprint) Rename(to string) *Command { return this.addCommand("rename", CommandOptions{To: to}) } // 修改表备注 func (this *Blueprint) ModifyComment(comment string) *Command { this.Comment = comment return this.addCommand("modifyComment", CommandOptions{}) } // 删除表 func (this *Blueprint) Drop() *Command { return this.addCommand("drop", CommandOptions{}) } // 删除表, 先判断再删除 func (this *Blueprint) DropIfExists() *Command { return this.addCommand("dropIfExists", CommandOptions{}) } func (this *Blueprint) ToSql(sc Schema) (statements []string) { this.addImpliedCommands(sc) for _, cmd := range this.commands { switch cmd.Type { case "create": statements = append(statements, sc.CompileCreate(this)...) case "add": statements = append(statements, sc.CompileAdd(this)...) case "change": statements = append(statements, sc.CompileChange(this)...) case "drop": statements = append(statements, sc.CompileDrop(this)...) case "dropIfExists": statements = append(statements, sc.CompileDropIfExists(this)...) case "dropColumn": statements = append(statements, sc.CompileDropColumn(this)...) case "rename": statements = append(statements, sc.CompileRename(this)...) case "modifyComment": statements = append(statements, sc.CompileModifyComment(this)...) } } return statements } // 判断是否是创建表 func (this *Blueprint) creating() bool { return lo.SomeBy(this.commands, func(item *Command) bool { if item.Type == "create" { return true } return false }) } func (this *Blueprint) addColumn(typ, name string, options *ColumnOptions) (definition *ColumnDefinition) { definition = &ColumnDefinition{Type: typ, Name: name} if options != nil { if options.Length > 0 { definition.Length = options.Length } if options.autoIncrement { definition.autoIncrement = true } if options.unsigned { definition.unsigned = true } if options.Total > 0 { definition.Total = options.Total } if options.Places > 0 { definition.Places = options.Places } if len(options.Allowed) > 0 { definition.Allowed = options.Allowed } if options.Precision > 0 { definition.Precision = options.Precision } if options.change { definition.change = true } } this.columns = append(this.columns, definition) return definition } func (this *Blueprint) addImpliedCommands(sc Schema) { if !this.creating() { if len(this.GetAddedColumns()) > 0 { this.commands = append([]*Command{this.createCommand("add", CommandOptions{})}, this.commands...) } if len(this.GetChangedColumns()) > 0 { this.commands = append([]*Command{this.createCommand("change", CommandOptions{})}, this.commands...) } } this.addFluentIndexes() } // 添加索引字段 func (this *Blueprint) addFluentIndexes() { for _, column := range this.columns { if column.primary { this.Primary(column.Name) continue } else if column.unique { this.Unique(column.Name) continue } else if column.index { this.Index(column.Name) continue } else if column.spatialIndex { this.SpatialIndex(column.Name) continue } } } func (this *Blueprint) addCommand(name string, options CommandOptions) (command *Command) { command = this.createCommand(name, options) this.commands = append(this.commands, command) return command } func (this *Blueprint) createCommand(name string, options CommandOptions) *Command { return &Command{Type: name, CommandOptions: options} } // 生成索引名称 func (this *Blueprint) generateIndexName(typ string, columns []string) string { return strings.ToLower(typ + "_" + strings.Join(columns, "_")) } func (this *Blueprint) GetAddedColumns() []*ColumnDefinition { return lo.Filter(this.columns, func(item *ColumnDefinition, _ int) bool { return !item.change }) } func (this *Blueprint) GetChangedColumns() []*ColumnDefinition { return lo.Filter(this.columns, func(item *ColumnDefinition, _ int) bool { return item.change }) } func (this *Blueprint) GetCommands() []*Command { return this.commands } func (this *Blueprint) GetTable() string { return this.table } // 命令类型 func (this *Command) Command() string { return this.Type }