feat: Oracle 驱动完整化——回调体系、版本感知、驱动抽象层与测试套件

- 新增驱动抽象层 driver_adapter(go-ora/godror 双驱动切换)
- 新增 Update/Delete/Query 回调,完善 RETURNING INTO 子句
- 修复 create.go 事务悬挂行 BUG、批量 RowsAffected、主键 WHERE 注入
- 版本感知体系:分页/自增/序列默认值/BOOLEAN/32k VARCHAR2 分级能力判定
- 11g 序列+触发器自增、ON UPDATE 触发器、默认值智能转换
- 升级 go-ora v2.8.19 → v2.9.0
- 单元测试+模块测试+集成测试共 96 个,真实 Oracle 11g 全绿
- 测试 DSN 密码移除,改为 ORACLE_DSN 环境变量注入
This commit is contained in:
2026-08-08 10:36:57 +08:00
parent 1f6c8a18a3
commit 9154afab6b
35 changed files with 4956 additions and 37 deletions
+372 -11
View File
@@ -16,6 +16,22 @@ type Migrator struct {
migrator.Migrator
}
// oracleDBVer 返回当前数据库版本号(用于版本感知的默认值处理)。
// m.Dialector 是 gorm.Dialector 接口,需断言为 oracle.Dialector 获取 DBVer。
func (m Migrator) oracleDBVer() string {
if d, ok := m.Dialector.(Dialector); ok {
return d.DBVer
}
return ""
}
// hasNEXTVALDefault 判断字段默认值是否为序列引用(.NEXTVAL)。
// 仅检查 DefaultValue 字符串,因为 string 类型字段的 DefaultValueInterface
// 会被 GORM 解析为字符串字面量,无法区分普通字符串与序列引用。
func hasNEXTVALDefault(field *schema.Field) bool {
return field != nil && strings.Contains(strings.ToUpper(field.DefaultValue), ".NEXTVAL")
}
func (m Migrator) CurrentDatabase() (name string) {
m.DB.Raw(
fmt.Sprintf(`SELECT ORA_DATABASE_NAME as "Current Database" FROM %s`, m.Dialector.(Dialector).DummyTableName()),
@@ -28,7 +44,178 @@ func (m Migrator) CreateTable(values ...interface{}) error {
m.TryQuotifyReservedWords(value)
m.TryRemoveOnUpdate(value)
}
return m.Migrator.CreateTable(values...)
// 先创建表
if err := m.Migrator.CreateTable(values...); err != nil {
return err
}
// 然后创建 ON UPDATE 触发器
for _, value := range values {
m.RunWithValue(value, func(stmt *gorm.Statement) error {
if stmt.Schema == nil {
return nil
}
for _, rel := range stmt.Schema.Relationships.Relations {
if err := m.CreateOnUpdateTrigger(value, rel); err != nil {
// 触发器创建失败不阻止表创建,只记录警告
// 可以选择忽略或记录日志
}
}
return nil
})
}
// Oracle 11g 不支持 IDENTITY 列,为自增主键创建序列 + BEFORE INSERT 触发器
for _, value := range values {
if err := m.createAutoIncrementSupport(value); err != nil {
return err
}
}
// Oracle 11g 下使用序列默认值(DEFAULT <seq>.NEXTVAL)的字段:
// 11g 的 DEFAULT 子句不允许引用序列 NEXTVALORA-0098412c 才支持),
// 因此建表后为这类字段创建 BEFORE INSERT 触发器实现等价语义;12c+ 无需。
dbVer := m.oracleDBVer()
if !supportsIdentity(dbVer) {
for _, value := range values {
if err := m.RunWithValue(value, func(stmt *gorm.Statement) error {
if stmt.Schema == nil {
return nil
}
for _, field := range stmt.Schema.Fields {
// 仅处理显式声明了序列默认值且非自增的字段。
// 自增主键的序列逻辑由 createAutoIncrementSupport 负责,
// 跳过以避免生成重复/冲突的 BEFORE INSERT 触发器。
if !field.HasDefaultValue || field.AutoIncrement || !hasNEXTVALDefault(field) {
continue
}
// 从 DefaultValue 提取序列名:取 ".NEXTVAL" 前的部分
seqName := extractSequenceNameFromDefault(field.DefaultValue)
if seqName == "" {
continue
}
if err := m.createSequenceDefaultTrigger(stmt, field, seqName); err != nil {
return err
}
}
return nil
}); err != nil {
return err
}
}
}
return nil
}
// sequenceName 返回自增主键对应的序列名
func (m Migrator) sequenceName(table string) string {
return fmt.Sprintf("SEQ_%s", table)
}
// triggerName 返回自增主键对应的触发器名
func (m Migrator) triggerName(table string) string {
return fmt.Sprintf("TRG_%s", table)
}
// createAutoIncrementSupport 为自增主键创建序列和 BEFORE INSERT 触发器(仅不支持 IDENTITY 的版本)
func (m Migrator) createAutoIncrementSupport(value interface{}) error {
// 12c+ 原生支持 IDENTITY 列,无需序列 + 触发器模拟
if d, ok := m.Dialector.(Dialector); ok && supportsIdentity(d.DBVer) {
return nil
}
return m.RunWithValue(value, func(stmt *gorm.Statement) error {
if stmt.Schema == nil {
return nil
}
for _, field := range stmt.Schema.Fields {
if !field.AutoIncrement {
continue
}
if field.DataType != schema.Int && field.DataType != schema.Uint {
continue
}
seqName := m.sequenceName(stmt.Table)
trgName := m.triggerName(stmt.Table)
// 创建序列
if err := m.DB.Exec(fmt.Sprintf("CREATE SEQUENCE %s START WITH 1 INCREMENT BY 1 NOCACHE", seqName)).Error; err != nil {
return err
}
// 创建 BEFORE INSERT 触发器:ID 为空时从序列取值
triggerSQL := fmt.Sprintf(`CREATE OR REPLACE TRIGGER %s
BEFORE INSERT ON %s
FOR EACH ROW
BEGIN
IF :NEW.%s IS NULL THEN
SELECT %s.NEXTVAL INTO :NEW.%s FROM DUAL;
END IF;
END;`, trgName, stmt.Table, field.DBName, seqName, field.DBName)
if err := m.DB.Exec(triggerSQL).Error; err != nil {
return err
}
}
return nil
})
}
// extractSequenceNameFromDefault 从序列默认值字符串中提取序列名。
// 例如 "SEQ_MY.NEXTVAL" → "SEQ_MY";同时兼容 GORM 对含括号默认值保持原文的
// 情况(如 "(SEQ_MY.NEXTVAL)")。
func extractSequenceNameFromDefault(defaultValue string) string {
v := strings.TrimSpace(defaultValue)
// 去掉可能的包裹括号
v = strings.TrimPrefix(v, "(")
v = strings.TrimSuffix(v, ")")
v = strings.TrimSpace(v)
idx := strings.Index(strings.ToUpper(v), ".NEXTVAL")
if idx <= 0 {
return ""
}
return strings.TrimSpace(v[:idx])
}
// createSequenceDefaultTrigger 为 11g 下使用序列默认值的字段创建 BEFORE INSERT 触发器。
// 11g 的 DEFAULT 子句不允许引用序列 NEXTVALORA-0098412c 才支持),
// 因此在建表后通过触发器实现等价语义:插入时列值为 NULL 则从序列取值回填。
// 触发器命名为 SEQDEF_TRG_<table>_<column>,避免与 autoIncrement 的 TRG_<table> 冲突。
func (m Migrator) createSequenceDefaultTrigger(stmt *gorm.Statement, field *schema.Field, seqName string) error {
trgName := fmt.Sprintf("SEQDEF_TRG_%s_%s", stmt.Table, field.DBName)
// Oracle 标识符最多 30 字符,超长时截断避免 ORA-00972
if len(trgName) > 30 {
trgName = trgName[:30]
}
triggerSQL := fmt.Sprintf(`CREATE OR REPLACE TRIGGER %s
BEFORE INSERT ON %s
FOR EACH ROW
BEGIN
IF :NEW.%s IS NULL THEN
SELECT %s.NEXTVAL INTO :NEW.%s FROM DUAL;
END IF;
END;`, trgName, stmt.Table, field.DBName, seqName, field.DBName)
return m.DB.Exec(triggerSQL).Error
}
// dropSequence 删除表对应的自增序列(不存在时忽略)
func (m Migrator) dropSequence(table string) error {
seqName := m.sequenceName(table)
var count int64
if err := m.DB.Raw("SELECT COUNT(*) FROM USER_SEQUENCES WHERE SEQUENCE_NAME = ?", seqName).Row().Scan(&count); err != nil {
return err
}
if count > 0 {
return m.DB.Exec("DROP SEQUENCE " + seqName).Error
}
return nil
}
func (m Migrator) DropTable(values ...interface{}) error {
@@ -38,7 +225,11 @@ func (m Migrator) DropTable(values ...interface{}) error {
tx := m.DB.Session(&gorm.Session{})
if m.HasTable(value) {
if err := m.RunWithValue(value, func(stmt *gorm.Statement) error {
return tx.Exec("DROP TABLE ? CASCADE CONSTRAINTS", clause.Table{Name: stmt.Table}).Error
if err := tx.Exec("DROP TABLE ? CASCADE CONSTRAINTS", clause.Table{Name: stmt.Table}).Error; err != nil {
return err
}
// 删除自增序列
return m.dropSequence(stmt.Table)
}); err != nil {
return err
}
@@ -81,8 +272,35 @@ func (m Migrator) ColumnTypes(value interface{}) ([]gorm.ColumnType, error) {
return err
}
// Oracle 返回大写列名,而模型字段 DBName 可能是小写(如 column:id)。
// 将列名映射回模型定义的 DBName 大小写,避免 AutoMigrate 误判列不存在。
upperToDBName := make(map[string]string, len(stmt.Schema.Fields))
for _, field := range stmt.Schema.Fields {
if field.DBName != "" {
upperToDBName[strings.ToUpper(field.DBName)] = field.DBName
}
}
for _, c := range rawColumnTypes {
columnTypes = append(columnTypes, migrator.ColumnType{SQLColumnType: c})
ct := migrator.ColumnType{SQLColumnType: c}
// 映射列名大小写
upperName := strings.ToUpper(c.Name())
if dbName, ok := upperToDBName[upperName]; ok {
ct.NameValue = sql.NullString{String: dbName, Valid: true}
}
// go-ora 未实现 RowsColumnTypeDatabaseTypeName,从数据字典获取真实数据类型,
// 避免 AutoMigrate 对每个非主键列都误判类型变化并触发 ALTER。
var dataType string
if err := m.DB.Raw(
"SELECT DATA_TYPE FROM USER_TAB_COLUMNS WHERE TABLE_NAME = ? AND UPPER(COLUMN_NAME) = ?",
stmt.Table, upperName,
).Row().Scan(&dataType); err == nil && dataType != "" {
ct.DataTypeValue = sql.NullString{String: dataType, Valid: true}
}
columnTypes = append(columnTypes, ct)
}
return
@@ -177,11 +395,10 @@ func (m Migrator) HasColumn(value interface{}, field string) bool {
return m.RunWithValue(value, func(stmt *gorm.Statement) error {
if stmt.Schema != nil && strings.Contains(stmt.Schema.Table, ".") {
ownertable := strings.Split(stmt.Schema.Table, ".")
return m.DB.Raw("SELECT COUNT(*) FROM ALL_TAB_COLUMNS WHERE OWNER = ? and TABLE_NAME = ? AND COLUMN_NAME = ?", ownertable[0], ownertable[1], field).Row().Scan(&count)
return m.DB.Raw("SELECT COUNT(*) FROM ALL_TAB_COLUMNS WHERE OWNER = ? AND TABLE_NAME = ? AND UPPER(COLUMN_NAME) = UPPER(?)", ownertable[0], ownertable[1], field).Row().Scan(&count)
} else {
return m.DB.Raw("SELECT COUNT(*) FROM USER_TAB_COLUMNS WHERE TABLE_NAME = ? AND COLUMN_NAME = ?", stmt.Table, field).Row().Scan(&count)
return m.DB.Raw("SELECT COUNT(*) FROM USER_TAB_COLUMNS WHERE TABLE_NAME = ? AND UPPER(COLUMN_NAME) = UPPER(?)", stmt.Table, field).Row().Scan(&count)
}
}) == nil && count > 0
}
@@ -191,9 +408,9 @@ func (m Migrator) AlterDataTypeOf(stmt *gorm.Statement, field *schema.Field) (ex
var nullable = ""
if stmt.Schema != nil && strings.Contains(stmt.Schema.Table, ".") {
ownertable := strings.Split(stmt.Schema.Table, ".")
m.DB.Raw("SELECT NULLABLE FROM ALL_TAB_COLUMNS WHERE OWNER = ? and TABLE_NAME = ? AND COLUMN_NAME = ?", ownertable[0], ownertable[1], field.DBName).Row().Scan(&nullable)
m.DB.Raw("SELECT NULLABLE FROM ALL_TAB_COLUMNS WHERE OWNER = ? AND TABLE_NAME = ? AND UPPER(COLUMN_NAME) = UPPER(?)", ownertable[0], ownertable[1], field.DBName).Row().Scan(&nullable)
} else {
m.DB.Raw("SELECT NULLABLE FROM USER_TAB_COLUMNS WHERE TABLE_NAME = ? AND COLUMN_NAME = ?", stmt.Table, field.DBName).Row().Scan(&nullable)
m.DB.Raw("SELECT NULLABLE FROM USER_TAB_COLUMNS WHERE TABLE_NAME = ? AND UPPER(COLUMN_NAME) = UPPER(?)", stmt.Table, field.DBName).Row().Scan(&nullable)
}
if field.NotNull && nullable == "Y" {
expr.SQL += " NOT NULL"
@@ -204,17 +421,71 @@ func (m Migrator) AlterDataTypeOf(stmt *gorm.Statement, field *schema.Field) (ex
}
if field.HasDefaultValue && (field.DefaultValueInterface != nil || field.DefaultValue != "") {
// 序列默认值(.NEXTVAL)优先走版本感知处理:
// string 字段的 DefaultValueInterface 会被 GORM 解析为字符串字面量(如 "SEQ_X.NEXTVAL"),
// 直接拼接会生成错误的字符串默认值而非序列引用,必须先识别出来。
if hasNEXTVALDefault(field) {
if dv := buildOracleDefault(m.oracleDBVer(), field.DefaultValue, field); dv != "" {
expr.SQL += " " + dv
}
return
}
if field.DefaultValueInterface != nil {
defaultStmt := &gorm.Statement{Vars: []interface{}{field.DefaultValueInterface}}
m.Dialector.BindVarTo(defaultStmt, defaultStmt, field.DefaultValueInterface)
expr.SQL += " DEFAULT " + m.Dialector.Explain(defaultStmt.SQL.String(), field.DefaultValueInterface)
} else if field.DefaultValue != "(-)" {
expr.SQL += " DEFAULT " + field.DefaultValue
// 使用 buildOracleDefault 进行智能转换(版本感知:11g 下 NEXTVAL 默认值
// 不能生成 DEFAULT 子句,返回空串则跳过,避免拼出非法 SQL)
if dv := buildOracleDefault(m.oracleDBVer(), field.DefaultValue, field); dv != "" {
expr.SQL += " " + dv
}
}
}
return
}
// FullDataTypeOf 返回字段的完整数据库类型(版本感知的默认值处理)。
// GORM 标准实现会把 DefaultValue 直接拼成 "DEFAULT xxx",对 11g 下引用序列的
// NEXTVAL 默认值会生成非法 SQLORA-00984),因此在此重写:
// 11g 下 NEXTVAL 默认值不生成 DEFAULT 子句,改由 CreateTable 流程在建表后
// 创建 BEFORE INSERT 触发器实现等价语义(见 createSequenceDefaultTrigger)。
func (m Migrator) FullDataTypeOf(field *schema.Field) (expr clause.Expr) {
expr.SQL = m.DataTypeOf(field)
if field.NotNull {
expr.SQL += " NOT NULL"
}
if field.HasDefaultValue && (field.DefaultValueInterface != nil || field.DefaultValue != "") {
// 序列默认值(.NEXTVAL)优先走版本感知处理:
// string 字段的 DefaultValueInterface 会被 GORM 解析为字符串字面量(如 "SEQ_X.NEXTVAL"),
// 直接拼接会生成错误的字符串默认值而非序列引用,必须先识别出来。
if hasNEXTVALDefault(field) {
dbVer := m.oracleDBVer()
if dv := buildOracleDefault(dbVer, field.DefaultValue, field); dv != "" {
expr.SQL += " " + dv
}
return
}
if field.DefaultValueInterface != nil {
defaultStmt := &gorm.Statement{Vars: []interface{}{field.DefaultValueInterface}}
m.Dialector.BindVarTo(defaultStmt, defaultStmt, field.DefaultValueInterface)
expr.SQL += " DEFAULT " + m.Dialector.Explain(defaultStmt.SQL.String(), field.DefaultValueInterface)
} else if field.DefaultValue != "(-)" {
// 版本感知:11g 下 NEXTVAL 默认值不能用 DEFAULT 子句(此处非 NEXTVAL
// 场景由 buildOracleDefault 正常生成 DEFAULT 子句)
if dv := buildOracleDefault(m.oracleDBVer(), field.DefaultValue, field); dv != "" {
expr.SQL += " " + dv
}
}
}
return
}
func (m Migrator) CreateConstraint(value interface{}, name string) error {
m.TryRemoveOnUpdate(value)
return m.Migrator.CreateConstraint(value, name)
@@ -263,11 +534,13 @@ func (m Migrator) HasIndex(value interface{}, name string) bool {
if idx := stmt.Schema.LookIndex(name); idx != nil {
name = idx.Name
}
// 索引名已是完整名称(如 IDX_TEST_USERS_EMAIL),直接大写后与 USER_INDEXES 中存储的名称比较,
// 不能再次通过 IndexName() 拼装,否则会得到错误的名字。
indexName := strings.ToUpper(name)
return m.DB.Raw(
"SELECT COUNT(*) FROM USER_INDEXES WHERE TABLE_NAME = ? AND INDEX_NAME = ?",
m.Migrator.DB.NamingStrategy.TableName(stmt.Table),
m.Migrator.DB.NamingStrategy.IndexName(stmt.Table, name),
indexName,
).Row().Scan(&count)
})
@@ -322,3 +595,91 @@ func (m Migrator) TryQuotifyReservedWords(values ...interface{}) error {
}
return nil
}
// CreateOnUpdateTrigger 创建 ON UPDATE 触发器
// Oracle 不支持原生的 ON UPDATE 外键操作,需要通过触发器模拟
func (m Migrator) CreateOnUpdateTrigger(value interface{}, rel *schema.Relationship) error {
if rel == nil {
return fmt.Errorf("relationship is nil")
}
constraint := rel.ParseConstraint()
if constraint == nil || constraint.OnUpdate == "" {
return nil
}
// 只处理 CASCADE 和 SET NULL
if constraint.OnUpdate != "CASCADE" && constraint.OnUpdate != "SET NULL" {
return nil
}
return m.RunWithValue(value, func(stmt *gorm.Statement) error {
triggerName := fmt.Sprintf("fk_trigger_%s_%s_%s",
stmt.Schema.Table,
rel.Field.DBName,
constraint.References[0].DBName,
)
var triggerSQL string
if constraint.OnUpdate == "CASCADE" {
// CASCADE: 当父表更新时,子表相应字段也更新
triggerSQL = fmt.Sprintf(`
CREATE OR REPLACE TRIGGER %s
AFTER UPDATE OF %s ON %s
FOR EACH ROW
BEGIN
UPDATE %s SET %s = :NEW.%s WHERE %s = :OLD.%s;
END;`,
triggerName,
constraint.References[0].DBName,
constraint.ReferenceSchema.Table,
stmt.Schema.Table,
rel.Field.DBName,
constraint.References[0].DBName,
rel.Field.DBName,
constraint.References[0].DBName,
)
} else if constraint.OnUpdate == "SET NULL" {
// SET NULL: 当父表更新时,子表相应字段设为 NULL
triggerSQL = fmt.Sprintf(`
CREATE OR REPLACE TRIGGER %s
AFTER UPDATE OF %s ON %s
FOR EACH ROW
BEGIN
UPDATE %s SET %s = NULL WHERE %s = :OLD.%s;
END;`,
triggerName,
constraint.References[0].DBName,
constraint.ReferenceSchema.Table,
stmt.Schema.Table,
rel.Field.DBName,
rel.Field.DBName,
constraint.References[0].DBName,
)
}
if triggerSQL != "" {
return m.DB.Exec(triggerSQL).Error
}
return nil
})
}
// DropOnUpdateTrigger 删除 ON UPDATE 触发器
func (m Migrator) DropOnUpdateTrigger(value interface{}, rel *schema.Relationship) error {
if rel == nil {
return fmt.Errorf("relationship is nil")
}
return m.RunWithValue(value, func(stmt *gorm.Statement) error {
triggerName := fmt.Sprintf("fk_trigger_%s_%s_%s",
stmt.Schema.Table,
rel.Field.DBName,
rel.Field.DBName,
)
return m.DB.Exec(fmt.Sprintf("DROP TRIGGER IF EXISTS %s", triggerName)).Error
})
}