diff --git a/clauses/helpers_test.go b/clauses/helpers_test.go new file mode 100644 index 0000000..979d5e7 --- /dev/null +++ b/clauses/helpers_test.go @@ -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() +} diff --git a/clauses/merge_test.go b/clauses/merge_test.go new file mode 100644 index 0000000..2110a16 --- /dev/null +++ b/clauses/merge_test.go @@ -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) + } +} diff --git a/clauses/returning_into.go b/clauses/returning_into.go index 0f2857b..83c5ab4 100644 --- a/clauses/returning_into.go +++ b/clauses/returning_into.go @@ -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 +} diff --git a/clauses/returning_into_test.go b/clauses/returning_into_test.go new file mode 100644 index 0000000..8bfcc08 --- /dev/null +++ b/clauses/returning_into_test.go @@ -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) + } +} diff --git a/clauses/when_matched.go b/clauses/when_matched.go index b44cfb7..eefd85e 100644 --- a/clauses/when_matched.go +++ b/clauses/when_matched.go @@ -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()) diff --git a/clauses/when_matched_test.go b/clauses/when_matched_test.go new file mode 100644 index 0000000..56d4a86 --- /dev/null +++ b/clauses/when_matched_test.go @@ -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) + } +} diff --git a/clauses/when_not_matched.go b/clauses/when_not_matched.go index 13b7641..f38b62e 100644 --- a/clauses/when_not_matched.go +++ b/clauses/when_not_matched.go @@ -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()) diff --git a/clauses/when_not_matched_test.go b/clauses/when_not_matched_test.go new file mode 100644 index 0000000..3e0c142 --- /dev/null +++ b/clauses/when_not_matched_test.go @@ -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) +} diff --git a/common.go b/common.go new file mode 100644 index 0000000..bfc274c --- /dev/null +++ b/common.go @@ -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 .NEXTVAL,直接生成 DEFAULT 子句 + if supportsIdentity(dbVer) { + // 去掉可能的包裹括号(GORM 对含括号的默认值保持原文, + // 如 "(SEQ_MY.NEXTVAL)"),生成标准的 DEFAULT .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) + } +} \ No newline at end of file diff --git a/common_test.go b/common_test.go new file mode 100644 index 0000000..5010271 --- /dev/null +++ b/common_test.go @@ -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 .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) + } + }) + } +} diff --git a/create.go b/create.go index 6025e4d..7b082a8 100644 --- a/create.go +++ b/create.go @@ -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 } } } diff --git a/delete.go b/delete.go new file mode 100644 index 0000000..7f42bf9 --- /dev/null +++ b/delete.go @@ -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) + } + } + }, + ) + } + } +} \ No newline at end of file diff --git a/driver_adapter/adapter.go b/driver_adapter/adapter.go new file mode 100644 index 0000000..0b36455 --- /dev/null +++ b/driver_adapter/adapter.go @@ -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 +} diff --git a/driver_adapter/adapter_test.go b/driver_adapter/adapter_test.go new file mode 100644 index 0000000..6b898f7 --- /dev/null +++ b/driver_adapter/adapter_test.go @@ -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") + } +} diff --git a/driver_adapter/godror.go b/driver_adapter/godror.go new file mode 100644 index 0000000..90ab8a6 --- /dev/null +++ b/driver_adapter/godror.go @@ -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{} + }) +} \ No newline at end of file diff --git a/driver_adapter/goora.go b/driver_adapter/goora.go new file mode 100644 index 0000000..9a3037d --- /dev/null +++ b/driver_adapter/goora.go @@ -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{} + }) +} \ No newline at end of file diff --git a/driver_adapter/goora_test.go b/driver_adapter/goora_test.go new file mode 100644 index 0000000..623ea5c --- /dev/null +++ b/driver_adapter/goora_test.go @@ -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 且无错误") + } +} diff --git a/go.mod b/go.mod index a17110b..7b430f0 100644 --- a/go.mod +++ b/go.mod @@ -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 ) diff --git a/migrator.go b/migrator.go index 93b4c1d..3aa3598 100644 --- a/migrator.go +++ b/migrator.go @@ -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 .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__,避免与 autoIncrement 的 TRG_
冲突。 +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 + }) +} diff --git a/migrator_test.go b/migrator_test.go new file mode 100644 index 0000000..b564a37 --- /dev/null +++ b/migrator_test.go @@ -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) diff --git a/namer_test.go b/namer_test.go new file mode 100644 index 0000000..9fc9a35 --- /dev/null +++ b/namer_test.go @@ -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") + } +} diff --git a/oracle.go b/oracle.go index f5a82ab..fe5f561 100644 --- a/oracle.go +++ b/oracle.go @@ -19,17 +19,65 @@ 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 - Conn gorm.ConnPool //*sql.DB - DefaultStringSize uint - DBName string - DBVer string + DriverName string + DSN string + Conn gorm.ConnPool //*sql.DB + 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: - sqlType = "TIMESTAMP WITH TIME ZONE" + 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) +} diff --git a/oracle_test.go b/oracle_test.go new file mode 100644 index 0000000..e6c0a60 --- /dev/null +++ b/oracle_test.go @@ -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 连接,跳过") +} diff --git a/query.go b/query.go new file mode 100644 index 0000000..5a9507e --- /dev/null +++ b/query.go @@ -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 +} \ No newline at end of file diff --git a/tests/create_test.go b/tests/create_test.go new file mode 100644 index 0000000..58c14ec --- /dev/null +++ b/tests/create_test.go @@ -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") + } +} diff --git a/tests/delete_test.go b/tests/delete_test.go new file mode 100644 index 0000000..7be8dce --- /dev/null +++ b/tests/delete_test.go @@ -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") + } +} diff --git a/tests/hook_test.go b/tests/hook_test.go new file mode 100644 index 0000000..03ecaf9 --- /dev/null +++ b/tests/hook_test.go @@ -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") + } +} diff --git a/tests/main_test.go b/tests/main_test.go new file mode 100644 index 0000000..e1db357 --- /dev/null +++ b/tests/main_test.go @@ -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) + } +} diff --git a/tests/migrate_test.go b/tests/migrate_test.go new file mode 100644 index 0000000..009d792 --- /dev/null +++ b/tests/migrate_test.go @@ -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)) + } +} diff --git a/tests/models.go b/tests/models.go new file mode 100644 index 0000000..aca0161 --- /dev/null +++ b/tests/models.go @@ -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" +} diff --git a/tests/query_test.go b/tests/query_test.go new file mode 100644 index 0000000..77fc83c --- /dev/null +++ b/tests/query_test.go @@ -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)) + } +} diff --git a/tests/sequence_test.go b/tests/sequence_test.go new file mode 100644 index 0000000..befb5da --- /dev/null +++ b/tests/sequence_test.go @@ -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 .NEXTVAL 子句) + if err := DB.AutoMigrate(&SeqDefaultViaDriverModel{}); err != nil { + t.Fatalf("failed to migrate: %v", err) + } + + // 验证序列默认值触发器已创建(命名 SEQDEF_TRG_
_, + // 避免与 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) +} diff --git a/tests/soft_delete_test.go b/tests/soft_delete_test.go new file mode 100644 index 0000000..e98e338 --- /dev/null +++ b/tests/soft_delete_test.go @@ -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) +} diff --git a/tests/update_test.go b/tests/update_test.go new file mode 100644 index 0000000..b9d1671 --- /dev/null +++ b/tests/update_test.go @@ -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) +} diff --git a/update.go b/update.go new file mode 100644 index 0000000..e4df566 --- /dev/null +++ b/update.go @@ -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) + } + } + }, + ) + } + } +} \ No newline at end of file