//go:build sqlite_vtable || vtable package vtab_test import ( "database/sql" "testing" "git.fsdpf.net/go/db/dialect/sqlite3/vtab" "github.com/stretchr/testify/suite" ) // ─── 内存表实现(用于测试)─────────────────────────────────── // // 模拟一张 users 表,数据存在内存 slice 中,支持完整 CRUD。 type userRow struct { id int64 name string age int64 } // usersModule 是虚拟表工厂 type usersModule struct{} func (m *usersModule) Create(args []string, declare func(string) error) (vtab.Table, error) { if err := declare("CREATE TABLE users(id INTEGER, name TEXT, age INTEGER)"); err != nil { return nil, err } return newUsersTable(), nil } func (m *usersModule) Connect(args []string, declare func(string) error) (vtab.Table, error) { return m.Create(args, declare) } // usersTable 是虚拟表实例,同时实现 WritableTable type usersTable struct { rows []userRow nextID int64 } func newUsersTable() *usersTable { return &usersTable{ nextID: 1, rows: []userRow{ {1, "Alice", 30}, {2, "Bob", 25}, {3, "Charlie", 35}, }, } } // BestIndex:不做任何约束下推,由 SQLite 全表扫描后过滤 func (t *usersTable) BestIndex(constraints []vtab.ConstraintInfo, _ []vtab.OrderByInfo) (*vtab.IndexOutput, error) { return &vtab.IndexOutput{ Used: make([]bool, len(constraints)), // 全部 false EstimatedCost: float64(len(t.rows)), EstimatedRows: float64(len(t.rows)), }, nil } func (t *usersTable) Open() (vtab.Cursor, error) { return &usersCursor{rows: t.rows}, nil } func (t *usersTable) Disconnect() error { return nil } func (t *usersTable) Destroy() error { return nil } // ── WritableTable ── func (t *usersTable) Insert(values []any) (int64, error) { // values: [id, name, age],id 可能为 nil(自动生成) id := t.nextID t.nextID++ if values[0] != nil { id = toInt64(values[0]) } name, _ := values[1].(string) age := toInt64(values[2]) t.rows = append(t.rows, userRow{id, name, age}) return id, nil } func (t *usersTable) Update(rowid any, values []any) error { rid := toInt64(rowid) for i, r := range t.rows { if r.id == rid { if values[1] != nil { t.rows[i].name, _ = values[1].(string) } if values[2] != nil { t.rows[i].age = toInt64(values[2]) } return nil } } return nil } func (t *usersTable) Delete(rowid any) error { rid := toInt64(rowid) for i, r := range t.rows { if r.id == rid { t.rows = append(t.rows[:i], t.rows[i+1:]...) return nil } } return nil } // ── Cursor ── type usersCursor struct { rows []userRow pos int } func (c *usersCursor) Filter(_ int, constraints []vtab.ConstraintInfo) error { var filtered []userRow for _, row := range c.rows { if rowMatchesAll(row, constraints) { filtered = append(filtered, row) } } c.rows = filtered c.pos = 0 return nil } func rowMatchesAll(row userRow, constraints []vtab.ConstraintInfo) bool { for _, c := range constraints { switch c.Column { case 0: // id v := toInt64(c.Value) switch c.Op { case vtab.OpEQ: if row.id != v { return false } case vtab.OpGT: if !(row.id > v) { return false } case vtab.OpGE: if !(row.id >= v) { return false } case vtab.OpLT: if !(row.id < v) { return false } case vtab.OpLE: if !(row.id <= v) { return false } } case 1: // name val, _ := c.Value.(string) if c.Op == vtab.OpEQ && row.name != val { return false } case 2: // age v := toInt64(c.Value) switch c.Op { case vtab.OpEQ: if row.age != v { return false } case vtab.OpGT: if !(row.age > v) { return false } case vtab.OpGE: if !(row.age >= v) { return false } case vtab.OpLT: if !(row.age < v) { return false } case vtab.OpLE: if !(row.age <= v) { return false } } } } return true } func (c *usersCursor) Next() error { c.pos++; return nil } func (c *usersCursor) EOF() bool { return c.pos >= len(c.rows) } func (c *usersCursor) Close() error { return nil } func (c *usersCursor) Rowid() (int64, error) { return c.rows[c.pos].id, nil } func (c *usersCursor) Column(col int) (any, error) { r := c.rows[c.pos] switch col { case 0: return r.id, nil case 1: return r.name, nil case 2: return r.age, nil } return nil, nil } func toInt64(v any) int64 { switch x := v.(type) { case int64: return x case int: return int64(x) case float64: return int64(x) } return 0 } // ─── 测试套件 ───────────────────────────────────────────────── func init() { vtab.Register("users_mod", &usersModule{}) } type VtabSuite struct { suite.Suite db *sql.DB } func (s *VtabSuite) SetupSuite() { db, err := sql.Open(vtab.DriverName, ":memory:") s.Require().NoError(err) s.db = db } func (s *VtabSuite) TearDownSuite() { s.db.Close() } // SetupTest 每个测试前重建虚拟表,保证初始数据一致 func (s *VtabSuite) SetupTest() { s.db.Exec(`DROP TABLE IF EXISTS users`) _, err := s.db.Exec(`CREATE VIRTUAL TABLE users USING users_mod()`) s.Require().NoError(err) } // ── SELECT ── func (s *VtabSuite) TestSelect_All() { rows, err := s.db.Query(`SELECT id, name, age FROM users ORDER BY id`) s.Require().NoError(err) defer rows.Close() var result []userRow for rows.Next() { var r userRow s.Require().NoError(rows.Scan(&r.id, &r.name, &r.age)) result = append(result, r) } s.Require().NoError(rows.Err()) s.Equal([]userRow{ {1, "Alice", 30}, {2, "Bob", 25}, {3, "Charlie", 35}, }, result) } func (s *VtabSuite) TestSelect_Where() { rows, err := s.db.Query(`SELECT name FROM users WHERE age > 28 ORDER BY id`) s.Require().NoError(err) defer rows.Close() var names []string for rows.Next() { var name string s.Require().NoError(rows.Scan(&name)) names = append(names, name) } s.Equal([]string{"Alice", "Charlie"}, names) } func (s *VtabSuite) TestSelect_IF() { var label string err := s.db.QueryRow(`SELECT IF(age >= 30, 'senior', 'junior') FROM users WHERE id = 2`).Scan(&label) s.Require().NoError(err) s.Equal("junior", label) } // ── INSERT ── func (s *VtabSuite) TestInsert() { _, err := s.db.Exec(`INSERT INTO users(name, age) VALUES(?, ?)`, "Dave", 28) s.Require().NoError(err) var count int s.Require().NoError(s.db.QueryRow(`SELECT COUNT(*) FROM users WHERE name = 'Dave'`).Scan(&count)) s.Equal(1, count) } // ── UPDATE ── func (s *VtabSuite) TestUpdate() { _, err := s.db.Exec(`UPDATE users SET age = ? WHERE name = ?`, 26, "Bob") s.Require().NoError(err) var age int64 s.Require().NoError(s.db.QueryRow(`SELECT age FROM users WHERE name = 'Bob'`).Scan(&age)) s.Equal(int64(26), age) } // ── DELETE ── func (s *VtabSuite) TestDelete() { _, err := s.db.Exec(`DELETE FROM users WHERE name = ?`, "Charlie") s.Require().NoError(err) var count int s.Require().NoError(s.db.QueryRow(`SELECT COUNT(*) FROM users WHERE name = 'Charlie'`).Scan(&count)) s.Equal(0, count) } // ── JOIN with real table ── // TestJoin_VtabAndRealTable 演示虚拟表与真实表的 JOIN。 // // 真实表 orders:每条订单记录 user_id 和 amount。 // 虚拟表 users:内存中的用户数据。 // 查询:统计每个用户的订单总金额,只返回有订单的用户。 func (s *VtabSuite) TestJoin_VtabAndRealTable() { // 建真实表并插入数据 _, err := s.db.Exec(` CREATE TABLE IF NOT EXISTS orders ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL, amount REAL NOT NULL )`) s.Require().NoError(err) defer s.db.Exec(`DROP TABLE IF EXISTS orders`) _, err = s.db.Exec(` INSERT INTO orders(user_id, amount) VALUES (1, 100.0), (1, 50.5), (2, 200.0), (3, 75.0), (3, 25.0)`) s.Require().NoError(err) // vtab users INNER JOIN real orders rows, err := s.db.Query(` SELECT u.name, SUM(o.amount) AS total FROM users AS u JOIN orders AS o ON o.user_id = u.id GROUP BY u.id, u.name ORDER BY u.id`) s.Require().NoError(err) defer rows.Close() type row struct { name string total float64 } var result []row for rows.Next() { var r row s.Require().NoError(rows.Scan(&r.name, &r.total)) result = append(result, r) } s.Require().NoError(rows.Err()) s.Equal([]row{ {"Alice", 150.5}, {"Bob", 200.0}, {"Charlie", 100.0}, }, result) } // TestJoin_VtabLeftJoinRealTable 演示 LEFT JOIN:列出所有用户及其订单数,没有订单的用户显示 0。 func (s *VtabSuite) TestJoin_VtabLeftJoinRealTable() { _, err := s.db.Exec(` CREATE TABLE IF NOT EXISTS orders ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL, amount REAL NOT NULL )`) s.Require().NoError(err) defer s.db.Exec(`DROP TABLE IF EXISTS orders`) // 只给 Alice 和 Bob 插入订单,Charlie 没有订单 _, err = s.db.Exec(` INSERT INTO orders(user_id, amount) VALUES (1, 100.0), (2, 200.0)`) s.Require().NoError(err) rows, err := s.db.Query(` SELECT u.name, COUNT(o.id) AS order_count FROM users AS u LEFT JOIN orders AS o ON o.user_id = u.id GROUP BY u.id, u.name ORDER BY u.id`) s.Require().NoError(err) defer rows.Close() type row struct { name string count int } var result []row for rows.Next() { var r row s.Require().NoError(rows.Scan(&r.name, &r.count)) result = append(result, r) } s.Require().NoError(rows.Err()) s.Equal([]row{ {"Alice", 1}, {"Bob", 1}, {"Charlie", 0}, // 没有订单,LEFT JOIN 保留 }, result) } func TestVtabSuite(t *testing.T) { suite.Run(t, new(VtabSuite)) }