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