From ac81f1ff0b32b1d4dcefbad62eebf33845344e8c Mon Sep 17 00:00:00 2001 From: what Date: Sat, 18 Apr 2026 13:41:41 +0800 Subject: [PATCH] =?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,