- 新增 dialect/sqlite3/vtab 包,提供干净的虚拟表接口(Module/Table/Cursor) - 适配层自动将 BestIndex 约束与 Filter 值绑定(ConstraintInfo.Value),无需手动编解码 IdxStr - 支持 OpLIMIT/OpOFFSET 约束下推 - 新增 SupportsDistinct 方言选项,控制 SELECT 级和表达式级 DISTINCT 生成 - sqlite3 方言注册 IF() 函数支持
366 lines
8.4 KiB
Go
366 lines
8.4 KiB
Go
//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))
|
||
}
|