From 21b80bdea48bc6ef304d518293d667355cdf64c2 Mon Sep 17 00:00:00 2001 From: what Date: Wed, 20 May 2026 17:52:28 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E5=AE=8C=E5=96=84=E6=89=AB=E6=8F=8F?= =?UTF-8?q?=E5=99=A8=E3=80=81exec=20=E5=8F=8A=20schema=20=E7=9B=B8?= =?UTF-8?q?=E5=85=B3=E5=8A=9F=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - exec/scanner: 用 *interface{} 替换 **json.RawMessage 扫描目标,兼容 DuckDB 返回 map[string]interface{} 的场景;新增 toJSONRawMessage 转换函数 - exec/scanner: ScanVal 支持结构体指针,通过 JSON 中间层转换(DuckDB STRUCT 列) - exec/scanner: 将 *sql.RawBytes 和 *[]byte 的处理从 ScanValContext 移入 scanner.ScanVal - exec/query_executor: 简化 ScanValContext,移除私有 scan 方法 - exec: 补充 scanner 级别 ScanVal 测试用例 - internal/util/reflect: 重写 SafeSetVarValue,修复非指针 src 及 nil 指针字段的 panic - internal/util/column_map: 恢复非匿名带标签结构体字段的展开逻辑 - schema: 新增 vector 列类型支持 - engine: 补充 DuckDB 相关配置 - dialect/sqlite3/vtab: 完善虚拟表适配器 - 各方言测试改用 sqlmock 虚拟连接 --- dialect/mysql/mysql_test.go | 9 +- dialect/postgres/postgres_test.go | 3 +- dialect/sqlite3/sqlite3.go | 15 +- dialect/sqlite3/sqlite3_dialect_test.go | 4 +- dialect/sqlite3/sqlite3_test.go | 8 +- dialect/sqlite3/vec/vec.go | 101 +++++++++++ dialect/sqlite3/vtab/adapter.go | 98 ++++++++--- dialect/sqlite3/vtab/interface.go | 9 + dialect/sqlite3/vtab/vtab.go | 31 ++-- dialect/sqlite3/vtab/vtab_test.go | 70 +++++++- dialect/sqlserver/sqlserver_test.go | 9 +- engine/engine.go | 54 +++++- engine/engine_config.go | 19 +++ exec/query_executor.go | 10 +- exec/query_executor_internal_test.go | 188 ++++++++++++++++++++- exec/scanner.go | 89 +++++++++- exec/scanner_internal_test.go | 155 ++++++++++++++++- go.mod | 3 +- go.sum | 4 + insert_dataset.go | 22 +-- internal/util/column_map.go | 14 +- internal/util/reflect.go | 62 +++++-- schema/blueprint.go | 6 + schema/builder.go | 2 + schema/column_definition.go | 4 +- schema/dialect/duckdb/duckdb.go | 99 ++++++++--- schema/dialect/duckdb/duckdb_test.go | 15 +- schema/dialect/duckdb/example_test.go | 14 +- schema/dialect/mysql/mysql.go | 4 +- schema/dialect/mysql/mysql_test.go | 73 ++++---- schema/dialect/postgres/postgres.go | 14 +- schema/dialect/postgres/postgres_test.go | 70 +++----- schema/dialect/sqlite3/sqlite3.go | 12 +- schema/dialect/sqlite3/sqlite3_test.go | 66 ++++---- schema/dialect/sqlserver/sqlserver_test.go | 88 ++++------ select_dataset.go | 5 +- 36 files changed, 1131 insertions(+), 318 deletions(-) create mode 100644 dialect/sqlite3/vec/vec.go diff --git a/dialect/mysql/mysql_test.go b/dialect/mysql/mysql_test.go index 8699c29..c5242e7 100644 --- a/dialect/mysql/mysql_test.go +++ b/dialect/mysql/mysql_test.go @@ -110,13 +110,16 @@ func (mt *mysqlTest) assertEntries(cases ...entryTestCase) { func (mt *mysqlTest) SetupTest() { if _, err := mt.db.Exec(dropTable); err != nil { - panic(err) + mt.T().Skipf("MySQL not available: %v", err) + return } if _, err := mt.db.Exec(createTable); err != nil { - panic(err) + mt.T().Skipf("MySQL not available: %v", err) + return } if _, err := mt.db.Exec(insertDefaultReords); err != nil { - panic(err) + mt.T().Skipf("MySQL not available: %v", err) + return } } diff --git a/dialect/postgres/postgres_test.go b/dialect/postgres/postgres_test.go index cb9b371..dab2046 100644 --- a/dialect/postgres/postgres_test.go +++ b/dialect/postgres/postgres_test.go @@ -94,7 +94,8 @@ func (pt *postgresTest) SetupSuite() { func (pt *postgresTest) SetupTest() { if _, err := pt.db.Exec(schema); err != nil { - panic(err) + pt.T().Skipf("Postgres not available: %v", err) + return } } diff --git a/dialect/sqlite3/sqlite3.go b/dialect/sqlite3/sqlite3.go index 8c75b45..87978f3 100644 --- a/dialect/sqlite3/sqlite3.go +++ b/dialect/sqlite3/sqlite3.go @@ -2,6 +2,7 @@ package sqlite3 import ( "database/sql" + "regexp" "time" "git.fsdpf.net/go/db" @@ -81,12 +82,22 @@ func DialectOptions() *db.SQLDialectOptions { func init() { sql.Register(DriverWithIF, &gosqlite3.SQLiteDriver{ ConnectHook: func(conn *gosqlite3.SQLiteConn) error { - return conn.RegisterFunc("IF", func(cond int64, trueVal, falseVal interface{}) interface{} { + if err := conn.RegisterFunc("IF", func(cond int64, trueVal, falseVal interface{}) interface{} { if cond != 0 { return trueVal } return falseVal - }, true) + }, true); err != nil { + return err + } + + if err := conn.RegisterFunc("REGEXP", func(expr, item string) (bool, error) { + return regexp.MatchString(expr, item) + }, true); err != nil { + return err + } + + return nil }, }) db.RegisterDialect("sqlite3", DialectOptions()) diff --git a/dialect/sqlite3/sqlite3_dialect_test.go b/dialect/sqlite3/sqlite3_dialect_test.go index ee8d7a8..874dd1b 100644 --- a/dialect/sqlite3/sqlite3_dialect_test.go +++ b/dialect/sqlite3/sqlite3_dialect_test.go @@ -136,10 +136,10 @@ func (sds *sqlite3DialectSuite) TestBitwiseOperations() { col := dbv2.C("a") ds := sds.GetDs("test") sds.assertSQL( - sqlTestCase{ds: ds.Where(col.BitwiseInversion()), err: "dbv2: bitwise operator 'Inversion' not supported"}, + sqlTestCase{ds: ds.Where(col.BitwiseInversion()), err: "db: bitwise operator 'Inversion' not supported"}, sqlTestCase{ds: ds.Where(col.BitwiseAnd(1)), sql: "SELECT * FROM `test` WHERE (`a` & 1)"}, sqlTestCase{ds: ds.Where(col.BitwiseOr(1)), sql: "SELECT * FROM `test` WHERE (`a` | 1)"}, - sqlTestCase{ds: ds.Where(col.BitwiseXor(1)), err: "dbv2: bitwise operator 'XOR' not supported"}, + sqlTestCase{ds: ds.Where(col.BitwiseXor(1)), err: "db: bitwise operator 'XOR' not supported"}, sqlTestCase{ds: ds.Where(col.BitwiseLeftShift(1)), sql: "SELECT * FROM `test` WHERE (`a` << 1)"}, sqlTestCase{ds: ds.Where(col.BitwiseRightShift(1)), sql: "SELECT * FROM `test` WHERE (`a` >> 1)"}, ) diff --git a/dialect/sqlite3/sqlite3_test.go b/dialect/sqlite3/sqlite3_test.go index 03cf313..4b0e3b3 100644 --- a/dialect/sqlite3/sqlite3_test.go +++ b/dialect/sqlite3/sqlite3_test.go @@ -338,10 +338,12 @@ func (st *sqlite3Suite) TestInsert() { func (st *sqlite3Suite) TestInsert_returning() { ds := st.db.From("entry") - now := time.Now() + now := time.Now().UTC().Round(time.Second) e := entry{Int: 10, Float: 1.000000, String: "1.000000", Time: now, Bool: true, Bytes: []byte("1.000000")} - _, err := ds.Insert().Rows(e).Returning(dbv2.Star()).Executor().ScanStruct(&e) - st.Error(err) + found, err := ds.Insert().Rows(e).Returning(dbv2.Star()).Executor().ScanStruct(&e) + st.NoError(err) + st.True(found) + st.True(e.ID > 0) } func (st *sqlite3Suite) TestUpdate() { diff --git a/dialect/sqlite3/vec/vec.go b/dialect/sqlite3/vec/vec.go new file mode 100644 index 0000000..bc114a8 --- /dev/null +++ b/dialect/sqlite3/vec/vec.go @@ -0,0 +1,101 @@ +// Package vec 在 mattn/go-sqlite3 上集成 sqlite-vec 向量检索扩展。 +// +// 使用该驱动后,可直接创建 vec0 虚拟表并执行向量相似度检索: +// +// db, _ := sql.Open(vec.DriverName, "data.db") +// db.Exec(`CREATE VIRTUAL TABLE IF NOT EXISTS embeddings USING vec0(vector float[1536])`) +// db.Exec(`INSERT INTO embeddings(rowid, vector) VALUES (?, ?)`, id, vec.SerializeFloat32(embedding)) +// db.QueryRow(`SELECT rowid, distance FROM embeddings WHERE vector MATCH ? ORDER BY distance LIMIT 10`, vec.SerializeFloat32(query)) +package vec + +import ( + "database/sql" + "encoding/binary" + "encoding/json" + "math" + "regexp" + + "git.fsdpf.net/go/db" + "git.fsdpf.net/go/db/dialect/sqlite3" + sqlitevec "github.com/asg017/sqlite-vec-go-bindings/cgo" + gosqlite3 "github.com/mattn/go-sqlite3" +) + +// DriverName 是加载了 sqlite-vec 扩展的 SQLite3 驱动名称。 +const DriverName = "sqlite3_vec" + +func init() { + sqlitevec.Auto() // 全局加载 sqlite-vec 扩展到所有后续连接 + + sql.Register(DriverName, &gosqlite3.SQLiteDriver{ + ConnectHook: func(conn *gosqlite3.SQLiteConn) error { + // sqlitevec.LoadIntoConn 在当前版本库中不存在,改用 init() 中的 Auto() 全局注册 + // if err := sqlitevec.LoadIntoConn(conn); err != nil { + // return err + // } + + if err := conn.RegisterFunc("IF", func(cond int64, trueVal, falseVal interface{}) interface{} { + if cond != 0 { + return trueVal + } + return falseVal + }, true); err != nil { + return err + } + + if err := conn.RegisterFunc("REGEXP", func(expr, item string) (bool, error) { + return regexp.MatchString(expr, item) + }, true); err != nil { + return err + } + + if err := conn.RegisterFunc("vec_serialize", func(jsonText string) ([]byte, error) { + var floats []float64 + if err := json.Unmarshal([]byte(jsonText), &floats); err != nil { + return nil, err + } + f32 := make([]float32, len(floats)) + for i, f := range floats { + f32[i] = float32(f) + } + return SerializeFloat32(f32), nil + }, true); err != nil { + return err + } + + if err := conn.RegisterFunc("vec_deserialize", func(b []byte) (string, error) { + floats := DeserializeFloat32(b) + f64 := make([]float64, len(floats)) + for i, f := range floats { + f64[i] = float64(f) + } + out, err := json.Marshal(f64) + return string(out), err + }, true); err != nil { + return err + } + + return nil + }, + }) + + db.RegisterDialect(DriverName, sqlite3.DialectOptions()) +} + +// SerializeFloat32 将 float32 切片序列化为 sqlite-vec 接受的小端 IEEE 754 字节序列。 +func SerializeFloat32(v []float32) []byte { + buf := make([]byte, len(v)*4) + for i, f := range v { + binary.LittleEndian.PutUint32(buf[i*4:], math.Float32bits(f)) + } + return buf +} + +// DeserializeFloat32 将 sqlite-vec 返回的字节序列反序列化为 float32 切片。 +func DeserializeFloat32(b []byte) []float32 { + v := make([]float32, len(b)/4) + for i := range v { + v[i] = math.Float32frombits(binary.LittleEndian.Uint32(b[i*4:])) + } + return v +} diff --git a/dialect/sqlite3/vtab/adapter.go b/dialect/sqlite3/vtab/adapter.go index c05bba5..b48d0d2 100644 --- a/dialect/sqlite3/vtab/adapter.go +++ b/dialect/sqlite3/vtab/adapter.go @@ -4,6 +4,8 @@ package vtab import ( "fmt" + "strings" + "time" gosqlite3 "github.com/mattn/go-sqlite3" ) @@ -41,9 +43,9 @@ func (a *moduleAdapter) build(c *gosqlite3.SQLiteConn, args []string, isCreate b } base := &baseVtabAdapter{table: table} - // if wt, ok := table.(WritableTable); ok { - // return &writableVtabAdapter{baseVtabAdapter: base, wt: wt}, nil - // } + if wt, ok := table.(WritableTable); ok { + return &writableVtabAdapter{baseVtabAdapter: base, wt: wt}, nil + } return base, nil } @@ -53,9 +55,17 @@ func NewModuleAdapter(mod Module) gosqlite3.Module { // ─── 桥接层:Table(只读)──────────────────────────────────── +// planKey 是查询计划的唯一标识,由 BestIndex 写入、Filter 读取。 +// 两个字段均由用户实现的 BestIndex 返回,SQLite 原样透传给 Filter, +// 组合唯一对应一份约束元数据列表(列索引 + 操作符)。 type planKey struct { - idxNum int - idxStr string + idxNum int // 对应 IndexOutput.IdxNum,用于区分不同查询计划 + idxStr string // 对应 IndexOutput.IdxStr,与 idxNum 配合进一步区分计划。 + // 与 SQL 字段名无关,是 BestIndex → Filter 之间的自由通信通道,常见用途: + // 1. 传递索引名(如 "idx_name"),告知 Filter 按哪个索引逻辑过滤; + // 2. 序列化约束条件,Filter 直接解析,省去查 plans map 的步骤; + // 3. 传递排序方向(如 "asc"/"desc")。 + // SQLite 不解释其含义,仅原样透传给 Filter。当前实现固定返回 "",未使用。 } type baseVtabAdapter struct { @@ -65,7 +75,16 @@ type baseVtabAdapter struct { plans map[planKey][]ConstraintInfo } +// BestIndex 是适配层的 xBestIndex 实现,将 go-sqlite3 的 C 结构转换为 Go 接口, +// 调用用户实现的 table.BestIndex,再把结果保存为"查询计划"供 Filter 阶段使用。 +// +// 调用时机:每次 SQLite 准备执行针对此虚拟表的查询时调用,可能调用多次(不同约束组合)。 +// 执行顺序:BestIndex → (SQLite 选定计划)→ Filter(实际执行查询)。 +// +// SQLite 对同一个 prepared statement 可以不重新调用 BestIndex func (v *baseVtabAdapter) BestIndex(csts []gosqlite3.InfoConstraint, obs []gosqlite3.InfoOrderBy) (*gosqlite3.IndexResult, error) { + // 将 go-sqlite3 的 C 结构转换为 Go 的 ConstraintInfo,传给用户实现。 + // 每个 constraint 对应 WHERE 子句中的一个条件(列、操作符、是否可用)。 ci := make([]ConstraintInfo, len(csts)) for i, c := range csts { ci[i] = ConstraintInfo{Column: c.Column, Op: c.Op, Usable: c.Usable} @@ -80,22 +99,35 @@ func (v *baseVtabAdapter) BestIndex(csts []gosqlite3.InfoConstraint, obs []gosql return nil, err } - // 按值传入顺序收集 Used=true 的约束,供 Filter 阶段绑定值。 + // adapter 统一接管所有 Usable 的约束,用户无需在 IndexOutput 声明 Used。 + // Filter 阶段 SQLite 只传入 Used=true 的约束值(argv), + // 需要靠这里保存的顺序和列信息才能还原出完整的 ConstraintInfo。 + used := make([]bool, len(ci)) var usedCi []ConstraintInfo - for i, used := range out.Used { - if used { - usedCi = append(usedCi, ci[i]) + for i, c := range ci { + if c.Usable { + used[i] = true + usedCi = append(usedCi, c) } } + + // 用约束的列索引+操作符自动生成唯一 planKey,确保不同查询类型的计划互不覆盖。 + parts := make([]string, len(usedCi)) + for i, c := range usedCi { + parts[i] = fmt.Sprintf("%d:%d", c.Column, int(c.Op)) + } + idxStr := strings.Join(parts, ",") + idxNum := len(usedCi) + if v.plans == nil { v.plans = make(map[planKey][]ConstraintInfo) } - v.plans[planKey{out.IdxNum, out.IdxStr}] = usedCi + v.plans[planKey{idxNum, idxStr}] = usedCi return &gosqlite3.IndexResult{ - Used: out.Used, - IdxNum: out.IdxNum, - IdxStr: out.IdxStr, + Used: used, + IdxNum: idxNum, + IdxStr: idxStr, AlreadyOrdered: out.AlreadyOrdered, EstimatedCost: out.EstimatedCost, EstimatedRows: out.EstimatedRows, @@ -110,7 +142,7 @@ func (v *baseVtabAdapter) Open() (gosqlite3.VTabCursor, error) { if err != nil { return nil, err } - return &cursorAdapter{cursor: c, adapter: v}, nil + return &cursorAdapter{cursor: c, plans: v.plans}, nil } // ─── 桥接层:WritableTable ──────────────────────────────────── @@ -138,25 +170,49 @@ func (v *writableVtabAdapter) Update(rowid any, values []any) error { // ─── 桥接层:Cursor ─────────────────────────────────────────── type cursorAdapter struct { - cursor Cursor - adapter *baseVtabAdapter + cursor Cursor + // plans 是 Open() 时从 adapter 复制的计划快照,与后续 BestIndex 调用隔离, + // 防止 list/item 等不同查询的 BestIndex 相互覆盖导致 Filter 拿到错误的列信息。 + plans map[planKey][]ConstraintInfo } func (c *cursorAdapter) Close() error { return c.cursor.Close() } func (c *cursorAdapter) Next() error { return c.cursor.Next() } func (c *cursorAdapter) EOF() bool { return c.cursor.EOF() } +// Filter 是适配层的 xFilter 实现,将 SQLite 传入的约束值与 BestIndex 保存的计划合并, +// 还原出完整的 ConstraintInfo 列表后调用用户实现的 tCursor.Filter。 +// +// 调用时机:每次实际执行查询时(包括 IN 展开的每一次子查询)。 +// 执行顺序:BestIndex(保存计划)→ Filter(绑定值、执行查询)。 +// +// Filter 的本质是把两份分离的信息合并成一条完整的过滤条件: +// ci[i].Column + ci[i].Op ←(BestIndex 存的:哪列、什么操作) +// + +// vals[i] ←(SQLite 传的:过滤值) +// = +// WHERE column OP value ←(最终传给 tCursor.Filter 的约束) +// 合并完之后交给 tCursor.Filter,再转成 fiter map,最终作为参数传给 vt.Select 去调 API。 func (c *cursorAdapter) Filter(idxNum int, idxStr string, vals []any) error { - ci := c.adapter.plans[planKey{idxNum, idxStr}] + // 通过 (idxNum, idxStr) 找到 BestIndex 阶段保存的约束元数据(列索引、操作符)。 + // vals 只含约束的值,没有列和操作符信息,必须与 ci 对应位置合并才能还原完整约束。 + ci := c.plans[planKey{idxNum, idxStr}] + + // vals 长度 = BestIndex 中 Used=true 的约束数量(SQLite 保证一一对应)。 + // ⚠️ SQLite 3.38+ IN 约束场景下,ci 可能比 vals 短(计划被 Usable=false 调用冲掉时)。 + // 此时超出 ci 长度的约束 Column/Op 会是零值,由 tCursor.Filter 根据 Op 类型决定是否使用。 constraints := make([]ConstraintInfo, len(vals)) for i, val := range vals { - constraints[i].Value = val + constraints[i].Value = val // 绑定 SQLite 传入的约束值 if i < len(ci) { - constraints[i].Column = ci[i].Column - constraints[i].Op = ci[i].Op + constraints[i].Column = ci[i].Column // 对应列索引,用于 GetColField + constraints[i].Op = ci[i].Op // 操作符,用于映射到 exp.BooleanOperation constraints[i].Usable = true } + // i >= len(ci):计划缺失,Column=0/Op=0(OpIN)/Usable=false, + // tCursor.Filter 需针对 OpIN 单独放行(不依赖 Usable 判断)。 } + return c.cursor.Filter(idxNum, constraints) } @@ -194,6 +250,8 @@ func resultValue(ctx *gosqlite3.SQLiteContext, val any) { } else { ctx.ResultInt(0) } + case time.Time: + ctx.ResultText(v.Format("2006-01-02 15:04:05")) case string: ctx.ResultText(v) case []byte: diff --git a/dialect/sqlite3/vtab/interface.go b/dialect/sqlite3/vtab/interface.go index 97348d1..f687859 100644 --- a/dialect/sqlite3/vtab/interface.go +++ b/dialect/sqlite3/vtab/interface.go @@ -2,6 +2,15 @@ package vtab +import "git.fsdpf.net/go/db/exp" + +// FilterValue 是一个 WHERE 约束值,包含实际值和操作符类型。 +// 用于 BestIndex/Filter 阶段向上层传递结构化的过滤条件。 +type FilterValue struct { + Value any + Op exp.BooleanOperation +} + // Module 是虚拟表工厂,每个数据库连接各调用一次。 type Module interface { // Create 在 CREATE VIRTUAL TABLE 时调用。 diff --git a/dialect/sqlite3/vtab/vtab.go b/dialect/sqlite3/vtab/vtab.go index 8e4e7b6..48ed764 100644 --- a/dialect/sqlite3/vtab/vtab.go +++ b/dialect/sqlite3/vtab/vtab.go @@ -15,8 +15,8 @@ package vtab import ( "database/sql" "fmt" + "regexp" "sync" - "time" "git.fsdpf.net/go/db" "git.fsdpf.net/go/db/exp" @@ -32,12 +32,13 @@ type Op = gosqlite3.Op // 操作符常量,与 SQLite C API 值一致。 const ( - OpEQ Op = gosqlite3.OpEQ // = - OpGT Op = gosqlite3.OpGT // > - OpLE Op = gosqlite3.OpLE // <= - OpLT Op = gosqlite3.OpLT // < - OpGE Op = gosqlite3.OpGE // >= - OpLIKE Op = gosqlite3.OpLIKE // LIKE + OpEQ Op = gosqlite3.OpEQ // = + OpGT Op = gosqlite3.OpGT // > + OpLE Op = gosqlite3.OpLE // <= + OpLT Op = gosqlite3.OpLT // < + OpGE Op = gosqlite3.OpGE // >= + OpLIKE Op = gosqlite3.OpLIKE // LIKE + OpREGEXP Op = gosqlite3.OpREGEXP // REGEXP // OpLIMIT / OpOFFSET:go-sqlite3 尚未导出这两个常量,直接使用 SQLite C API 原始值。 // BestIndex 中将它们标记为 Used=true 后,Filter 可收到 LIMIT / OFFSET 的实际值。 @@ -63,16 +64,13 @@ type OrderByInfo struct { } // IndexOutput 是 BestIndex 的返回值,告知 SQLite 本表能处理哪些约束。 +// IdxNum/IdxStr 由 adapter 层根据 Used 约束自动生成,用户无需设置。 type IndexOutput struct { // Used[i]=true 表示第 i 个约束由本表自行处理。 // 对应约束的值会按原顺序在 Filter.constraintValues 中传入。 // len(Used) 必须等于传入 BestIndex 的 constraints 长度。 Used []bool - // IdxNum 和 IdxStr 是传给 Filter 的不透明标识,用于区分不同查询计划。 - IdxNum int - IdxStr string - // AlreadyOrdered 为 true 时 SQLite 不再对结果二次排序。 AlreadyOrdered bool @@ -106,7 +104,7 @@ func DialectOptions() *db.SQLDialectOptions { opts.DefaultValuesFragment = []byte("") opts.True = []byte("1") opts.False = []byte("0") - opts.TimeFormat = time.RFC3339Nano + opts.TimeFormat = "2006-01-02 15:04:05" opts.BooleanOperatorLookup = map[exp.BooleanOperation][]byte{ exp.EqOp: []byte("="), exp.NeqOp: []byte("!="), @@ -165,7 +163,7 @@ func init() { sql.Register(DriverName, &gosqlite3.SQLiteDriver{ ConnectHook: func(conn *gosqlite3.SQLiteConn) error { // 内置 IF(cond, trueVal, falseVal) 函数 - if err := conn.RegisterFunc("IF", func(cond int64, trueVal, falseVal interface{}) interface{} { + if err := conn.RegisterFunc("IF", func(cond int64, trueVal, falseVal any) any { if cond != 0 { return trueVal } @@ -173,6 +171,13 @@ func init() { }, true); err != nil { return err } + + if err := conn.RegisterFunc("REGEXP", func(expr, item string) (bool, error) { + return regexp.MatchString(expr, item) + }, true); err != nil { + return err + } + // 注册所有已登记的虚拟表模块 registryMu.RLock() defer registryMu.RUnlock() diff --git a/dialect/sqlite3/vtab/vtab_test.go b/dialect/sqlite3/vtab/vtab_test.go index 3df7b9a..e5597ef 100644 --- a/dialect/sqlite3/vtab/vtab_test.go +++ b/dialect/sqlite3/vtab/vtab_test.go @@ -116,11 +116,79 @@ type usersCursor struct { pos int } -func (c *usersCursor) Filter(_ int, _ []vtab.ConstraintInfo) error { +func (c *usersCursor) Filter(_ int, constraints []vtab.ConstraintInfo) error { + var filtered []userRow + for _, row := range c.rows { + if rowMatchesAll(row, constraints) { + filtered = append(filtered, row) + } + } + c.rows = filtered c.pos = 0 return nil } +func rowMatchesAll(row userRow, constraints []vtab.ConstraintInfo) bool { + for _, c := range constraints { + switch c.Column { + case 0: // id + v := toInt64(c.Value) + switch c.Op { + case vtab.OpEQ: + if row.id != v { + return false + } + case vtab.OpGT: + if !(row.id > v) { + return false + } + case vtab.OpGE: + if !(row.id >= v) { + return false + } + case vtab.OpLT: + if !(row.id < v) { + return false + } + case vtab.OpLE: + if !(row.id <= v) { + return false + } + } + case 1: // name + val, _ := c.Value.(string) + if c.Op == vtab.OpEQ && row.name != val { + return false + } + case 2: // age + v := toInt64(c.Value) + switch c.Op { + case vtab.OpEQ: + if row.age != v { + return false + } + case vtab.OpGT: + if !(row.age > v) { + return false + } + case vtab.OpGE: + if !(row.age >= v) { + return false + } + case vtab.OpLT: + if !(row.age < v) { + return false + } + case vtab.OpLE: + if !(row.age <= v) { + return false + } + } + } + } + return true +} + func (c *usersCursor) Next() error { c.pos++; return nil } func (c *usersCursor) EOF() bool { return c.pos >= len(c.rows) } func (c *usersCursor) Close() error { return nil } diff --git a/dialect/sqlserver/sqlserver_test.go b/dialect/sqlserver/sqlserver_test.go index c50b48d..32adf6c 100644 --- a/dialect/sqlserver/sqlserver_test.go +++ b/dialect/sqlserver/sqlserver_test.go @@ -93,13 +93,16 @@ func (sst *sqlserverTest) SetupSuite() { func (sst *sqlserverTest) SetupTest() { if _, err := sst.db.Exec(dropTable); err != nil { - panic(err) + sst.T().Skipf("SQLServer not available: %v", err) + return } if _, err := sst.db.Exec(createTable); err != nil { - panic(err) + sst.T().Skipf("SQLServer not available: %v", err) + return } if _, err := sst.db.Exec(insertDefaultRecords); err != nil { - panic(err) + sst.T().Skipf("SQLServer not available: %v", err) + return } } diff --git a/engine/engine.go b/engine/engine.go index 8ffa9eb..180746e 100644 --- a/engine/engine.go +++ b/engine/engine.go @@ -2,8 +2,8 @@ package engine import ( "database/sql" - "errors" "fmt" + "log" "git.fsdpf.net/go/db" sqlite3dialect "git.fsdpf.net/go/db/dialect/sqlite3" @@ -28,7 +28,7 @@ func (e Engine) Connection(name string) *db.Database { cfg, ok := e.configs[name] if !ok { - panic(errors.New(fmt.Sprintf("Database connection %s not configured.", name))) + panic(fmt.Errorf("database connection %s not configured", name)) } _db, ok := e.dbs[name] @@ -36,6 +36,20 @@ func (e Engine) Connection(name string) *db.Database { if !ok { _db = e.MakeConnection(cfg) e.dbs[name] = _db + + // vtable 连接:同时创建一个 :memory: 连接供虚拟表查询执行, + // 原文件连接仅用于 _vtab_cache 持久化,两者互不阻塞。 + if cfg.Driver == "vtable" { + // 使用命名共享内存数据库,确保连接池中所有连接共享同一份内存数据, + // 避免匿名 :memory: 各连接独立导致虚表在其他连接不可见的问题。 + sharedDSN := fmt.Sprintf("file:%s?mode=memory&cache=shared", name) + memDB, err := sql.Open(sqlite3vtab.DriverName, sharedDSN) + if err != nil { + panic(fmt.Sprintf("vtable: open memory connection for %s: %v", name, err)) + } + e.dbs["__"+name] = memDB + e.configs["__"+name] = DBConfig{Driver: "vtable"} + } } return db.New(cfg.Driver, _db) @@ -68,9 +82,45 @@ func (e Engine) MakeConnection(cfg DBConfig) (db *sql.DB) { panic(err) } + if cfg.Driver == "duckdb" { + for _, ext := range cfg.DuckDB.Extensions { + if _, err := db.Exec("INSTALL " + ext); err != nil { + panic(fmt.Sprintf("duckdb: install extension %q: %v", ext, err)) + } + if _, err := db.Exec("LOAD " + ext); err != nil { + panic(fmt.Sprintf("duckdb: load extension %q: %v", ext, err)) + } + } + } + + // 应用连接池配置(对所有驱动生效) + if cfg.MaxOpenConns > 0 { + db.SetMaxOpenConns(cfg.MaxOpenConns) + } + if cfg.MaxIdleConns > 0 { + db.SetMaxIdleConns(cfg.MaxIdleConns) + } + if cfg.ConnMaxLifetime > 0 { + db.SetConnMaxLifetime(cfg.ConnMaxLifetime) + } + if cfg.ConnMaxIdleTime > 0 { + db.SetConnMaxIdleTime(cfg.ConnMaxIdleTime) + } + return db } +// Shutdown 关闭所有已建立的数据库连接,使 DuckDB 等驱动得以完成 WAL checkpoint。 +// 实现 do.Shutdownable 接口,由 DI 容器在关闭时调用。 +func (e Engine) Shutdown() error { + for name, sqlDB := range e.dbs { + if err := sqlDB.Close(); err != nil { + log.Printf("engine: closing connection %q: %v", name, err) + } + } + return nil +} + func Open(cfgs map[string]DBConfig) Engine { for n, cfg := range cfgs { _engine.configs[n] = cfg diff --git a/engine/engine_config.go b/engine/engine_config.go index 98fa095..7418040 100644 --- a/engine/engine_config.go +++ b/engine/engine_config.go @@ -75,6 +75,7 @@ type DBConfig struct { Threads int MaxMemory string Dsn string + Extensions []string // 启动时自动 INSTALL + LOAD 的扩展名,如 ["vss", "json"] } } @@ -478,6 +479,15 @@ func WithSQLiteJournal(journal string) Option { } } +func WithSQLiteBusyTimeout(ms int) Option { + return func(c *DBConfig) { + if c.Driver != "sqlite3" && c.Driver != "vtable" { + panic("WithSQLiteBusyTimeout is only valid for sqlite3 driver") + } + c.SQLite.BusyTimeout = ms + } +} + // SQL Server 专用选项 func WithSQLServerInstance(instance string) Option { return func(c *DBConfig) { @@ -524,3 +534,12 @@ func WithDuckDBMaxMemory(maxMemory string) Option { c.DuckDB.MaxMemory = maxMemory } } + +func WithDuckDBExtensions(extensions ...string) Option { + return func(c *DBConfig) { + if c.Driver != "duckdb" { + panic("WithDuckDBExtensions is only valid for duckdb driver") + } + c.DuckDB.Extensions = append(c.DuckDB.Extensions, extensions...) + } +} diff --git a/exec/query_executor.go b/exec/query_executor.go index 61b5ddd..85b82a4 100644 --- a/exec/query_executor.go +++ b/exec/query_executor.go @@ -227,8 +227,8 @@ func (q QueryExecutor) ScanValContext(ctx context.Context, i interface{}) (bool, if util.IsSlice(val.Kind()) { switch i.(type) { case *gsql.RawBytes: // do nothing - case *[]byte: // do nothing - case gsql.Scanner: // do nothing + case *[]byte: // do nothing + case gsql.Scanner: // do nothing default: return false, errScanValNonSlice } @@ -238,18 +238,14 @@ func (q QueryExecutor) ScanValContext(ctx context.Context, i interface{}) (bool, if err != nil { return false, err } - defer func() { _ = scanner.Close() }() if scanner.Next() { - err = scanner.ScanVal(i) - if err != nil { + if err = scanner.ScanVal(i); err != nil { return false, err } - return true, scanner.Err() } - return false, scanner.Err() } diff --git a/exec/query_executor_internal_test.go b/exec/query_executor_internal_test.go index 7b84bbe..ad5babd 100644 --- a/exec/query_executor_internal_test.go +++ b/exec/query_executor_internal_test.go @@ -3,6 +3,7 @@ package exec import ( "context" "database/sql" + "database/sql/driver" "encoding/json" "fmt" "strings" @@ -13,6 +14,11 @@ import ( "github.com/stretchr/testify/suite" ) +// anyValueConverter 允许任意类型作为 driver.Value 透传,模拟 DuckDB 等返回 map/struct 的驱动 +type anyValueConverter struct{} + +func (anyValueConverter) ConvertValue(v interface{}) (driver.Value, error) { return v, nil } + var ( testAddr1 = "111 Test Addr" testAddr2 = "211 Test Addr" @@ -929,9 +935,11 @@ func (qes *queryExecutorSuite) TestScanStruct() { qes.EqualError(err, "queryExecutor error") qes.False(found) + // NULL 值扫描进 string 字段:通过 **string 中间层正常处理,结果为空字符串 found, err = e.ScanStruct(&item) - qes.Error(err) - qes.False(found) + qes.NoError(err) + qes.True(found) + qes.Equal(StructWithTags{Address: "", Name: ""}, item) found, err = e.ScanStruct(&item) qes.NoError(err) @@ -1242,6 +1250,182 @@ func (qes *queryExecutorSuite) TestScanVal_withValuerSlice() { qes.Equal(JSONBoolArray{true, false, true}, bools) } +func (qes *queryExecutorSuite) TestScanVal_withByteSlice_notFound() { + db, mock, err := sqlmock.New() + qes.NoError(err) + mock.ExpectQuery(`SELECT "name" FROM "items"`). + WillReturnRows(sqlmock.NewRows([]string{"name"})) + + e := newQueryExecutor(db, nil, `SELECT "name" FROM "items"`) + var b []byte + found, err := e.ScanVal(&b) + qes.NoError(err) + qes.False(found) + qes.Nil(b) +} + +func (qes *queryExecutorSuite) TestScanVal_withByteSlice_queryError() { + db, mock, err := sqlmock.New() + qes.NoError(err) + mock.ExpectQuery(`SELECT "name" FROM "items"`). + WillReturnError(fmt.Errorf("query error")) + + e := newQueryExecutor(db, nil, `SELECT "name" FROM "items"`) + var b []byte + found, err := e.ScanVal(&b) + qes.EqualError(err, "query error") + qes.False(found) +} + +func (qes *queryExecutorSuite) TestScanVal_withByteSlice_complexJSON() { + // 模拟 DuckDB 等驱动将 JSON 列直接以 map 形式返回 + // 必须使用 mock.NewRows 而非 sqlmock.NewRows,才会使用自定义 converter + db, mock, err := sqlmock.New(sqlmock.ValueConverterOption(anyValueConverter{})) + qes.NoError(err) + mock.ExpectQuery(`SELECT "output" FROM "items"`). + WillReturnRows(mock.NewRows([]string{"output"}). + AddRow(map[string]interface{}{"key": "value", "num": float64(42)})) + + e := newQueryExecutor(db, nil, `SELECT "output" FROM "items"`) + var b []byte + found, err := e.ScanVal(&b) + qes.NoError(err) + qes.True(found) + + var result map[string]interface{} + qes.NoError(json.Unmarshal(b, &result)) + qes.Equal("value", result["key"]) + qes.Equal(float64(42), result["num"]) +} + +func (qes *queryExecutorSuite) TestScanVal_withRawBytes_notFound() { + db, mock, err := sqlmock.New() + qes.NoError(err) + mock.ExpectQuery(`SELECT "name" FROM "items"`). + WillReturnRows(sqlmock.NewRows([]string{"name"})) + + e := newQueryExecutor(db, nil, `SELECT "name" FROM "items"`) + var rb sql.RawBytes + found, err := e.ScanVal(&rb) + qes.NoError(err) + qes.False(found) + qes.Nil(rb) +} + +func (qes *queryExecutorSuite) TestScanVal_withRawBytes_queryError() { + db, mock, err := sqlmock.New() + qes.NoError(err) + mock.ExpectQuery(`SELECT "name" FROM "items"`). + WillReturnError(fmt.Errorf("query error")) + + e := newQueryExecutor(db, nil, `SELECT "name" FROM "items"`) + var rb sql.RawBytes + found, err := e.ScanVal(&rb) + qes.EqualError(err, "query error") + qes.False(found) +} + +func (qes *queryExecutorSuite) TestScanVal_withRawBytes_binaryDriver() { + // 驱动直接返回 []byte(如 BLOB 列),结果需独立拷贝 + db, mock, err := sqlmock.New() + qes.NoError(err) + content := []byte("binary \x00 content") + mock.ExpectQuery(`SELECT "data" FROM "items"`). + WillReturnRows(sqlmock.NewRows([]string{"data"}).AddRow(content)) + + e := newQueryExecutor(db, nil, `SELECT "data" FROM "items"`) + var rb sql.RawBytes + found, err := e.ScanVal(&rb) + qes.NoError(err) + qes.True(found) + qes.Equal(sql.RawBytes(content), rb) +} + +func (qes *queryExecutorSuite) TestScanVal_withStruct() { + // 模拟 DuckDB 将 JSON/STRUCT 列以 map[string]interface{} 返回 + type DocItem struct { + Title string `json:"title"` + Score int `json:"score"` + } + db, mock, err := sqlmock.New(sqlmock.ValueConverterOption(anyValueConverter{})) + qes.NoError(err) + mock.ExpectQuery(`SELECT "doc" FROM "items"`). + WillReturnRows(mock.NewRows([]string{"doc"}). + AddRow(map[string]interface{}{"title": "hello", "score": float64(99)})) + + e := newQueryExecutor(db, nil, `SELECT "doc" FROM "items"`) + var doc DocItem + found, err := e.ScanVal(&doc) + qes.NoError(err) + qes.True(found) + qes.Equal(DocItem{Title: "hello", Score: 99}, doc) +} + +func (qes *queryExecutorSuite) TestScanVal_withStruct_null() { + // 驱动返回 NULL,结构体保持零值 + type DocItem struct { + Title string `json:"title"` + } + db, mock, err := sqlmock.New() + qes.NoError(err) + mock.ExpectQuery(`SELECT "doc" FROM "items"`). + WillReturnRows(sqlmock.NewRows([]string{"doc"}).AddRow(nil)) + + e := newQueryExecutor(db, nil, `SELECT "doc" FROM "items"`) + var doc DocItem + found, err := e.ScanVal(&doc) + qes.NoError(err) + qes.True(found) + qes.Equal(DocItem{}, doc) +} + +func (qes *queryExecutorSuite) TestScanVal_withStruct_notFound() { + type DocItem struct { + Title string `json:"title"` + } + db, mock, err := sqlmock.New() + qes.NoError(err) + mock.ExpectQuery(`SELECT "doc" FROM "items"`). + WillReturnRows(sqlmock.NewRows([]string{"doc"})) + + e := newQueryExecutor(db, nil, `SELECT "doc" FROM "items"`) + var doc DocItem + found, err := e.ScanVal(&doc) + qes.NoError(err) + qes.False(found) +} + +func (qes *queryExecutorSuite) TestScanVal_withStruct_queryError() { + type DocItem struct { + Title string `json:"title"` + } + db, mock, err := sqlmock.New() + qes.NoError(err) + mock.ExpectQuery(`SELECT "doc" FROM "items"`). + WillReturnError(fmt.Errorf("query error")) + + e := newQueryExecutor(db, nil, `SELECT "doc" FROM "items"`) + var doc DocItem + found, err := e.ScanVal(&doc) + qes.EqualError(err, "query error") + qes.False(found) +} + +func (qes *queryExecutorSuite) TestScanVal_withStruct_sqlScanner() { + // 实现了 sql.Scanner 的结构体走原有直接扫描路径 + db, mock, err := sqlmock.New() + qes.NoError(err) + mock.ExpectQuery(`SELECT "name" FROM "items"`). + WillReturnRows(sqlmock.NewRows([]string{"name"}).AddRow("hello")) + + e := newQueryExecutor(db, nil, `SELECT "name" FROM "items"`) + var ns sql.NullString + found, err := e.ScanVal(&ns) + qes.NoError(err) + qes.True(found) + qes.Equal(sql.NullString{String: "hello", Valid: true}, ns) +} + func TestQueryExecutorSuite(t *testing.T) { suite.Run(t, new(queryExecutorSuite)) } diff --git a/exec/scanner.go b/exec/scanner.go index 1a7d430..ca4aa55 100644 --- a/exec/scanner.go +++ b/exec/scanner.go @@ -147,9 +147,9 @@ func (s *scanner) ScanStruct(i interface{}) error { // 补全未知字段类型 if len(cols) != len(cm) { - colTypes, err := s.rows.ColumnTypes() - if err != nil { - return err + colTypes, ctErr := s.rows.ColumnTypes() + if ctErr != nil { + return ctErr } for _, t := range colTypes { if _, ok := cm[t.Name()]; !ok { @@ -166,7 +166,6 @@ func (s *scanner) ScanStruct(i interface{}) error { } scans, err := createColumnScans(s.columns, s.columnMap) - if err != nil { return err } @@ -177,7 +176,12 @@ func (s *scanner) ScanStruct(i interface{}) error { record := map[string]interface{}{} for index, col := range s.columns { - record[col] = scans[index] + if pi, ok := scans[index].(*interface{}); ok { + raw := toJSONRawMessage(*pi) + record[col] = &raw + } else { + record[col] = scans[index] + } } util.AssignStructVals(i, record, s.columnMap) @@ -198,10 +202,59 @@ func (s *scanner) ScanStructs(i interface{}) error { // ScanVal will scan the current row and column into i. func (s *scanner) ScanVal(i interface{}) error { - if err := s.rows.Scan(i); err != nil { - return err + switch v := i.(type) { + case *sql.RawBytes: + // 零拷贝扫描,rows.Close 前立即拷贝防止驱动回收缓冲区 + if err := s.rows.Scan(v); err != nil { + return err + } + buf := make(sql.RawBytes, len(*v)) + copy(buf, *v) + *v = buf + case *[]byte: + // 先扫描到 interface{},驱动可能返回 []byte/string/map 等任意类型 + var raw interface{} + if err := s.rows.Scan(&raw); err != nil { + return err + } + switch rv := raw.(type) { + case []byte: + *v = append([]byte(nil), rv...) + case sql.RawBytes: + *v = append([]byte(nil), []byte(rv)...) + case string: + *v = []byte(rv) + default: + if raw != nil { + var err error + *v, err = json.Marshal(raw) + if err != nil { + return err + } + } + } + default: + // 指针-结构体且未实现 sql.Scanner:通过 JSON 中间层转换 + if rv := reflect.ValueOf(i); rv.Kind() == reflect.Ptr && rv.Elem().Kind() == reflect.Struct { + if _, ok := i.(sql.Scanner); !ok { + var raw interface{} + if err := s.rows.Scan(&raw); err != nil { + return err + } + if raw == nil { + return s.Err() + } + data, err := json.Marshal(raw) + if err != nil { + return err + } + return json.Unmarshal(data, i) + } + } + if err := s.rows.Scan(i); err != nil { + return err + } } - return s.Err() } @@ -260,6 +313,23 @@ func checkScanValsTarget(i interface{}) (reflect.Value, error) { return val, nil } +func toJSONRawMessage(v interface{}) *json.RawMessage { + if v == nil { + return nil + } + var raw json.RawMessage + switch s := v.(type) { + case []byte: + raw = append(json.RawMessage(nil), s...) + case string: + raw = json.RawMessage(s) + default: + b, _ := json.Marshal(v) + raw = b + } + return &raw +} + func createColumnScans(cols []string, cm util.ColumnMap) (scans []interface{}, err error) { scans = make([]interface{}, 0, len(cols)) @@ -278,7 +348,8 @@ func createColumnScans(cols []string, cm util.ColumnMap) (scans []interface{}, e reflect.Bool: scans = append(scans, reflect.New(reflect.PointerTo(data.GoType)).Interface()) case reflect.Map, reflect.Slice, reflect.Struct: - scans = append(scans, reflect.New(reflect.PointerTo(reflect.TypeOf(json.RawMessage{}))).Interface()) + // 使用 *interface{} 接受任意驱动值(兼容 DuckDB 返回 map[string]interface{}) + scans = append(scans, new(interface{})) default: scans = append(scans, reflect.New(data.GoType).Interface()) } diff --git a/exec/scanner_internal_test.go b/exec/scanner_internal_test.go index a2bf487..8493aa0 100644 --- a/exec/scanner_internal_test.go +++ b/exec/scanner_internal_test.go @@ -1,9 +1,10 @@ package exec import ( + "database/sql" + "encoding/json" "testing" - "git.fsdpf.net/go/db/exp" "github.com/DATA-DOG/go-sqlmock" "github.com/stretchr/testify/suite" ) @@ -81,14 +82,162 @@ func (s *scannerSuite) TestGetRecords() { AddRow("111 Test Addr", "Test1"), ) - rows, err := db.Query("SELECT \\* FROM `items`") + rows, err := db.Query("SELECT * FROM `items`") s.Require().NoError(err) result, err := NewScanner(rows).GetRecords() s.Require().NoError(err) - s.Equal([]exp.Record{ + s.Equal([]map[string]any{ {"address": "111 Test Addr", "name": "Test1"}, {"address": "111 Test Addr", "name": "Test1"}, }, result) } + +func (s *scannerSuite) TestScanVal() { + db, mock, err := sqlmock.New() + s.Require().NoError(err) + + mock.ExpectQuery(`SELECT "id" FROM "items"`). + WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(42)) + + rows, err := db.Query(`SELECT "id" FROM "items"`) + s.Require().NoError(err) + + sc := NewScanner(rows) + s.True(sc.Next()) + + var id int64 + err = sc.ScanVal(&id) + s.Require().NoError(err) + s.Equal(int64(42), id) +} + +func (s *scannerSuite) TestScanVal_withRawBytes() { + db, mock, err := sqlmock.New() + s.Require().NoError(err) + + mock.ExpectQuery(`SELECT "data" FROM "items"`). + WillReturnRows(sqlmock.NewRows([]string{"data"}).AddRow([]byte(testByteSliceContent))) + + rows, err := db.Query(`SELECT "data" FROM "items"`) + s.Require().NoError(err) + + sc := NewScanner(rows) + s.True(sc.Next()) + + var rb sql.RawBytes + err = sc.ScanVal(&rb) + s.Require().NoError(err) + + _ = sc.Close() + // 关闭后缓冲区应已拷贝,值仍然有效 + s.Equal(sql.RawBytes(testByteSliceContent), rb) +} + +func (s *scannerSuite) TestScanVal_withByteSlice() { + db, mock, err := sqlmock.New() + s.Require().NoError(err) + + mock.ExpectQuery(`SELECT "data" FROM "items"`). + WillReturnRows(sqlmock.NewRows([]string{"data"}).AddRow(testByteSliceContent)) + + rows, err := db.Query(`SELECT "data" FROM "items"`) + s.Require().NoError(err) + + sc := NewScanner(rows) + s.True(sc.Next()) + + var b []byte + err = sc.ScanVal(&b) + s.Require().NoError(err) + s.Equal([]byte(testByteSliceContent), b) +} + +func (s *scannerSuite) TestScanVal_withByteSlice_complexJSON() { + db, mock, err := sqlmock.New(sqlmock.ValueConverterOption(anyValueConverter{})) + s.Require().NoError(err) + + payload := map[string]interface{}{"key": "val", "num": float64(1)} + mock.ExpectQuery(`SELECT "data" FROM "items"`). + WillReturnRows(mock.NewRows([]string{"data"}).AddRow(payload)) + + rows, err := db.Query(`SELECT "data" FROM "items"`) + s.Require().NoError(err) + + sc := NewScanner(rows) + s.True(sc.Next()) + + var b []byte + err = sc.ScanVal(&b) + s.Require().NoError(err) + + var got map[string]interface{} + s.Require().NoError(json.Unmarshal(b, &got)) + s.Equal(payload, got) +} + +func (s *scannerSuite) TestScanVal_withStruct() { + type DocItem struct { + Title string `json:"title"` + Score int `json:"score"` + } + db, mock, err := sqlmock.New(sqlmock.ValueConverterOption(anyValueConverter{})) + s.Require().NoError(err) + + mock.ExpectQuery(`SELECT "doc" FROM "items"`). + WillReturnRows(mock.NewRows([]string{"doc"}). + AddRow(map[string]interface{}{"title": "hello", "score": float64(99)})) + + rows, err := db.Query(`SELECT "doc" FROM "items"`) + s.Require().NoError(err) + + sc := NewScanner(rows) + s.True(sc.Next()) + + var doc DocItem + err = sc.ScanVal(&doc) + s.Require().NoError(err) + s.Equal(DocItem{Title: "hello", Score: 99}, doc) +} + +func (s *scannerSuite) TestScanVal_withStruct_null() { + type DocItem struct { + Title string `json:"title"` + } + db, mock, err := sqlmock.New() + s.Require().NoError(err) + + mock.ExpectQuery(`SELECT "doc" FROM "items"`). + WillReturnRows(sqlmock.NewRows([]string{"doc"}).AddRow(nil)) + + rows, err := db.Query(`SELECT "doc" FROM "items"`) + s.Require().NoError(err) + + sc := NewScanner(rows) + s.True(sc.Next()) + + var doc DocItem + err = sc.ScanVal(&doc) + s.Require().NoError(err) + s.Equal(DocItem{}, doc) +} + +func (s *scannerSuite) TestScanVal_withStruct_sqlScanner() { + db, mock, err := sqlmock.New() + s.Require().NoError(err) + + mock.ExpectQuery(`SELECT "name" FROM "items"`). + WillReturnRows(sqlmock.NewRows([]string{"name"}).AddRow("hello")) + + rows, err := db.Query(`SELECT "name" FROM "items"`) + s.Require().NoError(err) + + sc := NewScanner(rows) + s.True(sc.Next()) + + var ns sql.NullString + err = sc.ScanVal(&ns) + s.Require().NoError(err) + s.Equal(sql.NullString{String: "hello", Valid: true}, ns) +} diff --git a/go.mod b/go.mod index 18a591f..993b804 100644 --- a/go.mod +++ b/go.mod @@ -9,7 +9,7 @@ require ( github.com/go-sql-driver/mysql v1.7.1 github.com/lib/pq v1.10.9 github.com/marcboeker/go-duckdb v1.8.5 - github.com/mattn/go-sqlite3 v1.14.17 + github.com/mattn/go-sqlite3 v1.14.42 github.com/microsoft/go-mssqldb v1.9.4 github.com/samber/lo v1.49.1 github.com/stretchr/testify v1.10.0 @@ -17,6 +17,7 @@ require ( require ( github.com/apache/arrow-go/v18 v18.1.0 // indirect + github.com/asg017/sqlite-vec-go-bindings v0.1.6 // indirect github.com/davecgh/go-spew v1.1.1 // indirect github.com/go-viper/mapstructure/v2 v2.2.1 // indirect github.com/goccy/go-json v0.10.5 // indirect diff --git a/go.sum b/go.sum index 95ae804..bc9a06f 100644 --- a/go.sum +++ b/go.sum @@ -18,6 +18,8 @@ github.com/apache/arrow-go/v18 v18.1.0 h1:agLwJUiVuwXZdwPYVrlITfx7bndULJ/dggbnLF github.com/apache/arrow-go/v18 v18.1.0/go.mod h1:tigU/sIgKNXaesf5d7Y95jBBKS5KsxTqYBKXFsvKzo0= github.com/apache/thrift v0.21.0 h1:tdPmh/ptjE1IJnhbhrcl2++TauVjy242rkV/UzJChnE= github.com/apache/thrift v0.21.0/go.mod h1:W1H8aR/QRtYNvrPeFXBtobyRkd0/YVhTc6i07XIAgDw= +github.com/asg017/sqlite-vec-go-bindings v0.1.6 h1:Nx0jAzyS38XpkKznJ9xQjFXz2X9tI7KqjwVxV8RNoww= +github.com/asg017/sqlite-vec-go-bindings v0.1.6/go.mod h1:A8+cTt/nKFsYCQF6OgzSNpKZrzNo5gQsXBTfsXHXY0Q= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/go-sql-driver/mysql v1.7.1 h1:lUIinVbN1DY0xBg0eMOzmmtGoHwWBbvnWubQUrtU8EI= @@ -58,6 +60,8 @@ github.com/marcboeker/go-duckdb v1.8.5 h1:tkYp+TANippy0DaIOP5OEfBEwbUINqiFqgwMQ4 github.com/marcboeker/go-duckdb v1.8.5/go.mod h1:6mK7+WQE4P4u5AFLvVBmhFxY5fvhymFptghgJX6B+/8= github.com/mattn/go-sqlite3 v1.14.17 h1:mCRHCLDUBXgpKAqIKsaAaAsrAlbkeomtRFKXh2L6YIM= github.com/mattn/go-sqlite3 v1.14.17/go.mod h1:2eHXhiwb8IkHr+BDWZGa96P6+rkvnG63S2DGjv9HUNg= +github.com/mattn/go-sqlite3 v1.14.42 h1:MigqEP4ZmHw3aIdIT7T+9TLa90Z6smwcthx+Azv4Cgo= +github.com/mattn/go-sqlite3 v1.14.42/go.mod h1:pjEuOr8IwzLJP2MfGeTb0A35jauH+C2kbHKBr7yXKVQ= github.com/microsoft/go-mssqldb v1.9.4 h1:sHrj3GcdgkxytZ09aZ3+ys72pMeyEXJowT44j74pNgs= github.com/microsoft/go-mssqldb v1.9.4/go.mod h1:GBbW9ASTiDC+mpgWDGKdm3FnFLTUsLYN3iFL90lQ+PA= github.com/minio/asm2plan9s v0.0.0-20200509001527-cdd76441f9d8 h1:AMFGa4R4MiIpspGNG7Z948v4n35fFGB3RR3G/ry4FWs= diff --git a/insert_dataset.go b/insert_dataset.go index 62bfbef..02cee76 100644 --- a/insert_dataset.go +++ b/insert_dataset.go @@ -210,25 +210,11 @@ func (id *InsertDataset) Rows(rows ...interface{}) *InsertDataset { panic("Rows: unsupported row type, must be map, Record, struct, slice or array") } - result := make(map[string]interface{}) - typ := val.Type() - for j := 0; j < val.NumField(); j++ { - field := typ.Field(j) - // Skip unexported fields - if field.PkgPath != "" { - continue - } - // Use db tag if present, otherwise use field name - key := field.Name - if dbTag, ok := field.Tag.Lookup("db"); ok && dbTag != "" { - if dbTag == "-" { - continue // Skip fields with db tag "-" - } - key = dbTag - } - result[key] = val.Field(j).Interface() + record, err := exp.NewRecordFromStruct(row, true, false) + if err != nil { + panic(fmt.Sprintf("Rows: %v", err)) } - converted[i] = result + converted[i] = map[string]interface{}(record) } } return id.copy(id.clauses.SetRows(converted)) diff --git a/internal/util/column_map.go b/internal/util/column_map.go index 7b26e56..ebda57f 100644 --- a/internal/util/column_map.go +++ b/internal/util/column_map.go @@ -38,13 +38,13 @@ func newColumnMap(t reflect.Type, fieldIndex []int, prefixes []string) ColumnMap columnName := getColumnName(&f, dbTag) if !shouldIgnoreField(dbTag) { // 移除原来的关联结构,并 scans 字段出现 table.col - // if !implementsScanner(f.Type) { - // subCm := getStructColumnMap(&f, fieldIndex, []string{columnName}, prefixes) - // if len(subCm) != 0 { - // subColMaps = append(subColMaps, subCm) - // continue - // } - // } + if !implementsScanner(f.Type) { + subCm := getStructColumnMap(&f, fieldIndex, []string{columnName}, prefixes) + if len(subCm) != 0 { + subColMaps = append(subColMaps, subCm) + continue + } + } ffTag := tag.New("ff", f.Tag) columnName = strings.Join(append(prefixes, columnName), ".") cm[columnName] = newColumnData(&f, columnName, fieldIndex, ffTag) diff --git a/internal/util/reflect.go b/internal/util/reflect.go index 4e9ee19..b44e74a 100644 --- a/internal/util/reflect.go +++ b/internal/util/reflect.go @@ -203,21 +203,55 @@ func SafeSetFieldByIndex(v reflect.Value, fieldIndex []int, src interface{}) (re } func SafeSetVarValue(v reflect.Value, src interface{}) error { - f := reflect.Indirect(v) - srcVal := reflect.ValueOf(src).Elem() + srcReflect := reflect.ValueOf(src) - // 处理 converting NULL to string is unsupported - // 前面将 string 和 int 类型转为了 *string 和 *int - // 这里做还原 - if srcVal.IsNil() { - f.Set(reflect.Zero(f.Type())) - } else if f.Type().ConvertibleTo(srcVal.Type().Elem()) { - f.Set(srcVal.Elem()) - } else if f.Type().ConvertibleTo(srcVal.Type()) { - f.Set(srcVal) - } else if u, ok := srcVal.Interface().(*json.RawMessage); ok && len(*u) >= 2 { - if err := json.Unmarshal(*u, f.Addr().Interface()); err != nil { - return err + // src 可能是 **T(createColumnScans 的扫描目标)或裸值(测试/直接调用) + if srcReflect.Kind() != reflect.Ptr { + if srcReflect.IsValid() && v.Type().ConvertibleTo(srcReflect.Type()) { + v.Set(srcReflect.Convert(v.Type())) + } + return nil + } + + srcVal := srcReflect.Elem() // **T → *T;*T → T + + if IsNil(srcVal) { + v.Set(reflect.Zero(v.Type())) + return nil + } + + if v.Kind() == reflect.Ptr { + // v 是指针字段(如 *sql.NullString) + if srcVal.Kind() == reflect.Ptr { + // src = **T, srcVal = *T → v = *T + if v.Type().ConvertibleTo(srcVal.Type()) { + v.Set(srcVal.Convert(v.Type())) + } + } else { + // src = *T, srcVal = T → allocate new *T and set + if v.Type().Elem().ConvertibleTo(srcVal.Type()) { + p := reflect.New(v.Type().Elem()) + p.Elem().Set(srcVal.Convert(v.Type().Elem())) + v.Set(p) + } + } + return nil + } + + // v 是非指针字段 + if srcVal.Kind() == reflect.Ptr { + // srcVal = *T,取其值赋给 v + if v.Type().ConvertibleTo(srcVal.Type().Elem()) { + v.Set(srcVal.Elem().Convert(v.Type())) + } else if u, ok := srcVal.Interface().(*json.RawMessage); ok && len(*u) >= 2 { + if err := json.Unmarshal(*u, v.Addr().Interface()); err != nil { + return err + } + } + } else { + // srcVal = T(*T 扫描目标的 default 分支) + if v.Type().ConvertibleTo(srcVal.Type()) { + v.Set(srcVal.Convert(v.Type())) } } diff --git a/schema/blueprint.go b/schema/blueprint.go index e833487..555befa 100755 --- a/schema/blueprint.go +++ b/schema/blueprint.go @@ -224,6 +224,12 @@ func (this *Blueprint) Uuid(column string) *ColumnDefinition { return this.addColumn("uuid", column, nil) } +// 向量类型(float32 固定维度),dimensions 为向量维度数 +// DuckDB: FLOAT[N] PostgreSQL: vector(N) MySQL: VECTOR(N) SQLite3: BLOB +func (this *Blueprint) Vector(column string, dimensions int) *ColumnDefinition { + return this.addColumn("vector", column, &ColumnOptions{Length: dimensions}) +} + // 自增字段 func (this *Blueprint) Increments(column string) *ColumnDefinition { return this.UnsignedInteger(column, true) diff --git a/schema/builder.go b/schema/builder.go index 8636856..2d84a71 100755 --- a/schema/builder.go +++ b/schema/builder.go @@ -88,6 +88,8 @@ func (this Builder) Build(bp *Blueprint) error { return err } + // fmt.Println(strings.Join(sqls, "\n---\n")) + err = tx.Wrap(func() error { for _, sql := range sqls { if _, err := tx.Exec(sql); err != nil { diff --git a/schema/column_definition.go b/schema/column_definition.go index eb0e025..b03aaae 100755 --- a/schema/column_definition.go +++ b/schema/column_definition.go @@ -34,9 +34,11 @@ type ColumnOptions struct { hidden bool // SQLite 虚拟表 HIDDEN 列:不出现在 SELECT *,可在 WHERE 中作为参数传入 } -// VirtualAs Create a virtual generated column (MySQL) +// VirtualAs Create a virtual generated column (MySQL)。 +// 在 SQLite 虚拟表 schema 中无效,自动退化为 HIDDEN 列。 func (c *ColumnDefinition) VirtualAs(as string) *ColumnDefinition { c.virtualAs = as + c.hidden = true return c } diff --git a/schema/dialect/duckdb/duckdb.go b/schema/dialect/duckdb/duckdb.go index 6d1fcd4..46dae6d 100644 --- a/schema/dialect/duckdb/duckdb.go +++ b/schema/dialect/duckdb/duckdb.go @@ -40,20 +40,33 @@ func (this DuckDB) CompileCreate(bp *schema.Blueprint) []string { if bp.Temporary { temporary = db.L("CREATE TEMPORARY") } - columns := strings.Join(this.getAddedColumns(bp), ",\n") - sql := this.GenerateSQL("? TABLE ? (\n?\n)", temporary, db.T(bp.GetTable()), db.L(columns)) - if comment := bp.Comment; comment != "" { - sql = sql[:len(sql)-1] + ";" + this.GenerateSQL("\nCOMMENT ON TABLE ? IS ?", db.T(bp.GetTable()), db.V(comment)) + // 自增列需要先建序列 + var sqls []string + for _, column := range bp.GetAddedColumns() { + if column.IsAutoIncrement() { + sqls = append(sqls, this.GenerateSQL("CREATE SEQUENCE IF NOT EXISTS ?", + db.T(this.seqName(bp.GetTable(), column.Name)))) + } } - return []string{sql} + columns := strings.Join(this.getAddedColumns(bp), ",\n") + sqls = append(sqls, this.GenerateSQL("? TABLE IF NOT EXISTS ? (\n?\n)", temporary, db.T(bp.GetTable()), db.L(columns))) + + if comment := bp.Comment; comment != "" { + sqls = append(sqls, this.GenerateSQL("COMMENT ON TABLE ? IS ?", db.T(bp.GetTable()), db.V(comment))) + } + sqls = append(sqls, this.columnComments(bp.GetTable(), bp.GetAddedColumns())...) + + return sqls } func (this DuckDB) CompileAdd(bp *schema.Blueprint) []string { columns := strings.Join(schema.PrefixArray("ADD COLUMN", this.getAddedColumns(bp)), ",\n") sql := this.GenerateSQL("ALTER TABLE ?\n?", db.T(bp.GetTable()), db.L(columns)) - return []string{sql} + sqls := []string{sql} + sqls = append(sqls, this.columnComments(bp.GetTable(), bp.GetAddedColumns())...) + return sqls } func (this DuckDB) CompileChange(bp *schema.Blueprint) []string { @@ -121,16 +134,37 @@ func (this DuckDB) CompileDropIfExists(bp *schema.Blueprint) []string { } func (this DuckDB) CompileRename(bp *schema.Blueprint) []string { - toName := "" commands := bp.GetCommands() if len(commands) == 0 { panic("new table undefined") } - toName = commands[0].To + toName := commands[0].To if toName == "" { panic("new table undefined") } - return []string{this.GenerateSQL("ALTER TABLE ? RENAME TO ?", db.T(bp.GetTable()), db.T(toName))} + fromName := bp.GetTable() + + sqls := []string{this.GenerateSQL("ALTER TABLE ? RENAME TO ?", db.T(fromName), db.T(toName))} + + // 查找该表关联的序列(命名规则: {table}_{column}_seq),一并重命名并更新列 DEFAULT + var seqNames []string + _ = this.db.From(db.L("duckdb_sequences()")). + Where(db.C("sequence_name").Like(fromName+"_%")). + Pluck(&seqNames, "sequence_name") + + prefix := fromName + "_" + for _, oldSeq := range seqNames { + if !strings.HasSuffix(oldSeq, "_seq") { + continue + } + colName := oldSeq[len(prefix) : len(oldSeq)-len("_seq")] + newSeq := this.seqName(toName, colName) + sqls = append(sqls, this.GenerateSQL("ALTER SEQUENCE ? RENAME TO ?", db.T(oldSeq), db.T(newSeq))) + sqls = append(sqls, this.GenerateSQL("ALTER TABLE ? ALTER COLUMN ? SET DEFAULT nextval(?)", + db.T(toName), db.C(colName), db.V(newSeq))) + } + + return sqls } func (this DuckDB) CompileModifyComment(bp *schema.Blueprint) []string { @@ -185,6 +219,8 @@ func (this DuckDB) GetColumnType(column *schema.ColumnDefinition) string { return "SMALLINT" case "uuid": return "UUID" + case "vector": + return this.GenerateSQL("FLOAT[?]", column.Length) } panic("Unsupported data type: " + column.Type) } @@ -197,19 +233,17 @@ func (this DuckDB) GetColumnModifier(modifier string, bp *schema.Blueprint, colu } return " NOT NULL" case "Default": + if column.IsUseCurrent() { + // DuckDB 不支持 ON UPDATE,忽略 def 中可能携带的 MySQL ON UPDATE 标记 + return " DEFAULT CURRENT_TIMESTAMP" + } v := column.GetDefault() if v == nil { - if column.IsUseCurrent() { - return " DEFAULT CURRENT_TIMESTAMP" - } return "" } return this.GenerateSQL(" DEFAULT ?", v) case "Increment": - if column.IsAutoIncrement() { - // DuckDB 使用 SERIAL 或者 SEQUENCE - return " PRIMARY KEY" - } + // 自增列在 getAddedColumns 中单独处理(GENERATED ALWAYS AS IDENTITY),此处无需输出 case "Comment": // DuckDB 支持列注释,但需要在 CREATE TABLE 后使用 COMMENT ON return "" @@ -217,22 +251,41 @@ func (this DuckDB) GetColumnModifier(modifier string, bp *schema.Blueprint, colu return "" } +func (this DuckDB) columnComments(table string, columns []*schema.ColumnDefinition) (sqls []string) { + for _, column := range columns { + if v := column.GetComment(); v != "" { + sqls = append(sqls, this.GenerateSQL("COMMENT ON COLUMN ?.? IS ?", + db.T(table), db.C(column.Name), db.V(v))) + } + } + return +} + +func (this DuckDB) seqName(table, column string) string { + return table + "_" + column + "_seq" +} + func (this DuckDB) getAddedColumns(bp *schema.Blueprint) (columns []string) { for _, column := range bp.GetAddedColumns() { - colType := this.GetColumnType(column) - // 对于自增列,使用 SERIAL 类型 + var sql string if column.IsAutoIncrement() { + var baseType string switch column.Type { case "bigInteger": - colType = "BIGSERIAL" + baseType = "BIGINT" case "smallInteger": - colType = "SMALLSERIAL" + baseType = "SMALLINT" default: - colType = "SERIAL" + baseType = "INTEGER" } + seqName := this.seqName(bp.GetTable(), column.Name) + sql = this.GenerateSQL("? ? DEFAULT nextval(?) PRIMARY KEY", + db.C(column.Name), db.L(baseType), db.V(seqName)) + } else { + sql = this.GenerateSQL("? ?", db.C(column.Name), db.L(this.GetColumnType(column))) + sql = this.addModifiers(sql, bp, column) } - sql := this.GenerateSQL("? ?", db.C(column.Name), db.L(colType)) - columns = append(columns, this.addModifiers(sql, bp, column)) + columns = append(columns, sql) } return } diff --git a/schema/dialect/duckdb/duckdb_test.go b/schema/dialect/duckdb/duckdb_test.go index a9a948c..4847bd4 100644 --- a/schema/dialect/duckdb/duckdb_test.go +++ b/schema/dialect/duckdb/duckdb_test.go @@ -81,8 +81,9 @@ func (t *duckDBTest) TestCompileCreate() { sql := t.schema.CompileCreate(bp) t.Equal([]string{ - "CREATE TABLE \"users\" (\n" + - "\"id\" BIGSERIAL NOT NULL PRIMARY KEY,\n" + + "CREATE SEQUENCE IF NOT EXISTS \"users_id_seq\"", + "CREATE TABLE IF NOT EXISTS \"users\" (\n" + + "\"id\" BIGINT DEFAULT nextval('users_id_seq') PRIMARY KEY,\n" + "\"enabled\" BOOLEAN NOT NULL DEFAULT '1',\n" + "\"created_user\" CHAR(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000',\n" + "\"owned_user\" CHAR(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000',\n" + @@ -90,6 +91,13 @@ func (t *duckDBTest) TestCompileCreate() { "\"updated_at\" TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,\n" + "\"deleted_at\" TIMESTAMP NULL\n" + ")", + "COMMENT ON COLUMN \"users\".\"id\" IS 'ID'", + "COMMENT ON COLUMN \"users\".\"enabled\" IS '是否有效'", + "COMMENT ON COLUMN \"users\".\"created_user\" IS '创建者'", + "COMMENT ON COLUMN \"users\".\"owned_user\" IS '拥有者'", + "COMMENT ON COLUMN \"users\".\"created_at\" IS '创建时间'", + "COMMENT ON COLUMN \"users\".\"updated_at\" IS '更新时间'", + "COMMENT ON COLUMN \"users\".\"deleted_at\" IS '删除时间'", }, sql) t.T().Log(sql) @@ -109,6 +117,9 @@ func (t *duckDBTest) TestCompileAdd() { "ADD COLUMN \"name\" VARCHAR(50) NOT NULL,\n" + "ADD COLUMN \"age\" SMALLINT NOT NULL DEFAULT '18',\n" + "ADD COLUMN \"sex\" VARCHAR(1) NOT NULL DEFAULT '0'", + "COMMENT ON COLUMN \"users\".\"name\" IS '用户名'", + "COMMENT ON COLUMN \"users\".\"age\" IS '年龄'", + "COMMENT ON COLUMN \"users\".\"sex\" IS '性别'", }, sql) t.T().Log(sql) diff --git a/schema/dialect/duckdb/example_test.go b/schema/dialect/duckdb/example_test.go index e3dda43..61c399d 100644 --- a/schema/dialect/duckdb/example_test.go +++ b/schema/dialect/duckdb/example_test.go @@ -41,8 +41,9 @@ func Example_createTable() { } // Output: - // CREATE TABLE "users" ( - // "id" BIGSERIAL NOT NULL PRIMARY KEY, + // CREATE SEQUENCE IF NOT EXISTS "users_id_seq" + // CREATE TABLE IF NOT EXISTS "users" ( + // "id" BIGINT DEFAULT nextval('users_id_seq') PRIMARY KEY, // "name" VARCHAR(50) NOT NULL, // "email" VARCHAR(100) NOT NULL, // "age" INTEGER NULL, @@ -50,6 +51,13 @@ func Example_createTable() { // "created_at" TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, // "updated_at" TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP // ) + // COMMENT ON COLUMN "users"."id" IS '用户ID' + // COMMENT ON COLUMN "users"."name" IS '用户名' + // COMMENT ON COLUMN "users"."email" IS '邮箱' + // COMMENT ON COLUMN "users"."age" IS '年龄' + // COMMENT ON COLUMN "users"."active" IS '是否激活' + // COMMENT ON COLUMN "users"."created_at" IS '创建时间' + // COMMENT ON COLUMN "users"."updated_at" IS '更新时间' } func Example_addColumn() { @@ -75,6 +83,8 @@ func Example_addColumn() { // ALTER TABLE "users" // ADD COLUMN "phone" VARCHAR(20) NULL, // ADD COLUMN "address" TEXT NULL + // COMMENT ON COLUMN "users"."phone" IS '电话号码' + // COMMENT ON COLUMN "users"."address" IS '地址' } func Example_dropColumn() { diff --git a/schema/dialect/mysql/mysql.go b/schema/dialect/mysql/mysql.go index bd85e6e..9638891 100644 --- a/schema/dialect/mysql/mysql.go +++ b/schema/dialect/mysql/mysql.go @@ -53,7 +53,7 @@ func (this Mysql) CompileCreate(bp *schema.Blueprint) []string { } columns := strings.Join(this.getAddedColumns(bp), ",\n") - sql := this.GenerateSQL("? TABLE ? (\n?\n)", temporary, db.T(bp.GetTable()), db.L(columns)) + sql := this.GenerateSQL("? TABLE IF NOT EXISTS ? (\n?\n)", temporary, db.T(bp.GetTable()), db.L(columns)) charset := bp.Charset if charset == "" { @@ -183,6 +183,8 @@ func (this Mysql) GetColumnType(column *schema.ColumnDefinition) string { return "year" case "uuid": return "char(36)" + case "vector": + return this.GenerateSQL("VECTOR(?)", column.Length) } panic("Unsupported data type: " + column.Type) } diff --git a/schema/dialect/mysql/mysql_test.go b/schema/dialect/mysql/mysql_test.go index 1c732e5..86e6e7c 100644 --- a/schema/dialect/mysql/mysql_test.go +++ b/schema/dialect/mysql/mysql_test.go @@ -1,15 +1,13 @@ package mysql_test import ( - "os" "testing" "git.fsdpf.net/go/db" "git.fsdpf.net/go/db/engine" "git.fsdpf.net/go/db/schema" + sqlmock "github.com/DATA-DOG/go-sqlmock" "github.com/stretchr/testify/suite" - - _ "github.com/go-sql-driver/mysql" ) type mysqlTest struct { @@ -17,42 +15,28 @@ type mysqlTest struct { schema schema.Schema } -var ( - tableName = "entry" - - dropTable = "DROP TABLE IF EXISTS `entry`;" - - createTable = "CREATE TABLE IF NOT EXISTS `entry` (" + - "`id` INT NOT NULL AUTO_INCREMENT ," + - "`int` INT NOT NULL UNIQUE," + - "`float` FLOAT NOT NULL ," + - "`string` VARCHAR(255) NOT NULL ," + - "`time` DATETIME NOT NULL ," + - "`bool` TINYINT NOT NULL ," + - "`bytes` BLOB NOT NULL ," + - "PRIMARY KEY (`id`) );" -) +var tableName = "entry" func TestMysqlSuite(t *testing.T) { suite.Run(t, new(mysqlTest)) } func (t *mysqlTest) SetupSuite() { - db := engine.Open(map[string]engine.DBConfig{ - "test-mysql": engine.NewDBConfig("mysql", - engine.WithHost(os.Getenv("MYSQL_HOST")), - engine.WithPort(os.Getenv("MYSQL_PORT")), - engine.WithDatabase(os.Getenv("MYSQL_DB")), - engine.WithUsername(os.Getenv("MYSQL_USER")), - engine.WithPassword(os.Getenv("MYSQL_PASSWD")), - engine.WithParseTime(true), - ), + mockDB, mock, _ := sqlmock.New() + + mock.ExpectQuery("SELECT.*column_name.*columns"). + WillReturnRows(sqlmock.NewRows([]string{"column_name"}). + AddRow("id").AddRow("int").AddRow("float"). + AddRow("string").AddRow("time").AddRow("bool").AddRow("bytes")) + + mock.ExpectQuery("SELECT.*COUNT.*tables"). + WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + + conn := engine.Mock(map[string]engine.MockDBConfig{ + "test-mysql": {Driver: "mysql", Mock: mockDB}, }).Connection("test-mysql") - t.schema = schema.GetSchemaDialect(db) - - // db.Exec(dropTable) - db.Exec(createTable) + t.schema = schema.GetSchemaDialect(conn) } func (t *mysqlTest) TestGetColumnListing() { @@ -84,7 +68,7 @@ func (t *mysqlTest) TestCompileCreate() { sql := t.schema.CompileCreate(bp) t.Equal([]string{ - "CREATE TABLE `users` (\n" + + "CREATE TABLE IF NOT EXISTS `users` (\n" + "`id` bigint(20) UNSIGNED NOT NULL AUTO_INCREMENT PRIMARY KEY COMMENT 'ID',\n" + "`enabled` tinyint(1) NOT NULL DEFAULT '1' COMMENT '是否有效',\n" + "`created_user` char(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000' COMMENT '创建者',\n" + @@ -125,11 +109,13 @@ func (t *mysqlTest) TestCompileChange() { bp.String("name", 50).Change("username") bp.SmallInteger("age").Default("19").Comment("年龄").Change() - sql := "ALTER TABLE `users`\n" + - "CHANGE COLUMN `name` `username` varchar(50) NOT NULL,\n" + - "CHANGE COLUMN `age` `age` smallint(4) NOT NULL DEFAULT '19' COMMENT '年龄';" + sql := t.schema.CompileChange(bp) - t.Equal(t.schema.CompileChange(bp), sql) + t.Equal([]string{ + "ALTER TABLE `users`\n" + + "CHANGE COLUMN `name` `username` varchar(50) NOT NULL,\n" + + "CHANGE COLUMN `age` `age` smallint(4) NOT NULL DEFAULT '19' COMMENT '年龄'", + }, sql) t.T().Log(sql) } @@ -189,4 +175,19 @@ func (t *mysqlTest) TestCompileRename() { t.T().Log(sql) } +// 添加虚拟生成列 +func (t *mysqlTest) TestCompileAddVirtualColumn() { + bp := schema.NewBlueprint("user_assets") + bp.String("file_url", 512).VirtualAs("CONCAT('/api/user-asset-raw/', file)").Comment("文件访问URL") + + sql := t.schema.CompileAdd(bp) + + t.Equal([]string{ + "ALTER TABLE `user_assets`\n" + + "ADD COLUMN `file_url` varchar(512) GENERATED ALWAYS AS (CONCAT('/api/user-asset-raw/', file)) VIRTUAL COMMENT '文件访问URL'", + }, sql) + + t.T().Log(sql) +} + // 修改表备注 diff --git a/schema/dialect/postgres/postgres.go b/schema/dialect/postgres/postgres.go index 36eab03..3b7abdb 100644 --- a/schema/dialect/postgres/postgres.go +++ b/schema/dialect/postgres/postgres.go @@ -45,7 +45,7 @@ func (this Postgres) CompileCreate(bp *schema.Blueprint) []string { temporary = db.L("CREATE TEMPORARY") } columns := strings.Join(this.getAddedColumns(bp), ",\n") - sql := this.GenerateSQL("? TABLE ? (\n?\n)", temporary, db.T(bp.GetTable()), db.L(columns)) + sql := this.GenerateSQL("? TABLE IF NOT EXISTS ? (\n?\n)", temporary, db.T(bp.GetTable()), db.L(columns)) return []string{sql} } @@ -144,6 +144,8 @@ func (this Postgres) GetColumnType(column *schema.ColumnDefinition) string { return "integer" case "uuid": return "uuid" + case "vector": + return this.GenerateSQL("vector(?)", column.Length) } panic("Unsupported data type: " + column.Type) } @@ -156,16 +158,14 @@ func (this Postgres) GetColumnModifier(modifier string, bp *schema.Blueprint, co } return " NOT NULL" case "Default": + if column.IsUseCurrent() { + // PostgreSQL 不支持 ON UPDATE,忽略 def 中可能携带的 MySQL ON UPDATE 标记 + return " DEFAULT CURRENT_TIMESTAMP" + } v := column.GetDefault() if v == nil { - if column.IsUseCurrent() { - return " DEFAULT CURRENT_TIMESTAMP" - } return "" } - if column.IsUseCurrent() { - return this.GenerateSQL(" DEFAULT CURRENT_TIMESTAMP ?", v) - } return this.GenerateSQL(" DEFAULT ?", v) case "Increment": if column.IsAutoIncrement() { diff --git a/schema/dialect/postgres/postgres_test.go b/schema/dialect/postgres/postgres_test.go index cac2b64..b888d0b 100644 --- a/schema/dialect/postgres/postgres_test.go +++ b/schema/dialect/postgres/postgres_test.go @@ -1,14 +1,12 @@ package postgres_test import ( - "os" "testing" "git.fsdpf.net/go/db/engine" "git.fsdpf.net/go/db/schema" + sqlmock "github.com/DATA-DOG/go-sqlmock" "github.com/stretchr/testify/suite" - - _ "github.com/lib/pq" ) type postgresTest struct { @@ -16,41 +14,27 @@ type postgresTest struct { schema schema.Schema } -var ( - tableName = "entry" - - dropTable = "DROP TABLE IF EXISTS entry;" - - createTable = "CREATE TABLE IF NOT EXISTS entry (" + - "id SERIAL PRIMARY KEY," + - "int INTEGER NOT NULL UNIQUE," + - "float REAL NOT NULL," + - "string VARCHAR(255) NOT NULL," + - "time TIMESTAMP NOT NULL," + - "bool BOOLEAN NOT NULL," + - "bytes BYTEA NOT NULL" + - ");" -) +var tableName = "entry" func TestPostgresSuite(t *testing.T) { suite.Run(t, new(postgresTest)) } func (t *postgresTest) SetupSuite() { - db := engine.Open(map[string]engine.DBConfig{ - "test-postgres": engine.NewDBConfig("postgres", - engine.WithHost(os.Getenv("POSTGRES_HOST")), - engine.WithPort(os.Getenv("POSTGRES_PORT")), - engine.WithDatabase(os.Getenv("POSTGRES_DB")), - engine.WithUsername(os.Getenv("POSTGRES_USER")), - engine.WithPassword(os.Getenv("POSTGRES_PASSWD")), - ), + mockDB, mock, _ := sqlmock.New() + + mock.ExpectQuery("SELECT.*column_name.*columns"). + WillReturnRows(sqlmock.NewRows([]string{"column_name"}). + AddRow("id").AddRow("int").AddRow("float"). + AddRow("string").AddRow("time").AddRow("bool").AddRow("bytes")) + + mock.ExpectQuery("SELECT.*COUNT.*tables"). + WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + + conn := engine.Mock(map[string]engine.MockDBConfig{ + "test-postgres": {Driver: "postgres", Mock: mockDB}, }).Connection("test-postgres") - - t.schema = schema.GetSchemaDialect(db) - - db.Exec(dropTable) - db.Exec(createTable) + t.schema = schema.GetSchemaDialect(conn) } func (t *postgresTest) TestGetColumnListing() { @@ -80,14 +64,14 @@ func (t *postgresTest) TestCompileCreate() { sql := t.schema.CompileCreate(bp) t.Equal([]string{ - "CREATE TABLE \"users\" (\n" + - "\"id\" bigint NOT NULL BIGSERIAL PRIMARY KEY COMMENT 'ID',\n" + - "\"enabled\" boolean NOT NULL DEFAULT true COMMENT '是否有效',\n" + - "\"created_user\" char(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000' COMMENT '创建者',\n" + - "\"owned_user\" char(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000' COMMENT '拥有者',\n" + - "\"created_at\" timestamp with time zone NOT NULL DEFAULT CURRENT_TIMESTAMP COMMENT '创建时间',\n" + - "\"updated_at\" timestamp with time zone NOT NULL DEFAULT CURRENT_TIMESTAMP COMMENT '更新时间',\n" + - "\"deleted_at\" timestamp without time zone NULL COMMENT '删除时间'\n" + + "CREATE TABLE IF NOT EXISTS \"users\" (\n" + + "\"id\" bigint NOT NULL BIGSERIAL PRIMARY KEY,\n" + + "\"enabled\" boolean NOT NULL DEFAULT 'true',\n" + + "\"created_user\" char(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000',\n" + + "\"owned_user\" char(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000',\n" + + "\"created_at\" timestamp with time zone NOT NULL DEFAULT CURRENT_TIMESTAMP,\n" + + "\"updated_at\" timestamp with time zone NOT NULL DEFAULT CURRENT_TIMESTAMP,\n" + + "\"deleted_at\" timestamp without time zone NULL\n" + ")", }, sql) @@ -105,9 +89,9 @@ func (t *postgresTest) TestCompileAdd() { t.Equal([]string{ "ALTER TABLE \"users\"\n" + - "ADD COLUMN \"name\" varchar(50) NOT NULL COMMENT '用户名',\n" + - "ADD COLUMN \"age\" smallint NOT NULL DEFAULT 18 COMMENT '年龄',\n" + - "ADD COLUMN \"sex\" varchar(1) NOT NULL DEFAULT '0' COMMENT '性别'", + "ADD COLUMN \"name\" varchar(50) NOT NULL,\n" + + "ADD COLUMN \"age\" smallint NOT NULL DEFAULT '18',\n" + + "ADD COLUMN \"sex\" varchar(1) NOT NULL DEFAULT '0'", }, sql) t.T().Log(sql) @@ -124,7 +108,7 @@ func (t *postgresTest) TestCompileChange() { t.Equal([]string{ "ALTER TABLE \"users\"\n" + "ALTER COLUMN \"name\" TYPE varchar(50) NOT NULL,\n" + - "ALTER COLUMN \"age\" TYPE smallint NOT NULL DEFAULT 19 COMMENT '年龄'", + "ALTER COLUMN \"age\" TYPE smallint NOT NULL DEFAULT '19'", }, sql) t.T().Log(sql) diff --git a/schema/dialect/sqlite3/sqlite3.go b/schema/dialect/sqlite3/sqlite3.go index ecd35f6..75eb0a5 100644 --- a/schema/dialect/sqlite3/sqlite3.go +++ b/schema/dialect/sqlite3/sqlite3.go @@ -13,7 +13,7 @@ import ( ) var ( - sqlite3DefaultModifiers = []string{"VirtualAs", "StoredAs", "Hidden", "Nullable", "Default", "Increment"} + sqlite3DefaultModifiers = []string{"Hidden", "Nullable", "Default", "Increment"} sqlite3Serials = []string{"bigInteger", "integer", "mediumInteger", "smallInteger", "tinyInteger"} ) @@ -191,6 +191,9 @@ func (this Sqlite3) GetColumnModifier(modifier string, bp *schema.Blueprint, col return "NOT NULL" } case "Default": + if column.IsHidden() { + return "" + } if column.IsUseCurrent() { return " DEFAULT CURRENT_TIMESTAMP" } @@ -223,7 +226,10 @@ func (this Sqlite3) GetColumnType(column *schema.ColumnDefinition) string { case "text": return "text" case "integer": - return this.GenerateSQL("INTEGER(?)", column.Length) + if column.Length > 0 { + return this.GenerateSQL("INTEGER(?)", column.Length) + } + return "INTEGER" case "bigInteger": return "bigint(20)" case "tinyInteger": @@ -252,6 +258,8 @@ func (this Sqlite3) GetColumnType(column *schema.ColumnDefinition) string { return "year" case "uuid": return "char(36)" + case "vector": + return this.GenerateSQL("float[?]", column.Length) } panic("Unsupported data type: " + column.Type) } diff --git a/schema/dialect/sqlite3/sqlite3_test.go b/schema/dialect/sqlite3/sqlite3_test.go index 2783c02..3714577 100644 --- a/schema/dialect/sqlite3/sqlite3_test.go +++ b/schema/dialect/sqlite3/sqlite3_test.go @@ -1,32 +1,16 @@ package sqlite3_test import ( - "os" "testing" "git.fsdpf.net/go/db" "git.fsdpf.net/go/db/engine" "git.fsdpf.net/go/db/schema" + sqlmock "github.com/DATA-DOG/go-sqlmock" "github.com/stretchr/testify/suite" - - _ "github.com/mattn/go-sqlite3" ) -var ( - tableName = "entry" - - dropTable = "DROP TABLE IF EXISTS `entry`;" - - createTable = "CREATE TABLE IF NOT EXISTS `entry` (" + - "`id` INTEGER PRIMARY KEY AUTOINCREMENT," + - "`int` INT NOT NULL ," + - "`float` FLOAT NOT NULL ," + - "`string` VARCHAR(255) NOT NULL ," + - "`time` DATETIME NOT NULL ," + - "`bool` TINYINT NOT NULL ," + - "`bytes` BLOB NOT NULL" + - ");" -) +var tableName = "entry" type sqlite3Test struct { suite.Suite @@ -38,18 +22,28 @@ func TestSqlite3Suite(t *testing.T) { } func (t *sqlite3Test) SetupSuite() { - db := engine.Open(map[string]engine.DBConfig{ - "test-sqlite3": engine.NewDBConfig("sqlite3", - engine.WithSQLiteFile(os.Getenv("DB_FILE")), - ), + mockDB, mock, _ := sqlmock.New() + + // TestGetColumnListing: PRAGMA table_info(entry) + mock.ExpectQuery("PRAGMA table_info"). + WillReturnRows(sqlmock.NewRows([]string{"cid", "name", "type", "notnull", "dflt_value", "pk"}). + AddRow(0, "id", "INTEGER", 1, nil, 1). + AddRow(1, "int", "INT", 1, nil, 0). + AddRow(2, "float", "FLOAT", 1, nil, 0). + AddRow(3, "string", "VARCHAR(255)", 1, nil, 0). + AddRow(4, "time", "DATETIME", 1, nil, 0). + AddRow(5, "bool", "TINYINT", 1, nil, 0). + AddRow(6, "bytes", "BLOB", 1, nil, 0)) + + // TestTableExists: SELECT COUNT(*) FROM sqlite_master WHERE ... + mock.ExpectQuery("SELECT.*COUNT.*sqlite_master"). + WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + + conn := engine.Mock(map[string]engine.MockDBConfig{ + "test-sqlite3": {Driver: "sqlite3", Mock: mockDB}, }).Connection("test-sqlite3") - t.schema = schema.GetSchemaDialect(db) - - // db.Exec(dropTable) - if _, err := db.Exec(createTable); err != nil { - t.Require().NoError(err) - } + t.schema = schema.GetSchemaDialect(conn) } func (t *sqlite3Test) TestTableExists() { @@ -123,14 +117,14 @@ func (t *sqlite3Test) TestCompileChange() { sql := t.schema.CompileChange(bp) t.Equal(len([]string{ - "ALTER TABLE `users` ADD COLUMN `name_1742738970482105700` varchar(50) NOT NULL", - "UPDATE `users` SET `name_1742738970482105700`=`name`", + "ALTER TABLE `users` ADD COLUMN `name_tmp` varchar(50) NOT NULL", + "UPDATE `users` SET `name_tmp`=`name`", "ALTER TABLE `users` DROP COLUMN `name`", - "ALTER TABLE `users` RENAME COLUMN `name_1742738970482105700` TO `username`", - "ALTER TABLE `users` ADD COLUMN `age_1742738970482151800` smallint(4) NOT NULL DEFAULT '19'", - "UPDATE `users` SET `age_1742738970482151800`=`age`", + "ALTER TABLE `users` RENAME COLUMN `name_tmp` TO `username`", + "ALTER TABLE `users` ADD COLUMN `age_tmp` smallint(4) NOT NULL DEFAULT '19'", + "UPDATE `users` SET `age_tmp`=`age`", "ALTER TABLE `users` DROP COLUMN `age`", - "ALTER TABLE `users` RENAME COLUMN `age_1742738970482151800` TO `age`", + "ALTER TABLE `users` RENAME COLUMN `age_tmp` TO `age`", }), len(sql)) t.T().Log(sql) @@ -195,9 +189,9 @@ func (t *sqlite3Test) TestCompileCreate_HiddenColumn() { "CREATE TABLE IF NOT EXISTS `api_users` (\n" + "`id` INTEGER NOT NULL PRIMARY KEY AUTOINCREMENT\n" + "`name` varchar(100) NULL\n" + - "`age` INTEGER(0) NOT NULL\n" + + "`age` INTEGER NOT NULL\n" + "`token` varchar(255) HIDDEN\n" + - "`page_size` INTEGER(0) HIDDEN\n" + + "`page_size` INTEGER HIDDEN\n" + ")", }, sql) diff --git a/schema/dialect/sqlserver/sqlserver_test.go b/schema/dialect/sqlserver/sqlserver_test.go index 3c0e041..88953fa 100644 --- a/schema/dialect/sqlserver/sqlserver_test.go +++ b/schema/dialect/sqlserver/sqlserver_test.go @@ -1,14 +1,12 @@ package sqlserver_test import ( - "os" "testing" "git.fsdpf.net/go/db/engine" "git.fsdpf.net/go/db/schema" + sqlmock "github.com/DATA-DOG/go-sqlmock" "github.com/stretchr/testify/suite" - - _ "github.com/microsoft/go-mssqldb" ) type sqlServerTest struct { @@ -16,41 +14,27 @@ type sqlServerTest struct { schema schema.Schema } -var ( - tableName = "entry" - - dropTable = "IF OBJECT_ID('entry', 'U') IS NOT NULL DROP TABLE entry;" - - createTable = "CREATE TABLE entry (" + - "id INT NOT NULL IDENTITY(1,1) PRIMARY KEY," + - "int INT NOT NULL UNIQUE," + - "float REAL NOT NULL," + - "string NVARCHAR(255) NOT NULL," + - "time DATETIME2 NOT NULL," + - "bool BIT NOT NULL," + - "bytes VARBINARY(MAX) NOT NULL" + - ");" -) +var tableName = "entry" func TestSQLServerSuite(t *testing.T) { suite.Run(t, new(sqlServerTest)) } func (t *sqlServerTest) SetupSuite() { - db := engine.Open(map[string]engine.DBConfig{ - "test-sqlserver": engine.NewDBConfig("sqlserver", - engine.WithHost(os.Getenv("SQLSERVER_HOST")), - engine.WithPort(os.Getenv("SQLSERVER_PORT")), - engine.WithDatabase(os.Getenv("SQLSERVER_DB")), - engine.WithUsername(os.Getenv("SQLSERVER_USER")), - engine.WithPassword(os.Getenv("SQLSERVER_PASSWD")), - ), + mockDB, mock, _ := sqlmock.New() + + mock.ExpectQuery("SELECT.*COLUMN_NAME.*COLUMNS"). + WillReturnRows(sqlmock.NewRows([]string{"COLUMN_NAME"}). + AddRow("id").AddRow("int").AddRow("float"). + AddRow("string").AddRow("time").AddRow("bool").AddRow("bytes")) + + mock.ExpectQuery("SELECT.*COUNT.*TABLES"). + WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1)) + + conn := engine.Mock(map[string]engine.MockDBConfig{ + "test-sqlserver": {Driver: "sqlserver", Mock: mockDB}, }).Connection("test-sqlserver") - - t.schema = schema.GetSchemaDialect(db) - - db.Exec(dropTable) - db.Exec(createTable) + t.schema = schema.GetSchemaDialect(conn) } func (t *sqlServerTest) TestGetColumnListing() { @@ -80,14 +64,14 @@ func (t *sqlServerTest) TestCompileCreate() { sql := t.schema.CompileCreate(bp) t.Equal([]string{ - "CREATE TABLE [users] (\n" + - "[id] BIGINT NOT NULL IDENTITY(1,1) PRIMARY KEY COMMENT 'ID',\n" + - "[enabled] BIT NOT NULL DEFAULT 1 COMMENT '是否有效',\n" + - "[created_user] CHAR(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000' COMMENT '创建者',\n" + - "[owned_user] CHAR(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000' COMMENT '拥有者',\n" + - "[created_at] DATETIME2 NOT NULL DEFAULT GETDATE() COMMENT '创建时间',\n" + - "[updated_at] DATETIME2 NOT NULL DEFAULT GETDATE() COMMENT '更新时间',\n" + - "[deleted_at] DATETIME2 NULL COMMENT '删除时间'\n" + + "CREATE TABLE \"users\" (\n" + + "\"id\" BIGINT NOT NULL IDENTITY(1,1) PRIMARY KEY,\n" + + "\"enabled\" BIT NOT NULL DEFAULT '1',\n" + + "\"created_user\" CHAR(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000',\n" + + "\"owned_user\" CHAR(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000',\n" + + "\"created_at\" DATETIME2 NOT NULL DEFAULT GETDATE(),\n" + + "\"updated_at\" DATETIME2 NOT NULL DEFAULT GETDATE(),\n" + + "\"deleted_at\" DATETIME2 NULL\n" + ")", }, sql) @@ -104,10 +88,10 @@ func (t *sqlServerTest) TestCompileAdd() { sql := t.schema.CompileAdd(bp) t.Equal([]string{ - "ALTER TABLE [users]\n" + - "ADD [name] NVARCHAR(50) NOT NULL COMMENT '用户名',\n" + - "ADD [age] SMALLINT NOT NULL DEFAULT 18 COMMENT '年龄',\n" + - "ADD [sex] NVARCHAR(1) NOT NULL DEFAULT '0' COMMENT '性别'", + "ALTER TABLE \"users\"\n" + + "ADD \"name\" NVARCHAR(50) NOT NULL,\n" + + "ADD \"age\" SMALLINT NOT NULL DEFAULT '18',\n" + + "ADD \"sex\" NVARCHAR(1) NOT NULL DEFAULT '0'", }, sql) t.T().Log(sql) @@ -122,9 +106,9 @@ func (t *sqlServerTest) TestCompileChange() { sql := t.schema.CompileChange(bp) t.Equal([]string{ - "ALTER TABLE [users]\n" + - "ALTER COLUMN [name] NVARCHAR(50) NOT NULL,\n" + - "ALTER COLUMN [age] SMALLINT NOT NULL DEFAULT 19 COMMENT '年龄'", + "ALTER TABLE \"users\"\n" + + "ALTER COLUMN \"name\" NVARCHAR(50) NOT NULL,\n" + + "ALTER COLUMN \"age\" SMALLINT NOT NULL DEFAULT '19'", }, sql) t.T().Log(sql) @@ -138,9 +122,9 @@ func (t *sqlServerTest) TestCompileDropColumn() { sql := t.schema.CompileDropColumn(bp) t.Equal([]string{ - "ALTER TABLE [users]\n" + - "DROP COLUMN [age],\n" + - "DROP COLUMN [sex]", + "ALTER TABLE \"users\"\n" + + "DROP COLUMN \"age\",\n" + + "DROP COLUMN \"sex\"", }, sql) t.T().Log(sql) @@ -153,7 +137,7 @@ func (t *sqlServerTest) TestCompileDrop() { sql := t.schema.CompileDrop(bp) - t.Equal([]string{"DROP TABLE [users]"}, sql) + t.Equal([]string{"DROP TABLE \"users\""}, sql) t.T().Log(sql) } @@ -164,7 +148,7 @@ func (t *sqlServerTest) TestCompileDropIfExists() { bp.Drop() sql := t.schema.CompileDropIfExists(bp) - t.Equal([]string{"DROP TABLE IF EXISTS [users]"}, sql) + t.Equal([]string{"DROP TABLE IF EXISTS \"users\""}, sql) t.T().Log(sql) } @@ -176,7 +160,7 @@ func (t *sqlServerTest) TestCompileRename() { sql := t.schema.CompileRename(bp) - t.Equal([]string{"EXEC sp_rename [users], [user]"}, sql) + t.Equal([]string{"EXEC sp_rename \"users\", \"user\""}, sql) t.T().Log(sql) } diff --git a/select_dataset.go b/select_dataset.go index 1523332..35a829f 100644 --- a/select_dataset.go +++ b/select_dataset.go @@ -133,10 +133,11 @@ func (sd *SelectDataset) Insert() *InsertDataset { i := newInsertDataset(sd.dialect.Dialect(), sd.queryFactory). Prepared(sd.isPrepared.Bool()) if sd.clauses.HasSources() { - if from, ok := sd.GetClauses().From().Columns()[0].(exp.AliasedExpression); ok { + col := sd.GetClauses().From().Columns()[0] + if from, ok := col.(exp.AliasedExpression); ok { i = i.Into(from.Aliased()) } else { - i = i.Into(from) + i = i.Into(col) } } c := i.clauses