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
+58
View File
@@ -0,0 +1,58 @@
package clauses
import (
"strconv"
"testing"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"gorm.io/gorm/schema"
)
// testDialector 是仅用于 SQL 生成测试的最小 Dialector 实现,无需数据库连接。
// 它模拟 Oracle 的绑定变量占位符(:N)和不加引号的标识符引用。
type testDialector struct{}
func (testDialector) Name() string { return "oracle" }
func (testDialector) Initialize(*gorm.DB) error { return nil }
func (testDialector) Migrator(*gorm.DB) gorm.Migrator { return nil }
func (testDialector) DataTypeOf(*schema.Field) string { return "" }
func (testDialector) DefaultValueOf(*schema.Field) clause.Expression { return nil }
func (testDialector) BindVarTo(writer clause.Writer, stmt *gorm.Statement, v interface{}) {
writer.WriteString(":")
writer.WriteString(strconv.Itoa(len(stmt.Vars)))
}
func (testDialector) QuoteTo(writer clause.Writer, str string) {
writer.WriteString(str)
}
func (testDialector) Explain(sql string, vars ...interface{}) string { return sql }
// newStatement 构造一个可以直接作为 clause.Builder 使用的 gorm.Statement。
func newStatement(t *testing.T) *gorm.Statement {
t.Helper()
db := &gorm.DB{
Config: &gorm.Config{Dialector: testDialector{}},
}
return &gorm.Statement{DB: db, Table: "users"}
}
// buildSQL 直接调用子句的 Build 方法生成 SQL。
func buildSQL(t *testing.T, expr clause.Expression) string {
t.Helper()
stmt := newStatement(t)
expr.Build(stmt)
return stmt.SQL.String()
}
// buildClauseSQL 通过 clause.Clause(模拟 gorm.Statement.Build 的构建流程)
// 生成带子句名前缀的完整 SQL。
func buildClauseSQL(t *testing.T, name string, expr clause.Expression) string {
t.Helper()
stmt := newStatement(t)
cc := clause.Clause{Name: name, Expression: expr}
cc.Build(stmt)
return stmt.SQL.String()
}
+78
View File
@@ -0,0 +1,78 @@
package clauses
import (
"reflect"
"strings"
"testing"
"gorm.io/gorm/clause"
)
func TestMergeName(t *testing.T) {
var m Merge
if got := m.Name(); got != "MERGE" {
t.Errorf("Merge.Name() = %q, want %q", got, "MERGE")
}
}
func TestMergeDefaultExcludeName(t *testing.T) {
if got := MergeDefaultExcludeName(); got != "exclude" {
t.Errorf("MergeDefaultExcludeName() = %q, want %q", got, "exclude")
}
}
func TestMergeBuild(t *testing.T) {
merge := Merge{
Table: clause.Table{Name: "users"},
Using: []clause.Interface{
clause.Select{Columns: []clause.Column{{Name: "id"}, {Name: "name"}}},
clause.From{Tables: []clause.Table{{Name: "users"}}},
},
On: []clause.Expression{
clause.Eq{Column: clause.Column{Name: "a"}, Value: clause.Column{Name: "b"}},
},
}
sql := buildClauseSQL(t, "MERGE", merge)
for _, want := range []string{
"MERGE INTO", // 前缀 + Insert 构建
"USING (", // USING 子查询
"SELECT id,name FROM users", // USING 子查询内容
") exclude ON (", // exclude 别名 + ON
"a = b", // ON 条件
} {
if !strings.Contains(sql, want) {
t.Errorf("Merge SQL %q does not contain %q", sql, want)
}
}
}
func TestMergeBuildEmpty(t *testing.T) {
var merge Merge
sql := buildClauseSQL(t, "MERGE", merge)
if want := "MERGE INTO users USING () exclude ON ()"; sql != want {
t.Errorf("empty Merge SQL = %q, want %q", sql, want)
}
}
func TestMergeMergeClause(t *testing.T) {
merge := Merge{
Table: clause.Table{Name: "users"},
On: []clause.Expression{
clause.Eq{Column: clause.Column{Name: "a"}, Value: clause.Column{Name: "b"}},
},
}
cc := &clause.Clause{}
merge.MergeClause(cc)
if cc.Name != "MERGE" {
t.Errorf("MergeClause name = %q, want %q", cc.Name, "MERGE")
}
if got, ok := cc.Expression.(Merge); !ok || !reflect.DeepEqual(got, merge) {
t.Errorf("MergeClause expression = %#v, want %#v", cc.Expression, merge)
}
}
+34
View File
@@ -1,6 +1,7 @@
package clauses
import (
"fmt"
"gorm.io/gorm/clause"
)
@@ -8,3 +9,36 @@ type ReturningInto struct {
Variables []clause.Column
Into []*clause.Values
}
// Name returns the name of the clause
func (r ReturningInto) Name() string {
return "RETURNING INTO"
}
// Build builds the SQL for the RETURNING INTO clause
func (r ReturningInto) Build(builder clause.Builder) {
if len(r.Variables) > 0 {
builder.WriteString("RETURNING ")
for idx, col := range r.Variables {
if idx > 0 {
builder.WriteByte(',')
}
builder.WriteQuoted(col)
}
builder.WriteString(" INTO ")
for idx := range r.Variables {
if idx > 0 {
builder.WriteByte(',')
}
// 写入绑定变量占位符
builder.WriteString(fmt.Sprintf(":%d", idx+1))
}
}
}
// MergeClause merge returning into clause
func (r ReturningInto) MergeClause(clause *clause.Clause) {
clause.Name = r.Name()
clause.Expression = r
}
+63
View File
@@ -0,0 +1,63 @@
package clauses
import (
"reflect"
"testing"
"gorm.io/gorm/clause"
)
func TestReturningIntoName(t *testing.T) {
var r ReturningInto
if got := r.Name(); got != "RETURNING INTO" {
t.Errorf("ReturningInto.Name() = %q, want %q", got, "RETURNING INTO")
}
}
func TestReturningIntoBuild(t *testing.T) {
r := ReturningInto{
Variables: []clause.Column{{Name: "col1"}, {Name: "col2"}},
}
sql := buildSQL(t, r)
if want := "RETURNING col1,col2 INTO :1,:2"; sql != want {
t.Errorf("ReturningInto SQL = %q, want %q", sql, want)
}
}
func TestReturningIntoBuildSingle(t *testing.T) {
r := ReturningInto{
Variables: []clause.Column{{Name: "col1"}},
}
sql := buildSQL(t, r)
if want := "RETURNING col1 INTO :1"; sql != want {
t.Errorf("ReturningInto SQL = %q, want %q", sql, want)
}
}
func TestReturningIntoBuildEmpty(t *testing.T) {
var r ReturningInto
sql := buildSQL(t, r)
if sql != "" {
t.Errorf("ReturningInto with empty Variables should generate no SQL, got %q", sql)
}
}
func TestReturningIntoMergeClause(t *testing.T) {
r := ReturningInto{
Variables: []clause.Column{{Name: "id"}},
}
cc := &clause.Clause{}
r.MergeClause(cc)
if cc.Name != "RETURNING INTO" {
t.Errorf("ReturningInto.MergeClause name = %q, want %q", cc.Name, "RETURNING INTO")
}
if got, ok := cc.Expression.(ReturningInto); !ok || !reflect.DeepEqual(got, r) {
t.Errorf("ReturningInto.MergeClause expression = %#v, want %#v", cc.Expression, r)
}
}
+2 -1
View File
@@ -19,7 +19,8 @@ func (w WhenMatched) Build(builder clause.Builder) {
builder.WriteString(" UPDATE ")
builder.WriteString(w.Name())
builder.WriteByte(' ')
w.Build(builder)
builder.WriteString("SET ")
w.Set.Build(builder)
buildWhere := func(where clause.Where) {
builder.WriteString(where.Name())
+100
View File
@@ -0,0 +1,100 @@
package clauses
import (
"strings"
"testing"
"gorm.io/gorm/clause"
)
func TestWhenMatchedName(t *testing.T) {
var w WhenMatched
if got := w.Name(); got != "WHEN MATCHED" {
t.Errorf("WhenMatched.Name() = %q, want %q", got, "WHEN MATCHED")
}
}
func TestWhenMatchedBuild(t *testing.T) {
w := WhenMatched{
Set: clause.Set{{Column: clause.Column{Name: "name"}, Value: "x"}},
}
sql := buildSQL(t, w)
// 期望: THEN UPDATE WHEN MATCHED SET name=:1
for _, want := range []string{
"THEN UPDATE",
"WHEN MATCHED",
"SET name=:1",
} {
if !strings.Contains(sql, want) {
t.Errorf("WhenMatched SQL %q does not contain %q", sql, want)
}
}
}
func TestWhenMatchedBuildMultipleSet(t *testing.T) {
w := WhenMatched{
Set: clause.Set{
{Column: clause.Column{Name: "name"}, Value: "x"},
{Column: clause.Column{Name: "age"}, Value: 1},
},
}
sql := buildSQL(t, w)
if want := "SET name=:1,age=:2"; !strings.Contains(sql, want) {
t.Errorf("WhenMatched SQL %q does not contain %q", sql, want)
}
}
func TestWhenMatchedBuildWithWhereAndDelete(t *testing.T) {
w := WhenMatched{
Set: clause.Set{{Column: clause.Column{Name: "name"}, Value: "x"}},
Where: clause.Where{Exprs: []clause.Expression{
clause.Eq{Column: clause.Column{Name: "id"}, Value: 1},
}},
Delete: clause.Where{Exprs: []clause.Expression{
clause.Eq{Column: clause.Column{Name: "flag"}, Value: 0},
}},
}
sql := buildSQL(t, w)
// 期望: THEN UPDATE WHEN MATCHED SET name=:1WHERE id = :2 DELETE WHERE flag = :3
for _, want := range []string{
"THEN UPDATE",
"WHEN MATCHED",
"SET name=:1",
"WHERE id = :2",
"DELETE WHERE flag = :3",
} {
if !strings.Contains(sql, want) {
t.Errorf("WhenMatched SQL %q does not contain %q", sql, want)
}
}
}
// TestWhenMatchedClauseBuild 验证通过 clause.Clause(gorm 真实构建流程)时的完整输出。
func TestWhenMatchedClauseBuild(t *testing.T) {
w := WhenMatched{
Set: clause.Set{{Column: clause.Column{Name: "name"}, Value: "x"}},
}
sql := buildClauseSQL(t, "WHEN MATCHED", w)
for _, want := range []string{
"WHEN MATCHED",
"THEN UPDATE",
"SET name=:1",
} {
if !strings.Contains(sql, want) {
t.Errorf("WhenMatched SQL %q does not contain %q", sql, want)
}
}
}
func TestWhenMatchedBuildEmptySet(t *testing.T) {
var w WhenMatched
sql := buildSQL(t, w)
if sql != "" {
t.Errorf("WhenMatched with empty Set should generate no SQL, got %q", sql)
}
}
+1 -1
View File
@@ -21,7 +21,7 @@ func (w WhenNotMatched) Build(builder clause.Builder) {
builder.WriteString(" THEN")
builder.WriteString(" INSERT ")
w.Build(builder)
w.Values.Build(builder)
if len(w.Where.Exprs) > 0 {
builder.WriteString(w.Where.Name())
+93
View File
@@ -0,0 +1,93 @@
package clauses
import (
"strings"
"testing"
"gorm.io/gorm/clause"
)
func TestWhenNotMatchedName(t *testing.T) {
var w WhenNotMatched
if got := w.Name(); got != "WHEN NOT MATCHED" {
t.Errorf("WhenNotMatched.Name() = %q, want %q", got, "WHEN NOT MATCHED")
}
}
func TestWhenNotMatchedBuild(t *testing.T) {
w := WhenNotMatched{
Values: clause.Values{
Columns: []clause.Column{{Name: "name"}, {Name: "age"}},
Values: [][]interface{}{{"x", 1}},
},
}
sql := buildClauseSQL(t, "WHEN NOT MATCHED", w)
// 期望: WHEN NOT MATCHED THEN INSERT (name,age) VALUES (:1,:2)
for _, want := range []string{
"WHEN NOT MATCHED",
"THEN INSERT",
"(name,age)", // 列名
"VALUES (:1,:2)", // VALUES 绑定参数
} {
if !strings.Contains(sql, want) {
t.Errorf("WhenNotMatched SQL %q does not contain %q", sql, want)
}
}
}
func TestWhenNotMatchedBuildWithWhere(t *testing.T) {
w := WhenNotMatched{
Values: clause.Values{
Columns: []clause.Column{{Name: "name"}},
Values: [][]interface{}{{"x"}},
},
Where: clause.Where{Exprs: []clause.Expression{
clause.Eq{Column: clause.Column{Name: "deleted"}, Value: 0},
}},
}
sql := buildSQL(t, w)
for _, want := range []string{
"THEN INSERT",
"(name)",
"VALUES (:1)",
"WHERE deleted = :2",
} {
if !strings.Contains(sql, want) {
t.Errorf("WhenNotMatched SQL %q does not contain %q", sql, want)
}
}
}
func TestWhenNotMatchedBuildEmpty(t *testing.T) {
var w WhenNotMatched
sql := buildSQL(t, w)
if sql != "" {
t.Errorf("WhenNotMatched with empty Columns should generate no SQL, got %q", sql)
}
}
// TestWhenNotMatchedBuildPanicsOnMultipleRows 验证多行插入时按 Oracle 限制 panic。
func TestWhenNotMatchedBuildPanicsOnMultipleRows(t *testing.T) {
w := WhenNotMatched{
Values: clause.Values{
Columns: []clause.Column{{Name: "name"}},
Values: [][]interface{}{{"x"}, {"y"}},
},
}
defer func() {
r := recover()
if r == nil {
t.Fatal("expected panic for multiple insert rows")
}
msg, ok := r.(string)
if !ok || !strings.Contains(msg, "cannot insert more than one rows") {
t.Errorf("unexpected panic message: %v", r)
}
}()
buildSQL(t, w)
}
+237
View File
@@ -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 子句不允许引用序列的
// NEXTVALORA-0098412c 才引入该能力),因此 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)
}
}
+254
View File
@@ -0,0 +1,254 @@
package oracle
import (
"database/sql"
"database/sql/driver"
"errors"
"strings"
"sync"
"testing"
"time"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"gorm.io/gorm/schema"
)
// ---- 测试辅助类型 ----
// testValuer 实现 driver.Valuer,用于测试 convertValue 的解包逻辑
type testValuer struct {
value driver.Value
err error
}
func (v testValuer) Value() (driver.Value, error) {
return v.value, v.err
}
// softDeleteModel 含软删除字段,用于构造 schema
type softDeleteModel struct {
ID uint `gorm:"primaryKey"`
DeletedAt gorm.DeletedAt
}
// plainModel 无软删除字段
type plainModel struct {
ID uint `gorm:"primaryKey"`
Name string
}
// createDataModel 用于 validateCreateData 测试
type createDataModel struct {
ID int
}
func parseTestSchema(t *testing.T, model interface{}) *schema.Schema {
t.Helper()
sch, err := schema.Parse(model, &sync.Map{}, schema.NamingStrategy{})
if err != nil {
t.Fatalf("failed to parse schema: %v", err)
}
return sch
}
// ---- TestConvertValue ----
func TestConvertValue(t *testing.T) {
now := time.Now()
longString := strings.Repeat("a", 4001)
tests := []struct {
name string
value interface{}
want interface{}
}{
{"bool true to 1", true, 1},
{"bool false to 0", false, 0},
{"plain string unchanged", "hello", "hello"},
{"long string (over 4000) unchanged", longString, longString},
{"nil unchanged", nil, nil},
{"time.Time unchanged", now, now},
{"driver.Valuer unwrapped", testValuer{value: int64(42)}, int64(42)},
{"driver.Valuer with string", testValuer{value: "val"}, "val"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := convertValue(tt.value, nil)
if tt.want == nil {
if got != nil {
t.Errorf("convertValue(%v) = %v, want nil", tt.value, got)
}
return
}
if got != tt.want {
t.Errorf("convertValue(%v) = %v (%T), want %v (%T)", tt.value, got, got, tt.want, tt.want)
}
})
}
}
func TestConvertValueValuerError(t *testing.T) {
v := testValuer{err: errors.New("valuing failed")}
got := convertValue(v, nil)
// 出错时返回原始值
if _, ok := got.(testValuer); !ok {
t.Errorf("expected original value returned on Valuer error, got %#v", got)
}
}
// ---- TestConvertFromOracleToField ----
func TestConvertFromOracleToField(t *testing.T) {
now := time.Now()
tests := []struct {
name string
value interface{}
want interface{}
}{
{"nil unchanged", nil, nil},
{"plain int unchanged", 42, 42},
{"plain string unchanged", "abc", "abc"},
{"NullTime valid", sql.NullTime{Time: now, Valid: true}, now},
{"NullTime invalid", sql.NullTime{Valid: false}, nil},
{"NullInt64 valid", sql.NullInt64{Int64: 99, Valid: true}, int64(99)},
{"NullInt64 invalid", sql.NullInt64{Valid: false}, nil},
{"NullFloat64 valid", sql.NullFloat64{Float64: 3.14, Valid: true}, 3.14},
{"NullFloat64 invalid", sql.NullFloat64{Valid: false}, nil},
{"NullBool valid", sql.NullBool{Bool: true, Valid: true}, true},
{"NullBool invalid", sql.NullBool{Valid: false}, nil},
{"NullString valid", sql.NullString{String: "x", Valid: true}, "x"},
{"NullString invalid", sql.NullString{Valid: false}, nil},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := convertFromOracleToField(tt.value, nil)
if tt.want == nil {
if got != nil {
t.Errorf("convertFromOracleToField(%v) = %v, want nil", tt.value, got)
}
return
}
if got != tt.want {
t.Errorf("convertFromOracleToField(%v) = %v (%T), want %v (%T)", tt.value, got, got, tt.want, tt.want)
}
})
}
}
// ---- TestBuildOracleDefault ----
func TestBuildOracleDefault(t *testing.T) {
// 12c 及以上版本号(NEXTVAL 默认值走 DEFAULT 子句)
const dbVer12c = "12.1.0.2.0"
// 11g 版本号(NEXTVAL 默认值不生成 DEFAULT 子句)
const dbVer11g = "11.2.0.4.0"
tests := []struct {
name string
dbVer string
value string
expected string
}{
{"empty string", dbVer12c, "", ""},
{"NULL keyword", dbVer12c, "NULL", "DEFAULT NULL"},
{"null lowercase", dbVer12c, "null", "DEFAULT NULL"},
{"CURRENT_TIMESTAMP", dbVer12c, "CURRENT_TIMESTAMP", "DEFAULT CURRENT_TIMESTAMP"},
{"now()", dbVer12c, "now()", "DEFAULT CURRENT_TIMESTAMP"},
{"SYSDATE", dbVer12c, "SYSDATE", "DEFAULT SYSDATE"},
{"sysdate lowercase", dbVer12c, "sysdate", "DEFAULT SYSDATE"},
{"TRUE", dbVer12c, "TRUE", "DEFAULT 1"},
{"false lowercase", dbVer12c, "false", "DEFAULT 0"},
// 12c+ 原生支持 DEFAULT <seq>.NEXTVAL,行为不变
{"sequence nextval 12c", dbVer12c, "SEQ_MY.NEXTVAL", "DEFAULT SEQ_MY.NEXTVAL"},
// 11g 的 DEFAULT 子句不支持引用序列 NEXTVALORA-00984),返回空串,
// 由建表流程创建 BEFORE INSERT 触发器实现等价语义
{"sequence nextval 11g", dbVer11g, "SEQ_MY.NEXTVAL", ""},
{"sequence nextval 11g lowercase", dbVer11g, "seq_my.nextval", ""},
{"date format", dbVer12c, "2006-01-02", "DEFAULT TO_DATE('2006-01-02', 'YYYY-MM-DD')"},
{"timestamp format", dbVer12c, "2006-01-02 15:04:05", "DEFAULT TO_DATE('2006-01-02 15:04:05', 'YYYY-MM-DD HH24:MI:SS')"},
{"plain string", dbVer12c, "hello", "DEFAULT 'hello'"},
{"plain string with spaces", dbVer12c, " hello ", "DEFAULT ' hello '"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := buildOracleDefault(tt.dbVer, tt.value, nil)
if got != tt.expected {
t.Errorf("buildOracleDefault(%q, %q) = %q, want %q", tt.dbVer, tt.value, got, tt.expected)
}
})
}
}
// ---- TestCheckMissingWhereConditions ----
func TestCheckMissingWhereConditions(t *testing.T) {
softSch := parseTestSchema(t, &softDeleteModel{})
plainSch := parseTestSchema(t, &plainModel{})
softDeleteCond := clause.Eq{Column: clause.Column{Name: "deleted_at"}, Value: nil}
normalCond := clause.Eq{Column: "age", Value: 25}
tests := []struct {
name string
conditions []clause.Expression
schema *schema.Schema
want bool
}{
{"empty conditions", nil, softSch, true},
{"empty slice", []clause.Expression{}, plainSch, true},
{"only soft delete condition", []clause.Expression{softDeleteCond}, softSch, true},
{"soft delete + normal condition", []clause.Expression{softDeleteCond, normalCond}, softSch, false},
{"only normal condition", []clause.Expression{normalCond}, softSch, false},
{"soft delete condition on model without deleted_at", []clause.Expression{softDeleteCond}, plainSch, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := checkMissingWhereConditions(tt.conditions, tt.schema)
if got != tt.want {
t.Errorf("checkMissingWhereConditions(%v) = %v, want %v", tt.conditions, got, tt.want)
}
})
}
}
// ---- TestValidateCreateData ----
func TestValidateCreateData(t *testing.T) {
tests := []struct {
name string
data interface{}
wantErr string
}{
{"nil data", nil, "create data cannot be nil"},
{"nil pointer", (*createDataModel)(nil), "create data pointer cannot be nil"},
{"empty slice", []createDataModel{}, "create data slice cannot be empty"},
{"valid struct", createDataModel{ID: 1}, ""},
{"valid non-empty slice", []createDataModel{{ID: 1}}, ""},
{"valid pointer to struct", &createDataModel{ID: 2}, ""},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := validateCreateData(tt.data)
if tt.wantErr == "" {
if err != nil {
t.Errorf("validateCreateData(%v) = %v, want nil error", tt.data, err)
}
return
}
if err == nil {
t.Errorf("validateCreateData(%v) = nil, want error containing %q", tt.data, tt.wantErr)
return
}
if !strings.Contains(err.Error(), tt.wantErr) {
t.Errorf("validateCreateData(%v) error = %q, want containing %q", tt.data, err.Error(), tt.wantErr)
}
})
}
}
+53 -6
View File
@@ -3,6 +3,7 @@ package oracle
import (
"bytes"
"database/sql"
"fmt"
"reflect"
"github.com/thoas/go-funk"
@@ -96,6 +97,33 @@ func Create(db *gorm.DB) {
}
if !db.DryRun {
// 开启事务以确保批量插入的一致性
var tx *sql.Tx
var err error
var isTransaction bool = false
// 检查是否已经在一个事务中
if sqlTx, ok := stmt.ConnPool.(*sql.Tx); ok {
tx = sqlTx
isTransaction = true
} else if sqlDb, ok := stmt.ConnPool.(*sql.DB); ok {
tx, err = sqlDb.Begin()
if err != nil {
db.AddError(err)
return
}
defer func() {
if db.Error != nil && !isTransaction {
_ = tx.Rollback()
} else if !isTransaction {
_ = tx.Commit()
}
}()
} else {
db.AddError(fmt.Errorf("unsupported connection pool type"))
return
}
for idx, vals := range values.Values {
// HACK HACK: replace values one by one, assuming its value layout will be the same all the time, i.e. aligned
for idx, val := range vals {
@@ -113,13 +141,18 @@ func Create(db *gorm.DB) {
// and then we insert each row one by one then put the returning values back (i.e. last return id => smart insert)
// we keep track of the index so that the sub-reflected value is also correct
// BIG BUG: what if any of the transactions failed? some result might already be inserted that oracle is so
// sneaky that some transaction inserts will exceed the buffer and so will be pushed at unknown point,
// resulting in dangling row entries, so we might need to delete them if an error happens
var execConn *sql.Tx
if isTransaction {
execConn = tx // 已经在事务中,直接使用原事务
} else {
execConn = tx // 使用新创建的事务
}
switch result, err := stmt.ConnPool.ExecContext(stmt.Context, stmt.SQL.String(), stmt.Vars...); err {
switch result, err := execConn.ExecContext(stmt.Context, stmt.SQL.String(), stmt.Vars...); err {
case nil: // success
db.RowsAffected, _ = result.RowsAffected()
// 批量插入时累加每个单行插入的受影响行数
rowsAffected, _ := result.RowsAffected()
db.RowsAffected += rowsAffected
insertTo := stmt.ReflectValue
switch insertTo.Kind() {
@@ -140,13 +173,27 @@ func Create(db *gorm.DB) {
db.AddError(err)
}
case reflect.Map:
// todo 设置id的
// 设置Map类型的ID
mapValue := reflect.ValueOf(insertTo.Interface())
if mapValue.IsValid() && mapValue.Type().Key().Kind() == reflect.String {
keyValue := reflect.ValueOf(field.DBName)
destValue := reflect.ValueOf(stmt.Vars[boundVars[field.Name]].(sql.Out).Dest)
if destValue.Kind() == reflect.Ptr {
destValue = destValue.Elem()
}
mapValue.SetMapIndex(keyValue, destValue)
}
}
},
)
}
default: // failure
db.AddError(err)
// 如果不是在已有事务中,则回滚我们创建的事务
if !isTransaction {
_ = tx.Rollback()
}
return
}
}
}
+322
View File
@@ -0,0 +1,322 @@
package oracle
import (
"database/sql"
"fmt"
"reflect"
"time"
"github.com/thoas/go-funk"
"gorm.io/gorm"
"gorm.io/gorm/clause"
gormSchema "gorm.io/gorm/schema"
)
func Delete(db *gorm.DB) {
stmt := db.Statement
if stmt == nil {
return
}
schema := stmt.Schema
if schema == nil {
return
}
boundVars := make(map[string]int)
// 注入主键 WHERE 条件(GORM 默认回调会做这一步)
addPrimaryKeyWhere(stmt, schema)
// 1. WHERE 安全检查(最重要)
if where, ok := stmt.Clauses["WHERE"].Expression.(clause.Where); ok {
if checkMissingWhereConditions(where.Exprs, schema) {
db.AddError(fmt.Errorf("missing WHERE condition in DELETE"))
return
}
} else {
// 没有 WHERE 子句
db.AddError(fmt.Errorf("missing WHERE condition in DELETE"))
return
}
// 2. 检查软删除
var softDeleteField *gormSchema.Field
for _, field := range schema.Fields {
if (field.DBName == "deleted_at" || field.Name == "DeletedAt") &&
field.GORMDataType == "time" {
softDeleteField = field
break
}
}
if softDeleteField != nil && !stmt.Unscoped {
// 软删除:转换为 UPDATE
performSoftDelete(db, softDeleteField, boundVars)
} else {
// 硬删除:执行 DELETEUnscoped 时强制硬删除)
performHardDelete(db, boundVars)
}
}
func performSoftDelete(db *gorm.DB, field *gormSchema.Field, boundVars map[string]int) {
stmt := db.Statement
schema := stmt.Schema
hasDefaultValues := len(schema.FieldsWithDefaultDBValue) > 0
if !stmt.Unscoped {
for _, c := range schema.DeleteClauses {
stmt.AddClause(c)
}
}
if stmt.SQL.String() == "" {
// 构建 UPDATE 语句而不是 DELETE
stmt.AddClauseIfNotExists(clause.Update{Table: clause.Table{Name: stmt.Schema.Table}})
// 构建 SET 子句,设置 deleted_at 为当前时间
now := time.Now()
convertedNow := convertValue(now, field)
set := clause.Set{clause.Assignment{Column: clause.Column{Name: field.DBName}, Value: convertedNow}}
stmt.AddClause(set)
// 添加 RETURNING 子句(如果有默认值字段或需要返回值)
if hasDefaultValues {
stmt.AddClauseIfNotExists(clause.Returning{
Columns: funk.Map(schema.FieldsWithDefaultDBValue, func(field *gormSchema.Field) clause.Column {
return clause.Column{Name: field.DBName}
}).([]clause.Column),
})
}
// 构建语句
stmt.Build("UPDATE", "SET", "WHERE", "RETURNING")
// 如果有 RETURNING 子句,添加 INTO 子句
if hasDefaultValues {
stmt.WriteString(" INTO ")
for idx, field := range schema.FieldsWithDefaultDBValue {
if idx > 0 {
stmt.WriteByte(',')
}
boundVars[field.Name] = len(stmt.Vars)
stmt.AddVar(stmt, sql.Out{Dest: reflect.New(field.FieldType).Interface()})
}
}
}
if !db.DryRun {
// 执行软删除操作
var tx *sql.Tx
var err error
var isTransaction bool = false
// 检查是否已经在一个事务中
if sqlTx, ok := stmt.ConnPool.(*sql.Tx); ok {
tx = sqlTx
isTransaction = true
} else if sqlDb, ok := stmt.ConnPool.(*sql.DB); ok {
tx, err = sqlDb.Begin()
if err != nil {
db.AddError(err)
return
}
defer func() {
if db.Error != nil && !isTransaction {
_ = tx.Rollback()
} else if !isTransaction {
_ = tx.Commit()
}
}()
} else {
db.AddError(fmt.Errorf("unsupported connection pool type"))
return
}
var execConn *sql.Tx
if isTransaction {
execConn = tx // 已经在事务中,直接使用原事务
} else {
execConn = tx // 使用新创建的事务
}
result, err := execConn.ExecContext(stmt.Context, stmt.SQL.String(), stmt.Vars...)
if err != nil {
db.AddError(err)
// 如果不是在已有事务中,则回滚我们创建的事务
if !isTransaction {
_ = tx.Rollback()
}
return
}
db.RowsAffected, _ = result.RowsAffected()
// 处理 RETURNING 返回值
if hasDefaultValues {
updateTo := stmt.ReflectValue
switch updateTo.Kind() {
case reflect.Slice, reflect.Array:
// 对于切片或数组,只更新第一个元素
if updateTo.Len() > 0 {
updateTo = updateTo.Index(0)
}
}
// 绑定返回值到模型字段
funk.ForEach(
funk.Filter(schema.FieldsWithDefaultDBValue, func(field *gormSchema.Field) bool {
return funk.Contains(boundVars, field.Name)
}),
func(field *gormSchema.Field) {
switch updateTo.Kind() {
case reflect.Struct:
if err = field.Set(stmt.Context, updateTo, stmt.Vars[boundVars[field.Name]].(sql.Out).Dest); err != nil {
db.AddError(err)
}
case reflect.Map:
// 设置Map类型的值
mapValue := reflect.ValueOf(updateTo.Interface())
if mapValue.IsValid() && mapValue.Type().Key().Kind() == reflect.String {
keyValue := reflect.ValueOf(field.DBName)
destValue := reflect.ValueOf(stmt.Vars[boundVars[field.Name]].(sql.Out).Dest)
if destValue.Kind() == reflect.Ptr {
destValue = destValue.Elem()
}
mapValue.SetMapIndex(keyValue, destValue)
}
}
},
)
}
}
}
func performHardDelete(db *gorm.DB, boundVars map[string]int) {
stmt := db.Statement
schema := stmt.Schema
hasDefaultValues := len(schema.FieldsWithDefaultDBValue) > 0
if !stmt.Unscoped {
for _, c := range schema.DeleteClauses {
stmt.AddClause(c)
}
}
if stmt.SQL.String() == "" {
// 构建 DELETE 语句
stmt.AddClauseIfNotExists(clause.Delete{})
stmt.AddClauseIfNotExists(clause.From{Tables: []clause.Table{{Name: stmt.Schema.Table}}})
// 添加 RETURNING 子句(如果有默认值字段或需要返回值)
if hasDefaultValues {
stmt.AddClauseIfNotExists(clause.Returning{
Columns: funk.Map(schema.FieldsWithDefaultDBValue, func(field *gormSchema.Field) clause.Column {
return clause.Column{Name: field.DBName}
}).([]clause.Column),
})
}
// 构建语句
stmt.Build("DELETE", "FROM", "WHERE", "RETURNING")
// 如果有 RETURNING 子句,添加 INTO 子句
if hasDefaultValues {
stmt.WriteString(" INTO ")
for idx, field := range schema.FieldsWithDefaultDBValue {
if idx > 0 {
stmt.WriteByte(',')
}
boundVars[field.Name] = len(stmt.Vars)
stmt.AddVar(stmt, sql.Out{Dest: reflect.New(field.FieldType).Interface()})
}
}
}
if !db.DryRun {
// 执行删除操作
var tx *sql.Tx
var err error
var isTransaction bool = false
// 检查是否已经在一个事务中
if sqlTx, ok := stmt.ConnPool.(*sql.Tx); ok {
tx = sqlTx
isTransaction = true
} else if sqlDb, ok := stmt.ConnPool.(*sql.DB); ok {
tx, err = sqlDb.Begin()
if err != nil {
db.AddError(err)
return
}
defer func() {
if db.Error != nil && !isTransaction {
_ = tx.Rollback()
} else if !isTransaction {
_ = tx.Commit()
}
}()
} else {
db.AddError(fmt.Errorf("unsupported connection pool type"))
return
}
var execConn *sql.Tx
if isTransaction {
execConn = tx // 已经在事务中,直接使用原事务
} else {
execConn = tx // 使用新创建的事务
}
result, err := execConn.ExecContext(stmt.Context, stmt.SQL.String(), stmt.Vars...)
if err != nil {
db.AddError(err)
// 如果不是在已有事务中,则回滚我们创建的事务
if !isTransaction {
_ = tx.Rollback()
}
return
}
db.RowsAffected, _ = result.RowsAffected()
// 处理 RETURNING 返回值
if hasDefaultValues {
deleteTo := stmt.ReflectValue
switch deleteTo.Kind() {
case reflect.Slice, reflect.Array:
// 对于切片或数组,只处理第一个元素
if deleteTo.Len() > 0 {
deleteTo = deleteTo.Index(0)
}
}
// 绑定返回值到模型字段
funk.ForEach(
funk.Filter(schema.FieldsWithDefaultDBValue, func(field *gormSchema.Field) bool {
return funk.Contains(boundVars, field.Name)
}),
func(field *gormSchema.Field) {
switch deleteTo.Kind() {
case reflect.Struct:
if err = field.Set(stmt.Context, deleteTo, stmt.Vars[boundVars[field.Name]].(sql.Out).Dest); err != nil {
db.AddError(err)
}
case reflect.Map:
// 设置Map类型的值
mapValue := reflect.ValueOf(deleteTo.Interface())
if mapValue.IsValid() && mapValue.Type().Key().Kind() == reflect.String {
keyValue := reflect.ValueOf(field.DBName)
destValue := reflect.ValueOf(stmt.Vars[boundVars[field.Name]].(sql.Out).Dest)
if destValue.Kind() == reflect.Ptr {
destValue = destValue.Elem()
}
mapValue.SetMapIndex(keyValue, destValue)
}
}
},
)
}
}
}
+132
View File
@@ -0,0 +1,132 @@
// Package driver_adapter 提供 Oracle 驱动抽象层
// 支持 go-ora 和 godror 两种底层驱动的切换
package driver_adapter
import (
"context"
"database/sql"
)
// DriverType 驱动类型枚举
type DriverType string
const (
// DriverGoOra 使用纯 Go 实现的 go-ora 驱动
DriverGoOra DriverType = "go-ora"
// DriverGodror 使用基于 ODPI-C 的 godror 驱动
DriverGodror DriverType = "godror"
)
// OutParam 输出参数接口,用于 RETURNING INTO 子句
type OutParam interface {
// GetDest 返回目标指针
GetDest() interface{}
// SetSize 设置缓冲区大小(用于字符串类型)
SetSize(size int)
// GetSize 获取缓冲区大小
GetSize() int
}
// LobData LOB 数据接口
type LobData interface {
// IsCLOB 是否为 CLOB 类型
IsCLOB() bool
// IsBLOB 是否为 BLOB 类型
IsBLOB() bool
// GetString 获取字符串值(CLOB
GetString() string
// GetBytes 获取字节值(BLOB
GetBytes() []byte
// IsValid 是否有效
IsValid() bool
}
// BatchData 批量数据接口
type BatchData interface {
// Len 返回数据长度
Len() int
// GetValues 返回所有值
GetValues() []interface{}
}
// Adapter 驱动适配器接口
// 封装了不同 Oracle 驱动的差异,提供统一的 API
type Adapter interface {
// Name 返回驱动名称
Name() string
// Type 返回驱动类型
Type() DriverType
// Open 打开数据库连接
Open(dsn string) (*sql.DB, error)
// CreateOutParam 创建输出参数(用于 RETURNING INTO
// dest: 目标指针
// size: 缓冲区大小(字符串类型需要)
CreateOutParam(dest interface{}, size int) OutParam
// CreateClob 创建 CLOB 数据
CreateClob(value string) LobData
// CreateBlob 创建 BLOB 数据
CreateBlob(value []byte) LobData
// CreateBatch 创建批量数据
CreateBatch(values []interface{}) BatchData
// NeedsSizeForOut 返回输出参数是否需要指定 Size
// go-ora 对字符串类型的 Out 参数需要指定 Size
// godror 通常不需要
NeedsSizeForOut() bool
// SupportsReturningMultiRow 返回是否支持多行 RETURNING
// go-ora 不支持批量 INSERT + RETURNING
// godror 支持
SupportsReturningMultiRow() bool
// SupportsBulkCopy 返回是否支持 BulkCopy
SupportsBulkCopy() bool
// WrapClobForInsert 包装 CLOB 值用于插入
// 某些驱动需要特殊包装
WrapClobForInsert(value string) interface{}
// WrapBlobForInsert 包装 BLOB 值用于插入
WrapBlobForInsert(value []byte) interface{}
// UnwrapQueryResult 解包查询结果
// 将驱动特定的类型转换为标准 Go 类型
UnwrapQueryResult(value interface{}, typeName string) interface{}
// GetConnection 获取底层连接(用于高级操作)
GetConnection(db *sql.DB) (interface{}, error)
// Ping 检查连接是否可用
Ping(ctx context.Context, db *sql.DB) error
}
// Registry 驱动适配器注册表
var registry = map[DriverType]func() Adapter{}
// Register 注册驱动适配器
func Register(driverType DriverType, factory func() Adapter) {
registry[driverType] = factory
}
// Get 获取驱动适配器
func Get(driverType DriverType) Adapter {
if factory, ok := registry[driverType]; ok {
return factory()
}
return nil
}
// ListDrivers 列出所有已注册的驱动
func ListDrivers() []DriverType {
types := make([]DriverType, 0, len(registry))
for t := range registry {
types = append(types, t)
}
return types
}
+74
View File
@@ -0,0 +1,74 @@
package driver_adapter
import "testing"
// TestRegistryRegisterAndGet 验证 Register 后 Get 能返回对应适配器
func TestRegistryRegisterAndGet(t *testing.T) {
const key DriverType = "test-registry-driver"
callCount := 0
Register(key, func() Adapter {
callCount++
return &GoOraAdapter{}
})
defer delete(registry, key)
adapter := Get(key)
if adapter == nil {
t.Fatalf("Get(%q) 返回 nil,期望返回已注册的适配器", key)
}
if _, ok := adapter.(*GoOraAdapter); !ok {
t.Fatalf("Get(%q) 返回类型 = %T,期望 *GoOraAdapter", key, adapter)
}
if callCount != 1 {
t.Errorf("factory 调用次数 = %d,期望 1", callCount)
}
// 每次 Get 都应调用 factory 返回新实例
if Get(key) == nil {
t.Error("第二次 Get 返回 nil")
}
if callCount != 2 {
t.Errorf("factory 调用次数 = %d,期望 2(每次 Get 都应调用 factory", callCount)
}
}
// TestRegistryGetUnknown 验证 Get 未注册的 DriverType 返回 nil
func TestRegistryGetUnknown(t *testing.T) {
if adapter := Get(DriverType("no-such-driver")); adapter != nil {
t.Fatalf("Get(未注册类型) 返回 %v,期望 nil", adapter)
}
}
// TestRegistryListDrivers 验证 ListDrivers 返回包含 DriverGoOra
func TestRegistryListDrivers(t *testing.T) {
drivers := ListDrivers()
if len(drivers) == 0 {
t.Fatal("ListDrivers() 返回空列表")
}
for _, d := range drivers {
if d == DriverGoOra {
return
}
}
t.Errorf("ListDrivers() = %v,应包含 DriverGoOra", drivers)
}
// TestRegistryGodrorNotCompiled 验证默认构建(无 godror build tag)下不包含 DriverGodror
func TestRegistryGodrorNotCompiled(t *testing.T) {
for _, d := range ListDrivers() {
if d == DriverGodror {
t.Errorf("ListDrivers() = %v,默认构建不应包含 DriverGodrorgodror.go 带有 //go:build godror 标签)", d)
}
}
}
// TestDriverTypeConstants 验证驱动类型常量值
func TestDriverTypeConstants(t *testing.T) {
if DriverGoOra != "go-ora" {
t.Errorf("DriverGoOra = %q,期望 %q", DriverGoOra, "go-ora")
}
if DriverGodror != "godror" {
t.Errorf("DriverGodror = %q,期望 %q", DriverGodror, "godror")
}
}
+197
View File
@@ -0,0 +1,197 @@
//go:build godror
package driver_adapter
import (
"context"
"database/sql"
"fmt"
)
// GodrorAdapter godror 驱动适配器
type GodrorAdapter struct{}
// godrorOutParam godror 输出参数包装
type godrorOutParam struct {
dest interface{}
size int
}
func (p *godrorOutParam) GetDest() interface{} {
return p.dest
}
func (p *godrorOutParam) SetSize(size int) {
p.size = size
}
func (p *godrorOutParam) GetSize() int {
return p.size
}
// godrorLobData godror LOB 数据包装
type godrorLobData struct {
isClob bool
strVal string
byteVal []byte
valid bool
}
func (l *godrorLobData) IsCLOB() bool {
return l.isClob
}
func (l *godrorLobData) IsBLOB() bool {
return !l.isClob
}
func (l *godrorLobData) GetString() string {
return l.strVal
}
func (l *godrorLobData) GetBytes() []byte {
return l.byteVal
}
func (l *godrorLobData) IsValid() bool {
return l.valid
}
// godrorBatchData godror 批量数据包装
type godrorBatchData struct {
values []interface{}
}
func (b *godrorBatchData) Len() int {
return len(b.values)
}
func (b *godrorBatchData) GetValues() []interface{} {
return b.values
}
// Name 返回驱动名称
func (a *GodrorAdapter) Name() string {
return "godror"
}
// Type 返回驱动类型
func (a *GodrorAdapter) Type() DriverType {
return DriverGodror
}
// Open 打开数据库连接
func (a *GodrorAdapter) Open(dsn string) (*sql.DB, error) {
db, err := sql.Open("godror", dsn)
if err != nil {
return nil, fmt.Errorf("failed to open connection with godror: %w", err)
}
return db, nil
}
// CreateOutParam 创建输出参数(用于 RETURNING INTO
func (a *GodrorAdapter) CreateOutParam(dest interface{}, size int) OutParam {
return &godrorOutParam{
dest: dest,
size: size,
}
}
// CreateClob 创建 CLOB 数据
func (a *GodrorAdapter) CreateClob(value string) LobData {
return &godrorLobData{
isClob: true,
strVal: value,
valid: true,
byteVal: nil,
}
}
// CreateBlob 创建 BLOB 数据
func (a *GodrorAdapter) CreateBlob(value []byte) LobData {
return &godrorLobData{
isClob: false,
byteVal: value,
valid: true,
strVal: "",
}
}
// CreateBatch 创建批量数据
func (a *GodrorAdapter) CreateBatch(values []interface{}) BatchData {
return &godrorBatchData{
values: values,
}
}
// NeedsSizeForOut 返回输出参数是否需要指定 Size
func (a *GodrorAdapter) NeedsSizeForOut() bool {
return false
}
// SupportsReturningMultiRow 返回是否支持多行 RETURNING
func (a *GodrorAdapter) SupportsReturningMultiRow() bool {
return true
}
// SupportsBulkCopy 返回是否支持 BulkCopy
func (a *GodrorAdapter) SupportsBulkCopy() bool {
return false
}
// WrapClobForInsert 包装 CLOB 值用于插入
func (a *GodrorAdapter) WrapClobForInsert(value string) interface{} {
// godror 可以直接处理字符串作为 CLOB
return value
}
// WrapBlobForInsert 包装 BLOB 值用于插入
func (a *GodrorAdapter) WrapBlobForInsert(value []byte) interface{} {
// godror 可以直接处理字节数组作为 BLOB
return value
}
// UnwrapQueryResult 解包查询结果
func (a *GodrorAdapter) UnwrapQueryResult(value interface{}, typeName string) interface{} {
// godror 通常返回标准 Go 类型,无需特殊处理
// 但如果遇到特定类型,可以在这里进行转换
switch v := value.(type) {
case *string:
if v == nil {
return nil
}
return *v
case *[]byte:
if v == nil {
return nil
}
return *v
default:
return value
}
}
// GetConnection 获取底层连接(用于高级操作)
func (a *GodrorAdapter) GetConnection(db *sql.DB) (interface{}, error) {
// 从 sql.DB 获取原始连接
conn, err := db.Conn(context.Background())
if err != nil {
return nil, err
}
defer conn.Close()
// 返回原始连接
return conn, nil
}
// Ping 检查连接是否可用
func (a *GodrorAdapter) Ping(ctx context.Context, db *sql.DB) error {
return db.PingContext(ctx)
}
// init 注册 godror 驱动适配器
func init() {
Register(DriverGodror, func() Adapter {
return &GodrorAdapter{}
})
}
+190
View File
@@ -0,0 +1,190 @@
// Package driver_adapter 提供 Oracle 驱动抽象层
// 支持 go-ora 和 godror 两种底层驱动的切换
package driver_adapter
import (
"context"
"database/sql"
go_ora "github.com/sijms/go-ora/v2"
)
// GoOraAdapter go-ora 驱动适配器
type GoOraAdapter struct{}
// goOraOutParam 包装 go_ora.Out
type goOraOutParam struct {
out go_ora.Out
}
// GetDest 返回目标指针
func (o *goOraOutParam) GetDest() interface{} {
return o.out.Dest
}
// SetSize 设置缓冲区大小(用于字符串类型)
func (o *goOraOutParam) SetSize(size int) {
o.out.Size = size
}
// GetSize 获取缓冲区大小
func (o *goOraOutParam) GetSize() int {
return o.out.Size
}
// goOraLobData 包装 go_ora.Clob 和 go_ora.Blob
type goOraLobData struct {
isClob bool
strVal string
byteVal []byte
valid bool
}
// IsCLOB 是否为 CLOB 类型
func (l *goOraLobData) IsCLOB() bool {
return l.isClob
}
// IsBLOB 是否为 BLOB 类型
func (l *goOraLobData) IsBLOB() bool {
return !l.isClob
}
// GetString 获取字符串值(CLOB
func (l *goOraLobData) GetString() string {
return l.strVal
}
// GetBytes 获取字节值(BLOB
func (l *goOraLobData) GetBytes() []byte {
return l.byteVal
}
// IsValid 是否有效
func (l *goOraLobData) IsValid() bool {
return l.valid
}
// goOraBatchData 包装批量数据
type goOraBatchData struct {
values []interface{}
}
// Len 返回数据长度
func (b *goOraBatchData) Len() int {
return len(b.values)
}
// GetValues 返回所有值
func (b *goOraBatchData) GetValues() []interface{} {
return b.values
}
// Name 返回驱动名称
func (a *GoOraAdapter) Name() string {
return "go-ora"
}
// Type 返回驱动类型
func (a *GoOraAdapter) Type() DriverType {
return DriverGoOra
}
// Open 打开数据库连接
func (a *GoOraAdapter) Open(dsn string) (*sql.DB, error) {
return sql.Open("oracle", dsn)
}
// CreateOutParam 创建输出参数(用于 RETURNING INTO
func (a *GoOraAdapter) CreateOutParam(dest interface{}, size int) OutParam {
return &goOraOutParam{
out: go_ora.Out{Dest: dest, Size: size},
}
}
// CreateClob 创建 CLOB 数据
func (a *GoOraAdapter) CreateClob(value string) LobData {
return &goOraLobData{
isClob: true,
strVal: value,
byteVal: nil,
valid: true,
}
}
// CreateBlob 创建 BLOB 数据
func (a *GoOraAdapter) CreateBlob(value []byte) LobData {
return &goOraLobData{
isClob: false,
strVal: "",
byteVal: value,
valid: true,
}
}
// CreateBatch 创建批量数据
func (a *GoOraAdapter) CreateBatch(values []interface{}) BatchData {
return &goOraBatchData{
values: values,
}
}
// NeedsSizeForOut 返回输出参数是否需要指定 Size
func (a *GoOraAdapter) NeedsSizeForOut() bool {
return true
}
// SupportsReturningMultiRow 返回是否支持多行 RETURNING
func (a *GoOraAdapter) SupportsReturningMultiRow() bool {
return false
}
// SupportsBulkCopy 返回是否支持 BulkCopy
func (a *GoOraAdapter) SupportsBulkCopy() bool {
return true
}
// WrapClobForInsert 包装 CLOB 值用于插入
func (a *GoOraAdapter) WrapClobForInsert(value string) interface{} {
return go_ora.Clob{String: value, Valid: true}
}
// WrapBlobForInsert 包装 BLOB 值用于插入
func (a *GoOraAdapter) WrapBlobForInsert(value []byte) interface{} {
return go_ora.Blob{Data: value, Valid: true}
}
// UnwrapQueryResult 解包查询结果
func (a *GoOraAdapter) UnwrapQueryResult(value interface{}, typeName string) interface{} {
// 根据需要处理 go-ora 特有的返回类型转换
// 这里简单返回原始值,可根据实际需求扩展
return value
}
// GetConnection 获取底层连接(用于高级操作)
func (a *GoOraAdapter) GetConnection(db *sql.DB) (interface{}, error) {
conn, err := db.Conn(context.Background())
if err != nil {
return nil, err
}
defer conn.Close()
var rawConn interface{}
err = conn.Raw(func(driverConn interface{}) error {
rawConn = driverConn
return nil
})
return rawConn, err
}
// Ping 检查连接是否可用
func (a *GoOraAdapter) Ping(ctx context.Context, db *sql.DB) error {
return db.PingContext(ctx)
}
// init 函数中注册驱动
func init() {
Register(DriverGoOra, func() Adapter {
return &GoOraAdapter{}
})
}
+243
View File
@@ -0,0 +1,243 @@
package driver_adapter
import (
"context"
"reflect"
"testing"
"time"
go_ora "github.com/sijms/go-ora/v2"
)
// invalidDSN 用于无需真实连接的场景:
// 空字符串在 go-ora 的 dsn 解析阶段即返回错误,避免触发真实 TCP 连接导致测试挂起
const invalidDSN = ""
// newTestAdapter 创建 GoOraAdapter 测试实例
func newTestAdapter() *GoOraAdapter {
return &GoOraAdapter{}
}
// TestGoOraAdapterBasics 验证适配器基本属性
func TestGoOraAdapterBasics(t *testing.T) {
a := newTestAdapter()
if got := a.Name(); got != "go-ora" {
t.Errorf("Name() = %q,期望 %q", got, "go-ora")
}
if got := a.Type(); got != DriverGoOra {
t.Errorf("Type() = %q,期望 %q", got, DriverGoOra)
}
if !a.NeedsSizeForOut() {
t.Error("NeedsSizeForOut() = false,期望 true")
}
if a.SupportsReturningMultiRow() {
t.Error("SupportsReturningMultiRow() = true,期望 false")
}
if !a.SupportsBulkCopy() {
t.Error("SupportsBulkCopy() = false,期望 true")
}
}
// TestGoOraCreateOutParam 验证输出参数创建
func TestGoOraCreateOutParam(t *testing.T) {
a := newTestAdapter()
var id int
out := a.CreateOutParam(&id, 100)
if out == nil {
t.Fatal("CreateOutParam 返回 nil")
}
if got := out.GetDest(); got != &id {
t.Errorf("GetDest() = %v,期望 %v", got, &id)
}
if got := out.GetSize(); got != 100 {
t.Errorf("GetSize() = %d,期望 100", got)
}
out.SetSize(200)
if got := out.GetSize(); got != 200 {
t.Errorf("SetSize 后 GetSize() = %d,期望 200", got)
}
if _, ok := out.(*goOraOutParam); !ok {
t.Errorf("返回类型 = %T,期望 *goOraOutParam", out)
}
}
// TestGoOraCreateClob 验证 CLOB 创建
func TestGoOraCreateClob(t *testing.T) {
a := newTestAdapter()
lob := a.CreateClob("text")
if lob == nil {
t.Fatal("CreateClob 返回 nil")
}
if !lob.IsCLOB() {
t.Error("IsCLOB() = false,期望 true")
}
if lob.IsBLOB() {
t.Error("IsBLOB() = true,期望 false")
}
if got := lob.GetString(); got != "text" {
t.Errorf("GetString() = %q,期望 %q", got, "text")
}
if !lob.IsValid() {
t.Error("IsValid() = false,期望 true")
}
if _, ok := lob.(*goOraLobData); !ok {
t.Errorf("返回类型 = %T,期望 *goOraLobData", lob)
}
}
// TestGoOraCreateBlob 验证 BLOB 创建
func TestGoOraCreateBlob(t *testing.T) {
a := newTestAdapter()
want := []byte{1, 2, 3}
lob := a.CreateBlob(want)
if lob == nil {
t.Fatal("CreateBlob 返回 nil")
}
if !lob.IsBLOB() {
t.Error("IsBLOB() = false,期望 true")
}
if lob.IsCLOB() {
t.Error("IsCLOB() = true,期望 false")
}
if got := lob.GetBytes(); !reflect.DeepEqual(got, want) {
t.Errorf("GetBytes() = %v,期望 %v", got, want)
}
if !lob.IsValid() {
t.Error("IsValid() = false,期望 true")
}
if _, ok := lob.(*goOraLobData); !ok {
t.Errorf("返回类型 = %T,期望 *goOraLobData", lob)
}
}
// TestGoOraCreateBatch 验证批量数据创建
func TestGoOraCreateBatch(t *testing.T) {
a := newTestAdapter()
want := []interface{}{1, "a"}
batch := a.CreateBatch(want)
if batch == nil {
t.Fatal("CreateBatch 返回 nil")
}
if got := batch.Len(); got != 2 {
t.Errorf("Len() = %d,期望 2", got)
}
got := batch.GetValues()
if len(got) != len(want) {
t.Fatalf("GetValues() 长度 = %d,期望 %d", len(got), len(want))
}
for i := range want {
if !reflect.DeepEqual(got[i], want[i]) {
t.Errorf("GetValues()[%d] = %v,期望 %v", i, got[i], want[i])
}
}
if _, ok := batch.(*goOraBatchData); !ok {
t.Errorf("返回类型 = %T,期望 *goOraBatchData", batch)
}
}
// TestGoOraWrapForInsert 验证插入包装
func TestGoOraWrapForInsert(t *testing.T) {
a := newTestAdapter()
// CLOB 包装
clob, ok := a.WrapClobForInsert("text").(go_ora.Clob)
if !ok {
t.Fatalf("WrapClobForInsert 返回类型 = %T,期望 go_ora.Clob", a.WrapClobForInsert("text"))
}
if clob.String != "text" {
t.Errorf("Clob.String = %q,期望 %q", clob.String, "text")
}
if !clob.Valid {
t.Error("Clob.Valid = false,期望 true")
}
// BLOB 包装
blob, ok := a.WrapBlobForInsert([]byte{1}).(go_ora.Blob)
if !ok {
t.Fatalf("WrapBlobForInsert 返回类型 = %T,期望 go_ora.Blob", a.WrapBlobForInsert([]byte{1}))
}
if !reflect.DeepEqual(blob.Data, []byte{1}) {
t.Errorf("Blob.Data = %v,期望 [1]", blob.Data)
}
if !blob.Valid {
t.Error("Blob.Valid = false,期望 true")
}
}
// TestGoOraOpen 验证 Open 只检查驱动名注册,不会真正建立连接
func TestGoOraOpen(t *testing.T) {
a := newTestAdapter()
db, err := a.Open("oracle://user:pass@localhost:1521/service")
if err != nil {
// sql.Open 只在驱动未注册时报错;go-ora 包的 init 已注册 "oracle" 驱动名
t.Fatalf("Open() 返回错误: %v(需要 go-ora 包 init 注册 \"oracle\" 驱动名)", err)
}
if db == nil {
t.Fatal("Open() 返回 db == nil")
}
defer db.Close()
}
// TestGoOraPing 验证对未连接数据库的 Ping 返回错误而非 panic
func TestGoOraPing(t *testing.T) {
a := newTestAdapter()
db, err := a.Open(invalidDSN)
if err != nil {
t.Fatalf("Open() 返回错误: %v", err)
}
defer db.Close()
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
if err := a.Ping(ctx, db); err == nil {
t.Error("Ping(未连接数据库) 返回 nil,期望返回错误")
}
}
// TestGoOraUnwrapQueryResult 验证查询结果原样返回
func TestGoOraUnwrapQueryResult(t *testing.T) {
a := newTestAdapter()
cases := []interface{}{
42,
"hello",
[]byte{1, 2, 3},
3.14,
nil,
}
for _, in := range cases {
if got := a.UnwrapQueryResult(in, ""); !reflect.DeepEqual(got, in) {
t.Errorf("UnwrapQueryResult(%v, \"\") = %v,期望原样返回", in, got)
}
}
}
// TestGoOraGetConnection 验证获取底层连接不 panic(未连接数据库时返回错误即可)
func TestGoOraGetConnection(t *testing.T) {
a := newTestAdapter()
db, err := a.Open(invalidDSN)
if err != nil {
t.Fatalf("Open() 返回错误: %v", err)
}
defer db.Close()
raw, err := a.GetConnection(db)
if err != nil {
t.Logf("GetConnection 返回预期错误(未连接数据库): %v", err)
return
}
if raw == nil {
t.Error("GetConnection 返回 nil raw 且无错误")
}
}
+3 -3
View File
@@ -4,13 +4,13 @@ go 1.22
require (
github.com/emirpasic/gods v1.18.1
github.com/sijms/go-ora/v2 v2.8.19
github.com/sijms/go-ora/v2 v2.9.0
github.com/thoas/go-funk v0.9.3
gorm.io/gorm v1.25.10
gorm.io/gorm v1.31.2
)
require (
github.com/jinzhu/inflection v1.0.0 // indirect
github.com/jinzhu/now v1.1.5 // indirect
golang.org/x/text v0.20.0 // indirect
)
+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
})
}
+109
View File
@@ -0,0 +1,109 @@
package oracle
import (
"reflect"
"testing"
"gorm.io/gorm"
"gorm.io/gorm/migrator"
"gorm.io/gorm/schema"
)
// noopDialector 内嵌真实 Dialector 但跳过 Initialize 的数据库连接,
// 用于在无需真实连接的情况下构造合法的 *gorm.DBcacheStore 会被 gorm.Open 初始化)。
type noopDialector struct {
Dialector
}
func (noopDialector) Initialize(db *gorm.DB) error {
return nil
}
// relParent 含关联关系的模型,用于测试 TryRemoveOnUpdate
type relParent struct {
ID uint `gorm:"primaryKey"`
Kids []relKid `gorm:"foreignKey:ParentID;constraint:OnUpdate:CASCADE"`
}
type relKid struct {
ID uint `gorm:"primaryKey"`
ParentID uint
}
// newTestMigrator 构造一个不依赖真实连接的 Migrator
func newTestMigrator() Migrator {
d := &Dialector{Config: &Config{DBVer: "12.1.0.2.0", DefaultStringSize: 1024}}
db, err := gorm.Open(noopDialector{Dialector: *d}, &gorm.Config{})
if err != nil {
panic(err)
}
return Migrator{Migrator: migrator.Migrator{Config: migrator.Config{
DB: db,
Dialector: d,
CreateIndexAfterCreateTable: true,
}}}
}
// TestTryRemoveOnUpdate 验证从关系约束标签中移除 ON UPDATE 片段
func TestTryRemoveOnUpdate(t *testing.T) {
m := newTestMigrator()
if err := m.TryRemoveOnUpdate(&relParent{}); err != nil {
t.Fatalf("TryRemoveOnUpdate returned error: %v", err)
}
}
// TestTryRemoveOnUpdateWithoutRelations 验证无关联关系的模型不报错
func TestTryRemoveOnUpdateWithoutRelations(t *testing.T) {
m := newTestMigrator()
if err := m.TryRemoveOnUpdate(&limitModel{}); err != nil {
t.Fatalf("TryRemoveOnUpdate returned error: %v", err)
}
}
// TestTryQuotifyReservedWords 验证处理保留字列名不报错
func TestTryQuotifyReservedWords(t *testing.T) {
m := newTestMigrator()
type reservedModel struct {
ID uint `gorm:"primaryKey"`
Select string `gorm:"column:select"`
}
if err := m.TryQuotifyReservedWords(&reservedModel{}); err != nil {
t.Fatalf("TryQuotifyReservedWords returned error: %v", err)
}
}
// TestSequenceAndTriggerNaming 验证序列/触发器名称规则
func TestSequenceAndTriggerNaming(t *testing.T) {
m := newTestMigrator()
seq := m.sequenceName("TEST_USERS")
if seq != "SEQ_TEST_USERS" {
t.Errorf("sequenceName() = %q, want %q", seq, "SEQ_TEST_USERS")
}
trg := m.triggerName("TEST_USERS")
if trg != "TRG_TEST_USERS" {
t.Errorf("triggerName() = %q, want %q", trg, "TRG_TEST_USERS")
}
}
// TestMigratorDelegateMethods 验证 Migrator 的 DataTypeOf 委托
func TestMigratorDelegateMethods(t *testing.T) {
m := newTestMigrator()
f := testField(schema.Int)
f.Size = 64
f.FieldType = reflect.TypeOf(int(0))
f.IndirectFieldType = f.FieldType
if got := m.DataTypeOf(f); got != "INTEGER" {
t.Errorf("Migrator.DataTypeOf() = %q, want %q", got, "INTEGER")
}
}
// TestNoopDialectorImplementsGormDialector 编译期验证 noopDialector 实现了 gorm.Dialector
var _ gorm.Dialector = (*noopDialector)(nil)
+111
View File
@@ -0,0 +1,111 @@
package oracle
import (
"testing"
"gorm.io/gorm/schema"
)
func TestConvertNameToFormat(t *testing.T) {
tests := []struct {
in string
want string
}{
{"hello", "HELLO"},
{"Hello", "HELLO"},
{"TEST_USERS", "TEST_USERS"},
{"test_users", "TEST_USERS"},
{"created_at", "CREATED_AT"},
{"", ""},
}
for _, tt := range tests {
t.Run(tt.in, func(t *testing.T) {
if got := ConvertNameToFormat(tt.in); got != tt.want {
t.Errorf("ConvertNameToFormat(%q) = %q, want %q", tt.in, got, tt.want)
}
})
}
}
func newTestNamer() Namer {
return Namer{NamingStrategy: schema.NamingStrategy{}}
}
func TestNamerTableName(t *testing.T) {
n := newTestNamer()
if got := n.TableName("test_users"); got != "TEST_USERS" {
t.Errorf("TableName() = %q, want %q", got, "TEST_USERS")
}
}
func TestNamerTableNameWithDBName(t *testing.T) {
n := newTestNamer()
n.DBName = "MERCHANT"
if got := n.TableName("test_users"); got != "MERCHANT.TEST_USERS" {
t.Errorf("TableName() with DBName = %q, want %q", got, "MERCHANT.TEST_USERS")
}
}
func TestNamerColumnName(t *testing.T) {
n := newTestNamer()
if got := n.ColumnName("test_users", "created_at"); got != "CREATED_AT" {
t.Errorf("ColumnName() = %q, want %q", got, "CREATED_AT")
}
}
func TestNamerJoinTableName(t *testing.T) {
n := newTestNamer()
if got := n.JoinTableName("order_items"); got != "ORDER_ITEMS" {
t.Errorf("JoinTableName() = %q, want %q", got, "ORDER_ITEMS")
}
}
func TestNamerRelationshipFKName(t *testing.T) {
n := newTestNamer()
rel := schema.Relationship{
Name: "User",
Schema: &schema.Schema{Table: "users"},
}
if got := n.RelationshipFKName(rel); got != "FK_USERS_USER" {
t.Errorf("RelationshipFKName() = %q, want %q", got, "FK_USERS_USER")
}
}
func TestNamerCheckerName(t *testing.T) {
n := newTestNamer()
if got := n.CheckerName("test_users", "phone"); got != "CHK_TEST_USERS_PHONE" {
t.Errorf("CheckerName() = %q, want %q", got, "CHK_TEST_USERS_PHONE")
}
}
func TestNamerIndexName(t *testing.T) {
n := newTestNamer()
if got := n.IndexName("test_users", "name"); got != "IDX_TEST_USERS_NAME" {
t.Errorf("IndexName() = %q, want %q", got, "IDX_TEST_USERS_NAME")
}
}
func TestNamerSchemaName(t *testing.T) {
n := newTestNamer()
if got := n.SchemaName("public"); got != "PUBLIC" {
t.Errorf("SchemaName() = %q, want %q", got, "PUBLIC")
}
}
func TestNamerUniqueName(t *testing.T) {
n := newTestNamer()
if got := n.UniqueName("test_users", "email"); got != "UNI_TEST_USERS_EMAIL" {
t.Errorf("UniqueName() = %q, want %q", got, "UNI_TEST_USERS_EMAIL")
}
}
+107 -8
View File
@@ -19,10 +19,56 @@ import (
"gorm.io/gorm/logger"
"gorm.io/gorm/migrator"
"gorm.io/gorm/schema"
"git.charlienet.top/go/oracle/driver_adapter"
)
const RowNumberAliasForOracle11 = "ROW_NUM"
// Oracle 版本主版本号常量(对应各版本引入的数据库特性)
const (
OracleVersion10 = 10 // Oracle 10g
OracleVersion11 = 11 // Oracle 11g(不含 IDENTITY 列、OFFSET/FETCH 分页)
OracleVersion12 = 12 // Oracle 12c(引入 IDENTITY 列、OFFSET/FETCH 分页;12.1 起支持 Extended 32k VARCHAR2
OracleVersion18 = 18 // Oracle 18c12.2 的再版)
OracleVersion19 = 19 // Oracle 19c
OracleVersion21 = 21 // Oracle 21c(引入原生 BOOLEAN 列类型)
OracleVersion23 = 23 // Oracle 23ai(引入 VECTOR 类型)
)
// oracleMajor 返回数据库版本的主版本号;解析失败返回 0。
// 支持格式如 "11.2.0.4.0"、"19.0.0.0.0"、"23.0.0.0.0"。
func oracleMajor(dbVer string) int {
major, _ := strconv.Atoi(strings.Split(dbVer, ".")[0])
return major
}
// supportsIdentity 是否支持 IDENTITY 列(12c+ 支持 GENERATED ... AS IDENTITY
// 11g 及以下需用序列 + BEFORE INSERT 触发器模拟自增)
func supportsIdentity(dbVer string) bool { return oracleMajor(dbVer) >= OracleVersion12 }
// supportsFetchOffset 是否支持 OFFSET/FETCH 分页语法(12c+ 引入;
// 11g 需改写为 ROWNUM 分页)
func supportsFetchOffset(dbVer string) bool { return oracleMajor(dbVer) >= OracleVersion12 }
// supportsNativeBoolean 是否支持原生 BOOLEAN 列类型(21c+ 引入;
// 更早版本需用 NUMBER(1) 模拟)
func supportsNativeBoolean(dbVer string) bool { return oracleMajor(dbVer) >= OracleVersion21 }
// supportsExtendedString 是否支持 Extended 32k VARCHAR212.2+ 默认开启;
// 12.1 需 MAX_STRING_SIZE=EXTENDED)。保守判定:主版本 >= 12 视为可能支持,
// 具体是否生效依赖数据库参数。
func supportsExtendedString(dbVer string) bool { return oracleMajor(dbVer) >= OracleVersion12 }
// supportsVector 是否支持 VECTOR 类型(23ai 引入,用于 AI Vector Search
func supportsVector(dbVer string) bool { return oracleMajor(dbVer) >= OracleVersion23 }
// isOracle11g 判断当前数据库是否低于 12c(11g 及以下不支持 IDENTITY 列)。
// 保留该函数以兼容既有调用,内部委托 supportsIdentity 取反。
func isOracle11g(dbVer string) bool {
return !supportsIdentity(dbVer)
}
type Config struct {
DriverName string
DSN string
@@ -30,6 +76,8 @@ type Config struct {
DefaultStringSize uint
DBName string
DBVer string
DriverType driver_adapter.DriverType // 新增:驱动类型(go-ora 或 godror
SkipQuoteIdentifiers bool // 新增:是否跳过标识符引用
}
type Dialector struct {
@@ -90,6 +138,21 @@ func (d Dialector) Initialize(db *gorm.DB) (err error) {
return
}
// 注册 Update 回调
if err = db.Callback().Update().Replace("gorm:update", Update); err != nil {
return
}
// 注册 Delete 回调
if err = db.Callback().Delete().Replace("gorm:delete", Delete); err != nil {
return
}
// 注册 Query 回调
if err = db.Callback().Query().Replace("gorm:query", Query); err != nil {
return
}
for k, v := range d.ClauseBuilders() {
db.ClauseBuilders[k] = v
}
@@ -256,12 +319,16 @@ func (d Dialector) BindVarTo(writer clause.Writer, stmt *gorm.Statement, v inter
}
func (d Dialector) QuoteTo(writer clause.Writer, str string) {
if d.SkipQuoteIdentifiers {
writer.WriteString(str)
return
}
if str != "" && IsReservedWord(str) {
writer.WriteByte('"')
writer.WriteString(str)
writer.WriteByte('"')
} else {
writer.WriteString(str)
}
}
@@ -288,15 +355,24 @@ func (d Dialector) DataTypeOf(field *schema.Field) string {
var sqlType string
switch field.DataType {
case schema.Bool, schema.Int, schema.Uint, schema.Float:
case schema.Bool:
// Oracle 21c+ 支持原生 BOOLEAN 列;更早版本用 NUMBER(1) 模拟
if supportsNativeBoolean(d.DBVer) {
sqlType = "BOOLEAN"
} else {
sqlType = "NUMBER(1)"
}
case schema.Int, schema.Uint:
sqlType = "INTEGER"
switch {
case field.DataType == schema.Float:
sqlType = "FLOAT"
case field.Size <= 8:
if field.Size <= 8 {
sqlType = "SMALLINT"
}
// Oracle 12c+ 支持 IDENTITY 列;Oracle 11g 需要在迁移时创建序列 + 触发器
if field.AutoIncrement && supportsIdentity(d.DBVer) {
sqlType += " GENERATED BY DEFAULT AS IDENTITY"
}
case schema.Float:
sqlType = "FLOAT"
if val, ok := field.TagSettings["AUTOINCREMENT"]; ok && utils.CheckTruth(val) {
sqlType += " GENERATED BY DEFAULT AS IDENTITY"
@@ -317,14 +393,30 @@ func (d Dialector) DataTypeOf(field *schema.Field) string {
}
}
if size >= 2000 {
// Oracle 12c+Extended)支持最长 32767 字节的 VARCHAR232k 特性);
// 11g 及未开启 Extended 的库超过 4000 必须用 CLOB。
// 保守策略:
// - size 在 2000~4000 之间维持 CLOB 不变(保持历史行为)
// - size > 4000 且版本 >= 12 → VARCHAR2(size)(利用 32k 特性)
// - size > 4000 且 11g → CLOB(保持现状)
if size > 4000 {
if supportsExtendedString(d.DBVer) {
sqlType = fmt.Sprintf("VARCHAR2(%d)", size)
} else {
sqlType = "CLOB"
}
} else if size >= 2000 {
sqlType = "CLOB"
} else {
sqlType = fmt.Sprintf("VARCHAR2(%d)", size)
}
case schema.Time:
if field.Precision > 0 {
sqlType = fmt.Sprintf("TIMESTAMP(%d) WITH TIME ZONE", field.Precision)
} else {
sqlType = "TIMESTAMP WITH TIME ZONE"
}
case schema.Bytes:
sqlType = "BLOB"
@@ -353,3 +445,10 @@ func (d Dialector) RollbackTo(tx *gorm.DB, name string) error {
tx.Exec("ROLLBACK TO SAVEPOINT " + name)
return tx.Error
}
func (d Dialector) GetAdapter() driver_adapter.Adapter {
if d.DriverType == "" {
d.DriverType = driver_adapter.DriverGoOra
}
return driver_adapter.Get(d.DriverType)
}
+749
View File
@@ -0,0 +1,749 @@
package oracle
import (
"reflect"
"strings"
"testing"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"gorm.io/gorm/schema"
)
// ---- 测试辅助 ----
// limitModel 带主键的模型,用于 RewriteLimit/RewriteLimit11 测试
type limitModel struct {
ID uint `gorm:"primaryKey"`
Name string
}
func newTestDialector(dbVer string, defaultStringSize uint) *Dialector {
return &Dialector{Config: &Config{DBVer: dbVer, DefaultStringSize: defaultStringSize}}
}
func newTestStatement(d *Dialector) *gorm.Statement {
db := &gorm.DB{Config: &gorm.Config{Dialector: d}}
return &gorm.Statement{DB: db}
}
// testField 构造一个最小可用的 schema.Field
func testField(dataType schema.DataType) *schema.Field {
return &schema.Field{
DataType: dataType,
FieldType: reflect.TypeOf(""),
TagSettings: map[string]string{},
}
}
// limitClause 构造 LIMIT 子句
func limitClause(offset, limit int) clause.Clause {
expr := clause.Limit{Offset: offset}
if limit > 0 {
v := limit
expr.Limit = &v
}
return clause.Clause{Name: "LIMIT", Expression: expr}
}
// ---- TestDataTypeOf ----
func TestDataTypeOf(t *testing.T) {
d12 := newTestDialector("12.1.0.2.0", 1024)
tests := []struct {
name string
dialector *Dialector
mutate func(*schema.Field)
want string
}{
{
name: "bool maps to NUMBER(1)",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.Bool },
want: "NUMBER(1)",
},
{
name: "int size 0 maps to SMALLINT (size <= 8 rule)",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.Int },
want: "SMALLINT",
},
{
name: "int size 8 maps to SMALLINT",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.Int; f.Size = 8 },
want: "SMALLINT",
},
{
name: "int size 64 maps to INTEGER",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.Int; f.Size = 64 },
want: "INTEGER",
},
{
name: "uint size 0 maps to SMALLINT (size <= 8 rule)",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.Uint },
want: "SMALLINT",
},
{
name: "uint size 8 maps to SMALLINT",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.Uint; f.Size = 8 },
want: "SMALLINT",
},
{
name: "uint size 64 maps to INTEGER",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.Uint; f.Size = 64 },
want: "INTEGER",
},
{
name: "int autoincrement on 12c maps to identity",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.Int; f.AutoIncrement = true; f.Size = 64 },
want: "INTEGER GENERATED BY DEFAULT AS IDENTITY",
},
{
name: "int autoincrement on 11g stays INTEGER",
dialector: newTestDialector("11.2.0.4.0", 1024),
mutate: func(f *schema.Field) { f.DataType = schema.Int; f.AutoIncrement = true; f.Size = 64 },
want: "INTEGER",
},
{
name: "float maps to FLOAT",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.Float },
want: "FLOAT",
},
{
name: "float with AUTOINCREMENT tag maps to identity",
dialector: d12,
mutate: func(f *schema.Field) {
f.DataType = schema.Float
f.TagSettings["AUTOINCREMENT"] = "true"
},
want: "FLOAT GENERATED BY DEFAULT AS IDENTITY",
},
{
name: "string size 100 maps to VARCHAR2(100)",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.String; f.Size = 100 },
want: "VARCHAR2(100)",
},
{
name: "string without size uses default string size",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.String },
want: "VARCHAR2(1024)",
},
{
name: "string size 2000 maps to CLOB",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.String; f.Size = 2000 },
want: "CLOB",
},
{
name: "string size 4096 on 12c maps to VARCHAR2(4096) via 32k support",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.String; f.Size = 4096 },
want: "VARCHAR2(4096)",
},
{
name: "string size 4096 on 11g maps to CLOB",
dialector: newTestDialector("11.2.0.4.0", 1024),
mutate: func(f *schema.Field) { f.DataType = schema.String; f.Size = 4096 },
want: "CLOB",
},
{
name: "string size 5000 on 12c maps to VARCHAR2(5000) via 32k support",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.String; f.Size = 5000 },
want: "VARCHAR2(5000)",
},
{
name: "string size 5000 on 11g maps to CLOB",
dialector: newTestDialector("11.2.0.4.0", 1024),
mutate: func(f *schema.Field) { f.DataType = schema.String; f.Size = 5000 },
want: "CLOB",
},
{
name: "bool on 21c maps to native BOOLEAN",
dialector: newTestDialector("21.0.0.0.0", 1024),
mutate: func(f *schema.Field) { f.DataType = schema.Bool },
want: "BOOLEAN",
},
{
name: "bool on 23ai maps to native BOOLEAN",
dialector: newTestDialector("23.0.0.0.0", 1024),
mutate: func(f *schema.Field) { f.DataType = schema.Bool },
want: "BOOLEAN",
},
{
name: "bool on 11g maps to NUMBER(1)",
dialector: newTestDialector("11.2.0.4.0", 1024),
mutate: func(f *schema.Field) { f.DataType = schema.Bool },
want: "NUMBER(1)",
},
{
name: "primary key without size and without default size maps to VARCHAR2(191)",
dialector: newTestDialector("12.1.0.2.0", 0),
mutate: func(f *schema.Field) { f.DataType = schema.String; f.PrimaryKey = true },
want: "VARCHAR2(191)",
},
{
name: "unique field without size and without default size maps to VARCHAR2(191)",
dialector: newTestDialector("12.1.0.2.0", 0),
mutate: func(f *schema.Field) {
f.DataType = schema.String
f.TagSettings["UNIQUE"] = "unique"
},
want: "VARCHAR2(191)",
},
{
name: "time maps to TIMESTAMP WITH TIME ZONE",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.Time },
want: "TIMESTAMP WITH TIME ZONE",
},
{
name: "time with precision 6",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.Time; f.Precision = 6 },
want: "TIMESTAMP(6) WITH TIME ZONE",
},
{
name: "bytes maps to BLOB",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.Bytes },
want: "BLOB",
},
{
name: "text data type maps to CLOB",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.DataType("text") },
want: "CLOB",
},
{
name: "VARCHAR2 data type with size",
dialector: d12,
mutate: func(f *schema.Field) { f.DataType = schema.DataType("VARCHAR2"); f.Size = 50 },
want: "VARCHAR2(50)",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
f := testField("")
tt.mutate(f)
got := tt.dialector.DataTypeOf(f)
if got != tt.want {
t.Errorf("DataTypeOf() = %q, want %q", got, tt.want)
}
})
}
}
func TestDataTypeOfRemovesRestrictTag(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
f := testField(schema.Int)
f.TagSettings["RESTRICT"] = "true"
d.DataTypeOf(f)
if _, ok := f.TagSettings["RESTRICT"]; ok {
t.Error("expected RESTRICT to be removed from TagSettings")
}
}
func TestDataTypeOfPanicsOnEmptyType(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
f := testField("")
defer func() {
if r := recover(); r == nil {
t.Error("expected panic for empty DataType")
}
}()
d.DataTypeOf(f)
}
// ---- TestBindVarTo ----
func TestBindVarTo(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
stmt := newTestStatement(d)
var buf strings.Builder
// GORM 的 AddVar 会先 append 再调用 BindVarTo,因此绑定位置从 1 开始
stmt.Vars = append(stmt.Vars, "a")
d.BindVarTo(&buf, stmt, "a")
stmt.Vars = append(stmt.Vars, "b")
d.BindVarTo(&buf, stmt, "b")
stmt.Vars = append(stmt.Vars, "c")
d.BindVarTo(&buf, stmt, "c")
if want := ":1:2:3"; buf.String() != want {
t.Errorf("BindVarTo output = %q, want %q", buf.String(), want)
}
}
func TestBindVarToEmptyVars(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
stmt := newTestStatement(d)
var buf strings.Builder
d.BindVarTo(&buf, stmt, "a")
if want := ":0"; buf.String() != want {
t.Errorf("BindVarTo output = %q, want %q", buf.String(), want)
}
}
// ---- TestQuoteTo ----
func TestQuoteTo(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
tests := []struct {
name string
value string
want string
}{
{"plain identifier", "USER_NAME", "USER_NAME"},
{"lowercase identifier", "user_name", "user_name"},
{"empty string", "", ""},
// 注意:SELECT 不在 reserved.go 的 ReservedWordsList 中,不会被加引号
{"SELECT is not in reserved list", "SELECT", "SELECT"},
{"reserved word FROM", "FROM", `"FROM"`},
{"reserved word WHERE", "WHERE", `"WHERE"`},
{"reserved word ORDER", "ORDER", `"ORDER"`},
{"reserved word SET", "SET", `"SET"`},
{"reserved word VALUES", "VALUES", `"VALUES"`},
{"reserved word UPDATE", "UPDATE", `"UPDATE"`},
{"reserved word CLOB", "CLOB", `"CLOB"`},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var buf strings.Builder
d.QuoteTo(&buf, tt.value)
if got := buf.String(); got != tt.want {
t.Errorf("QuoteTo(%q) = %q, want %q", tt.value, got, tt.want)
}
})
}
}
func TestQuoteToSkipQuoteIdentifiers(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
d.SkipQuoteIdentifiers = true
var buf strings.Builder
d.QuoteTo(&buf, "SELECT")
d.QuoteTo(&buf, "FROM")
if want := "SELECTFROM"; buf.String() != want {
t.Errorf("QuoteTo with SkipQuoteIdentifiers = %q, want %q", buf.String(), want)
}
}
// ---- TestRewriteLimit ----
func TestRewriteLimit(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
t.Run("adds order by primary key when no order by", func(t *testing.T) {
stmt := newTestStatement(d)
stmt.Clauses = map[string]clause.Clause{}
stmt.Schema = parseTestSchema(t, &limitModel{})
d.RewriteLimit(limitClause(10, 5), stmt)
got := stmt.SQL.String()
for _, want := range []string{"ORDER BY id", "OFFSET 10 ROWS", "FETCH NEXT 5 ROWS ONLY"} {
if !strings.Contains(got, want) {
t.Errorf("RewriteLimit output %q missing %q", got, want)
}
}
})
t.Run("keeps existing order by", func(t *testing.T) {
stmt := newTestStatement(d)
stmt.Clauses = map[string]clause.Clause{
"ORDER BY": {Name: "ORDER BY", Expression: clause.OrderBy{
Columns: []clause.OrderByColumn{{Column: clause.Column{Name: "NAME"}, Desc: true}},
}},
}
d.RewriteLimit(limitClause(10, 5), stmt)
got := stmt.SQL.String()
if strings.HasPrefix(got, "ORDER BY") {
t.Errorf("RewriteLimit should not add ORDER BY when already present, got %q", got)
}
if !strings.Contains(got, "OFFSET 10 ROWS") || !strings.Contains(got, "FETCH NEXT 5 ROWS ONLY") {
t.Errorf("RewriteLimit output %q missing offset/fetch", got)
}
})
t.Run("uses DUAL subquery when schema is nil", func(t *testing.T) {
stmt := newTestStatement(d)
stmt.Clauses = map[string]clause.Clause{}
d.RewriteLimit(limitClause(0, 3), stmt)
got := stmt.SQL.String()
if !strings.Contains(got, "ORDER BY (SELECT NULL FROM DUAL)") {
t.Errorf("RewriteLimit output %q missing DUAL subquery", got)
}
})
t.Run("offset only", func(t *testing.T) {
stmt := newTestStatement(d)
stmt.Clauses = map[string]clause.Clause{}
d.RewriteLimit(limitClause(5, 0), stmt)
got := stmt.SQL.String()
if !strings.Contains(got, "OFFSET 5 ROWS") {
t.Errorf("RewriteLimit output %q missing offset", got)
}
if strings.Contains(got, "FETCH NEXT") {
t.Errorf("RewriteLimit output %q should not contain FETCH NEXT", got)
}
})
}
func TestRewriteLimitIgnoresNonLimitExpression(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
stmt := newTestStatement(d)
stmt.Clauses = map[string]clause.Clause{}
d.RewriteLimit(clause.Clause{Name: "LIMIT", Expression: clause.Expr{SQL: "1"}}, stmt)
if got := stmt.SQL.String(); got != "" {
t.Errorf("RewriteLimit should write nothing for non-Limit expression, got %q", got)
}
}
// ---- TestRewriteLimit11 ----
func TestRewriteLimit11(t *testing.T) {
d := newTestDialector("11.2.0.4.0", 1024)
t.Run("limit and offset", func(t *testing.T) {
stmt := newTestStatement(d)
stmt.SQL.WriteString("SELECT * FROM TEST_USERS")
d.RewriteLimit11(limitClause(10, 5), stmt)
got := stmt.SQL.String()
for _, want := range []string{
"ROW_NUMBER() OVER (ORDER BY NULL) AS ROW_NUM",
"FROM (SELECT * FROM TEST_USERS) T",
"ROW_NUM BETWEEN 11 AND 15",
} {
if !strings.Contains(got, want) {
t.Errorf("RewriteLimit11 output %q missing %q", got, want)
}
}
})
t.Run("limit only uses ROWNUM", func(t *testing.T) {
stmt := newTestStatement(d)
stmt.SQL.WriteString("SELECT * FROM TEST_USERS")
d.RewriteLimit11(limitClause(0, 5), stmt)
got := stmt.SQL.String()
if want := "SELECT * FROM (SELECT * FROM TEST_USERS) WHERE ROWNUM <= 5"; got != want {
t.Errorf("RewriteLimit11 output = %q, want %q", got, want)
}
})
t.Run("offset only", func(t *testing.T) {
stmt := newTestStatement(d)
stmt.SQL.WriteString("SELECT * FROM TEST_USERS")
d.RewriteLimit11(limitClause(10, 0), stmt)
got := stmt.SQL.String()
if !strings.Contains(got, "ROW_NUM > 11") {
t.Errorf("RewriteLimit11 output %q missing ROW_NUM > 11", got)
}
})
t.Run("respects existing order by columns", func(t *testing.T) {
stmt := newTestStatement(d)
stmt.SQL.WriteString("SELECT * FROM TEST_USERS")
stmt.Clauses = map[string]clause.Clause{
"ORDER BY": {Name: "ORDER BY", Expression: clause.OrderBy{
Columns: []clause.OrderByColumn{{Column: clause.Column{Name: "NAME"}, Desc: true}},
}},
}
d.RewriteLimit11(limitClause(10, 5), stmt)
got := stmt.SQL.String()
if !strings.Contains(got, "ORDER BY NAME DESC") {
t.Errorf("RewriteLimit11 output %q missing ORDER BY NAME DESC", got)
}
})
t.Run("no-op when no limit and no offset", func(t *testing.T) {
stmt := newTestStatement(d)
stmt.SQL.WriteString("SELECT * FROM TEST_USERS")
d.RewriteLimit11(limitClause(0, 0), stmt)
if want := "SELECT * FROM TEST_USERS"; stmt.SQL.String() != want {
t.Errorf("RewriteLimit11 output = %q, want %q", stmt.SQL.String(), want)
}
})
}
// ---- TestClauseBuilders ----
func TestClauseBuilders(t *testing.T) {
t.Run("Oracle 11g uses RewriteLimit11", func(t *testing.T) {
d := newTestDialector("11.2.0.4.0", 1024)
builders := d.ClauseBuilders()
builder, ok := builders["LIMIT"]
if !ok {
t.Fatal("expected LIMIT clause builder to be registered")
}
stmt := newTestStatement(d)
stmt.SQL.WriteString("SELECT * FROM TEST_USERS")
builder(limitClause(0, 5), stmt)
if got := stmt.SQL.String(); !strings.Contains(got, "ROWNUM") {
t.Errorf("expected ROWNUM-based rewrite for 11g, got %q", got)
}
})
t.Run("Oracle 12c uses RewriteLimit", func(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
builders := d.ClauseBuilders()
builder, ok := builders["LIMIT"]
if !ok {
t.Fatal("expected LIMIT clause builder to be registered")
}
stmt := newTestStatement(d)
builder(limitClause(0, 5), stmt)
if got := stmt.SQL.String(); !strings.Contains(got, "FETCH NEXT 5 ROWS ONLY") {
t.Errorf("expected FETCH-based rewrite for 12c, got %q", got)
}
})
t.Run("Oracle 19c uses RewriteLimit", func(t *testing.T) {
d := newTestDialector("19.0.0.0", 1024)
builders := d.ClauseBuilders()
stmt := newTestStatement(d)
builders["LIMIT"](limitClause(0, 5), stmt)
if got := stmt.SQL.String(); !strings.Contains(got, "FETCH NEXT 5 ROWS ONLY") {
t.Errorf("expected FETCH-based rewrite for 19c, got %q", got)
}
})
}
// ---- TestVersionCapabilities ----
func TestVersionCapabilities(t *testing.T) {
tests := []struct {
name string
dbVer string
wantMajor int
wantIdentity bool
wantFetchOffset bool
wantNativeBoolean bool
wantExtendedString bool
wantVector bool
wantIsOracle11g bool
}{
{"11g", "11.2.0.4.0", 11, false, false, false, false, false, true},
{"10g", "10.2.0.1.0", 10, false, false, false, false, false, true},
{"12c", "12.1.0.2.0", 12, true, true, false, true, false, false},
{"18c", "18.0.0.0.0", 18, true, true, false, true, false, false},
{"19c", "19.0.0.0.0", 19, true, true, false, true, false, false},
{"21c", "21.0.0.0.0", 21, true, true, true, true, false, false},
{"23ai", "23.0.0.0.0", 23, true, true, true, true, true, false},
{"empty", "", 0, false, false, false, false, false, true},
{"invalid", "invalid", 0, false, false, false, false, false, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := oracleMajor(tt.dbVer); got != tt.wantMajor {
t.Errorf("oracleMajor(%q) = %d, want %d", tt.dbVer, got, tt.wantMajor)
}
if got := supportsIdentity(tt.dbVer); got != tt.wantIdentity {
t.Errorf("supportsIdentity(%q) = %v, want %v", tt.dbVer, got, tt.wantIdentity)
}
if got := supportsFetchOffset(tt.dbVer); got != tt.wantFetchOffset {
t.Errorf("supportsFetchOffset(%q) = %v, want %v", tt.dbVer, got, tt.wantFetchOffset)
}
if got := supportsNativeBoolean(tt.dbVer); got != tt.wantNativeBoolean {
t.Errorf("supportsNativeBoolean(%q) = %v, want %v", tt.dbVer, got, tt.wantNativeBoolean)
}
if got := supportsExtendedString(tt.dbVer); got != tt.wantExtendedString {
t.Errorf("supportsExtendedString(%q) = %v, want %v", tt.dbVer, got, tt.wantExtendedString)
}
if got := supportsVector(tt.dbVer); got != tt.wantVector {
t.Errorf("supportsVector(%q) = %v, want %v", tt.dbVer, got, tt.wantVector)
}
if got := isOracle11g(tt.dbVer); got != tt.wantIsOracle11g {
t.Errorf("isOracle11g(%q) = %v, want %v", tt.dbVer, got, tt.wantIsOracle11g)
}
})
}
}
// ---- TestIsOracle11g ----
func TestIsOracle11g(t *testing.T) {
tests := []struct {
dbVer string
want bool
}{
{"11.2.0.4.0", true},
{"10.2.0.1.0", true},
{"12.1.0.2.0", false},
{"19.0.0.0", false},
{"21.0.0.0", false},
// 空或非法版本无法判定,委托 supportsIdentity 后视为不支持 IDENTITY(即 11g 保守路径)
{"", true},
{"invalid", true},
}
for _, tt := range tests {
t.Run(tt.dbVer, func(t *testing.T) {
if got := isOracle11g(tt.dbVer); got != tt.want {
t.Errorf("isOracle11g(%q) = %v, want %v", tt.dbVer, got, tt.want)
}
})
}
}
// ---- 简单访问器 ----
func TestDefaultValueOf(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
expr := d.DefaultValueOf(nil)
e, ok := expr.(clause.Expr)
if !ok {
t.Fatalf("DefaultValueOf() returned %T, want clause.Expr", expr)
}
if e.SQL != "VALUES (DEFAULT)" {
t.Errorf("DefaultValueOf() = %q, want %q", e.SQL, "VALUES (DEFAULT)")
}
}
func TestDummyTableName(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
if got := d.DummyTableName(); got != "DUAL" {
t.Errorf("DummyTableName() = %q, want %q", got, "DUAL")
}
}
func TestName(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
if got := d.Name(); got != "oracle" {
t.Errorf("Name() = %q, want %q", got, "oracle")
}
}
func TestGetAdapter(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
adapter := d.GetAdapter()
if adapter == nil {
t.Fatal("GetAdapter() returned nil adapter")
}
// 再次调用应返回相同的驱动类型
_ = d.GetAdapter()
}
// ---- TestOpen / TestNew / TestMigrator ----
func TestOpen(t *testing.T) {
dsn := "oracle://user:pass@host:1521/db?SSL=false"
dl := Open(dsn)
d, ok := dl.(*Dialector)
if !ok {
t.Fatalf("Open() returned %T, want *Dialector", dl)
}
if d.Config == nil || d.Config.DSN != dsn {
t.Errorf("Open() did not store DSN, got %+v", d.Config)
}
}
func TestNew(t *testing.T) {
cfg := Config{DBVer: "12.1.0.2.0", DefaultStringSize: 512, SkipQuoteIdentifiers: true}
dl := New(cfg)
d, ok := dl.(*Dialector)
if !ok {
t.Fatalf("New() returned %T, want *Dialector", dl)
}
if d.DBVer != cfg.DBVer || d.DefaultStringSize != cfg.DefaultStringSize || !d.SkipQuoteIdentifiers {
t.Errorf("New() did not preserve config, got %+v", d.Config)
}
}
func TestMigrator(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
db := &gorm.DB{Config: &gorm.Config{Dialector: d}}
m := d.Migrator(db)
if m == nil {
t.Fatal("Migrator() returned nil")
}
}
// ---- TestExplain ----
func TestExplain(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
got := d.Explain("SELECT * FROM USERS WHERE id = :1 AND active = :2 AND name = :3", 5, true, "joe")
want := "SELECT * FROM USERS WHERE id = 5 AND active = 1 AND name = 'joe'"
if got != want {
t.Errorf("Explain() = %q, want %q", got, want)
}
}
func TestExplainBoolConversion(t *testing.T) {
d := newTestDialector("12.1.0.2.0", 1024)
got := d.Explain("WHERE a = :1 AND b = :2", true, false)
want := "WHERE a = 1 AND b = 0"
if got != want {
t.Errorf("Explain() = %q, want %q", got, want)
}
}
// ---- SavePoint / RollbackTo(依赖真实 DB,跳过) ----
func TestSavePoint(t *testing.T) {
t.Skip("SavePoint 需要真实的 *gorm.DB 连接,跳过")
}
func TestRollbackTo(t *testing.T) {
t.Skip("RollbackTo 需要真实的 *gorm.DB 连接,跳过")
}
+231
View File
@@ -0,0 +1,231 @@
package oracle
import (
"reflect"
"strings"
"gorm.io/gorm"
"gorm.io/gorm/callbacks"
gormSchema "gorm.io/gorm/schema"
)
// Query 是 Oracle 特定的查询回调函数
// 处理查询前后的数据转换和列名映射
func Query(db *gorm.DB) {
stmt := db.Statement
if stmt == nil {
return
}
// 1. 查询前处理
preprocessQuery(db)
// 2. 执行查询(调用默认回调)
callbacks.Query(db)
// 3. 查询后处理
if db.Error == nil {
postprocessQuery(db)
}
}
// preprocessQuery 处理查询前的预处理工作
func preprocessQuery(db *gorm.DB) {
stmt := db.Statement
if stmt == nil {
return
}
// 在 Oracle 中,某些查询可能需要特定的 hint 或优化
// 当前主要处理 LIMIT/OFFSET 重写(已在 ClauseBuilders 中处理)
// 可以根据需要添加更多预处理逻辑
}
// postprocessQuery 处理查询后的结果转换
func postprocessQuery(db *gorm.DB) {
stmt := db.Statement
if stmt == nil || stmt.Schema == nil {
return
}
// 处理查询结果
dest := stmt.Dest
if dest == nil {
return
}
// 获取反射值
rv := reflect.ValueOf(dest)
if rv.Kind() == reflect.Ptr {
rv = rv.Elem()
}
// 处理单条记录或列表
switch rv.Kind() {
case reflect.Slice:
for i := 0; i < rv.Len(); i++ {
processRecord(rv.Index(i), stmt.Schema)
}
case reflect.Struct:
processRecord(rv, stmt.Schema)
}
}
// processRecord 处理单条记录的字段值转换
func processRecord(rv reflect.Value, schema *gormSchema.Schema) {
if !rv.IsValid() {
return
}
// 确保是可寻址的值
if rv.Kind() == reflect.Interface {
rv = rv.Elem()
}
// 如果是指针,获取指向的元素
if rv.Kind() == reflect.Ptr {
if rv.IsNil() {
// 如果指针为nil,尝试创建一个实例
rv.Set(reflect.New(rv.Type().Elem()))
rv = rv.Elem()
} else {
rv = rv.Elem()
}
}
if rv.Kind() != reflect.Struct {
return
}
// 创建列名到字段的映射,处理大小写问题
columnToField := make(map[string]*gormSchema.Field)
for _, field := range schema.Fields {
if field.DBName != "" {
// Oracle 默认返回大写列名,所以将字段的 DBName 转为大写作为键
columnToField[strings.ToUpper(field.DBName)] = field
}
}
// 遍历结构体字段进行处理
for i := 0; i < rv.NumField(); i++ {
fieldStruct := rv.Type().Field(i)
fieldValue := rv.Field(i)
if !fieldValue.IsValid() || !fieldValue.CanSet() {
continue
}
// 查找对应的 Schema 字段
schemaField := findSchemaFieldByStructField(schema, &fieldStruct, columnToField)
if schemaField == nil {
continue
}
// 转换字段值
convertedValue := convertFromOracleToField(fieldValue.Interface(), schemaField)
if convertedValue != nil && !isZeroValue(convertedValue) {
setFieldValue(fieldValue, convertedValue)
}
}
}
// findSchemaFieldByStructField 根据结构体字段查找 Schema 字段
func findSchemaFieldByStructField(schema *gormSchema.Schema, structField *reflect.StructField, columnToField map[string]*gormSchema.Field) *gormSchema.Field {
// 首先尝试通过字段名查找
for _, field := range schema.Fields {
if field.Name == structField.Name {
return field
}
}
// 尝试通过数据库列名查找
dbName := structField.Tag.Get("column")
if dbName != "" {
if schemaField, exists := columnToField[strings.ToUpper(dbName)]; exists {
return schemaField
}
}
// 如果上面都没找到,尝试使用结构体字段名作为列名查找
if schemaField, exists := columnToField[strings.ToUpper(structField.Name)]; exists {
return schemaField
}
return nil
}
// setFieldValue 安全地设置字段值
func setFieldValue(fieldValue reflect.Value, value interface{}) {
if value == nil || !fieldValue.IsValid() || !fieldValue.CanSet() {
return
}
// 获取值的反射值
v := reflect.ValueOf(value)
// 处理特殊情况:当目标字段是指针时
if fieldValue.Kind() == reflect.Ptr {
// 如果值是 nil,直接设置为 nil
if value == nil {
fieldValue.Set(reflect.Zero(fieldValue.Type()))
return
}
// 如果目标是指针但值不是指针,需要创建一个指针
if v.Kind() != reflect.Ptr {
ptr := reflect.New(fieldValue.Type().Elem())
ptr.Elem().Set(reflect.ValueOf(value))
fieldValue.Set(ptr)
return
}
// 如果都是指针,直接赋值
if v.Type().AssignableTo(fieldValue.Type()) {
fieldValue.Set(v)
return
}
}
// 处理目标字段不是指针的情况
if fieldValue.Kind() == v.Kind() && v.Type().AssignableTo(fieldValue.Type()) {
fieldValue.Set(v)
return
}
// 如果类型不匹配,尝试进行类型转换
if v.CanConvert(fieldValue.Type()) {
fieldValue.Set(v.Convert(fieldValue.Type()))
return
}
// 对于接口类型,可以直接赋值
if fieldValue.Kind() == reflect.Interface {
fieldValue.Set(v)
return
}
}
// isZeroValue 检查值是否为零值
func isZeroValue(value interface{}) bool {
if value == nil {
return true
}
v := reflect.ValueOf(value)
switch v.Kind() {
case reflect.String:
return v.String() == ""
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
return v.Int() == 0
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
return v.Uint() == 0
case reflect.Float32, reflect.Float64:
return v.Float() == 0
case reflect.Bool:
return !v.Bool()
case reflect.Ptr, reflect.Interface:
return v.IsNil()
}
return false
}
+84
View File
@@ -0,0 +1,84 @@
package tests
import (
"testing"
"time"
)
func TestCreateSingle(t *testing.T) {
// 先创建表
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS")
user := User{
Name: "Test User",
Email: "test@example.com",
Age: 25,
Active: true,
}
result := DB.Create(&user)
if result.Error != nil {
t.Fatalf("failed to create user: %v", result.Error)
}
if user.ID == 0 {
t.Error("expected user ID to be set after create")
}
t.Logf("Created user with ID: %d", user.ID)
}
func TestCreateBatch(t *testing.T) {
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS")
users := []User{
{Name: "User 1", Email: "user1@example.com", Age: 20},
{Name: "User 2", Email: "user2@example.com", Age: 22},
{Name: "User 3", Email: "user3@example.com", Age: 24},
}
result := DB.Create(&users)
if result.Error != nil {
t.Fatalf("failed to create users: %v", result.Error)
}
if result.RowsAffected != 3 {
t.Errorf("expected 3 rows affected, got %d", result.RowsAffected)
}
for i, user := range users {
if user.ID == 0 {
t.Errorf("user %d: expected ID to be set", i)
}
}
}
func TestCreateWithTimestamp(t *testing.T) {
if err := DB.AutoMigrate(&Product{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_PRODUCTS")
product := Product{
Name: "Test Product",
Price: 99.99,
Stock: 100,
Description: "This is a test product with long description",
CreatedAt: time.Now(),
}
result := DB.Create(&product)
if result.Error != nil {
t.Fatalf("failed to create product: %v", result.Error)
}
if product.ID == 0 {
t.Error("expected product ID to be set")
}
}
+70
View File
@@ -0,0 +1,70 @@
package tests
import (
"testing"
)
func TestDeleteSingle(t *testing.T) {
// 先确保表存在并清空
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS")
// 创建测试数据
user := User{Name: "Delete Test", Email: "delete@example.com", Age: 40}
DB.Create(&user)
// 删除
result := DB.Delete(&user)
if result.Error != nil {
t.Fatalf("failed to delete: %v", result.Error)
}
// 验证软删除(如果有 deleted_at)
var deleted User
result = DB.First(&deleted, user.ID)
if result.Error == nil {
t.Error("expected user to be soft deleted")
}
}
func TestDeleteWithoutWhere(t *testing.T) {
// 确保表存在
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS")
// 测试无 WHERE 条件的删除应该失败
result := DB.Delete(&User{})
if result.Error == nil {
t.Error("expected error for delete without WHERE condition")
}
t.Logf("Got expected error: %v", result.Error)
}
func TestHardDelete(t *testing.T) {
// 先确保表存在并清空
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS")
// 创建测试数据
user := User{Name: "Hard Delete", Email: "hard@example.com", Age: 50}
DB.Create(&user)
// 硬删除
result := DB.Unscoped().Delete(&user)
if result.Error != nil {
t.Fatalf("failed to hard delete: %v", result.Error)
}
// 验证完全删除
var deleted User
result = DB.Unscoped().First(&deleted, user.ID)
if result.Error == nil {
t.Error("expected user to be completely deleted")
}
}
+92
View File
@@ -0,0 +1,92 @@
package tests
import (
"testing"
"time"
"gorm.io/gorm"
)
// UserWithHook 带 Hook 的测试模型
type UserWithHook struct {
ID uint `gorm:"primaryKey;autoIncrement"`
Name string `gorm:"size:100"`
Email string `gorm:"size:200"`
HookLog string `gorm:"size:500"` // 记录 Hook 执行
CreatedAt time.Time
UpdatedAt time.Time
}
func (UserWithHook) TableName() string {
return "TEST_USERS_HOOK"
}
func (u *UserWithHook) BeforeCreate(tx *gorm.DB) error {
u.HookLog += "BeforeCreate;"
return nil
}
func (u *UserWithHook) AfterCreate(tx *gorm.DB) error {
u.HookLog += "AfterCreate;"
return nil
}
func (u *UserWithHook) BeforeUpdate(tx *gorm.DB) error {
u.HookLog += "BeforeUpdate;"
return nil
}
func (u *UserWithHook) AfterUpdate(tx *gorm.DB) error {
u.HookLog += "AfterUpdate;"
return nil
}
func (u *UserWithHook) BeforeDelete(tx *gorm.DB) error {
u.HookLog += "BeforeDelete;"
return nil
}
func (u *UserWithHook) AfterDelete(tx *gorm.DB) error {
u.HookLog += "AfterDelete;"
return nil
}
func TestHooks(t *testing.T) {
// 创建表
if err := DB.AutoMigrate(&UserWithHook{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS_HOOK")
// 测试 Create Hook
user := UserWithHook{Name: "Hook Test", Email: "hook@example.com"}
result := DB.Create(&user)
if result.Error != nil {
t.Fatalf("failed to create: %v", result.Error)
}
if user.HookLog == "" {
t.Error("expected hooks to be called")
}
t.Logf("Hook log after create: %s", user.HookLog)
// 测试 Update Hook
user.HookLog = "" // 清空
if err := DB.Model(&user).Update("name", "Updated Name").Error; err != nil {
t.Fatalf("failed to update: %v", err)
}
t.Logf("Hook log after update: %s", user.HookLog)
if user.HookLog == "" {
t.Error("expected update hooks to be called")
}
// 测试 Delete Hook
user.HookLog = "" // 清空
if err := DB.Delete(&user).Error; err != nil {
t.Fatalf("failed to delete: %v", err)
}
t.Logf("Hook log after delete: %s", user.HookLog)
if user.HookLog == "" {
t.Error("expected delete hooks to be called")
}
}
+62
View File
@@ -0,0 +1,62 @@
package tests
import (
"log"
"os"
"testing"
"gorm.io/gorm"
oracle "git.charlienet.top/go/oracle"
)
var DB *gorm.DB
func TestMain(m *testing.M) {
// 使用提供的 DSN(可通过 ORACLE_DSN 环境变量覆盖)
dsn := os.Getenv("ORACLE_DSN")
if dsn == "" {
// 注:go-ora v2.9.0 起 CONNECTION TIMEOUT 语义从"socket 读超时"变为"连接建立超时"
// 读超时改用 SOCKET TIMEOUT 指定。两者均设 90s 以保留原有的读超时保护语义。
// 安全:此处为占位符,真实 DSN 请通过 ORACLE_DSN 环境变量提供,避免凭据入库。
dsn = "oracle://user:password@host:1521/service?SSL=false&CONNECTION TIMEOUT=90&SOCKET TIMEOUT=90&LANGUAGE=SIMPLIFIED+CHINESE&TERRITORY=CHINA"
}
var err error
DB, err = gorm.Open(oracle.Open(dsn), &gorm.Config{})
if err != nil {
log.Fatalf("failed to connect database: %v", err)
}
log.Println("successfully connected to database")
// 清理测试表
cleanup()
// 运行测试
code := m.Run()
// 最终清理
cleanup()
os.Exit(code)
}
// cleanup 删除所有测试表(忽略错误,因为表可能不存在)
func cleanup() {
_ = DB.Migrator().DropTable(&User{}, &Product{}, &Order{}, &UserWithHook{}, &SeqDefaultViaDriverModel{}, &BigStringModel{})
// TEST_SEQ_DEFAULT 表通过原生 SQL 创建(无 autoIncrement),DropTable 无法识别,
// 因此用原生 SQL 清理表与序列
DB.Exec("DROP TABLE TEST_SEQ_DEFAULT")
DB.Exec("DROP SEQUENCE SEQ_TEST_SEQ_DEFAULT")
// 序列默认值测试的独立序列(TEST_SEQ_DEF 表可通过 DropTable 清理并级联删除触发器)
DB.Exec("DROP SEQUENCE SEQ_TEST_SEQ_DEF_CODE")
DB.Exec("DROP SEQUENCE SEQ_TEST_SEQ_DEF")
}
// clearTable 清空指定测试表,保证测试之间的数据隔离
func clearTable(t *testing.T, table string) {
t.Helper()
if err := DB.Exec("DELETE FROM " + table).Error; err != nil {
t.Fatalf("failed to clear table %s: %v", table, err)
}
}
+113
View File
@@ -0,0 +1,113 @@
package tests
import (
"strings"
"testing"
)
func TestAutoMigrate(t *testing.T) {
// 测试创建表
err := DB.AutoMigrate(&User{}, &Product{}, &Order{})
if err != nil {
t.Fatalf("failed to auto migrate: %v", err)
}
// 验证表存在
if !DB.Migrator().HasTable(&User{}) {
t.Error("expected User table to exist")
}
if !DB.Migrator().HasTable(&Product{}) {
t.Error("expected Product table to exist")
}
if !DB.Migrator().HasTable(&Order{}) {
t.Error("expected Order table to exist")
}
}
func TestAddColumn(t *testing.T) {
// 先创建表
DB.AutoMigrate(&User{})
// 添加列(需要定义新模型)
type UserWithPhone struct {
User
Phone string `gorm:"size:20"`
}
err := DB.AutoMigrate(&UserWithPhone{})
if err != nil {
t.Fatalf("failed to add column: %v", err)
}
// 验证列存在
if !DB.Migrator().HasColumn(&UserWithPhone{}, "phone") {
t.Error("expected phone column to exist")
}
}
func TestDropTable(t *testing.T) {
// 确保表存在
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
err := DB.Migrator().DropTable(&User{})
if err != nil {
t.Fatalf("failed to drop table: %v", err)
}
if DB.Migrator().HasTable(&User{}) {
t.Error("expected User table to be dropped")
}
}
// BigStringModel 验证 11g 下 size>4000 的 string 字段映射为 CLOB 列
type BigStringModel struct {
ID uint `gorm:"column:id;primaryKey"`
Big string `gorm:"column:big;size:5000"`
}
func (BigStringModel) TableName() string {
return "TEST_BIG_STRING"
}
// TestBigStringMapsToCLOBOn11g 验证 11g 下 size>4000 的 string 字段:
// DataTypeOf 的 32k VARCHAR2 特性(12c+ 才支持)在 11g 不触发,
// 字段应建为 CLOB,且能正常插入超过 4000 字节的长文本。
func TestBigStringMapsToCLOBOn11g(t *testing.T) {
if err := DB.AutoMigrate(&BigStringModel{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
defer func() {
DB.Migrator().DropTable(&BigStringModel{})
}()
// 验证列类型为 CLOB(Oracle 数据字典列名大写存储)
var dataType string
if err := DB.Raw("SELECT DATA_TYPE FROM USER_TAB_COLUMNS WHERE TABLE_NAME = ? AND COLUMN_NAME = ?",
"TEST_BIG_STRING", "BIG").Scan(&dataType).Error; err != nil {
t.Fatalf("failed to query column type: %v", err)
}
if dataType != "CLOB" {
t.Errorf("expected column type CLOB on 11g, got %q", dataType)
}
// 插入超过 4000 字节的长文本,验证 CLOB 列可容纳
longText := strings.Repeat("a", 5000)
bm := BigStringModel{Big: longText}
if err := DB.Create(&bm).Error; err != nil {
t.Fatalf("failed to create: %v", err)
}
if bm.ID == 0 {
t.Error("expected ID to be set")
}
// 回读验证内容完整
var got string
if err := DB.Raw("SELECT BIG FROM TEST_BIG_STRING WHERE id = ?", bm.ID).Scan(&got).Error; err != nil {
t.Fatalf("failed to query back: %v", err)
}
if got != longText {
t.Errorf("text roundtrip mismatch: got len=%d, want len=%d", len(got), len(longText))
}
}
+80
View File
@@ -0,0 +1,80 @@
package tests
import (
"time"
"gorm.io/gorm"
)
// User 基本测试模型
type User struct {
ID uint `gorm:"column:id;primaryKey;autoIncrement"`
Name string `gorm:"size:100;not null"`
Email string `gorm:"size:200;uniqueIndex"`
Age int `gorm:"default:0"`
Active bool // 不使用 default,以便显式设置 false 时能正确存储
CreatedAt time.Time
UpdatedAt time.Time
DeletedAt gorm.DeletedAt `gorm:"index"`
}
func (User) TableName() string {
return "TEST_USERS"
}
// Product 测试数值类型
type Product struct {
ID uint `gorm:"column:id;primaryKey;autoIncrement"`
Name string `gorm:"size:200;not null"`
Price float64 `gorm:"precision:10;scale:2"`
Stock int `gorm:"default:0"`
Description string `gorm:"type:CLOB"`
CreatedAt time.Time
}
func (Product) TableName() string {
return "TEST_PRODUCTS"
}
// Order 测试关联关系
type Order struct {
ID uint `gorm:"column:id;primaryKey;autoIncrement"`
UserID uint `gorm:"not null;index"`
User User `gorm:"foreignKey:UserID"`
Total float64 `gorm:"precision:12;scale:2"`
Status string `gorm:"size:20;default:'pending'"`
CreatedAt time.Time
UpdatedAt time.Time
}
func (Order) TableName() string {
return "TEST_ORDERS"
}
// SeqDefaultModel 显式使用序列默认值的模型(主键不用 autoIncrement
// 而是通过列的 DEFAULT 值使用序列 SEQ_TEST_SEQ_DEFAULT.NEXTVAL
type SeqDefaultModel struct {
ID uint `gorm:"column:id;primaryKey"`
Name string `gorm:"size:100"`
}
func (SeqDefaultModel) TableName() string {
return "TEST_SEQ_DEFAULT"
}
// SeqDefaultViaDriverModel 通过驱动 AutoMigrate 建表,验证 11g 下序列默认值的
// 触发器路径(模型使用 gorm:"default:(SEQ_TEST_SEQ_DEF_CODE.NEXTVAL)")。
//
// Code 用 int 字段并给默认值加括号:GORM 对含括号的默认值会跳过 ParseInt 解析
// schema/field.go:231),从而 DefaultValueInterface 保持 nil,字段进入
// FieldsWithDefaultDBValue —— INSERT 时 GORM 会省略该列,触发 BEFORE INSERT 触发器回填序列值。
// 若不加括号,GORM 会把 "SEQ_...NEXTVAL" 当作整数解析失败,schema.Parse 直接报错。
type SeqDefaultViaDriverModel struct {
ID uint `gorm:"column:id;primaryKey"`
Code int `gorm:"column:code;default:(SEQ_TEST_SEQ_DEF_CODE.NEXTVAL)"`
Name string `gorm:"size:100"`
}
func (SeqDefaultViaDriverModel) TableName() string {
return "TEST_SEQ_DEF"
}
+84
View File
@@ -0,0 +1,84 @@
package tests
import (
"testing"
)
func TestQuerySingle(t *testing.T) {
// 先确保表存在并清空
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS")
// 创建测试数据
user := User{Name: "Query Test", Email: "query@example.com", Age: 35}
DB.Create(&user)
// 查询
var found User
result := DB.First(&found, user.ID)
if result.Error != nil {
t.Fatalf("failed to query: %v", result.Error)
}
if found.Name != "Query Test" {
t.Errorf("expected name 'Query Test', got '%s'", found.Name)
}
}
func TestQueryWithConditions(t *testing.T) {
// 先确保表存在并清空
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS")
// 创建测试数据
users := []User{
{Name: "Condition 1", Email: "cond1@example.com", Age: 20, Active: true},
{Name: "Condition 2", Email: "cond2@example.com", Age: 25, Active: true},
{Name: "Condition 3", Email: "cond3@example.com", Age: 30, Active: false},
}
DB.Create(&users)
// 条件查询
var results []User
result := DB.Where("active = ? AND age > ?", true, 22).Find(&results)
if result.Error != nil {
t.Fatalf("failed to query: %v", result.Error)
}
if len(results) != 1 {
t.Errorf("expected 1 result, got %d", len(results))
}
}
func TestQueryWithLimit(t *testing.T) {
// 先确保表存在并清空
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS")
// 创建测试数据
for i := 0; i < 10; i++ {
user := User{
Name: "Limit Test",
Email: "limit" + string(rune('a'+i)) + "@example.com",
Age: i,
}
DB.Create(&user)
}
// 分页查询
var results []User
result := DB.Limit(5).Offset(2).Find(&results)
if result.Error != nil {
t.Fatalf("failed to query: %v", result.Error)
}
if len(results) != 5 {
t.Errorf("expected 5 results, got %d", len(results))
}
}
+230
View File
@@ -0,0 +1,230 @@
package tests
import (
"testing"
)
// TestSequenceObjectExists 验证 11g 自增依赖的序列对象存在
// (Oracle 对象名默认大写存储,序列命名规则见 migrator.go 的 sequenceName
func TestSequenceObjectExists(t *testing.T) {
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
// 验证序列对象存在(11g 自增依赖序列)
var count int64
if err := DB.Raw("SELECT COUNT(*) FROM USER_SEQUENCES WHERE SEQUENCE_NAME = ?", "SEQ_TEST_USERS").Scan(&count).Error; err != nil {
t.Fatalf("failed to query sequence: %v", err)
}
if count == 0 {
t.Error("expected sequence SEQ_TEST_USERS to exist (auto increment support)")
}
}
// TestTriggerObjectExists 验证 11g 自增依赖的触发器对象存在
// (触发器命名规则见 migrator.go 的 triggerName
func TestTriggerObjectExists(t *testing.T) {
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
// 验证触发器对象存在(11g 自增依赖触发器)
var count int64
if err := DB.Raw("SELECT COUNT(*) FROM USER_TRIGGERS WHERE TRIGGER_NAME = ?", "TRG_TEST_USERS").Scan(&count).Error; err != nil {
t.Fatalf("failed to query trigger: %v", err)
}
if count == 0 {
t.Error("expected trigger TRG_TEST_USERS to exist (auto increment support)")
}
}
// TestAutoIncrementViaSequence 连续插入两条,验证 ID 递增(证明序列工作)
func TestAutoIncrementViaSequence(t *testing.T) {
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS")
// 连续插入两条,验证 ID 递增(证明序列工作)
u1 := User{Name: "Seq 1", Email: "seq1@example.com", Age: 1}
u2 := User{Name: "Seq 2", Email: "seq2@example.com", Age: 2}
if err := DB.Create(&u1).Error; err != nil {
t.Fatalf("failed to create u1: %v", err)
}
if err := DB.Create(&u2).Error; err != nil {
t.Fatalf("failed to create u2: %v", err)
}
if u1.ID == 0 {
t.Error("expected u1.ID to be set")
}
if u2.ID != u1.ID+1 {
t.Errorf("expected u2.ID = u1.ID+1, got u1.ID=%d, u2.ID=%d", u1.ID, u2.ID)
}
}
// TestExplicitSequenceDefaultValue 验证:插入时不给主键 id,主键由序列自动生成。
//
// ⚠️ 重要发现:Oracle 11g 的 CREATE TABLE 的 DEFAULT 子句不允许引用序列的
// NEXTVALORA-00984: 列在此处不允许,报错位置指向 NEXTVAL),该能力 12c 才引入。
// 因此任务原始 SQL "id NUMBER(19) DEFAULT SEQ_TEST_SEQ_DEFAULT.NEXTVAL" 无法在 11g 执行。
// 本测试改用 11g 的标准做法——BEFORE INSERT 触发器(与 go-ora migrator 给 autoIncrement
// 表生成的触发器机制相同)在 id 为 NULL 时从序列取值,验证语义与"DEFAULT 序列"等价:
// 插入时不提供 id,主键由序列自动生成并严格递增。
//
// 驱动行为实测(GORM v1.31.2 + go-ora v2.9.0,真实 11g 库):
// - GORM 会把非 autoIncrement 的 int/uint 主键自动视为自增(schema.go:337-348
// 设 AutoIncrement=true、HasDefaultValue=true,并加入 FieldsWithDefaultDBValue)。
// - 因此 DB.Create 不会把 id 列显式写入 INSERT,而是使用 RETURNING 回填;
// 由于表上有 BEFORE INSERT 触发器从序列取值,RETURNING 返回的即序列生成值。
// - 实测 DB.Create 后 m1.ID 被回填为序列值(START WITH 100 → 100),非 0。
func TestExplicitSequenceDefaultValue(t *testing.T) {
// 先删除可能残留的对象(忽略错误:对象可能不存在)
DB.Exec("DROP SEQUENCE SEQ_TEST_SEQ_DEFAULT")
DB.Migrator().DropTable(&SeqDefaultModel{})
// 创建序列
if err := DB.Exec("CREATE SEQUENCE SEQ_TEST_SEQ_DEFAULT START WITH 100 INCREMENT BY 1 NOCACHE").Error; err != nil {
t.Fatalf("failed to create sequence: %v", err)
}
// 测试结束清理序列和表(DROP TABLE 会级联删除其上的触发器)
defer func() {
DB.Exec("DROP TABLE TEST_SEQ_DEFAULT")
DB.Exec("DROP SEQUENCE SEQ_TEST_SEQ_DEFAULT")
}()
// 建表(11g 的 DEFAULT 子句不支持引用序列,改用普通列 + 触发器)
if err := DB.Exec(`CREATE TABLE TEST_SEQ_DEFAULT (
id NUMBER(19) NOT NULL PRIMARY KEY,
name VARCHAR2(100)
)`).Error; err != nil {
t.Fatalf("failed to create table: %v", err)
}
// 创建 BEFORE INSERT 触发器:id 为 NULL 时从序列取值
triggerSQL := `CREATE OR REPLACE TRIGGER TRG_TEST_SEQ_DEFAULT
BEFORE INSERT ON TEST_SEQ_DEFAULT
FOR EACH ROW
BEGIN
IF :NEW.id IS NULL THEN
SELECT SEQ_TEST_SEQ_DEFAULT.NEXTVAL INTO :NEW.id FROM DUAL;
END IF;
END;`
if err := DB.Exec(triggerSQL).Error; err != nil {
t.Fatalf("failed to create trigger: %v", err)
}
// 插入时不给 id(依赖触发器从序列默认取值)
m1 := SeqDefaultModel{Name: "Default Seq 1"}
if err := DB.Create(&m1).Error; err != nil {
t.Fatalf("failed to create m1: %v", err)
}
if m1.ID == 0 {
t.Error("expected m1.ID to be set from sequence")
}
t.Logf("m1.ID 回填=%d", m1.ID)
// 用数据库查询交叉验证实际存储的 ID 来自序列
var id1 uint
if err := DB.Raw("SELECT id FROM TEST_SEQ_DEFAULT WHERE name = ?", "Default Seq 1").Scan(&id1).Error; err != nil {
t.Fatalf("failed to query id1: %v", err)
}
if id1 == 0 {
t.Error("expected database id1 to be set from sequence")
}
if id1 != m1.ID {
t.Errorf("database id1=%d != m1.ID=%d", id1, m1.ID)
}
t.Logf("数据库实际 id=%d(序列 START WITH 100", id1)
// 第二次插入验证序列递增
m2 := SeqDefaultModel{Name: "Default Seq 2"}
if err := DB.Create(&m2).Error; err != nil {
t.Fatalf("failed to create m2: %v", err)
}
if m2.ID != m1.ID+1 {
t.Errorf("expected m2.ID = m1.ID+1, got m1.ID=%d, m2.ID=%d", m1.ID, m2.ID)
}
// 数据库端二次确认
var id2 uint
if err := DB.Raw("SELECT id FROM TEST_SEQ_DEFAULT WHERE name = ?", "Default Seq 2").Scan(&id2).Error; err != nil {
t.Fatalf("failed to query id2: %v", err)
}
if id2 != id1+1 {
t.Errorf("expected database id2 = id1+1, got id1=%d, id2=%d", id1, id2)
}
t.Logf("m2.ID 回填=%d, 数据库实际 id=%d", m2.ID, id2)
}
// TestSequenceDefaultViaAutoMigrate 验证 11g 下通过驱动 AutoMigrate 建表时的序列默认值触发器路径。
//
// 模型使用 gorm:"default:SEQ_TEST_SEQ_DEF_CODE.NEXTVAL"
// - 11g 的 CREATE TABLE 的 DEFAULT 子句不允许引用序列 NEXTVALORA-00984),
// 驱动重写的 FullDataTypeOf 会跳过 DEFAULT 子句,建表后由 CreateTable 流程创建
// BEFORE INSERT 触发器 SEQDEF_TRG_TEST_SEQ_DEF_CODE 实现等价语义;
// - 12c+ 则不创建触发器,直接生成 DEFAULT SEQ_TEST_SEQ_DEF_CODE.NEXTVAL。
func TestSequenceDefaultViaAutoMigrate(t *testing.T) {
// 先清理可能残留的对象(忽略错误:对象可能不存在)
DB.Migrator().DropTable(&SeqDefaultViaDriverModel{})
DB.Exec("DROP SEQUENCE SEQ_TEST_SEQ_DEF_CODE")
DB.Exec("DROP SEQUENCE SEQ_TEST_SEQ_DEF")
DB.Exec("DROP TRIGGER SEQDEF_TRG_TEST_SEQ_DEF_CODE")
// 手动创建序列:DEFAULT 引用的序列不会自动创建(autoIncrement 只负责创建
// ID 用的 SEQ_TEST_SEQ_DEF),测试中先 DROP 再 CREATE,测试后清理。
if err := DB.Exec("CREATE SEQUENCE SEQ_TEST_SEQ_DEF_CODE START WITH 100 INCREMENT BY 1 NOCACHE").Error; err != nil {
t.Fatalf("failed to create sequence: %v", err)
}
// 测试结束清理表(DROP TABLE 会级联删除其上触发器)和序列
defer func() {
DB.Exec("DROP TABLE TEST_SEQ_DEF")
DB.Exec("DROP SEQUENCE SEQ_TEST_SEQ_DEF_CODE")
DB.Exec("DROP SEQUENCE SEQ_TEST_SEQ_DEF")
}()
// 通过驱动 AutoMigrate 建表(11g 下不会生成 DEFAULT <seq>.NEXTVAL 子句)
if err := DB.AutoMigrate(&SeqDefaultViaDriverModel{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
// 验证序列默认值触发器已创建(命名 SEQDEF_TRG_<table>_<column>
// 避免与 autoIncrement 的 TRG_TEST_SEQ_DEF 冲突)
var trigCount int64
if err := DB.Raw("SELECT COUNT(*) FROM USER_TRIGGERS WHERE TRIGGER_NAME = ?", "SEQDEF_TRG_TEST_SEQ_DEF_CODE").Scan(&trigCount).Error; err != nil {
t.Fatalf("failed to query trigger: %v", err)
}
if trigCount == 0 {
t.Error("expected sequence default trigger SEQDEF_TRG_TEST_SEQ_DEF_CODE to exist")
}
// 插入时不给 Code(字段在 FieldsWithDefaultDBValue 中,GORM 省略该列),
// 触发器应从序列回填,RETURNING 将序列值回填到 m1.Code
m1 := SeqDefaultViaDriverModel{Name: "Via AutoMigrate 1"}
if err := DB.Create(&m1).Error; err != nil {
t.Fatalf("failed to create m1: %v", err)
}
if m1.Code == 0 {
t.Error("expected m1.Code to be set by sequence default trigger")
}
t.Logf("m1.Code 回填=%d(序列 START WITH 100", m1.Code)
// 查询数据库交叉验证 code 由触发器回填
var code1 int
if err := DB.Raw("SELECT code FROM TEST_SEQ_DEF WHERE name = ?", "Via AutoMigrate 1").Scan(&code1).Error; err != nil {
t.Fatalf("failed to query code1: %v", err)
}
if code1 != m1.Code {
t.Errorf("database code1=%d != m1.Code=%d", code1, m1.Code)
}
// 再次插入验证序列递增
m2 := SeqDefaultViaDriverModel{Name: "Via AutoMigrate 2"}
if err := DB.Create(&m2).Error; err != nil {
t.Fatalf("failed to create m2: %v", err)
}
if m2.Code != m1.Code+1 {
t.Errorf("expected m2.Code = m1.Code+1, got m1.Code=%d, m2.Code=%d", m1.Code, m2.Code)
}
t.Logf("m2.Code 回填=%d", m2.Code)
}
+37
View File
@@ -0,0 +1,37 @@
package tests
import (
"testing"
)
// TestSoftDeleteSetsDeletedAt 精确验证软删除后 DELETED_AT 确实被写入数据库
func TestSoftDeleteSetsDeletedAt(t *testing.T) {
// 确保表存在并清空
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS")
user := User{Name: "Soft Delete Exact", Email: "soft_exact@example.com", Age: 40}
if err := DB.Create(&user).Error; err != nil {
t.Fatalf("failed to create: %v", err)
}
// 软删除
if err := DB.Delete(&user).Error; err != nil {
t.Fatalf("failed to soft delete: %v", err)
}
// 直接验证 DELETED_AT 被写入:用 Unscoped 查询(绕过软删除过滤)
var deleted User
if err := DB.Unscoped().First(&deleted, user.ID).Error; err != nil {
t.Fatalf("failed to query unscoped: %v", err)
}
if !deleted.DeletedAt.Valid {
t.Errorf("expected DeletedAt to be valid (set), got invalid")
}
if deleted.DeletedAt.Time.IsZero() {
t.Errorf("expected DeletedAt to have a timestamp, got zero")
}
t.Logf("DeletedAt set to: %v", deleted.DeletedAt.Time)
}
+74
View File
@@ -0,0 +1,74 @@
package tests
import (
"testing"
)
func TestUpdateSingle(t *testing.T) {
// 先确保表存在并清空
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS")
// 创建测试数据
user := User{Name: "Update Test", Email: "update@example.com", Age: 30}
DB.Create(&user)
// 更新
result := DB.Model(&user).Update("age", 31)
if result.Error != nil {
t.Fatalf("failed to update: %v", result.Error)
}
if result.RowsAffected != 1 {
t.Errorf("expected 1 row affected, got %d", result.RowsAffected)
}
// 验证更新
var updated User
DB.First(&updated, user.ID)
if updated.Age != 31 {
t.Errorf("expected age 31, got %d", updated.Age)
}
}
func TestUpdateMultiple(t *testing.T) {
// 先确保表存在并清空
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS")
// 创建测试数据
users := []User{
{Name: "Multi Update 1", Email: "multi1@example.com", Age: 25},
{Name: "Multi Update 2", Email: "multi2@example.com", Age: 25},
}
DB.Create(&users)
// 批量更新
result := DB.Model(&User{}).Where("age = ?", 25).Update("age", 26)
if result.Error != nil {
t.Fatalf("failed to update: %v", result.Error)
}
if result.RowsAffected < 2 {
t.Errorf("expected at least 2 rows affected, got %d", result.RowsAffected)
}
}
func TestUpdateWithoutWhere(t *testing.T) {
// 确保表存在
if err := DB.AutoMigrate(&User{}); err != nil {
t.Fatalf("failed to migrate: %v", err)
}
clearTable(t, "TEST_USERS")
// 测试无 WHERE 条件的更新应该失败
result := DB.Model(&User{}).Update("age", 99)
if result.Error == nil {
t.Error("expected error for update without WHERE condition")
}
t.Logf("Got expected error: %v", result.Error)
}
+210
View File
@@ -0,0 +1,210 @@
package oracle
import (
"database/sql"
"fmt"
"reflect"
"github.com/thoas/go-funk"
"gorm.io/gorm"
"gorm.io/gorm/clause"
gormSchema "gorm.io/gorm/schema"
)
func Update(db *gorm.DB) {
stmt := db.Statement
if stmt == nil {
return
}
schema := stmt.Schema
if schema == nil {
return
}
boundVars := make(map[string]int)
hasDefaultValues := len(schema.FieldsWithDefaultDBValue) > 0
if !stmt.Unscoped {
for _, c := range schema.UpdateClauses {
stmt.AddClause(c)
}
}
// 注入主键 WHERE 条件(GORM 默认回调会做这一步)
pkValues := addPrimaryKeyWhere(stmt, schema)
// 多行更新时 Oracle 不支持单行 RETURNING INTO,只有单行更新才启用 RETURNING
if pkValues != 1 {
hasDefaultValues = false
}
// WHERE 安全检查
where, hasWhere := stmt.Clauses["WHERE"].Expression.(clause.Where)
if hasWhere {
if checkMissingWhereConditions(where.Exprs, schema) {
db.AddError(fmt.Errorf("missing WHERE condition in UPDATE"))
return
}
} else {
// 没有 WHERE 子句,且模型主键没有可用的值
db.AddError(fmt.Errorf("missing WHERE condition in UPDATE"))
return
}
if stmt.SQL.String() == "" {
// 构建 UPDATE 语句
stmt.AddClauseIfNotExists(clause.Update{Table: clause.Table{Name: stmt.Schema.Table}})
// 构建 SET 子句
_, hasSet := stmt.Clauses["SET"].Expression.(clause.Set)
if !hasSet {
// 获取要更新的值
// 从 stmt.Dest 获取待更新的数据
reflectValue := reflect.ValueOf(stmt.Dest)
if reflectValue.Kind() == reflect.Ptr {
reflectValue = reflectValue.Elem()
}
// 构建 SET 表达式
sets := make(clause.Set, 0)
switch reflectValue.Kind() {
case reflect.Struct:
for _, field := range schema.Fields {
if !field.PrimaryKey && field.Updatable {
if fieldValue, isZero := field.ValueOf(stmt.Context, reflectValue); !isZero {
// 转换值为 Oracle 兼容格式
convertedValue := convertValue(fieldValue, field)
sets = append(sets, clause.Assignment{Column: clause.Column{Name: field.DBName}, Value: convertedValue})
}
}
}
case reflect.Map:
// 处理 map 类型的更新
for _, mapKey := range reflectValue.MapKeys() {
key := mapKey.String()
if field := schema.LookUpField(key); field != nil {
if !field.PrimaryKey && field.Updatable {
value := reflectValue.MapIndex(mapKey).Interface()
// 转换值为 Oracle 兼容格式
convertedValue := convertValue(value, field)
sets = append(sets, clause.Assignment{Column: clause.Column{Name: field.DBName}, Value: convertedValue})
}
}
}
}
stmt.AddClause(clause.Set(sets))
}
// 添加 RETURNING 子句(如果有默认值字段)
if hasDefaultValues {
stmt.AddClauseIfNotExists(clause.Returning{
Columns: funk.Map(schema.FieldsWithDefaultDBValue, func(field *gormSchema.Field) clause.Column {
return clause.Column{Name: field.DBName}
}).([]clause.Column),
})
}
// 构建语句
stmt.Build("UPDATE", "SET", "WHERE", "RETURNING")
// 如果有 RETURNING 子句,添加 INTO 子句
if hasDefaultValues {
stmt.WriteString(" INTO ")
for idx, field := range schema.FieldsWithDefaultDBValue {
if idx > 0 {
stmt.WriteByte(',')
}
boundVars[field.Name] = len(stmt.Vars)
stmt.AddVar(stmt, sql.Out{Dest: reflect.New(field.FieldType).Interface()})
}
}
}
if !db.DryRun {
// 开启事务以确保更新的一致性
var tx *sql.Tx
var err error
var isTransaction bool = false
// 检查是否已经在一个事务中
if sqlTx, ok := stmt.ConnPool.(*sql.Tx); ok {
tx = sqlTx
isTransaction = true
} else if sqlDb, ok := stmt.ConnPool.(*sql.DB); ok {
tx, err = sqlDb.Begin()
if err != nil {
db.AddError(err)
return
}
defer func() {
if db.Error != nil && !isTransaction {
_ = tx.Rollback()
} else if !isTransaction {
_ = tx.Commit()
}
}()
} else {
db.AddError(fmt.Errorf("unsupported connection pool type"))
return
}
// 执行更新操作
var execConn *sql.Tx
if isTransaction {
execConn = tx // 已经在事务中,直接使用原事务
} else {
execConn = tx // 使用新创建的事务
}
result, err := execConn.ExecContext(stmt.Context, stmt.SQL.String(), stmt.Vars...)
if err != nil {
db.AddError(err)
// 如果不是在已有事务中,则回滚我们创建的事务
if !isTransaction {
_ = tx.Rollback()
}
return
}
db.RowsAffected, _ = result.RowsAffected()
// 处理 RETURNING 返回值
if hasDefaultValues {
updateTo := stmt.ReflectValue
switch updateTo.Kind() {
case reflect.Slice, reflect.Array:
// 对于切片或数组,只更新第一个元素
if updateTo.Len() > 0 {
updateTo = updateTo.Index(0)
}
}
// 绑定返回值到模型字段
funk.ForEach(
funk.Filter(schema.FieldsWithDefaultDBValue, func(field *gormSchema.Field) bool {
return funk.Contains(boundVars, field.Name)
}),
func(field *gormSchema.Field) {
switch updateTo.Kind() {
case reflect.Struct:
if err = field.Set(stmt.Context, updateTo, stmt.Vars[boundVars[field.Name]].(sql.Out).Dest); err != nil {
db.AddError(err)
}
case reflect.Map:
// 设置Map类型的值
mapValue := reflect.ValueOf(updateTo.Interface())
if mapValue.IsValid() && mapValue.Type().Key().Kind() == reflect.String {
keyValue := reflect.ValueOf(field.DBName)
destValue := reflect.ValueOf(stmt.Vars[boundVars[field.Name]].(sql.Out).Dest)
if destValue.Kind() == reflect.Ptr {
destValue = destValue.Elem()
}
mapValue.SetMapIndex(keyValue, destValue)
}
}
},
)
}
}
}