all repos

viye @ 91ebee3

my shot at reimplementing xiki
2 files changed, 27 insertions(+), 33 deletions(-)
refactor sql tool tests
Author: Oleksandr Smirnov olexsmir@gmail.com
Committed at: 2026-08-31 13:52:14 +0300
Authored at: 2026-08-30 21:09:47 +0300
Change ID: ymvtnlmvtllzumlrlntqxntxkszpqktw
Parent: cf8a742
M internal/plumbing/sql/sql.go
···
        16
        16
         

      
        17
        17
         type Tool struct{}

      
        18
        18
         

      
        19
        
        -func (Tool) Name() string { return "sql(select, insert, update, delete)" }

      
        
        19
        +func (Tool) Name() string { return "sql(select, insert, update, delete, etc...)" }

      
        20
        20
         func (Tool) Match(c *viye.Context) bool {

      
        21
        21
         	switch strings.ToLower(c.Cmd) {

      
        22
        
        -	case "select", "explain":

      
        
        22
        +	case "select", "explain", "pragma":

      
        23
        23
         		return len(c.Args) >= 1

      
        24
        24
         	case "insert":

      
        25
        25
         		return len(c.Args) >= 2 && strings.EqualFold(c.Args[0], "into")

      ···
        29
        29
         		return len(c.Args) >= 2 && strings.EqualFold(c.Args[0], "from")

      
        30
        30
         	case "create", "drop", "alter":

      
        31
        31
         		return len(c.Args) >= 2

      
        32
        
        -	case "pragma":

      
        33
        
        -		return len(c.Args) >= 1

      
        34
        32
         	case "analyze":

      
        35
        33
         		return true

      
        36
        34
         	}

      
M internal/plumbing/sql/sql_test.go
···
        1
        1
         package sql

      
        2
        2
         

      
        3
        3
         import (

      
        4
        
        -	"context"

      
        5
        4
         	"database/sql"

      
        6
        5
         	"strings"

      
        7
        6
         	"testing"

      ···
        10
        9
         	"olexsmir.xyz/x/is"

      
        11
        10
         )

      
        12
        11
         

      
        13
        
        -func TestMatchSQL(t *testing.T) {

      
        
        12
        +func TestMatch(t *testing.T) {

      
        14
        13
         	for tname, tt := range map[string]struct {

      
        15
        14
         		cmd  string

      
        16
        15
         		args []string

      ···
        44
        43
         

      
        45
        44
         func TestRunQuery_select(t *testing.T) {

      
        46
        45
         	db := newTestDB(t)

      
        47
        
        -	ctx := context.Background()

      
        48
        46
         	_, err := db.Exec("insert into users values (1, 'alice', 30), (2, 'bob', NULL)")

      
        49
        47
         	is.Err(t, err, nil)

      
        50
        48
         

      ···
        56
        54
         		"select * from users order by id",

      
        57
        55
         		"SELECT * FROM users ORDER BY id",

      
        58
        56
         	} {

      
        59
        
        -		out, err := runQuery(ctx, db, query)

      
        
        57
        +		out, err := runQuery(t.Context(), db, query)

      
        60
        58
         		is.Err(t, err, nil)

      
        61
        59
         		is.Equal(t, want, out)

      
        62
        60
         	}

      ···
        64
        62
         

      
        65
        63
         func TestRunQuery_truncation(t *testing.T) {

      
        66
        64
         	db := newTestDB(t)

      
        67
        
        -	_, err := db.Exec("with recursive c(x) as (select 1 union all select x+1 from c where x < 105) insert into users select x, 'u'||x, x from c")

      
        
        65
        +	_, err := db.ExecContext(t.Context(), "with recursive c(x) as (select 1 union all select x+1 from c where x < 105) insert into users select x, 'u'||x, x from c")

      
        68
        66
         	is.Err(t, err, nil)

      
        69
        67
         

      
        70
        
        -	out, err := runQuery(context.Background(), db, "select * from users")

      
        
        68
        +	out, err := runQuery(t.Context(), db, "select * from users")

      
        71
        69
         	is.Err(t, err, nil)

      
        72
        70
         	if !strings.Contains(out, "... 100+ more row(s), output truncated at 100 rows") {

      
        73
        71
         		t.Errorf("output not truncated:\n%s", out)

      
        74
        72
         	}

      
        75
        73
         }

      
        76
        74
         

      
        77
        
        -func TestWithLimit(t *testing.T) {

      
        78
        
        -	for _, tt := range []struct{ in, want string }{

      
        79
        
        -		{"select * from users;", "select * from users limit 101"},

      
        80
        
        -		{"select * from users limit 10", "select * from users limit 10"},

      
        81
        
        -		{"select * from users limit 10 offset 5", "select * from users limit 10 offset 5"},

      
        82
        
        -		{"pragma table_info(users)", "pragma table_info(users)"},

      
        83
        
        -		{"update users set age = 1", "update users set age = 1"},

      
        84
        
        -	} {

      
        85
        
        -		is.Equal(t, tt.want, withLimit(tt.in))

      
        86
        
        -	}

      
        87
        
        -}

      
        88
        
        -

      
        89
        75
         func TestRunQuery_write(t *testing.T) {

      
        90
        76
         	db := newTestDB(t)

      
        91
        
        -	ctx := context.Background()

      
        92
        
        -	_, err := db.Exec("insert into users values (1, 'alice', 30)")

      
        
        77
        +	_, err := db.ExecContext(t.Context(), "insert into users values (1, 'alice', 30)")

      
        93
        78
         	is.Err(t, err, nil)

      
        94
        79
         

      
        95
        80
         	for _, tt := range []struct{ query, want string }{

      ···
        97
        82
         		{"insert into users values (3, 'carol', 40)", "| 1 row(s) affected"},

      
        98
        83
         		{"delete from users where id = 3", "| 1 row(s) affected"},

      
        99
        84
         	} {

      
        100
        
        -		out, err := runQuery(ctx, db, tt.query)

      
        
        85
        +		out, err := runQuery(t.Context(), db, tt.query)

      
        101
        86
         		is.Err(t, err, nil)

      
        102
        87
         		is.Equal(t, tt.want, out)

      
        103
        88
         	}

      ···
        105
        90
         

      
        106
        91
         func TestUpdateFromBody(t *testing.T) {

      
        107
        92
         	db := newTestDB(t)

      
        108
        
        -	ctx := context.Background()

      
        109
        
        -	_, err := db.Exec("insert into users values (1, 'alice', 30)")

      
        
        93
        +	_, err := db.ExecContext(t.Context(), "insert into users values (1, 'alice', 30)")

      
        110
        94
         	is.Err(t, err, nil)

      
        111
        95
         

      
        112
        
        -	out, err := updateFromBody(ctx, db, "sqlite", "update users",

      
        
        96
        +	out, err := updateFromBody(t.Context(), db, "sqlite", "update users",

      
        113
        97
         		[]string{"id name age", "1 alice 31", "3 carol 40", "garbage"})

      
        114
        98
         	is.Err(t, err, nil)

      
        115
        99
         	is.Equal(t, "| 1 updated, 1 inserted, 1 skipped (bad column count)", out)

      
        116
        100
         

      
        117
        101
         	var count, aliceAge int

      
        118
        
        -	err = db.QueryRow("select count(*) from users").Scan(&count)

      
        
        102
        +	err = db.QueryRowContext(t.Context(), "select count(*) from users").Scan(&count)

      
        119
        103
         	is.Err(t, err, nil)

      
        120
        104
         	is.Equal(t, 2, count)

      
        121
        
        -	err = db.QueryRow("select age from users where id = 1").Scan(&aliceAge)

      
        
        105
        +

      
        
        106
        +	err = db.QueryRowContext(t.Context(), "select age from users where id = 1").Scan(&aliceAge)

      
        122
        107
         	is.Err(t, err, nil)

      
        123
        108
         	is.Equal(t, 31, aliceAge)

      
        124
        109
         }

      ···
        166
        151
         	}

      
        167
        152
         }

      
        168
        153
         

      
        
        154
        +func TestWithLimit(t *testing.T) {

      
        
        155
        +	for _, tt := range []struct{ in, want string }{

      
        
        156
        +		{"select * from users;", "select * from users limit 101"},

      
        
        157
        +		{"select * from users limit 10", "select * from users limit 10"},

      
        
        158
        +		{"select * from users limit 10 offset 5", "select * from users limit 10 offset 5"},

      
        
        159
        +		{"pragma table_info(users)", "pragma table_info(users)"},

      
        
        160
        +		{"update users set age = 1", "update users set age = 1"},

      
        
        161
        +	} {

      
        
        162
        +		is.Equal(t, tt.want, withLimit(tt.in))

      
        
        163
        +	}

      
        
        164
        +}

      
        
        165
        +

      
        169
        166
         func newTestDB(t *testing.T) *sql.DB {

      
        170
        167
         	t.Helper()

      
        171
        168
         	db, err := sql.Open("sqlite", ":memory:")

      
        172
        169
         	is.Err(t, err, nil)

      
        173
        
        -	// Single connection keeps the in-memory database shared across queries.

      
        174
        
        -	db.SetMaxOpenConns(1)

      
        
        170
        +	db.SetMaxOpenConns(1) // keep connection shared across queries.

      
        175
        171
         	t.Cleanup(func() { db.Close() })

      
        176
        
        -	_, err = db.Exec("create table users (id integer primary key, name text, age integer)")

      
        
        172
        +	_, err = db.ExecContext(t.Context(), "create table users (id integer primary key, name text, age integer)")

      
        177
        173
         	is.Err(t, err, nil)

      
        178
        174
         	return db

      
        179
        175
         }