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