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:
@@ -0,0 +1,237 @@
|
||||
package oracle
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"database/sql/driver"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
"gorm.io/gorm/schema"
|
||||
)
|
||||
|
||||
// convertValue 将 Go 值转换为 Oracle 兼容格式
|
||||
func convertValue(value interface{}, field *schema.Field) interface{} {
|
||||
if value == nil {
|
||||
return value
|
||||
}
|
||||
|
||||
switch v := value.(type) {
|
||||
case bool:
|
||||
if v {
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
case string:
|
||||
// 检查是否超长字符串,如果需要CLOB处理则保持原样
|
||||
if len(v) > 4000 {
|
||||
// 标记需要CLOB处理,这里只是示例,实际可能需要额外逻辑
|
||||
return v
|
||||
}
|
||||
return v
|
||||
case driver.Valuer:
|
||||
// 调用 Value() 方法解包
|
||||
val, err := v.Value()
|
||||
if err != nil {
|
||||
return value // 如果出错,返回原始值
|
||||
}
|
||||
return val
|
||||
case time.Time:
|
||||
return v
|
||||
default:
|
||||
return value
|
||||
}
|
||||
}
|
||||
|
||||
// convertFromOracleToField 将 Oracle 返回值转换为 Go 类型
|
||||
func convertFromOracleToField(value interface{}, field *schema.Field) interface{} {
|
||||
if value == nil {
|
||||
return value
|
||||
}
|
||||
|
||||
switch v := value.(type) {
|
||||
case sql.NullTime:
|
||||
if v.Valid {
|
||||
return v.Time
|
||||
}
|
||||
return nil
|
||||
case sql.NullInt64:
|
||||
if v.Valid {
|
||||
return v.Int64
|
||||
}
|
||||
return nil
|
||||
case sql.NullFloat64:
|
||||
if v.Valid {
|
||||
return v.Float64
|
||||
}
|
||||
return nil
|
||||
case sql.NullBool:
|
||||
if v.Valid {
|
||||
return v.Bool
|
||||
}
|
||||
return nil
|
||||
case sql.NullString:
|
||||
if v.Valid {
|
||||
return v.String
|
||||
}
|
||||
return nil
|
||||
default:
|
||||
return value
|
||||
}
|
||||
}
|
||||
|
||||
// validateCreateData 验证创建数据
|
||||
func validateCreateData(data interface{}) error {
|
||||
if data == nil {
|
||||
return fmt.Errorf("create data cannot be nil")
|
||||
}
|
||||
|
||||
rv := reflect.ValueOf(data)
|
||||
if rv.Kind() == reflect.Ptr {
|
||||
if rv.IsNil() {
|
||||
return fmt.Errorf("create data pointer cannot be nil")
|
||||
}
|
||||
rv = rv.Elem()
|
||||
}
|
||||
|
||||
if rv.Kind() == reflect.Slice {
|
||||
if rv.Len() == 0 {
|
||||
return fmt.Errorf("create data slice cannot be empty")
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// checkMissingWhereConditions 检查 WHERE 条件是否缺失
|
||||
func checkMissingWhereConditions(conditions []clause.Expression, schema *schema.Schema) bool {
|
||||
if len(conditions) == 0 {
|
||||
return true
|
||||
}
|
||||
count := 0
|
||||
for _, condition := range conditions {
|
||||
// 检查是否是软删除条件 (deleted_at IS NULL)
|
||||
if isSoftDeleteCondition(condition, schema) {
|
||||
count++
|
||||
}
|
||||
}
|
||||
|
||||
// 如果只有软删除条件或没有其他条件,则认为缺少WHERE条件
|
||||
return count >= len(conditions)
|
||||
}
|
||||
|
||||
// isSoftDeleteCondition 检查条件是否为软删除条件(deleted_at IS NULL)
|
||||
func isSoftDeleteCondition(condition clause.Expression, sch *schema.Schema) bool {
|
||||
// 查找是否有 deleted_at 字段
|
||||
var softDeleteField *schema.Field
|
||||
for _, field := range sch.Fields {
|
||||
if strings.EqualFold(field.DBName, "deleted_at") || field.Name == "DeletedAt" {
|
||||
softDeleteField = field
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if softDeleteField == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
// GORM 的软删除条件在 WHERE 中表现为对 deleted_at 列的相等/包含判断(值为空,构建为 IS NULL)
|
||||
switch cond := condition.(type) {
|
||||
case clause.Eq:
|
||||
return strings.EqualFold(columnNameOf(cond.Column), softDeleteField.DBName)
|
||||
case clause.IN:
|
||||
return strings.EqualFold(columnNameOf(cond.Column), softDeleteField.DBName)
|
||||
case clause.Neq:
|
||||
return strings.EqualFold(columnNameOf(cond.Column), softDeleteField.DBName)
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// columnNameOf 从 clause 表达式的 Column 字段中提取列名
|
||||
func columnNameOf(col interface{}) string {
|
||||
switch c := col.(type) {
|
||||
case clause.Column:
|
||||
return c.Name
|
||||
case string:
|
||||
return c
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// addPrimaryKeyWhere 根据模型的主键值注入 WHERE 条件。
|
||||
// GORM 默认的 Update/Delete 回调会在语句中注入主键条件,
|
||||
// 但驱动自定义回调替换了默认实现,因此需要手动补齐。
|
||||
// 返回注入的主键值数量(0 表示没有可用的主键值)。
|
||||
func addPrimaryKeyWhere(stmt *gorm.Statement, sch *schema.Schema) int {
|
||||
if stmt == nil || sch == nil {
|
||||
return 0
|
||||
}
|
||||
|
||||
_, queryValues := schema.GetIdentityFieldValuesMap(stmt.Context, stmt.ReflectValue, sch.PrimaryFields)
|
||||
column, values := schema.ToQueryValues(stmt.Table, sch.PrimaryFieldDBNames, queryValues)
|
||||
if len(values) > 0 {
|
||||
stmt.AddClause(clause.Where{Exprs: []clause.Expression{clause.IN{Column: column, Values: values}}})
|
||||
}
|
||||
return len(values)
|
||||
}
|
||||
|
||||
// buildOracleDefault 智能转换默认值
|
||||
// dbVer 用于版本感知的默认值处理:Oracle 11g 的 DEFAULT 子句不允许引用序列的
|
||||
// NEXTVAL(ORA-00984,12c 才引入该能力),因此 11g 下 NEXTVAL 分支返回空字符串,
|
||||
// 由调用方通过 BEFORE INSERT 触发器实现等价语义(见 migrator.createSequenceDefaultTrigger)。
|
||||
func buildOracleDefault(dbVer string, defaultValue string, field *schema.Field) string {
|
||||
if defaultValue == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
lowerVal := strings.ToLower(strings.TrimSpace(defaultValue))
|
||||
|
||||
switch lowerVal {
|
||||
case "null":
|
||||
return "DEFAULT NULL"
|
||||
case "current_timestamp", "now()":
|
||||
return "DEFAULT CURRENT_TIMESTAMP"
|
||||
case "sysdate":
|
||||
return "DEFAULT SYSDATE"
|
||||
case "true":
|
||||
return "DEFAULT 1"
|
||||
case "false":
|
||||
return "DEFAULT 0"
|
||||
default:
|
||||
// 检查是否为序列
|
||||
if strings.Contains(strings.ToUpper(defaultValue), ".NEXTVAL") {
|
||||
// 12c+ 原生支持 DEFAULT <seq>.NEXTVAL,直接生成 DEFAULT 子句
|
||||
if supportsIdentity(dbVer) {
|
||||
// 去掉可能的包裹括号(GORM 对含括号的默认值保持原文,
|
||||
// 如 "(SEQ_MY.NEXTVAL)"),生成标准的 DEFAULT <seq>.NEXTVAL
|
||||
seqExpr := strings.TrimSpace(defaultValue)
|
||||
seqExpr = strings.TrimPrefix(seqExpr, "(")
|
||||
seqExpr = strings.TrimSuffix(seqExpr, ")")
|
||||
return fmt.Sprintf("DEFAULT %s", strings.TrimSpace(seqExpr))
|
||||
}
|
||||
// 11g 不支持 DEFAULT 子句引用序列(ORA-00984),
|
||||
// 返回空字符串,由建表流程创建 BEFORE INSERT 触发器实现等价语义
|
||||
return ""
|
||||
}
|
||||
|
||||
// 检查日期格式 "2006-01-02"
|
||||
dateRegex := regexp.MustCompile(`^\d{4}-\d{2}-\d{2}$`)
|
||||
if dateRegex.MatchString(defaultValue) {
|
||||
return fmt.Sprintf("DEFAULT TO_DATE('%s', 'YYYY-MM-DD')", defaultValue)
|
||||
}
|
||||
|
||||
// 检查时间戳格式 "2006-01-02 15:04:05"
|
||||
timestampRegex := regexp.MustCompile(`^\d{4}-\d{2}-\d{2} \d{2}:\d{2}:\d{2}$`)
|
||||
if timestampRegex.MatchString(defaultValue) {
|
||||
return fmt.Sprintf("DEFAULT TO_DATE('%s', 'YYYY-MM-DD HH24:MI:SS')", defaultValue)
|
||||
}
|
||||
|
||||
// 普通字符串用单引号包围
|
||||
return fmt.Sprintf("DEFAULT '%s'", defaultValue)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user