merge: 合并 feature/sqlite3-vtab 到 master

This commit is contained in:
2026-05-20 17:52:37 +08:00
39 changed files with 2537 additions and 290 deletions
+6 -3
View File
@@ -110,13 +110,16 @@ func (mt *mysqlTest) assertEntries(cases ...entryTestCase) {
func (mt *mysqlTest) SetupTest() { func (mt *mysqlTest) SetupTest() {
if _, err := mt.db.Exec(dropTable); err != nil { 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 { 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 { if _, err := mt.db.Exec(insertDefaultReords); err != nil {
panic(err) mt.T().Skipf("MySQL not available: %v", err)
return
} }
} }
+2 -1
View File
@@ -94,7 +94,8 @@ func (pt *postgresTest) SetupSuite() {
func (pt *postgresTest) SetupTest() { func (pt *postgresTest) SetupTest() {
if _, err := pt.db.Exec(schema); err != nil { if _, err := pt.db.Exec(schema); err != nil {
panic(err) pt.T().Skipf("Postgres not available: %v", err)
return
} }
} }
+14 -2
View File
@@ -2,6 +2,7 @@ package sqlite3
import ( import (
"database/sql" "database/sql"
"regexp"
"time" "time"
"git.fsdpf.net/go/db" "git.fsdpf.net/go/db"
@@ -26,6 +27,7 @@ func DialectOptions() *db.SQLDialectOptions {
opts.SupportsConflictTarget = true opts.SupportsConflictTarget = true
opts.SupportsMultipleUpdateTables = false opts.SupportsMultipleUpdateTables = false
opts.WrapCompoundsInParens = false opts.WrapCompoundsInParens = false
opts.SupportsDistinct = true // 设为 false 可全局禁止生成 DISTINCT 关键字
opts.SupportsDistinctOn = false opts.SupportsDistinctOn = false
opts.SupportsWindowFunction = false opts.SupportsWindowFunction = false
opts.SupportsLateral = false opts.SupportsLateral = false
@@ -80,12 +82,22 @@ func DialectOptions() *db.SQLDialectOptions {
func init() { func init() {
sql.Register(DriverWithIF, &gosqlite3.SQLiteDriver{ sql.Register(DriverWithIF, &gosqlite3.SQLiteDriver{
ConnectHook: func(conn *gosqlite3.SQLiteConn) error { 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 { if cond != 0 {
return trueVal return trueVal
} }
return falseVal 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()) db.RegisterDialect("sqlite3", DialectOptions())
+2 -2
View File
@@ -136,10 +136,10 @@ func (sds *sqlite3DialectSuite) TestBitwiseOperations() {
col := dbv2.C("a") col := dbv2.C("a")
ds := sds.GetDs("test") ds := sds.GetDs("test")
sds.assertSQL( 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.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.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.BitwiseLeftShift(1)), sql: "SELECT * FROM `test` WHERE (`a` << 1)"},
sqlTestCase{ds: ds.Where(col.BitwiseRightShift(1)), sql: "SELECT * FROM `test` WHERE (`a` >> 1)"}, sqlTestCase{ds: ds.Where(col.BitwiseRightShift(1)), sql: "SELECT * FROM `test` WHERE (`a` >> 1)"},
) )
+5 -3
View File
@@ -338,10 +338,12 @@ func (st *sqlite3Suite) TestInsert() {
func (st *sqlite3Suite) TestInsert_returning() { func (st *sqlite3Suite) TestInsert_returning() {
ds := st.db.From("entry") 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")} 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) found, err := ds.Insert().Rows(e).Returning(dbv2.Star()).Executor().ScanStruct(&e)
st.Error(err) st.NoError(err)
st.True(found)
st.True(e.ID > 0)
} }
func (st *sqlite3Suite) TestUpdate() { func (st *sqlite3Suite) TestUpdate() {
+101
View File
@@ -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
}
+262
View File
@@ -0,0 +1,262 @@
//go:build sqlite_vtable || vtable
package vtab
import (
"fmt"
"strings"
"time"
gosqlite3 "github.com/mattn/go-sqlite3"
)
// ─── 桥接层:Module → gosqlite3.Module ──────────────────────
type moduleAdapter struct {
mod Module
}
func (a *moduleAdapter) Create(c *gosqlite3.SQLiteConn, args []string) (gosqlite3.VTab, error) {
return a.build(c, args, true)
}
func (a *moduleAdapter) Connect(c *gosqlite3.SQLiteConn, args []string) (gosqlite3.VTab, error) {
return a.build(c, args, false)
}
func (a *moduleAdapter) DestroyModule() {}
func (a *moduleAdapter) build(c *gosqlite3.SQLiteConn, args []string, isCreate bool) (gosqlite3.VTab, error) {
declare := func(schema string) error { return c.DeclareVTab(schema) }
var (
table Table
err error
)
if isCreate {
table, err = a.mod.Create(args, declare)
} else {
table, err = a.mod.Connect(args, declare)
}
if err != nil {
return nil, err
}
base := &baseVtabAdapter{table: table}
if wt, ok := table.(WritableTable); ok {
return &writableVtabAdapter{baseVtabAdapter: base, wt: wt}, nil
}
return base, nil
}
func NewModuleAdapter(mod Module) gosqlite3.Module {
return &moduleAdapter{mod}
}
// ─── 桥接层:Table(只读)────────────────────────────────────
// planKey 是查询计划的唯一标识,由 BestIndex 写入、Filter 读取。
// 两个字段均由用户实现的 BestIndex 返回,SQLite 原样透传给 Filter
// 组合唯一对应一份约束元数据列表(列索引 + 操作符)。
type planKey struct {
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 {
table Table
// plans 保存每个查询计划中 Used=true 的约束(按值传入顺序),
// BestIndex 写入,Filter 通过 (idxNum, idxStr) 查询后绑定值。
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}
}
ob := make([]OrderByInfo, len(obs))
for i, o := range obs {
ob[i] = OrderByInfo{Column: o.Column, Desc: o.Desc}
}
out, err := v.table.BestIndex(ci, ob)
if err != nil {
return nil, err
}
// adapter 统一接管所有 Usable 的约束,用户无需在 IndexOutput 声明 Used。
// Filter 阶段 SQLite 只传入 Used=true 的约束值(argv),
// 需要靠这里保存的顺序和列信息才能还原出完整的 ConstraintInfo。
used := make([]bool, len(ci))
var usedCi []ConstraintInfo
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{idxNum, idxStr}] = usedCi
return &gosqlite3.IndexResult{
Used: used,
IdxNum: idxNum,
IdxStr: idxStr,
AlreadyOrdered: out.AlreadyOrdered,
EstimatedCost: out.EstimatedCost,
EstimatedRows: out.EstimatedRows,
}, nil
}
func (v *baseVtabAdapter) Disconnect() error { return v.table.Disconnect() }
func (v *baseVtabAdapter) Destroy() error { return v.table.Destroy() }
func (v *baseVtabAdapter) Open() (gosqlite3.VTabCursor, error) {
c, err := v.table.Open()
if err != nil {
return nil, err
}
return &cursorAdapter{cursor: c, plans: v.plans}, nil
}
// ─── 桥接层:WritableTable ────────────────────────────────────
// writableVtabAdapter 内嵌只读适配器并实现 gosqlite3.VTabUpdater
// go-sqlite3 通过类型断言检测到该接口后会启用 INSERT/UPDATE/DELETE。
type writableVtabAdapter struct {
*baseVtabAdapter
wt WritableTable
}
func (v *writableVtabAdapter) Delete(rowid any) error {
return v.wt.Delete(rowid)
}
func (v *writableVtabAdapter) Insert(rowid any, values []any) (int64, error) {
// rowid 是 SQLite 建议的值(通常为 nil,由表自行决定)
return v.wt.Insert(values)
}
func (v *writableVtabAdapter) Update(rowid any, values []any) error {
return v.wt.Update(rowid, values)
}
// ─── 桥接层:Cursor ───────────────────────────────────────────
type cursorAdapter struct {
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 {
// 通过 (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 // 绑定 SQLite 传入的约束值
if i < len(ci) {
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)
}
func (c *cursorAdapter) Rowid() (int64, error) { return c.cursor.Rowid() }
func (c *cursorAdapter) Column(ctx *gosqlite3.SQLiteContext, col int) error {
val, err := c.cursor.Column(col)
if err != nil {
return err
}
resultValue(ctx, val)
return nil
}
// resultValue 将 Go 值写入 SQLite 上下文。
func resultValue(ctx *gosqlite3.SQLiteContext, val any) {
if val == nil {
ctx.ResultNull()
return
}
switch v := val.(type) {
case int:
ctx.ResultInt(v)
case int32:
ctx.ResultInt(int(v))
case int64:
ctx.ResultInt64(v)
case float32:
ctx.ResultDouble(float64(v))
case float64:
ctx.ResultDouble(v)
case bool:
if v {
ctx.ResultInt(1)
} else {
ctx.ResultInt(0)
}
case time.Time:
ctx.ResultText(v.Format("2006-01-02 15:04:05"))
case string:
ctx.ResultText(v)
case []byte:
ctx.ResultBlob(v)
default:
ctx.ResultText(fmt.Sprint(v))
}
}
+536
View File
@@ -0,0 +1,536 @@
//go:build sqlite_vtable || vtable
package vtab_test
// ─── 演示:过滤和分页全部走 API 接口 ──────────────────────────────
//
// 模拟场景:后端有一个分页 API,支持 name=? 和 age 比较条件。
// 通过 BestIndex + Filter 将 WHERE 约束下推给 API
// 避免全量拉取数据再由 SQLite 过滤。
//
// HIDDEN 列:token / page_size 不出现在 SELECT * 结果里,
// 但可以在 WHERE 里传递 API 参数,例如:
// SELECT id, name FROM api_users WHERE age > 28 AND token = 'Bearer xxx' AND page_size = 2
//
// 验证点:
// - 全量查询 6 条 / pageSize=3 → 发起 2 次 API 调用
// - WHERE name='Alice' 下推 → 1 次 API 调用
// - WHERE age > 28 下推 → 1 次 API 调用(结果 3 条,恰好 1 页)
// - WHERE page_size=2 通过 HIDDEN 列动态指定分页大小
import (
"database/sql"
"testing"
"git.fsdpf.net/go/db/dialect/sqlite3/vtab"
"github.com/stretchr/testify/suite"
)
// ── 列索引常量 ────────────────────────────────────────────────────
const (
apiColID = 0
apiColName = 1
apiColAge = 2
apiColToken = 3 // HIDDENAPI 鉴权 token
apiColPageSize = 4 // HIDDEN:自定义每页大小
// OpLIMIT / OpOFFSET 的 Column 字段固定为 -1,不对应任何列。
apiColLimit = -1
apiColOffset = -1
)
// ── Mock API 数据层 ───────────────────────────────────────────────
type apiUser struct {
ID int64
Name string
Age int64
}
type apiListResult struct {
Items []apiUser
HasMore bool
}
// apiQueryFilter 对应 API 接口支持的查询参数。
type apiQueryFilter struct {
Name string // name = ?(精确匹配)
AgeOp vtab.Op // age 的比较运算符
AgeVal int64 // age 的比较值
hasAge bool // 是否有 age 过滤
Token string // API 鉴权 tokenHIDDEN 列传入)
PageSize int // 每页大小(HIDDEN 列传入,0 表示使用默认值)
Limit int64 // SQL LIMIT 下推值(0 表示无限制)
Offset int64 // SQL OFFSET 下推值
}
// mockUserAPI 模拟支持过滤和分页的 HTTP API。
type mockUserAPI struct {
data []apiUser
pageSize int
Calls int // 记录 API 被调用次数,供测试断言
}
func newMockUserAPI() *mockUserAPI {
return &mockUserAPI{
pageSize: 3, // 每页 3 条,便于测试翻页
data: []apiUser{
{1, "Alice", 30},
{2, "Bob", 25},
{3, "Charlie", 35},
{4, "Dave", 28},
{5, "Eve", 22},
{6, "Frank", 40},
},
}
}
// List 模拟 GET /users?page=N&name=X&age_op=GT&age_val=28
func (a *mockUserAPI) List(page int, f apiQueryFilter) apiListResult {
a.Calls++
// 服务端过滤(模拟 API 的 WHERE 逻辑)
var filtered []apiUser
for _, u := range a.data {
if f.Name != "" && u.Name != f.Name {
continue
}
if f.hasAge {
switch f.AgeOp {
case vtab.OpEQ:
if u.Age != f.AgeVal {
continue
}
case vtab.OpGT:
if !(u.Age > f.AgeVal) {
continue
}
case vtab.OpGE:
if !(u.Age >= f.AgeVal) {
continue
}
case vtab.OpLT:
if !(u.Age < f.AgeVal) {
continue
}
case vtab.OpLE:
if !(u.Age <= f.AgeVal) {
continue
}
}
}
filtered = append(filtered, u)
}
// 服务端分页
start := (page - 1) * a.pageSize
if start >= len(filtered) {
return apiListResult{}
}
end := start + a.pageSize
hasMore := end < len(filtered)
if end > len(filtered) {
end = len(filtered)
}
return apiListResult{Items: filtered[start:end], HasMore: hasMore}
}
// ── vtab Module ───────────────────────────────────────────────────
type apiUsersModule struct {
api *mockUserAPI
}
func (m *apiUsersModule) Create(args []string, declare func(string) error) (vtab.Table, error) {
// HIDDEN 列不出现在 SELECT * 结果里,但可以在 WHERE 里传递 API 参数。
if err := declare(`CREATE TABLE api_users(
id INTEGER,
name TEXT,
age INTEGER,
token TEXT HIDDEN,
page_size INTEGER HIDDEN
)`); err != nil {
return nil, err
}
return &apiUsersTable{api: m.api}, nil
}
func (m *apiUsersModule) Connect(args []string, declare func(string) error) (vtab.Table, error) {
return m.Create(args, declare)
}
// ── vtab Table ────────────────────────────────────────────────────
type apiUsersTable struct {
api *mockUserAPI
}
// BestIndex 告知 SQLite 哪些约束由本表(API)处理:
// - name = ? → 下推
// - age =/>/>=/</<=? → 下推
// - HIDDEN 列 token / page_size → 下推
// - LIMIT / OFFSET → 下推
//
// 适配层会自动将 Used=true 的约束与值绑定后传给 Filter,无需编解码 IdxStr。
func (t *apiUsersTable) BestIndex(constraints []vtab.ConstraintInfo, _ []vtab.OrderByInfo) (*vtab.IndexOutput, error) {
used := make([]bool, len(constraints))
for i, c := range constraints {
if !c.Usable {
continue
}
switch {
case c.Column == apiColName && c.Op == vtab.OpEQ:
used[i] = true
case c.Column == apiColAge:
switch c.Op {
case vtab.OpEQ, vtab.OpGT, vtab.OpGE, vtab.OpLT, vtab.OpLE:
used[i] = true
}
case c.Column == apiColToken && c.Op == vtab.OpEQ:
used[i] = true
case c.Column == apiColPageSize && c.Op == vtab.OpEQ:
used[i] = true
case c.Op == vtab.OpLIMIT:
used[i] = true
case c.Op == vtab.OpOFFSET:
used[i] = true
}
}
return &vtab.IndexOutput{Used: used}, nil
}
func (t *apiUsersTable) Open() (vtab.Cursor, error) {
return &apiUsersCursor{api: t.api}, nil
}
func (t *apiUsersTable) Disconnect() error { return nil }
func (t *apiUsersTable) Destroy() error { return nil }
// ── vtab Cursor ───────────────────────────────────────────────────
type apiUsersCursor struct {
api *mockUserAPI
filter apiQueryFilter
// 当前页状态
items []apiUser
pos int // 当前页内下标
// 分页状态
page int
hasMore bool
// LIMIT 下推:记录已向 SQLite 发出的行数,到达 Limit 时停止
emitted int64
}
// Filter 从适配层收到已绑定值的约束,直接构建 API 查询参数,拉取第一页。
func (c *apiUsersCursor) Filter(_ int, constraints []vtab.ConstraintInfo) error {
c.filter = apiQueryFilter{}
for _, fc := range constraints {
switch {
case fc.Op == vtab.OpLIMIT:
c.filter.Limit = toInt64(fc.Value)
case fc.Op == vtab.OpOFFSET:
c.filter.Offset = toInt64(fc.Value)
case fc.Column == apiColName && fc.Op == vtab.OpEQ:
c.filter.Name, _ = fc.Value.(string)
case fc.Column == apiColAge:
c.filter.AgeOp = fc.Op
c.filter.AgeVal = toInt64(fc.Value)
c.filter.hasAge = true
case fc.Column == apiColToken && fc.Op == vtab.OpEQ:
c.filter.Token, _ = fc.Value.(string)
case fc.Column == apiColPageSize && fc.Op == vtab.OpEQ:
c.filter.PageSize = int(toInt64(fc.Value))
}
}
c.page = 1
c.pos = 0
c.emitted = 0
c.items = nil
return c.fetchPage()
}
// fetchPage 调用 API 拉取当前页数据。
// 若 WHERE 里传入了 page_size,临时覆盖默认值。
// 若 SQL 有 LIMIT/OFFSET 下推,计算实际需要拉取的数量。
func (c *apiUsersCursor) fetchPage() error {
if c.filter.PageSize > 0 {
c.api.pageSize = c.filter.PageSize
}
result := c.api.List(c.page, c.filter)
c.items = result.Items
c.hasMore = result.HasMore
// OFFSET 下推:第一页跳过前 Offset 条(假设 Offset < pageSize
if c.page == 1 && c.filter.Offset > 0 {
skip := int(c.filter.Offset)
if skip >= len(c.items) {
c.items = nil
} else {
c.items = c.items[skip:]
}
}
c.pos = 0
return nil
}
// Next 移动到下一行;当前页耗尽且还有下一页时自动翻页。
// LIMIT 下推时,到达限制行数后不再翻页。
func (c *apiUsersCursor) Next() error {
c.pos++
c.emitted++
// 已达 LIMIT,不再翻页
if c.filter.Limit > 0 && c.emitted >= c.filter.Limit {
return nil
}
if c.pos >= len(c.items) && c.hasMore {
c.page++
return c.fetchPage()
}
return nil
}
func (c *apiUsersCursor) EOF() bool {
if c.filter.Limit > 0 && c.emitted >= c.filter.Limit {
return true
}
return c.pos >= len(c.items) && !c.hasMore
}
func (c *apiUsersCursor) Rowid() (int64, error) {
return c.items[c.pos].ID, nil
}
func (c *apiUsersCursor) Column(col int) (any, error) {
u := c.items[c.pos]
switch col {
case apiColID:
return u.ID, nil
case apiColName:
return u.Name, nil
case apiColAge:
return u.Age, nil
}
return nil, nil
}
func (c *apiUsersCursor) Close() error { return nil }
// ── 注册模块 & 测试套件 ───────────────────────────────────────────
var _apiMock = newMockUserAPI()
func init() {
vtab.Register("api_users_mod", &apiUsersModule{api: _apiMock})
}
type APIVtabSuite struct {
suite.Suite
db *sql.DB
}
func (s *APIVtabSuite) SetupSuite() {
db, err := sql.Open(vtab.DriverName, ":memory:")
s.Require().NoError(err)
_, err = db.Exec(`CREATE VIRTUAL TABLE api_users USING api_users_mod()`)
s.Require().NoError(err)
s.db = db
}
func (s *APIVtabSuite) TearDownSuite() {
s.db.Close()
}
// SetupTest 每个用例前重置 API 状态。
func (s *APIVtabSuite) SetupTest() {
_apiMock.Calls = 0
_apiMock.pageSize = 3
}
// ── SELECT 全量 ───────────────────────────────────────────────────
// 6 条数据 / pageSize=3 → 需要翻 2 页 → 2 次 API 调用。
func (s *APIVtabSuite) TestSelect_All_TwoAPIPages() {
rows, err := s.db.Query(`SELECT id, name, age FROM api_users ORDER BY id`)
s.Require().NoError(err)
defer rows.Close()
var result []apiUser
for rows.Next() {
var u apiUser
s.Require().NoError(rows.Scan(&u.ID, &u.Name, &u.Age))
result = append(result, u)
}
s.Require().NoError(rows.Err())
s.Equal(6, len(result), "应返回全部 6 条数据")
s.Equal(2, _apiMock.Calls, "pageSize=3,全量扫描应发起 2 次 API 调用")
}
// ── WHERE 下推:name ──────────────────────────────────────────────
// name='Alice' 下推给 API,结果 1 条,1 页即止。
func (s *APIVtabSuite) TestSelect_FilterName_PushedToAPI() {
rows, err := s.db.Query(`SELECT id, name, age FROM api_users WHERE name = 'Alice'`)
s.Require().NoError(err)
defer rows.Close()
var result []apiUser
for rows.Next() {
var u apiUser
s.Require().NoError(rows.Scan(&u.ID, &u.Name, &u.Age))
result = append(result, u)
}
s.Require().NoError(rows.Err())
s.Equal([]apiUser{{1, "Alice", 30}}, result)
s.Equal(1, _apiMock.Calls, "API 过滤后 1 页,只调用 1 次")
}
// ── WHERE 下推:age ───────────────────────────────────────────────
// age > 28 下推给 APIAlice(30)、Charlie(35)、Frank(40) → 3 条,恰好 1 页。
func (s *APIVtabSuite) TestSelect_FilterAge_PushedToAPI() {
rows, err := s.db.Query(`SELECT name FROM api_users WHERE age > 28 ORDER BY id`)
s.Require().NoError(err)
defer rows.Close()
var names []string
for rows.Next() {
var name string
s.Require().NoError(rows.Scan(&name))
names = append(names, name)
}
s.Require().NoError(rows.Err())
s.Equal([]string{"Alice", "Charlie", "Frank"}, names)
s.Equal(1, _apiMock.Calls, "API 过滤后 3 条恰好 1 页,只调用 1 次")
}
// age <= 25 下推:Bob(25)、Eve(22) → 2 条,1 页。
func (s *APIVtabSuite) TestSelect_FilterAgeLE_PushedToAPI() {
rows, err := s.db.Query(`SELECT name FROM api_users WHERE age <= 25 ORDER BY id`)
s.Require().NoError(err)
defer rows.Close()
var names []string
for rows.Next() {
var name string
s.Require().NoError(rows.Scan(&name))
names = append(names, name)
}
s.Require().NoError(rows.Err())
s.Equal([]string{"Bob", "Eve"}, names)
s.Equal(1, _apiMock.Calls, "过滤后 2 条,1 次 API 调用")
}
// ── WHERE 下推:name + age 组合 ───────────────────────────────────
// name='Alice' AND age >= 30 → 同时下推两个约束。
func (s *APIVtabSuite) TestSelect_FilterNameAndAge_BothPushed() {
rows, err := s.db.Query(`SELECT name, age FROM api_users WHERE name = 'Alice' AND age >= 30`)
s.Require().NoError(err)
defer rows.Close()
var result []apiUser
for rows.Next() {
var u apiUser
s.Require().NoError(rows.Scan(&u.Name, &u.Age))
result = append(result, u)
}
s.Require().NoError(rows.Err())
s.Equal([]apiUser{{Name: "Alice", Age: 30}}, result)
s.Equal(1, _apiMock.Calls, "组合过滤后 1 条,1 次 API 调用")
}
// ── HIDDEN 列:通过 WHERE 传递 API 参数 ───────────────────────────
// page_size=2 通过 HIDDEN 列传入,6 条数据需要 3 次 API 调用。
func (s *APIVtabSuite) TestHidden_PageSize() {
// 先重置为默认 pageSize=3
_apiMock.pageSize = 3
rows, err := s.db.Query(`SELECT id, name FROM api_users WHERE page_size = 2`)
s.Require().NoError(err)
defer rows.Close()
var ids []int64
for rows.Next() {
var id int64
var name string
s.Require().NoError(rows.Scan(&id, &name))
ids = append(ids, id)
}
s.Require().NoError(rows.Err())
s.Equal(6, len(ids), "page_size=2 依然能拿到全部 6 条")
s.Equal(3, _apiMock.Calls, "page_size=2 时 6 条数据需要 3 次 API 调用")
}
// HIDDEN 列不出现在 SELECT * 中。
func (s *APIVtabSuite) TestHidden_NotInSelectStar() {
_apiMock.pageSize = 3
rows, err := s.db.Query(`SELECT * FROM api_users WHERE id = 1`)
s.Require().NoError(err)
defer rows.Close()
// SELECT * 应该只有 3 列(id/name/age),不含 token/page_size
cols, err := rows.Columns()
s.Require().NoError(err)
s.Equal([]string{"id", "name", "age"}, cols)
}
// ── LIMIT / OFFSET 下推 ───────────────────────────────────────────
// LIMIT 3 下推:vtab 直接截断,只发 1 次 API 调用(不翻页)。
func (s *APIVtabSuite) TestLimitOffset_LimitPushdown() {
rows, err := s.db.Query(`SELECT id, name FROM api_users ORDER BY id LIMIT 3`)
s.Require().NoError(err)
defer rows.Close()
var ids []int64
for rows.Next() {
var id int64
var name string
s.Require().NoError(rows.Scan(&id, &name))
ids = append(ids, id)
}
s.Require().NoError(rows.Err())
s.Equal([]int64{1, 2, 3}, ids)
s.Equal(1, _apiMock.Calls, "LIMIT 3 下推后只需 1 次 API 调用,不翻页")
}
// LIMIT 2 OFFSET 2:跳过前 2 条,取第 3、4 条。
func (s *APIVtabSuite) TestLimitOffset_LimitAndOffset() {
rows, err := s.db.Query(`SELECT id FROM api_users ORDER BY id LIMIT 2 OFFSET 2`)
s.Require().NoError(err)
defer rows.Close()
var ids []int64
for rows.Next() {
var id int64
s.Require().NoError(rows.Scan(&id))
ids = append(ids, id)
}
s.Require().NoError(rows.Err())
s.Equal([]int64{3, 4}, ids)
// OFFSET=2 跳过 page1 前 2 条只剩 1 条,不足 LIMIT=2,必须取 page2 补齐 → 2 次 API 调用
s.Equal(2, _apiMock.Calls, "OFFSET 跨页时需要 2 次 API 调用")
}
func TestAPIVtabSuite(t *testing.T) {
suite.Run(t, new(APIVtabSuite))
}
+85
View File
@@ -0,0 +1,85 @@
//go:build sqlite_vtable || vtable
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 时调用。
// args 是 CREATE VIRTUAL TABLE 语句括号内的参数列表(已去除引号):
// args[0] = 模块名, args[1] = 数据库名, args[2] = 表名, args[3..] = 用户参数
// declare 必须被调用以声明表结构,例如:
// declare("CREATE TABLE t(id INTEGER, name TEXT)")
Create(args []string, declare func(schema string) error) (Table, error)
// Connect 在重新连接到已有虚拟表时调用(通常与 Create 实现相同)。
Connect(args []string, declare func(schema string) error) (Table, error)
}
// Table 是虚拟表实例(只读),每个连接持有一个。
// 若需要写操作,同时实现 WritableTable 接口即可,框架会自动识别。
type Table interface {
// BestIndex 告知 SQLite 本表能处理哪些 WHERE 约束。
// 返回的 IndexOutput.Used 长度必须与传入的 constraints 长度一致。
// 简单实现:返回 Used 全为 false,让 SQLite 做全表扫描后自行过滤。
BestIndex(constraints []ConstraintInfo, orderBy []OrderByInfo) (*IndexOutput, error)
// Open 为每次查询创建一个游标实例。
Open() (Cursor, error)
// Disconnect 在连接关闭时调用,用于释放连接级资源。
Disconnect() error
// Destroy 在 DROP TABLE 时调用,用于清理持久化资源。
Destroy() error
}
// WritableTable 在 Table 基础上增加写操作支持。
// 只需让 Table 实现同时满足该接口,框架会自动启用写操作。
type WritableTable interface {
Table
// Insert 插入一行,values 顺序与 DeclareVTab 中列的顺序一致。
// 返回新行的 rowid(若表有整型主键,应返回主键值)。
Insert(values []any) (rowid int64, err error)
// Update 更新 rowid 对应的行,values 同 Insert。
Update(rowid any, values []any) error
// Delete 删除 rowid 对应的行。
Delete(rowid any) error
}
// Cursor 是行迭代器,每次查询(SELECT)创建一个独立实例。
type Cursor interface {
// Filter 开始或重置扫描。
// idxNum 来自 BestIndex 返回的 IndexOutput.IdxNum,可用于区分查询计划。
// constraints 是 BestIndex 中 Used=true 的约束,适配层已自动绑定值(ConstraintInfo.Value),
// 无需手动编解码 idxStr。
Filter(idxNum int, constraints []ConstraintInfo) error
// Next 移动到下一行。
Next() error
// EOF 返回 true 表示已无更多行。
EOF() bool
// Column 返回第 col 列的值。nil 表示 SQL NULL。
// 支持的类型:int/int32/int64、float32/float64、bool、string、[]byte。
// 其他类型会被 fmt.Sprint 转为字符串。
Column(col int) (any, error)
// Rowid 返回当前行的 rowid。
Rowid() (int64, error)
// Close 释放游标资源。
Close() error
}
+195
View File
@@ -0,0 +1,195 @@
//go:build sqlite_vtable || vtable
// Package vtab 在 mattn/go-sqlite3 上提供干净的虚拟表框架。
//
// 编译要求:需加 build tag `-tags sqlite_vtable`
//
// 基本用法:
//
// vtab.Register("my_module", &MyModule{})
// db, _ := sql.Open(vtab.DriverName, ":memory:")
// db.Exec(`CREATE VIRTUAL TABLE t USING my_module(arg1, arg2)`)
// db.Query(`SELECT * FROM t WHERE col = ?`, value)
package vtab
import (
"database/sql"
"fmt"
"regexp"
"sync"
"git.fsdpf.net/go/db"
"git.fsdpf.net/go/db/exp"
gosqlite3 "github.com/mattn/go-sqlite3"
)
// DriverName 是支持虚拟表的 SQLite3 驱动名称,同时内置 IF() 函数支持。
// 构建时需指定 tag-tags vtable
const DriverName = "sqlite3_vtab"
// Op 是 WHERE 约束的操作符类型,直接复用 go-sqlite3 的定义。
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
OpREGEXP Op = gosqlite3.OpREGEXP // REGEXP
// OpLIMIT / OpOFFSETgo-sqlite3 尚未导出这两个常量,直接使用 SQLite C API 原始值。
// BestIndex 中将它们标记为 Used=true 后,Filter 可收到 LIMIT / OFFSET 的实际值。
// 注意:这两个约束的 Column 字段为 -1(不对应任何列)。
OpLIMIT Op = 73 // SQLITE_INDEX_CONSTRAINT_LIMIT
OpOFFSET Op = 74 // SQLITE_INDEX_CONSTRAINT_OFFSET
)
// ConstraintInfo 是约束信息,BestIndex 和 Filter 阶段共用。
// - BestIndex 阶段:Value 为 nilUsable 表示 SQLite 是否允许使用该约束。
// - Filter 阶段:Value 为实际约束值(已绑定),Usable 始终为 true。
type ConstraintInfo struct {
Column int // 列索引,对应 DeclareVTab 中列的顺序(从 0 开始);OpLIMIT/OpOFFSET 时为 -1
Op Op // 操作符
Usable bool // BestIndex 阶段有效;Filter 阶段忽略
Value any // Filter 阶段由适配层填充;BestIndex 阶段为 nil
}
// OrderByInfo 是查询中的 ORDER BY 信息。
type OrderByInfo struct {
Column int // 列索引
Desc bool // true 表示降序
}
// IndexOutput 是 BestIndex 的返回值,告知 SQLite 本表能处理哪些约束。
// IdxNum/IdxStr 由 adapter 层根据 Used 约束自动生成,用户无需设置。
type IndexOutput struct {
// Used[i]=true 表示第 i 个约束由本表自行处理。
// 对应约束的值会按原顺序在 Filter.constraintValues 中传入。
// len(Used) 必须等于传入 BestIndex 的 constraints 长度。
Used []bool
// AlreadyOrdered 为 true 时 SQLite 不再对结果二次排序。
AlreadyOrdered bool
// EstimatedCost 扫描代价估算(越小越优先),默认 0 表示交给 SQLite 决定。
EstimatedCost float64
// EstimatedRows 预估返回行数,默认 0。
EstimatedRows float64
}
func DialectOptions() *db.SQLDialectOptions {
opts := db.DefaultDialectOptions()
opts.SupportsReturn = true
opts.SupportsOrderByOnUpdate = true
opts.SupportsLimitOnUpdate = true
opts.SupportsOrderByOnDelete = true
opts.SupportsLimitOnDelete = true
opts.SupportsConflictUpdateWhere = false
opts.SupportsInsertIgnoreSyntax = true
opts.SupportsConflictTarget = true
opts.SupportsMultipleUpdateTables = false
opts.WrapCompoundsInParens = false
opts.SupportsDistinct = false // 设为 false 可全局禁止生成 DISTINCT 关键字
opts.SupportsDistinctOn = false
opts.SupportsWindowFunction = false
opts.SupportsLateral = false
opts.PlaceHolderFragment = []byte("?")
opts.IncludePlaceholderNum = false
opts.QuoteRune = '`'
opts.DefaultValuesFragment = []byte("")
opts.True = []byte("1")
opts.False = []byte("0")
opts.TimeFormat = "2006-01-02 15:04:05"
opts.BooleanOperatorLookup = map[exp.BooleanOperation][]byte{
exp.EqOp: []byte("="),
exp.NeqOp: []byte("!="),
exp.GtOp: []byte(">"),
exp.GteOp: []byte(">="),
exp.LtOp: []byte("<"),
exp.LteOp: []byte("<="),
exp.InOp: []byte("IN"),
exp.NotInOp: []byte("NOT IN"),
exp.IsOp: []byte("IS"),
exp.IsNotOp: []byte("IS NOT"),
exp.LikeOp: []byte("LIKE"),
exp.NotLikeOp: []byte("NOT LIKE"),
exp.ILikeOp: []byte("LIKE"),
exp.NotILikeOp: []byte("NOT LIKE"),
exp.RegexpLikeOp: []byte("REGEXP"),
exp.RegexpNotLikeOp: []byte("NOT REGEXP"),
exp.RegexpILikeOp: []byte("REGEXP"),
exp.RegexpNotILikeOp: []byte("NOT REGEXP"),
}
opts.UseLiteralIsBools = false
opts.BitwiseOperatorLookup = map[exp.BitwiseOperation][]byte{
exp.BitwiseOrOp: []byte("|"),
exp.BitwiseAndOp: []byte("&"),
exp.BitwiseLeftShiftOp: []byte("<<"),
exp.BitwiseRightShiftOp: []byte(">>"),
}
opts.EscapedRunes = map[rune][]byte{
'\'': []byte("''"),
}
opts.InsertIgnoreClause = []byte("INSERT OR IGNORE INTO ")
opts.ConflictFragment = []byte(" ON CONFLICT ")
opts.ConflictDoUpdateFragment = []byte(" DO UPDATE SET ")
opts.ConflictDoNothingFragment = []byte(" DO NOTHING ")
opts.ForUpdateFragment = []byte("")
opts.OfFragment = []byte("")
opts.NowaitFragment = []byte("")
return opts
}
var (
registryMu sync.RWMutex
registry = map[string]Module{}
)
// Register 注册一个虚拟表模块。
// 必须在第一次打开数据库连接之前调用。
// 注册后,可在 SQL 中使用:CREATE VIRTUAL TABLE t USING moduleName(args...)
func Register(moduleName string, m Module) {
registryMu.Lock()
registry[moduleName] = m
registryMu.Unlock()
}
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 any) any {
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
}
// 注册所有已登记的虚拟表模块
registryMu.RLock()
defer registryMu.RUnlock()
for name, m := range registry {
if err := conn.CreateModule(name, &moduleAdapter{mod: m}); err != nil {
return fmt.Errorf("vtab: 注册模块 %q 失败: %w", name, err)
}
}
return nil
},
})
db.RegisterDialect("vtable", DialectOptions())
}
+433
View File
@@ -0,0 +1,433 @@
//go:build sqlite_vtable || vtable
package vtab_test
import (
"database/sql"
"testing"
"git.fsdpf.net/go/db/dialect/sqlite3/vtab"
"github.com/stretchr/testify/suite"
)
// ─── 内存表实现(用于测试)───────────────────────────────────
//
// 模拟一张 users 表,数据存在内存 slice 中,支持完整 CRUD。
type userRow struct {
id int64
name string
age int64
}
// usersModule 是虚拟表工厂
type usersModule struct{}
func (m *usersModule) Create(args []string, declare func(string) error) (vtab.Table, error) {
if err := declare("CREATE TABLE users(id INTEGER, name TEXT, age INTEGER)"); err != nil {
return nil, err
}
return newUsersTable(), nil
}
func (m *usersModule) Connect(args []string, declare func(string) error) (vtab.Table, error) {
return m.Create(args, declare)
}
// usersTable 是虚拟表实例,同时实现 WritableTable
type usersTable struct {
rows []userRow
nextID int64
}
func newUsersTable() *usersTable {
return &usersTable{
nextID: 1,
rows: []userRow{
{1, "Alice", 30},
{2, "Bob", 25},
{3, "Charlie", 35},
},
}
}
// BestIndex:不做任何约束下推,由 SQLite 全表扫描后过滤
func (t *usersTable) BestIndex(constraints []vtab.ConstraintInfo, _ []vtab.OrderByInfo) (*vtab.IndexOutput, error) {
return &vtab.IndexOutput{
Used: make([]bool, len(constraints)), // 全部 false
EstimatedCost: float64(len(t.rows)),
EstimatedRows: float64(len(t.rows)),
}, nil
}
func (t *usersTable) Open() (vtab.Cursor, error) {
return &usersCursor{rows: t.rows}, nil
}
func (t *usersTable) Disconnect() error { return nil }
func (t *usersTable) Destroy() error { return nil }
// ── WritableTable ──
func (t *usersTable) Insert(values []any) (int64, error) {
// values: [id, name, age]id 可能为 nil(自动生成)
id := t.nextID
t.nextID++
if values[0] != nil {
id = toInt64(values[0])
}
name, _ := values[1].(string)
age := toInt64(values[2])
t.rows = append(t.rows, userRow{id, name, age})
return id, nil
}
func (t *usersTable) Update(rowid any, values []any) error {
rid := toInt64(rowid)
for i, r := range t.rows {
if r.id == rid {
if values[1] != nil {
t.rows[i].name, _ = values[1].(string)
}
if values[2] != nil {
t.rows[i].age = toInt64(values[2])
}
return nil
}
}
return nil
}
func (t *usersTable) Delete(rowid any) error {
rid := toInt64(rowid)
for i, r := range t.rows {
if r.id == rid {
t.rows = append(t.rows[:i], t.rows[i+1:]...)
return nil
}
}
return nil
}
// ── Cursor ──
type usersCursor struct {
rows []userRow
pos int
}
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 }
func (c *usersCursor) Rowid() (int64, error) {
return c.rows[c.pos].id, nil
}
func (c *usersCursor) Column(col int) (any, error) {
r := c.rows[c.pos]
switch col {
case 0:
return r.id, nil
case 1:
return r.name, nil
case 2:
return r.age, nil
}
return nil, nil
}
func toInt64(v any) int64 {
switch x := v.(type) {
case int64:
return x
case int:
return int64(x)
case float64:
return int64(x)
}
return 0
}
// ─── 测试套件 ─────────────────────────────────────────────────
func init() {
vtab.Register("users_mod", &usersModule{})
}
type VtabSuite struct {
suite.Suite
db *sql.DB
}
func (s *VtabSuite) SetupSuite() {
db, err := sql.Open(vtab.DriverName, ":memory:")
s.Require().NoError(err)
s.db = db
}
func (s *VtabSuite) TearDownSuite() {
s.db.Close()
}
// SetupTest 每个测试前重建虚拟表,保证初始数据一致
func (s *VtabSuite) SetupTest() {
s.db.Exec(`DROP TABLE IF EXISTS users`)
_, err := s.db.Exec(`CREATE VIRTUAL TABLE users USING users_mod()`)
s.Require().NoError(err)
}
// ── SELECT ──
func (s *VtabSuite) TestSelect_All() {
rows, err := s.db.Query(`SELECT id, name, age FROM users ORDER BY id`)
s.Require().NoError(err)
defer rows.Close()
var result []userRow
for rows.Next() {
var r userRow
s.Require().NoError(rows.Scan(&r.id, &r.name, &r.age))
result = append(result, r)
}
s.Require().NoError(rows.Err())
s.Equal([]userRow{
{1, "Alice", 30},
{2, "Bob", 25},
{3, "Charlie", 35},
}, result)
}
func (s *VtabSuite) TestSelect_Where() {
rows, err := s.db.Query(`SELECT name FROM users WHERE age > 28 ORDER BY id`)
s.Require().NoError(err)
defer rows.Close()
var names []string
for rows.Next() {
var name string
s.Require().NoError(rows.Scan(&name))
names = append(names, name)
}
s.Equal([]string{"Alice", "Charlie"}, names)
}
func (s *VtabSuite) TestSelect_IF() {
var label string
err := s.db.QueryRow(`SELECT IF(age >= 30, 'senior', 'junior') FROM users WHERE id = 2`).Scan(&label)
s.Require().NoError(err)
s.Equal("junior", label)
}
// ── INSERT ──
func (s *VtabSuite) TestInsert() {
_, err := s.db.Exec(`INSERT INTO users(name, age) VALUES(?, ?)`, "Dave", 28)
s.Require().NoError(err)
var count int
s.Require().NoError(s.db.QueryRow(`SELECT COUNT(*) FROM users WHERE name = 'Dave'`).Scan(&count))
s.Equal(1, count)
}
// ── UPDATE ──
func (s *VtabSuite) TestUpdate() {
_, err := s.db.Exec(`UPDATE users SET age = ? WHERE name = ?`, 26, "Bob")
s.Require().NoError(err)
var age int64
s.Require().NoError(s.db.QueryRow(`SELECT age FROM users WHERE name = 'Bob'`).Scan(&age))
s.Equal(int64(26), age)
}
// ── DELETE ──
func (s *VtabSuite) TestDelete() {
_, err := s.db.Exec(`DELETE FROM users WHERE name = ?`, "Charlie")
s.Require().NoError(err)
var count int
s.Require().NoError(s.db.QueryRow(`SELECT COUNT(*) FROM users WHERE name = 'Charlie'`).Scan(&count))
s.Equal(0, count)
}
// ── JOIN with real table ──
// TestJoin_VtabAndRealTable 演示虚拟表与真实表的 JOIN。
//
// 真实表 orders:每条订单记录 user_id 和 amount。
// 虚拟表 users:内存中的用户数据。
// 查询:统计每个用户的订单总金额,只返回有订单的用户。
func (s *VtabSuite) TestJoin_VtabAndRealTable() {
// 建真实表并插入数据
_, err := s.db.Exec(`
CREATE TABLE IF NOT EXISTS orders (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
amount REAL NOT NULL
)`)
s.Require().NoError(err)
defer s.db.Exec(`DROP TABLE IF EXISTS orders`)
_, err = s.db.Exec(`
INSERT INTO orders(user_id, amount) VALUES
(1, 100.0),
(1, 50.5),
(2, 200.0),
(3, 75.0),
(3, 25.0)`)
s.Require().NoError(err)
// vtab users INNER JOIN real orders
rows, err := s.db.Query(`
SELECT u.name, SUM(o.amount) AS total
FROM users AS u
JOIN orders AS o ON o.user_id = u.id
GROUP BY u.id, u.name
ORDER BY u.id`)
s.Require().NoError(err)
defer rows.Close()
type row struct {
name string
total float64
}
var result []row
for rows.Next() {
var r row
s.Require().NoError(rows.Scan(&r.name, &r.total))
result = append(result, r)
}
s.Require().NoError(rows.Err())
s.Equal([]row{
{"Alice", 150.5},
{"Bob", 200.0},
{"Charlie", 100.0},
}, result)
}
// TestJoin_VtabLeftJoinRealTable 演示 LEFT JOIN:列出所有用户及其订单数,没有订单的用户显示 0。
func (s *VtabSuite) TestJoin_VtabLeftJoinRealTable() {
_, err := s.db.Exec(`
CREATE TABLE IF NOT EXISTS orders (
id INTEGER PRIMARY KEY AUTOINCREMENT,
user_id INTEGER NOT NULL,
amount REAL NOT NULL
)`)
s.Require().NoError(err)
defer s.db.Exec(`DROP TABLE IF EXISTS orders`)
// 只给 Alice 和 Bob 插入订单,Charlie 没有订单
_, err = s.db.Exec(`
INSERT INTO orders(user_id, amount) VALUES
(1, 100.0),
(2, 200.0)`)
s.Require().NoError(err)
rows, err := s.db.Query(`
SELECT u.name, COUNT(o.id) AS order_count
FROM users AS u
LEFT JOIN orders AS o ON o.user_id = u.id
GROUP BY u.id, u.name
ORDER BY u.id`)
s.Require().NoError(err)
defer rows.Close()
type row struct {
name string
count int
}
var result []row
for rows.Next() {
var r row
s.Require().NoError(rows.Scan(&r.name, &r.count))
result = append(result, r)
}
s.Require().NoError(rows.Err())
s.Equal([]row{
{"Alice", 1},
{"Bob", 1},
{"Charlie", 0}, // 没有订单,LEFT JOIN 保留
}, result)
}
func TestVtabSuite(t *testing.T) {
suite.Run(t, new(VtabSuite))
}
+6 -3
View File
@@ -93,13 +93,16 @@ func (sst *sqlserverTest) SetupSuite() {
func (sst *sqlserverTest) SetupTest() { func (sst *sqlserverTest) SetupTest() {
if _, err := sst.db.Exec(dropTable); err != nil { 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 { 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 { if _, err := sst.db.Exec(insertDefaultRecords); err != nil {
panic(err) sst.T().Skipf("SQLServer not available: %v", err)
return
} }
} }
+55 -2
View File
@@ -2,11 +2,12 @@ package engine
import ( import (
"database/sql" "database/sql"
"errors"
"fmt" "fmt"
"log"
"git.fsdpf.net/go/db" "git.fsdpf.net/go/db"
sqlite3dialect "git.fsdpf.net/go/db/dialect/sqlite3" sqlite3dialect "git.fsdpf.net/go/db/dialect/sqlite3"
sqlite3vtab "git.fsdpf.net/go/db/dialect/sqlite3/vtab"
) )
type Engine struct { type Engine struct {
@@ -27,7 +28,7 @@ func (e Engine) Connection(name string) *db.Database {
cfg, ok := e.configs[name] cfg, ok := e.configs[name]
if !ok { 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] _db, ok := e.dbs[name]
@@ -35,6 +36,20 @@ func (e Engine) Connection(name string) *db.Database {
if !ok { if !ok {
_db = e.MakeConnection(cfg) _db = e.MakeConnection(cfg)
e.dbs[name] = _db 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) return db.New(cfg.Driver, _db)
@@ -48,6 +63,8 @@ func (e Engine) MakeConnection(cfg DBConfig) (db *sql.DB) {
case "mysql": case "mysql":
case "sqlite3": case "sqlite3":
driverName = sqlite3dialect.DriverWithIF driverName = sqlite3dialect.DriverWithIF
case "vtable":
driverName = sqlite3vtab.DriverName
case "sqlserver": case "sqlserver":
case "postgres": case "postgres":
case "duckdb": case "duckdb":
@@ -65,9 +82,45 @@ func (e Engine) MakeConnection(cfg DBConfig) (db *sql.DB) {
panic(err) 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 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 { func Open(cfgs map[string]DBConfig) Engine {
for n, cfg := range cfgs { for n, cfg := range cfgs {
_engine.configs[n] = cfg _engine.configs[n] = cfg
+25 -5
View File
@@ -75,6 +75,7 @@ type DBConfig struct {
Threads int Threads int
MaxMemory string MaxMemory string
Dsn string Dsn string
Extensions []string // 启动时自动 INSTALL + LOAD 的扩展名,如 ["vss", "json"]
} }
} }
@@ -88,7 +89,7 @@ func (c *DBConfig) ToDSN() string {
return c.toMySQLDSN() return c.toMySQLDSN()
case "pgsql": case "pgsql":
return c.toPostgreSQLDSN() return c.toPostgreSQLDSN()
case "sqlite3": case "sqlite3", "vtable":
return c.toSQLiteDSN() return c.toSQLiteDSN()
case "sqlserver": case "sqlserver":
return c.toSQLServerDSN() return c.toSQLServerDSN()
@@ -328,6 +329,7 @@ func NewDBConfig(driver string, options ...Option) DBConfig {
"sqlite3": true, "sqlite3": true,
"sqlserver": true, "sqlserver": true,
"duckdb": true, "duckdb": true,
"vtable": true,
} }
if !validDrivers[driver] { if !validDrivers[driver] {
panic(fmt.Sprintf("Unsupported driver: %s", driver)) panic(fmt.Sprintf("Unsupported driver: %s", driver))
@@ -363,7 +365,7 @@ func NewDBConfig(driver string, options ...Option) DBConfig {
if config.Password == "" { if config.Password == "" {
panic(fmt.Sprintf("Password is required for %s driver", config.Driver)) panic(fmt.Sprintf("Password is required for %s driver", config.Driver))
} }
case "sqlite3": case "sqlite3", "vtable":
if config.SQLite.File == "" { if config.SQLite.File == "" {
panic("File is required for sqlite3 driver") panic("File is required for sqlite3 driver")
} }
@@ -429,7 +431,7 @@ func WithWriteHosts(hosts []string) Option {
} }
} }
// MySQL 专用选项 // WithMySQLCollation MySQL 专用选项
func WithMySQLCollation(collation string) Option { func WithMySQLCollation(collation string) Option {
return func(c *DBConfig) { return func(c *DBConfig) {
if c.Driver != "mysql" { if c.Driver != "mysql" {
@@ -461,7 +463,7 @@ func WithPgSslmode(sslmode string) Option {
// SQLite 专用选项 // SQLite 专用选项
func WithSQLiteFile(file string) Option { func WithSQLiteFile(file string) Option {
return func(c *DBConfig) { return func(c *DBConfig) {
if c.Driver != "sqlite3" { if c.Driver != "sqlite3" && c.Driver != "vtable" {
panic("WithSQLiteFile is only valid for sqlite3 driver") panic("WithSQLiteFile is only valid for sqlite3 driver")
} }
c.SQLite.File = file c.SQLite.File = file
@@ -470,13 +472,22 @@ func WithSQLiteFile(file string) Option {
func WithSQLiteJournal(journal string) Option { func WithSQLiteJournal(journal string) Option {
return func(c *DBConfig) { return func(c *DBConfig) {
if c.Driver != "sqlite3" { if c.Driver != "sqlite3" && c.Driver != "vtable" {
panic("WithSQLiteJournal is only valid for sqlite3 driver") panic("WithSQLiteJournal is only valid for sqlite3 driver")
} }
c.SQLite.Journal = journal c.SQLite.Journal = journal
} }
} }
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 专用选项 // SQL Server 专用选项
func WithSQLServerInstance(instance string) Option { func WithSQLServerInstance(instance string) Option {
return func(c *DBConfig) { return func(c *DBConfig) {
@@ -523,3 +534,12 @@ func WithDuckDBMaxMemory(maxMemory string) Option {
c.DuckDB.MaxMemory = maxMemory 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...)
}
}
+1 -5
View File
@@ -238,18 +238,14 @@ func (q QueryExecutor) ScanValContext(ctx context.Context, i interface{}) (bool,
if err != nil { if err != nil {
return false, err return false, err
} }
defer func() { _ = scanner.Close() }() defer func() { _ = scanner.Close() }()
if scanner.Next() { if scanner.Next() {
err = scanner.ScanVal(i) if err = scanner.ScanVal(i); err != nil {
if err != nil {
return false, err return false, err
} }
return true, scanner.Err() return true, scanner.Err()
} }
return false, scanner.Err() return false, scanner.Err()
} }
+186 -2
View File
@@ -3,6 +3,7 @@ package exec
import ( import (
"context" "context"
"database/sql" "database/sql"
"database/sql/driver"
"encoding/json" "encoding/json"
"fmt" "fmt"
"strings" "strings"
@@ -13,6 +14,11 @@ import (
"github.com/stretchr/testify/suite" "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 ( var (
testAddr1 = "111 Test Addr" testAddr1 = "111 Test Addr"
testAddr2 = "211 Test Addr" testAddr2 = "211 Test Addr"
@@ -929,9 +935,11 @@ func (qes *queryExecutorSuite) TestScanStruct() {
qes.EqualError(err, "queryExecutor error") qes.EqualError(err, "queryExecutor error")
qes.False(found) qes.False(found)
// NULL 值扫描进 string 字段:通过 **string 中间层正常处理,结果为空字符串
found, err = e.ScanStruct(&item) found, err = e.ScanStruct(&item)
qes.Error(err) qes.NoError(err)
qes.False(found) qes.True(found)
qes.Equal(StructWithTags{Address: "", Name: ""}, item)
found, err = e.ScanStruct(&item) found, err = e.ScanStruct(&item)
qes.NoError(err) qes.NoError(err)
@@ -1242,6 +1250,182 @@ func (qes *queryExecutorSuite) TestScanVal_withValuerSlice() {
qes.Equal(JSONBoolArray{true, false, true}, bools) 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) { func TestQueryExecutorSuite(t *testing.T) {
suite.Run(t, new(queryExecutorSuite)) suite.Run(t, new(queryExecutorSuite))
} }
+77 -6
View File
@@ -147,9 +147,9 @@ func (s *scanner) ScanStruct(i interface{}) error {
// 补全未知字段类型 // 补全未知字段类型
if len(cols) != len(cm) { if len(cols) != len(cm) {
colTypes, err := s.rows.ColumnTypes() colTypes, ctErr := s.rows.ColumnTypes()
if err != nil { if ctErr != nil {
return err return ctErr
} }
for _, t := range colTypes { for _, t := range colTypes {
if _, ok := cm[t.Name()]; !ok { if _, ok := cm[t.Name()]; !ok {
@@ -166,7 +166,6 @@ func (s *scanner) ScanStruct(i interface{}) error {
} }
scans, err := createColumnScans(s.columns, s.columnMap) scans, err := createColumnScans(s.columns, s.columnMap)
if err != nil { if err != nil {
return err return err
} }
@@ -177,8 +176,13 @@ func (s *scanner) ScanStruct(i interface{}) error {
record := map[string]interface{}{} record := map[string]interface{}{}
for index, col := range s.columns { for index, col := range s.columns {
if pi, ok := scans[index].(*interface{}); ok {
raw := toJSONRawMessage(*pi)
record[col] = &raw
} else {
record[col] = scans[index] record[col] = scans[index]
} }
}
util.AssignStructVals(i, record, s.columnMap) 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. // ScanVal will scan the current row and column into i.
func (s *scanner) ScanVal(i interface{}) error { func (s *scanner) ScanVal(i interface{}) error {
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 { if err := s.rows.Scan(i); err != nil {
return err return err
} }
}
return s.Err() return s.Err()
} }
@@ -260,6 +313,23 @@ func checkScanValsTarget(i interface{}) (reflect.Value, error) {
return val, nil 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) { func createColumnScans(cols []string, cm util.ColumnMap) (scans []interface{}, err error) {
scans = make([]interface{}, 0, len(cols)) scans = make([]interface{}, 0, len(cols))
@@ -278,7 +348,8 @@ func createColumnScans(cols []string, cm util.ColumnMap) (scans []interface{}, e
reflect.Bool: reflect.Bool:
scans = append(scans, reflect.New(reflect.PointerTo(data.GoType)).Interface()) scans = append(scans, reflect.New(reflect.PointerTo(data.GoType)).Interface())
case reflect.Map, reflect.Slice, reflect.Struct: 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: default:
scans = append(scans, reflect.New(data.GoType).Interface()) scans = append(scans, reflect.New(data.GoType).Interface())
} }
+152 -3
View File
@@ -1,9 +1,10 @@
package exec package exec
import ( import (
"database/sql"
"encoding/json"
"testing" "testing"
"git.fsdpf.net/go/db/exp"
"github.com/DATA-DOG/go-sqlmock" "github.com/DATA-DOG/go-sqlmock"
"github.com/stretchr/testify/suite" "github.com/stretchr/testify/suite"
) )
@@ -81,14 +82,162 @@ func (s *scannerSuite) TestGetRecords() {
AddRow("111 Test Addr", "Test1"), AddRow("111 Test Addr", "Test1"),
) )
rows, err := db.Query("SELECT \\* FROM `items`") rows, err := db.Query("SELECT * FROM `items`")
s.Require().NoError(err) s.Require().NoError(err)
result, err := NewScanner(rows).GetRecords() result, err := NewScanner(rows).GetRecords()
s.Require().NoError(err) 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"},
{"address": "111 Test Addr", "name": "Test1"}, {"address": "111 Test Addr", "name": "Test1"},
}, result) }, 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)
}
+2 -1
View File
@@ -9,7 +9,7 @@ require (
github.com/go-sql-driver/mysql v1.7.1 github.com/go-sql-driver/mysql v1.7.1
github.com/lib/pq v1.10.9 github.com/lib/pq v1.10.9
github.com/marcboeker/go-duckdb v1.8.5 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/microsoft/go-mssqldb v1.9.4
github.com/samber/lo v1.49.1 github.com/samber/lo v1.49.1
github.com/stretchr/testify v1.10.0 github.com/stretchr/testify v1.10.0
@@ -17,6 +17,7 @@ require (
require ( require (
github.com/apache/arrow-go/v18 v18.1.0 // indirect 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/davecgh/go-spew v1.1.1 // indirect
github.com/go-viper/mapstructure/v2 v2.2.1 // indirect github.com/go-viper/mapstructure/v2 v2.2.1 // indirect
github.com/goccy/go-json v0.10.5 // indirect github.com/goccy/go-json v0.10.5 // indirect
+4
View File
@@ -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/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 h1:tdPmh/ptjE1IJnhbhrcl2++TauVjy242rkV/UzJChnE=
github.com/apache/thrift v0.21.0/go.mod h1:W1H8aR/QRtYNvrPeFXBtobyRkd0/YVhTc6i07XIAgDw= 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 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= 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= 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/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 h1:mCRHCLDUBXgpKAqIKsaAaAsrAlbkeomtRFKXh2L6YIM=
github.com/mattn/go-sqlite3 v1.14.17/go.mod h1:2eHXhiwb8IkHr+BDWZGa96P6+rkvnG63S2DGjv9HUNg= 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 h1:sHrj3GcdgkxytZ09aZ3+ys72pMeyEXJowT44j74pNgs=
github.com/microsoft/go-mssqldb v1.9.4/go.mod h1:GBbW9ASTiDC+mpgWDGKdm3FnFLTUsLYN3iFL90lQ+PA= 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= github.com/minio/asm2plan9s v0.0.0-20200509001527-cdd76441f9d8 h1:AMFGa4R4MiIpspGNG7Z948v4n35fFGB3RR3G/ry4FWs=
+4 -18
View File
@@ -210,25 +210,11 @@ func (id *InsertDataset) Rows(rows ...interface{}) *InsertDataset {
panic("Rows: unsupported row type, must be map, Record, struct, slice or array") panic("Rows: unsupported row type, must be map, Record, struct, slice or array")
} }
result := make(map[string]interface{}) record, err := exp.NewRecordFromStruct(row, true, false)
typ := val.Type() if err != nil {
for j := 0; j < val.NumField(); j++ { panic(fmt.Sprintf("Rows: %v", err))
field := typ.Field(j)
// Skip unexported fields
if field.PkgPath != "" {
continue
} }
// Use db tag if present, otherwise use field name converted[i] = map[string]interface{}(record)
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()
}
converted[i] = result
} }
} }
return id.copy(id.clauses.SetRows(converted)) return id.copy(id.clauses.SetRows(converted))
+7 -7
View File
@@ -38,13 +38,13 @@ func newColumnMap(t reflect.Type, fieldIndex []int, prefixes []string) ColumnMap
columnName := getColumnName(&f, dbTag) columnName := getColumnName(&f, dbTag)
if !shouldIgnoreField(dbTag) { if !shouldIgnoreField(dbTag) {
// 移除原来的关联结构,并 scans 字段出现 table.col // 移除原来的关联结构,并 scans 字段出现 table.col
// if !implementsScanner(f.Type) { if !implementsScanner(f.Type) {
// subCm := getStructColumnMap(&f, fieldIndex, []string{columnName}, prefixes) subCm := getStructColumnMap(&f, fieldIndex, []string{columnName}, prefixes)
// if len(subCm) != 0 { if len(subCm) != 0 {
// subColMaps = append(subColMaps, subCm) subColMaps = append(subColMaps, subCm)
// continue continue
// } }
// } }
ffTag := tag.New("ff", f.Tag) ffTag := tag.New("ff", f.Tag)
columnName = strings.Join(append(prefixes, columnName), ".") columnName = strings.Join(append(prefixes, columnName), ".")
cm[columnName] = newColumnData(&f, columnName, fieldIndex, ffTag) cm[columnName] = newColumnData(&f, columnName, fieldIndex, ffTag)
+46 -12
View File
@@ -203,23 +203,57 @@ func SafeSetFieldByIndex(v reflect.Value, fieldIndex []int, src interface{}) (re
} }
func SafeSetVarValue(v reflect.Value, src interface{}) error { func SafeSetVarValue(v reflect.Value, src interface{}) error {
f := reflect.Indirect(v) srcReflect := reflect.ValueOf(src)
srcVal := reflect.ValueOf(src).Elem()
// 处理 converting NULL to string is unsupported // src 可能是 **TcreateColumnScans 的扫描目标)或裸值(测试/直接调用)
// 前面将 string 和 int 类型转为了 *string 和 *int if srcReflect.Kind() != reflect.Ptr {
// 这里做还原 if srcReflect.IsValid() && v.Type().ConvertibleTo(srcReflect.Type()) {
if srcVal.IsNil() { v.Set(srcReflect.Convert(v.Type()))
f.Set(reflect.Zero(f.Type())) }
} else if f.Type().ConvertibleTo(srcVal.Type().Elem()) { return nil
f.Set(srcVal.Elem()) }
} else if f.Type().ConvertibleTo(srcVal.Type()) {
f.Set(srcVal) 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 { } else if u, ok := srcVal.Interface().(*json.RawMessage); ok && len(*u) >= 2 {
if err := json.Unmarshal(*u, f.Addr().Interface()); err != nil { if err := json.Unmarshal(*u, v.Addr().Interface()); err != nil {
return err return err
} }
} }
} else {
// srcVal = T*T 扫描目标的 default 分支)
if v.Type().ConvertibleTo(srcVal.Type()) {
v.Set(srcVal.Convert(v.Type()))
}
}
return nil return nil
} }
+6
View File
@@ -224,6 +224,12 @@ func (this *Blueprint) Uuid(column string) *ColumnDefinition {
return this.addColumn("uuid", column, nil) 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 { func (this *Blueprint) Increments(column string) *ColumnDefinition {
return this.UnsignedInteger(column, true) return this.UnsignedInteger(column, true)
+2
View File
@@ -88,6 +88,8 @@ func (this Builder) Build(bp *Blueprint) error {
return err return err
} }
// fmt.Println(strings.Join(sqls, "\n---\n"))
err = tx.Wrap(func() error { err = tx.Wrap(func() error {
for _, sql := range sqls { for _, sql := range sqls {
if _, err := tx.Exec(sql); err != nil { if _, err := tx.Exec(sql); err != nil {
+15 -1
View File
@@ -31,11 +31,14 @@ type ColumnOptions struct {
primary bool // Add a primary index primary bool // Add a primary index
index bool // Add an index index bool // Add an index
spatialIndex bool // Add a spatial index spatialIndex bool // Add a spatial index
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 { func (c *ColumnDefinition) VirtualAs(as string) *ColumnDefinition {
c.virtualAs = as c.virtualAs = as
c.hidden = true
return c return c
} }
@@ -166,3 +169,14 @@ func (c *ColumnDefinition) IsUseCurrent() bool {
func (c *ColumnDefinition) GetRename() string { func (c *ColumnDefinition) GetRename() string {
return c.rename return c.rename
} }
// Hidden 将列标记为 SQLite 虚拟表 HIDDEN 列。
// HIDDEN 列不出现在 SELECT * 结果中,但可以通过 WHERE col = ? 向虚拟表传递参数。
func (c *ColumnDefinition) Hidden() *ColumnDefinition {
c.hidden = true
return c
}
func (c *ColumnDefinition) IsHidden() bool {
return c.hidden
}
+75 -22
View File
@@ -40,20 +40,33 @@ func (this DuckDB) CompileCreate(bp *schema.Blueprint) []string {
if bp.Temporary { if bp.Temporary {
temporary = db.L("CREATE 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 { func (this DuckDB) CompileAdd(bp *schema.Blueprint) []string {
columns := strings.Join(schema.PrefixArray("ADD COLUMN", this.getAddedColumns(bp)), ",\n") columns := strings.Join(schema.PrefixArray("ADD COLUMN", this.getAddedColumns(bp)), ",\n")
sql := this.GenerateSQL("ALTER TABLE ?\n?", db.T(bp.GetTable()), db.L(columns)) 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 { 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 { func (this DuckDB) CompileRename(bp *schema.Blueprint) []string {
toName := ""
commands := bp.GetCommands() commands := bp.GetCommands()
if len(commands) == 0 { if len(commands) == 0 {
panic("new table undefined") panic("new table undefined")
} }
toName = commands[0].To toName := commands[0].To
if toName == "" { if toName == "" {
panic("new table undefined") 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 { func (this DuckDB) CompileModifyComment(bp *schema.Blueprint) []string {
@@ -185,6 +219,8 @@ func (this DuckDB) GetColumnType(column *schema.ColumnDefinition) string {
return "SMALLINT" return "SMALLINT"
case "uuid": case "uuid":
return "UUID" return "UUID"
case "vector":
return this.GenerateSQL("FLOAT[?]", column.Length)
} }
panic("Unsupported data type: " + column.Type) panic("Unsupported data type: " + column.Type)
} }
@@ -197,19 +233,17 @@ func (this DuckDB) GetColumnModifier(modifier string, bp *schema.Blueprint, colu
} }
return " NOT NULL" return " NOT NULL"
case "Default": case "Default":
v := column.GetDefault()
if v == nil {
if column.IsUseCurrent() { if column.IsUseCurrent() {
// DuckDB 不支持 ON UPDATE,忽略 def 中可能携带的 MySQL ON UPDATE 标记
return " DEFAULT CURRENT_TIMESTAMP" return " DEFAULT CURRENT_TIMESTAMP"
} }
v := column.GetDefault()
if v == nil {
return "" return ""
} }
return this.GenerateSQL(" DEFAULT ?", v) return this.GenerateSQL(" DEFAULT ?", v)
case "Increment": case "Increment":
if column.IsAutoIncrement() { // 自增列在 getAddedColumns 中单独处理(GENERATED ALWAYS AS IDENTITY),此处无需输出
// DuckDB 使用 SERIAL 或者 SEQUENCE
return " PRIMARY KEY"
}
case "Comment": case "Comment":
// DuckDB 支持列注释,但需要在 CREATE TABLE 后使用 COMMENT ON // DuckDB 支持列注释,但需要在 CREATE TABLE 后使用 COMMENT ON
return "" return ""
@@ -217,22 +251,41 @@ func (this DuckDB) GetColumnModifier(modifier string, bp *schema.Blueprint, colu
return "" 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) { func (this DuckDB) getAddedColumns(bp *schema.Blueprint) (columns []string) {
for _, column := range bp.GetAddedColumns() { for _, column := range bp.GetAddedColumns() {
colType := this.GetColumnType(column) var sql string
// 对于自增列,使用 SERIAL 类型
if column.IsAutoIncrement() { if column.IsAutoIncrement() {
var baseType string
switch column.Type { switch column.Type {
case "bigInteger": case "bigInteger":
colType = "BIGSERIAL" baseType = "BIGINT"
case "smallInteger": case "smallInteger":
colType = "SMALLSERIAL" baseType = "SMALLINT"
default: 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, sql)
columns = append(columns, this.addModifiers(sql, bp, column))
} }
return return
} }
+13 -2
View File
@@ -81,8 +81,9 @@ func (t *duckDBTest) TestCompileCreate() {
sql := t.schema.CompileCreate(bp) sql := t.schema.CompileCreate(bp)
t.Equal([]string{ t.Equal([]string{
"CREATE TABLE \"users\" (\n" + "CREATE SEQUENCE IF NOT EXISTS \"users_id_seq\"",
"\"id\" BIGSERIAL NOT NULL PRIMARY KEY,\n" + "CREATE TABLE IF NOT EXISTS \"users\" (\n" +
"\"id\" BIGINT DEFAULT nextval('users_id_seq') PRIMARY KEY,\n" +
"\"enabled\" BOOLEAN NOT NULL DEFAULT '1',\n" + "\"enabled\" BOOLEAN NOT NULL DEFAULT '1',\n" +
"\"created_user\" CHAR(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000',\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" + "\"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" + "\"updated_at\" TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,\n" +
"\"deleted_at\" TIMESTAMP NULL\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) }, sql)
t.T().Log(sql) t.T().Log(sql)
@@ -109,6 +117,9 @@ func (t *duckDBTest) TestCompileAdd() {
"ADD COLUMN \"name\" VARCHAR(50) NOT NULL,\n" + "ADD COLUMN \"name\" VARCHAR(50) NOT NULL,\n" +
"ADD COLUMN \"age\" SMALLINT NOT NULL DEFAULT '18',\n" + "ADD COLUMN \"age\" SMALLINT NOT NULL DEFAULT '18',\n" +
"ADD COLUMN \"sex\" VARCHAR(1) NOT NULL DEFAULT '0'", "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) }, sql)
t.T().Log(sql) t.T().Log(sql)
+12 -2
View File
@@ -41,8 +41,9 @@ func Example_createTable() {
} }
// Output: // Output:
// CREATE TABLE "users" ( // CREATE SEQUENCE IF NOT EXISTS "users_id_seq"
// "id" BIGSERIAL NOT NULL PRIMARY KEY, // CREATE TABLE IF NOT EXISTS "users" (
// "id" BIGINT DEFAULT nextval('users_id_seq') PRIMARY KEY,
// "name" VARCHAR(50) NOT NULL, // "name" VARCHAR(50) NOT NULL,
// "email" VARCHAR(100) NOT NULL, // "email" VARCHAR(100) NOT NULL,
// "age" INTEGER NULL, // "age" INTEGER NULL,
@@ -50,6 +51,13 @@ func Example_createTable() {
// "created_at" TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, // "created_at" TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
// "updated_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() { func Example_addColumn() {
@@ -75,6 +83,8 @@ func Example_addColumn() {
// ALTER TABLE "users" // ALTER TABLE "users"
// ADD COLUMN "phone" VARCHAR(20) NULL, // ADD COLUMN "phone" VARCHAR(20) NULL,
// ADD COLUMN "address" TEXT NULL // ADD COLUMN "address" TEXT NULL
// COMMENT ON COLUMN "users"."phone" IS '电话号码'
// COMMENT ON COLUMN "users"."address" IS '地址'
} }
func Example_dropColumn() { func Example_dropColumn() {
+3 -1
View File
@@ -53,7 +53,7 @@ func (this Mysql) CompileCreate(bp *schema.Blueprint) []string {
} }
columns := strings.Join(this.getAddedColumns(bp), ",\n") 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 charset := bp.Charset
if charset == "" { if charset == "" {
@@ -183,6 +183,8 @@ func (this Mysql) GetColumnType(column *schema.ColumnDefinition) string {
return "year" return "year"
case "uuid": case "uuid":
return "char(36)" return "char(36)"
case "vector":
return this.GenerateSQL("VECTOR(?)", column.Length)
} }
panic("Unsupported data type: " + column.Type) panic("Unsupported data type: " + column.Type)
} }
+37 -36
View File
@@ -1,15 +1,13 @@
package mysql_test package mysql_test
import ( import (
"os"
"testing" "testing"
"git.fsdpf.net/go/db" "git.fsdpf.net/go/db"
"git.fsdpf.net/go/db/engine" "git.fsdpf.net/go/db/engine"
"git.fsdpf.net/go/db/schema" "git.fsdpf.net/go/db/schema"
sqlmock "github.com/DATA-DOG/go-sqlmock"
"github.com/stretchr/testify/suite" "github.com/stretchr/testify/suite"
_ "github.com/go-sql-driver/mysql"
) )
type mysqlTest struct { type mysqlTest struct {
@@ -17,42 +15,28 @@ type mysqlTest struct {
schema schema.Schema schema schema.Schema
} }
var ( var tableName = "entry"
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`) );"
)
func TestMysqlSuite(t *testing.T) { func TestMysqlSuite(t *testing.T) {
suite.Run(t, new(mysqlTest)) suite.Run(t, new(mysqlTest))
} }
func (t *mysqlTest) SetupSuite() { func (t *mysqlTest) SetupSuite() {
db := engine.Open(map[string]engine.DBConfig{ mockDB, mock, _ := sqlmock.New()
"test-mysql": engine.NewDBConfig("mysql",
engine.WithHost(os.Getenv("MYSQL_HOST")), mock.ExpectQuery("SELECT.*column_name.*columns").
engine.WithPort(os.Getenv("MYSQL_PORT")), WillReturnRows(sqlmock.NewRows([]string{"column_name"}).
engine.WithDatabase(os.Getenv("MYSQL_DB")), AddRow("id").AddRow("int").AddRow("float").
engine.WithUsername(os.Getenv("MYSQL_USER")), AddRow("string").AddRow("time").AddRow("bool").AddRow("bytes"))
engine.WithPassword(os.Getenv("MYSQL_PASSWD")),
engine.WithParseTime(true), 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") }).Connection("test-mysql")
t.schema = schema.GetSchemaDialect(db) t.schema = schema.GetSchemaDialect(conn)
// db.Exec(dropTable)
db.Exec(createTable)
} }
func (t *mysqlTest) TestGetColumnListing() { func (t *mysqlTest) TestGetColumnListing() {
@@ -84,7 +68,7 @@ func (t *mysqlTest) TestCompileCreate() {
sql := t.schema.CompileCreate(bp) sql := t.schema.CompileCreate(bp)
t.Equal([]string{ 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" + "`id` bigint(20) UNSIGNED NOT NULL AUTO_INCREMENT PRIMARY KEY COMMENT 'ID',\n" +
"`enabled` tinyint(1) NOT NULL DEFAULT '1' COMMENT '是否有效',\n" + "`enabled` tinyint(1) NOT NULL DEFAULT '1' COMMENT '是否有效',\n" +
"`created_user` char(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000' 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.String("name", 50).Change("username")
bp.SmallInteger("age").Default("19").Comment("年龄").Change() bp.SmallInteger("age").Default("19").Comment("年龄").Change()
sql := "ALTER TABLE `users`\n" + sql := t.schema.CompileChange(bp)
"CHANGE COLUMN `name` `username` varchar(50) NOT NULL,\n" +
"CHANGE COLUMN `age` `age` smallint(4) NOT NULL DEFAULT '19' COMMENT '年龄';"
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) t.T().Log(sql)
} }
@@ -189,4 +175,19 @@ func (t *mysqlTest) TestCompileRename() {
t.T().Log(sql) 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)
}
// 修改表备注 // 修改表备注
+6 -6
View File
@@ -45,7 +45,7 @@ func (this Postgres) CompileCreate(bp *schema.Blueprint) []string {
temporary = db.L("CREATE TEMPORARY") temporary = db.L("CREATE TEMPORARY")
} }
columns := strings.Join(this.getAddedColumns(bp), ",\n") 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} return []string{sql}
} }
@@ -144,6 +144,8 @@ func (this Postgres) GetColumnType(column *schema.ColumnDefinition) string {
return "integer" return "integer"
case "uuid": case "uuid":
return "uuid" return "uuid"
case "vector":
return this.GenerateSQL("vector(?)", column.Length)
} }
panic("Unsupported data type: " + column.Type) panic("Unsupported data type: " + column.Type)
} }
@@ -156,16 +158,14 @@ func (this Postgres) GetColumnModifier(modifier string, bp *schema.Blueprint, co
} }
return " NOT NULL" return " NOT NULL"
case "Default": case "Default":
v := column.GetDefault()
if v == nil {
if column.IsUseCurrent() { if column.IsUseCurrent() {
// PostgreSQL 不支持 ON UPDATE,忽略 def 中可能携带的 MySQL ON UPDATE 标记
return " DEFAULT CURRENT_TIMESTAMP" return " DEFAULT CURRENT_TIMESTAMP"
} }
v := column.GetDefault()
if v == nil {
return "" return ""
} }
if column.IsUseCurrent() {
return this.GenerateSQL(" DEFAULT CURRENT_TIMESTAMP ?", v)
}
return this.GenerateSQL(" DEFAULT ?", v) return this.GenerateSQL(" DEFAULT ?", v)
case "Increment": case "Increment":
if column.IsAutoIncrement() { if column.IsAutoIncrement() {
+27 -43
View File
@@ -1,14 +1,12 @@
package postgres_test package postgres_test
import ( import (
"os"
"testing" "testing"
"git.fsdpf.net/go/db/engine" "git.fsdpf.net/go/db/engine"
"git.fsdpf.net/go/db/schema" "git.fsdpf.net/go/db/schema"
sqlmock "github.com/DATA-DOG/go-sqlmock"
"github.com/stretchr/testify/suite" "github.com/stretchr/testify/suite"
_ "github.com/lib/pq"
) )
type postgresTest struct { type postgresTest struct {
@@ -16,41 +14,27 @@ type postgresTest struct {
schema schema.Schema schema schema.Schema
} }
var ( var tableName = "entry"
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" +
");"
)
func TestPostgresSuite(t *testing.T) { func TestPostgresSuite(t *testing.T) {
suite.Run(t, new(postgresTest)) suite.Run(t, new(postgresTest))
} }
func (t *postgresTest) SetupSuite() { func (t *postgresTest) SetupSuite() {
db := engine.Open(map[string]engine.DBConfig{ mockDB, mock, _ := sqlmock.New()
"test-postgres": engine.NewDBConfig("postgres",
engine.WithHost(os.Getenv("POSTGRES_HOST")), mock.ExpectQuery("SELECT.*column_name.*columns").
engine.WithPort(os.Getenv("POSTGRES_PORT")), WillReturnRows(sqlmock.NewRows([]string{"column_name"}).
engine.WithDatabase(os.Getenv("POSTGRES_DB")), AddRow("id").AddRow("int").AddRow("float").
engine.WithUsername(os.Getenv("POSTGRES_USER")), AddRow("string").AddRow("time").AddRow("bool").AddRow("bytes"))
engine.WithPassword(os.Getenv("POSTGRES_PASSWD")),
), 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") }).Connection("test-postgres")
t.schema = schema.GetSchemaDialect(conn)
t.schema = schema.GetSchemaDialect(db)
db.Exec(dropTable)
db.Exec(createTable)
} }
func (t *postgresTest) TestGetColumnListing() { func (t *postgresTest) TestGetColumnListing() {
@@ -80,14 +64,14 @@ func (t *postgresTest) TestCompileCreate() {
sql := t.schema.CompileCreate(bp) sql := t.schema.CompileCreate(bp)
t.Equal([]string{ t.Equal([]string{
"CREATE TABLE \"users\" (\n" + "CREATE TABLE IF NOT EXISTS \"users\" (\n" +
"\"id\" bigint NOT NULL BIGSERIAL PRIMARY KEY COMMENT 'ID',\n" + "\"id\" bigint NOT NULL BIGSERIAL PRIMARY KEY,\n" +
"\"enabled\" boolean NOT NULL DEFAULT true COMMENT '是否有效',\n" + "\"enabled\" boolean NOT NULL DEFAULT 'true',\n" +
"\"created_user\" char(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000' COMMENT '创建者',\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' COMMENT '拥有者',\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 COMMENT '创建时间',\n" + "\"created_at\" timestamp with time zone NOT NULL DEFAULT CURRENT_TIMESTAMP,\n" +
"\"updated_at\" timestamp with time zone NOT NULL DEFAULT CURRENT_TIMESTAMP COMMENT '更新时间',\n" + "\"updated_at\" timestamp with time zone NOT NULL DEFAULT CURRENT_TIMESTAMP,\n" +
"\"deleted_at\" timestamp without time zone NULL COMMENT '删除时间'\n" + "\"deleted_at\" timestamp without time zone NULL\n" +
")", ")",
}, sql) }, sql)
@@ -105,9 +89,9 @@ func (t *postgresTest) TestCompileAdd() {
t.Equal([]string{ t.Equal([]string{
"ALTER TABLE \"users\"\n" + "ALTER TABLE \"users\"\n" +
"ADD COLUMN \"name\" varchar(50) NOT NULL COMMENT '用户名',\n" + "ADD COLUMN \"name\" varchar(50) NOT NULL,\n" +
"ADD COLUMN \"age\" smallint NOT NULL DEFAULT 18 COMMENT '年龄',\n" + "ADD COLUMN \"age\" smallint NOT NULL DEFAULT '18',\n" +
"ADD COLUMN \"sex\" varchar(1) NOT NULL DEFAULT '0' COMMENT '性别'", "ADD COLUMN \"sex\" varchar(1) NOT NULL DEFAULT '0'",
}, sql) }, sql)
t.T().Log(sql) t.T().Log(sql)
@@ -124,7 +108,7 @@ func (t *postgresTest) TestCompileChange() {
t.Equal([]string{ t.Equal([]string{
"ALTER TABLE \"users\"\n" + "ALTER TABLE \"users\"\n" +
"ALTER COLUMN \"name\" TYPE varchar(50) NOT NULL,\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) }, sql)
t.T().Log(sql) t.T().Log(sql)
+20 -4
View File
@@ -13,7 +13,7 @@ import (
) )
var ( var (
sqlite3DefaultModifiers = []string{"VirtualAs", "StoredAs", "Nullable", "Default", "Increment"} sqlite3DefaultModifiers = []string{"Hidden", "Nullable", "Default", "Increment"}
sqlite3Serials = []string{"bigInteger", "integer", "mediumInteger", "smallInteger", "tinyInteger"} sqlite3Serials = []string{"bigInteger", "integer", "mediumInteger", "smallInteger", "tinyInteger"}
) )
@@ -178,14 +178,22 @@ func (this Sqlite3) GetColumnModifier(modifier string, bp *schema.Blueprint, col
if v := column.GetStoredAs(); v != "" { if v := column.GetStoredAs(); v != "" {
return this.GenerateSQL(" GENERATED ALWAYS AS (?) STORED", db.L(v)) return this.GenerateSQL(" GENERATED ALWAYS AS (?) STORED", db.L(v))
} }
case "Hidden":
if column.IsHidden() {
return "HIDDEN"
}
case "Nullable": case "Nullable":
if column.GetVirtualAs() == "" && column.GetStoredAs() == "" { // HIDDEN 列和生成列不加 NULL/NOT NULL
if column.GetVirtualAs() == "" && column.GetStoredAs() == "" && !column.IsHidden() {
if column.IsNullable() { if column.IsNullable() {
return "NULL" return "NULL"
} }
return "NOT NULL" return "NOT NULL"
} }
case "Default": case "Default":
if column.IsHidden() {
return ""
}
if column.IsUseCurrent() { if column.IsUseCurrent() {
return " DEFAULT CURRENT_TIMESTAMP" return " DEFAULT CURRENT_TIMESTAMP"
} }
@@ -218,7 +226,10 @@ func (this Sqlite3) GetColumnType(column *schema.ColumnDefinition) string {
case "text": case "text":
return "text" return "text"
case "integer": case "integer":
if column.Length > 0 {
return this.GenerateSQL("INTEGER(?)", column.Length) return this.GenerateSQL("INTEGER(?)", column.Length)
}
return "INTEGER"
case "bigInteger": case "bigInteger":
return "bigint(20)" return "bigint(20)"
case "tinyInteger": case "tinyInteger":
@@ -247,6 +258,8 @@ func (this Sqlite3) GetColumnType(column *schema.ColumnDefinition) string {
return "year" return "year"
case "uuid": case "uuid":
return "char(36)" return "char(36)"
case "vector":
return this.GenerateSQL("float[?]", column.Length)
} }
panic("Unsupported data type: " + column.Type) panic("Unsupported data type: " + column.Type)
} }
@@ -267,10 +280,13 @@ func (this Sqlite3) GenerateSQL(sql string, args ...any) string {
} }
func init() { func init() {
schema.RegisterDialect("sqlite3", func(db *db.Database) schema.Schema { sc := func(db *db.Database) schema.Schema {
return &Sqlite3{ return &Sqlite3{
db: db, db: db,
esg: sqlgen.NewExpressionSQLGenerator("sqlite3", sqlite3.DialectOptions()), esg: sqlgen.NewExpressionSQLGenerator("sqlite3", sqlite3.DialectOptions()),
} }
}) }
schema.RegisterDialect("sqlite3", sc)
schema.RegisterDialect("vtable", sc)
} }
+54 -34
View File
@@ -1,32 +1,16 @@
package sqlite3_test package sqlite3_test
import ( import (
"os"
"testing" "testing"
"git.fsdpf.net/go/db" "git.fsdpf.net/go/db"
"git.fsdpf.net/go/db/engine" "git.fsdpf.net/go/db/engine"
"git.fsdpf.net/go/db/schema" "git.fsdpf.net/go/db/schema"
sqlmock "github.com/DATA-DOG/go-sqlmock"
"github.com/stretchr/testify/suite" "github.com/stretchr/testify/suite"
_ "github.com/mattn/go-sqlite3"
) )
var ( var tableName = "entry"
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" +
");"
)
type sqlite3Test struct { type sqlite3Test struct {
suite.Suite suite.Suite
@@ -38,18 +22,28 @@ func TestSqlite3Suite(t *testing.T) {
} }
func (t *sqlite3Test) SetupSuite() { func (t *sqlite3Test) SetupSuite() {
db := engine.Open(map[string]engine.DBConfig{ mockDB, mock, _ := sqlmock.New()
"test-sqlite3": engine.NewDBConfig("sqlite3",
engine.WithSQLiteFile(os.Getenv("DB_FILE")), // 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") }).Connection("test-sqlite3")
t.schema = schema.GetSchemaDialect(db) t.schema = schema.GetSchemaDialect(conn)
// db.Exec(dropTable)
if _, err := db.Exec(createTable); err != nil {
t.Require().NoError(err)
}
} }
func (t *sqlite3Test) TestTableExists() { func (t *sqlite3Test) TestTableExists() {
@@ -123,14 +117,14 @@ func (t *sqlite3Test) TestCompileChange() {
sql := t.schema.CompileChange(bp) sql := t.schema.CompileChange(bp)
t.Equal(len([]string{ t.Equal(len([]string{
"ALTER TABLE `users` ADD COLUMN `name_1742738970482105700` varchar(50) NOT NULL", "ALTER TABLE `users` ADD COLUMN `name_tmp` varchar(50) NOT NULL",
"UPDATE `users` SET `name_1742738970482105700`=`name`", "UPDATE `users` SET `name_tmp`=`name`",
"ALTER TABLE `users` DROP COLUMN `name`", "ALTER TABLE `users` DROP COLUMN `name`",
"ALTER TABLE `users` RENAME COLUMN `name_1742738970482105700` TO `username`", "ALTER TABLE `users` RENAME COLUMN `name_tmp` TO `username`",
"ALTER TABLE `users` ADD COLUMN `age_1742738970482151800` smallint(4) NOT NULL DEFAULT '19'", "ALTER TABLE `users` ADD COLUMN `age_tmp` smallint(4) NOT NULL DEFAULT '19'",
"UPDATE `users` SET `age_1742738970482151800`=`age`", "UPDATE `users` SET `age_tmp`=`age`",
"ALTER TABLE `users` DROP COLUMN `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)) }), len(sql))
t.T().Log(sql) t.T().Log(sql)
@@ -178,6 +172,32 @@ func (t *sqlite3Test) TestCompileDropIfExists() {
t.T().Log(sql) t.T().Log(sql)
} }
// HIDDEN 列(虚拟表参数列)
func (t *sqlite3Test) TestCompileCreate_HiddenColumn() {
bp := schema.NewBlueprint("api_users")
bp.Create()
bp.BigIncrements("id").AutoIncrement()
bp.String("name", 100).Nullable()
bp.Integer("age")
bp.String("token", 255).Hidden().Nullable()
bp.Integer("page_size").Hidden().Nullable()
sql := t.schema.CompileCreate(bp)
t.Equal([]string{
"CREATE TABLE IF NOT EXISTS `api_users` (\n" +
"`id` INTEGER NOT NULL PRIMARY KEY AUTOINCREMENT\n" +
"`name` varchar(100) NULL\n" +
"`age` INTEGER NOT NULL\n" +
"`token` varchar(255) HIDDEN\n" +
"`page_size` INTEGER HIDDEN\n" +
")",
}, sql)
t.T().Log(sql)
}
// 表重命名 // 表重命名
func (t *sqlite3Test) TestCompileRename() { func (t *sqlite3Test) TestCompileRename() {
bp := schema.NewBlueprint("users") bp := schema.NewBlueprint("users")
+36 -52
View File
@@ -1,14 +1,12 @@
package sqlserver_test package sqlserver_test
import ( import (
"os"
"testing" "testing"
"git.fsdpf.net/go/db/engine" "git.fsdpf.net/go/db/engine"
"git.fsdpf.net/go/db/schema" "git.fsdpf.net/go/db/schema"
sqlmock "github.com/DATA-DOG/go-sqlmock"
"github.com/stretchr/testify/suite" "github.com/stretchr/testify/suite"
_ "github.com/microsoft/go-mssqldb"
) )
type sqlServerTest struct { type sqlServerTest struct {
@@ -16,41 +14,27 @@ type sqlServerTest struct {
schema schema.Schema schema schema.Schema
} }
var ( var tableName = "entry"
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" +
");"
)
func TestSQLServerSuite(t *testing.T) { func TestSQLServerSuite(t *testing.T) {
suite.Run(t, new(sqlServerTest)) suite.Run(t, new(sqlServerTest))
} }
func (t *sqlServerTest) SetupSuite() { func (t *sqlServerTest) SetupSuite() {
db := engine.Open(map[string]engine.DBConfig{ mockDB, mock, _ := sqlmock.New()
"test-sqlserver": engine.NewDBConfig("sqlserver",
engine.WithHost(os.Getenv("SQLSERVER_HOST")), mock.ExpectQuery("SELECT.*COLUMN_NAME.*COLUMNS").
engine.WithPort(os.Getenv("SQLSERVER_PORT")), WillReturnRows(sqlmock.NewRows([]string{"COLUMN_NAME"}).
engine.WithDatabase(os.Getenv("SQLSERVER_DB")), AddRow("id").AddRow("int").AddRow("float").
engine.WithUsername(os.Getenv("SQLSERVER_USER")), AddRow("string").AddRow("time").AddRow("bool").AddRow("bytes"))
engine.WithPassword(os.Getenv("SQLSERVER_PASSWD")),
), 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") }).Connection("test-sqlserver")
t.schema = schema.GetSchemaDialect(conn)
t.schema = schema.GetSchemaDialect(db)
db.Exec(dropTable)
db.Exec(createTable)
} }
func (t *sqlServerTest) TestGetColumnListing() { func (t *sqlServerTest) TestGetColumnListing() {
@@ -80,14 +64,14 @@ func (t *sqlServerTest) TestCompileCreate() {
sql := t.schema.CompileCreate(bp) sql := t.schema.CompileCreate(bp)
t.Equal([]string{ t.Equal([]string{
"CREATE TABLE [users] (\n" + "CREATE TABLE \"users\" (\n" +
"[id] BIGINT NOT NULL IDENTITY(1,1) PRIMARY KEY COMMENT 'ID',\n" + "\"id\" BIGINT NOT NULL IDENTITY(1,1) PRIMARY KEY,\n" +
"[enabled] BIT NOT NULL DEFAULT 1 COMMENT '是否有效',\n" + "\"enabled\" BIT NOT NULL DEFAULT '1',\n" +
"[created_user] CHAR(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000' COMMENT '创建者',\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' COMMENT '拥有者',\n" + "\"owned_user\" CHAR(36) NOT NULL DEFAULT '00000000-0000-0000-0000-000000000000',\n" +
"[created_at] DATETIME2 NOT NULL DEFAULT GETDATE() COMMENT '创建时间',\n" + "\"created_at\" DATETIME2 NOT NULL DEFAULT GETDATE(),\n" +
"[updated_at] DATETIME2 NOT NULL DEFAULT GETDATE() COMMENT '更新时间',\n" + "\"updated_at\" DATETIME2 NOT NULL DEFAULT GETDATE(),\n" +
"[deleted_at] DATETIME2 NULL COMMENT '删除时间'\n" + "\"deleted_at\" DATETIME2 NULL\n" +
")", ")",
}, sql) }, sql)
@@ -104,10 +88,10 @@ func (t *sqlServerTest) TestCompileAdd() {
sql := t.schema.CompileAdd(bp) sql := t.schema.CompileAdd(bp)
t.Equal([]string{ t.Equal([]string{
"ALTER TABLE [users]\n" + "ALTER TABLE \"users\"\n" +
"ADD [name] NVARCHAR(50) NOT NULL COMMENT '用户名',\n" + "ADD \"name\" NVARCHAR(50) NOT NULL,\n" +
"ADD [age] SMALLINT NOT NULL DEFAULT 18 COMMENT '年龄',\n" + "ADD \"age\" SMALLINT NOT NULL DEFAULT '18',\n" +
"ADD [sex] NVARCHAR(1) NOT NULL DEFAULT '0' COMMENT '性别'", "ADD \"sex\" NVARCHAR(1) NOT NULL DEFAULT '0'",
}, sql) }, sql)
t.T().Log(sql) t.T().Log(sql)
@@ -122,9 +106,9 @@ func (t *sqlServerTest) TestCompileChange() {
sql := t.schema.CompileChange(bp) sql := t.schema.CompileChange(bp)
t.Equal([]string{ t.Equal([]string{
"ALTER TABLE [users]\n" + "ALTER TABLE \"users\"\n" +
"ALTER COLUMN [name] NVARCHAR(50) NOT NULL,\n" + "ALTER COLUMN \"name\" NVARCHAR(50) NOT NULL,\n" +
"ALTER COLUMN [age] SMALLINT NOT NULL DEFAULT 19 COMMENT '年龄'", "ALTER COLUMN \"age\" SMALLINT NOT NULL DEFAULT '19'",
}, sql) }, sql)
t.T().Log(sql) t.T().Log(sql)
@@ -138,9 +122,9 @@ func (t *sqlServerTest) TestCompileDropColumn() {
sql := t.schema.CompileDropColumn(bp) sql := t.schema.CompileDropColumn(bp)
t.Equal([]string{ t.Equal([]string{
"ALTER TABLE [users]\n" + "ALTER TABLE \"users\"\n" +
"DROP COLUMN [age],\n" + "DROP COLUMN \"age\",\n" +
"DROP COLUMN [sex]", "DROP COLUMN \"sex\"",
}, sql) }, sql)
t.T().Log(sql) t.T().Log(sql)
@@ -153,7 +137,7 @@ func (t *sqlServerTest) TestCompileDrop() {
sql := t.schema.CompileDrop(bp) sql := t.schema.CompileDrop(bp)
t.Equal([]string{"DROP TABLE [users]"}, sql) t.Equal([]string{"DROP TABLE \"users\""}, sql)
t.T().Log(sql) t.T().Log(sql)
} }
@@ -164,7 +148,7 @@ func (t *sqlServerTest) TestCompileDropIfExists() {
bp.Drop() bp.Drop()
sql := t.schema.CompileDropIfExists(bp) 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) t.T().Log(sql)
} }
@@ -176,7 +160,7 @@ func (t *sqlServerTest) TestCompileRename() {
sql := t.schema.CompileRename(bp) 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) t.T().Log(sql)
} }
+3 -2
View File
@@ -133,10 +133,11 @@ func (sd *SelectDataset) Insert() *InsertDataset {
i := newInsertDataset(sd.dialect.Dialect(), sd.queryFactory). i := newInsertDataset(sd.dialect.Dialect(), sd.queryFactory).
Prepared(sd.isPrepared.Bool()) Prepared(sd.isPrepared.Bool())
if sd.clauses.HasSources() { 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()) i = i.Into(from.Aliased())
} else { } else {
i = i.Into(from) i = i.Into(col)
} }
} }
c := i.clauses c := i.clauses
+9
View File
@@ -559,6 +559,15 @@ func (esg *expressionSQLGenerator) literalExpressionSQL(b sb.SQLBuilder, literal
// //
// COUNT(I("a")) -> COUNT("a") // COUNT(I("a")) -> COUNT("a")
func (esg *expressionSQLGenerator) sqlFunctionExpressionSQL(b sb.SQLBuilder, sqlFunc exp.SQLFunctionExpression) { func (esg *expressionSQLGenerator) sqlFunctionExpressionSQL(b sb.SQLBuilder, sqlFunc exp.SQLFunctionExpression) {
if sqlFunc.Name() == "DISTINCT" && !esg.dialectOptions.SupportsDistinct {
for i, arg := range sqlFunc.Args() {
if i > 0 {
b.WriteRunes(esg.dialectOptions.CommaRune, esg.dialectOptions.SpaceRune)
}
esg.Generate(b, arg)
}
return
}
b.WriteStrings(sqlFunc.Name()) b.WriteStrings(sqlFunc.Name())
esg.Generate(b, sqlFunc.Args()) esg.Generate(b, sqlFunc.Args())
} }
+3
View File
@@ -34,6 +34,8 @@ type (
SupportsWithCTERecursive bool SupportsWithCTERecursive bool
// Set to true if multiple tables are supported in UPDATE statement. (DEFAULT=true) // Set to true if multiple tables are supported in UPDATE statement. (DEFAULT=true)
SupportsMultipleUpdateTables bool SupportsMultipleUpdateTables bool
// Set to true if DISTINCT is supported (DEFAULT=true)
SupportsDistinct bool
// Set to true if DISTINCT ON is supported (DEFAULT=true) // Set to true if DISTINCT ON is supported (DEFAULT=true)
SupportsDistinctOn bool SupportsDistinctOn bool
// Set to true if LATERAL queries are supported (DEFAULT=true) // Set to true if LATERAL queries are supported (DEFAULT=true)
@@ -420,6 +422,7 @@ func DefaultDialectOptions() *SQLDialectOptions {
SupportsConflictTarget: true, SupportsConflictTarget: true,
SupportsWithCTE: true, SupportsWithCTE: true,
SupportsWithCTERecursive: true, SupportsWithCTERecursive: true,
SupportsDistinct: true,
SupportsDistinctOn: true, SupportsDistinctOn: true,
WrapCompoundsInParens: true, WrapCompoundsInParens: true,
SupportsWindowFunction: true, SupportsWindowFunction: true,