feat: Oracle 驱动完整化——回调体系、版本感知、驱动抽象层与测试套件
- 新增驱动抽象层 driver_adapter(go-ora/godror 双驱动切换) - 新增 Update/Delete/Query 回调,完善 RETURNING INTO 子句 - 修复 create.go 事务悬挂行 BUG、批量 RowsAffected、主键 WHERE 注入 - 版本感知体系:分页/自增/序列默认值/BOOLEAN/32k VARCHAR2 分级能力判定 - 11g 序列+触发器自增、ON UPDATE 触发器、默认值智能转换 - 升级 go-ora v2.8.19 → v2.9.0 - 单元测试+模块测试+集成测试共 96 个,真实 Oracle 11g 全绿 - 测试 DSN 密码移除,改为 ORACLE_DSN 环境变量注入
This commit is contained in:
@@ -0,0 +1,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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -0,0 +1,237 @@
|
||||
package oracle
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"database/sql/driver"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
"gorm.io/gorm/schema"
|
||||
)
|
||||
|
||||
// convertValue 将 Go 值转换为 Oracle 兼容格式
|
||||
func convertValue(value interface{}, field *schema.Field) interface{} {
|
||||
if value == nil {
|
||||
return value
|
||||
}
|
||||
|
||||
switch v := value.(type) {
|
||||
case bool:
|
||||
if v {
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
case string:
|
||||
// 检查是否超长字符串,如果需要CLOB处理则保持原样
|
||||
if len(v) > 4000 {
|
||||
// 标记需要CLOB处理,这里只是示例,实际可能需要额外逻辑
|
||||
return v
|
||||
}
|
||||
return v
|
||||
case driver.Valuer:
|
||||
// 调用 Value() 方法解包
|
||||
val, err := v.Value()
|
||||
if err != nil {
|
||||
return value // 如果出错,返回原始值
|
||||
}
|
||||
return val
|
||||
case time.Time:
|
||||
return v
|
||||
default:
|
||||
return value
|
||||
}
|
||||
}
|
||||
|
||||
// convertFromOracleToField 将 Oracle 返回值转换为 Go 类型
|
||||
func convertFromOracleToField(value interface{}, field *schema.Field) interface{} {
|
||||
if value == nil {
|
||||
return value
|
||||
}
|
||||
|
||||
switch v := value.(type) {
|
||||
case sql.NullTime:
|
||||
if v.Valid {
|
||||
return v.Time
|
||||
}
|
||||
return nil
|
||||
case sql.NullInt64:
|
||||
if v.Valid {
|
||||
return v.Int64
|
||||
}
|
||||
return nil
|
||||
case sql.NullFloat64:
|
||||
if v.Valid {
|
||||
return v.Float64
|
||||
}
|
||||
return nil
|
||||
case sql.NullBool:
|
||||
if v.Valid {
|
||||
return v.Bool
|
||||
}
|
||||
return nil
|
||||
case sql.NullString:
|
||||
if v.Valid {
|
||||
return v.String
|
||||
}
|
||||
return nil
|
||||
default:
|
||||
return value
|
||||
}
|
||||
}
|
||||
|
||||
// validateCreateData 验证创建数据
|
||||
func validateCreateData(data interface{}) error {
|
||||
if data == nil {
|
||||
return fmt.Errorf("create data cannot be nil")
|
||||
}
|
||||
|
||||
rv := reflect.ValueOf(data)
|
||||
if rv.Kind() == reflect.Ptr {
|
||||
if rv.IsNil() {
|
||||
return fmt.Errorf("create data pointer cannot be nil")
|
||||
}
|
||||
rv = rv.Elem()
|
||||
}
|
||||
|
||||
if rv.Kind() == reflect.Slice {
|
||||
if rv.Len() == 0 {
|
||||
return fmt.Errorf("create data slice cannot be empty")
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// checkMissingWhereConditions 检查 WHERE 条件是否缺失
|
||||
func checkMissingWhereConditions(conditions []clause.Expression, schema *schema.Schema) bool {
|
||||
if len(conditions) == 0 {
|
||||
return true
|
||||
}
|
||||
count := 0
|
||||
for _, condition := range conditions {
|
||||
// 检查是否是软删除条件 (deleted_at IS NULL)
|
||||
if isSoftDeleteCondition(condition, schema) {
|
||||
count++
|
||||
}
|
||||
}
|
||||
|
||||
// 如果只有软删除条件或没有其他条件,则认为缺少WHERE条件
|
||||
return count >= len(conditions)
|
||||
}
|
||||
|
||||
// isSoftDeleteCondition 检查条件是否为软删除条件(deleted_at IS NULL)
|
||||
func isSoftDeleteCondition(condition clause.Expression, sch *schema.Schema) bool {
|
||||
// 查找是否有 deleted_at 字段
|
||||
var softDeleteField *schema.Field
|
||||
for _, field := range sch.Fields {
|
||||
if strings.EqualFold(field.DBName, "deleted_at") || field.Name == "DeletedAt" {
|
||||
softDeleteField = field
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if softDeleteField == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
// GORM 的软删除条件在 WHERE 中表现为对 deleted_at 列的相等/包含判断(值为空,构建为 IS NULL)
|
||||
switch cond := condition.(type) {
|
||||
case clause.Eq:
|
||||
return strings.EqualFold(columnNameOf(cond.Column), softDeleteField.DBName)
|
||||
case clause.IN:
|
||||
return strings.EqualFold(columnNameOf(cond.Column), softDeleteField.DBName)
|
||||
case clause.Neq:
|
||||
return strings.EqualFold(columnNameOf(cond.Column), softDeleteField.DBName)
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// columnNameOf 从 clause 表达式的 Column 字段中提取列名
|
||||
func columnNameOf(col interface{}) string {
|
||||
switch c := col.(type) {
|
||||
case clause.Column:
|
||||
return c.Name
|
||||
case string:
|
||||
return c
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// addPrimaryKeyWhere 根据模型的主键值注入 WHERE 条件。
|
||||
// GORM 默认的 Update/Delete 回调会在语句中注入主键条件,
|
||||
// 但驱动自定义回调替换了默认实现,因此需要手动补齐。
|
||||
// 返回注入的主键值数量(0 表示没有可用的主键值)。
|
||||
func addPrimaryKeyWhere(stmt *gorm.Statement, sch *schema.Schema) int {
|
||||
if stmt == nil || sch == nil {
|
||||
return 0
|
||||
}
|
||||
|
||||
_, queryValues := schema.GetIdentityFieldValuesMap(stmt.Context, stmt.ReflectValue, sch.PrimaryFields)
|
||||
column, values := schema.ToQueryValues(stmt.Table, sch.PrimaryFieldDBNames, queryValues)
|
||||
if len(values) > 0 {
|
||||
stmt.AddClause(clause.Where{Exprs: []clause.Expression{clause.IN{Column: column, Values: values}}})
|
||||
}
|
||||
return len(values)
|
||||
}
|
||||
|
||||
// buildOracleDefault 智能转换默认值
|
||||
// dbVer 用于版本感知的默认值处理:Oracle 11g 的 DEFAULT 子句不允许引用序列的
|
||||
// NEXTVAL(ORA-00984,12c 才引入该能力),因此 11g 下 NEXTVAL 分支返回空字符串,
|
||||
// 由调用方通过 BEFORE INSERT 触发器实现等价语义(见 migrator.createSequenceDefaultTrigger)。
|
||||
func buildOracleDefault(dbVer string, defaultValue string, field *schema.Field) string {
|
||||
if defaultValue == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
lowerVal := strings.ToLower(strings.TrimSpace(defaultValue))
|
||||
|
||||
switch lowerVal {
|
||||
case "null":
|
||||
return "DEFAULT NULL"
|
||||
case "current_timestamp", "now()":
|
||||
return "DEFAULT CURRENT_TIMESTAMP"
|
||||
case "sysdate":
|
||||
return "DEFAULT SYSDATE"
|
||||
case "true":
|
||||
return "DEFAULT 1"
|
||||
case "false":
|
||||
return "DEFAULT 0"
|
||||
default:
|
||||
// 检查是否为序列
|
||||
if strings.Contains(strings.ToUpper(defaultValue), ".NEXTVAL") {
|
||||
// 12c+ 原生支持 DEFAULT <seq>.NEXTVAL,直接生成 DEFAULT 子句
|
||||
if supportsIdentity(dbVer) {
|
||||
// 去掉可能的包裹括号(GORM 对含括号的默认值保持原文,
|
||||
// 如 "(SEQ_MY.NEXTVAL)"),生成标准的 DEFAULT <seq>.NEXTVAL
|
||||
seqExpr := strings.TrimSpace(defaultValue)
|
||||
seqExpr = strings.TrimPrefix(seqExpr, "(")
|
||||
seqExpr = strings.TrimSuffix(seqExpr, ")")
|
||||
return fmt.Sprintf("DEFAULT %s", strings.TrimSpace(seqExpr))
|
||||
}
|
||||
// 11g 不支持 DEFAULT 子句引用序列(ORA-00984),
|
||||
// 返回空字符串,由建表流程创建 BEFORE INSERT 触发器实现等价语义
|
||||
return ""
|
||||
}
|
||||
|
||||
// 检查日期格式 "2006-01-02"
|
||||
dateRegex := regexp.MustCompile(`^\d{4}-\d{2}-\d{2}$`)
|
||||
if dateRegex.MatchString(defaultValue) {
|
||||
return fmt.Sprintf("DEFAULT TO_DATE('%s', 'YYYY-MM-DD')", defaultValue)
|
||||
}
|
||||
|
||||
// 检查时间戳格式 "2006-01-02 15:04:05"
|
||||
timestampRegex := regexp.MustCompile(`^\d{4}-\d{2}-\d{2} \d{2}:\d{2}:\d{2}$`)
|
||||
if timestampRegex.MatchString(defaultValue) {
|
||||
return fmt.Sprintf("DEFAULT TO_DATE('%s', 'YYYY-MM-DD HH24:MI:SS')", defaultValue)
|
||||
}
|
||||
|
||||
// 普通字符串用单引号包围
|
||||
return fmt.Sprintf("DEFAULT '%s'", defaultValue)
|
||||
}
|
||||
}
|
||||
+254
@@ -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 子句不支持引用序列 NEXTVAL(ORA-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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
// 硬删除:执行 DELETE(Unscoped 时强制硬删除)
|
||||
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)
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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,默认构建不应包含 DriverGodror(godror.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")
|
||||
}
|
||||
}
|
||||
@@ -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{}
|
||||
})
|
||||
}
|
||||
@@ -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{}
|
||||
})
|
||||
}
|
||||
@@ -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 且无错误")
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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 子句不允许引用序列 NEXTVAL(ORA-00984,12c 才支持),
|
||||
// 因此建表后为这类字段创建 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 子句不允许引用序列 NEXTVAL(ORA-00984,12c 才支持),
|
||||
// 因此在建表后通过触发器实现等价语义:插入时列值为 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 默认值会生成非法 SQL(ORA-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
|
||||
})
|
||||
}
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
package oracle
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/migrator"
|
||||
"gorm.io/gorm/schema"
|
||||
)
|
||||
|
||||
// noopDialector 内嵌真实 Dialector 但跳过 Initialize 的数据库连接,
|
||||
// 用于在无需真实连接的情况下构造合法的 *gorm.DB(cacheStore 会被 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
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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 18c(12.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 VARCHAR2(12.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 字节的 VARCHAR2(32k 特性);
|
||||
// 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
@@ -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 连接,跳过")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
@@ -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 子句不允许引用序列的
|
||||
// NEXTVAL(ORA-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 子句不允许引用序列 NEXTVAL(ORA-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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user