all repos

viye @ 91ebee3cd3f7b85bdeab49233d3a4a15c998c06b

my shot at reimplementing xiki

viye/internal/plumbing/sql/sql_test.go (view raw)

Oleksandr Smirnov Oleksandr Smirnov
olexsmir@gmail.com
refactor sql tool tests, 1 month ago
1
package sql
2
3
import (
4
	"database/sql"
5
	"strings"
6
	"testing"
7
8
	"github.com/olexsmir/viye/internal/viye"
9
	"olexsmir.xyz/x/is"
10
)
11
12
func TestMatch(t *testing.T) {
13
	for tname, tt := range map[string]struct {
14
		cmd  string
15
		args []string
16
		want bool
17
	}{
18
		"select":                    {"select", []string{"*", "from", "users"}, true},
19
		"select with no expression": {"select", nil, false},
20
		"insert":                    {"insert", []string{"into", "users"}, true},
21
		"insert without into":       {"insert", []string{"users"}, false},
22
		"update":                    {"update", []string{"users", "set", "x=1"}, true},
23
		"update without set":        {"update", []string{"users", "x=1"}, false},
24
		"delete":                    {"delete", []string{"from", "users"}, true},
25
		"delete missing table":      {"delete", []string{"from"}, false},
26
		"create":                    {"create", []string{"table", "t"}, true},
27
		"create incomplete":         {"create", []string{"table"}, false},
28
		"drop":                      {"drop", []string{"table", "t"}, true},
29
		"alter":                     {"alter", []string{"table", "t"}, true},
30
		"pragma":                    {"pragma", []string{"table_info(t)"}, true},
31
		"explain":                   {"explain", []string{"select", "1"}, true},
32
		"analyze alone":             {"analyze", nil, true},
33
		"uppercase":                 {"SELECT", []string{"*", "from", "t"}, true},
34
		"not sql":                   {"curl", []string{"https://x"}, false},
35
	} {
36
		t.Run(tname, func(t *testing.T) {
37
			if got := (Tool{}).Match(&viye.Context{Cmd: tt.cmd, Args: tt.args}); got != tt.want {
38
				t.Errorf("Match(%q, %v) = %v, want %v", tt.cmd, tt.args, got, tt.want)
39
			}
40
		})
41
	}
42
}
43
44
func TestRunQuery_select(t *testing.T) {
45
	db := newTestDB(t)
46
	_, err := db.Exec("insert into users values (1, 'alice', 30), (2, 'bob', NULL)")
47
	is.Err(t, err, nil)
48
49
	want := `: id  name   age
50
: 1   alice  30
51
: 2   bob    NULL
52
`
53
	for _, query := range []string{
54
		"select * from users order by id",
55
		"SELECT * FROM users ORDER BY id",
56
	} {
57
		out, err := runQuery(t.Context(), db, query)
58
		is.Err(t, err, nil)
59
		is.Equal(t, want, out)
60
	}
61
}
62
63
func TestRunQuery_truncation(t *testing.T) {
64
	db := newTestDB(t)
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")
66
	is.Err(t, err, nil)
67
68
	out, err := runQuery(t.Context(), db, "select * from users")
69
	is.Err(t, err, nil)
70
	if !strings.Contains(out, "... 100+ more row(s), output truncated at 100 rows") {
71
		t.Errorf("output not truncated:\n%s", out)
72
	}
73
}
74
75
func TestRunQuery_write(t *testing.T) {
76
	db := newTestDB(t)
77
	_, err := db.ExecContext(t.Context(), "insert into users values (1, 'alice', 30)")
78
	is.Err(t, err, nil)
79
80
	for _, tt := range []struct{ query, want string }{
81
		{"update users set age = 40 where id = 1", "| 1 row(s) affected"},
82
		{"insert into users values (3, 'carol', 40)", "| 1 row(s) affected"},
83
		{"delete from users where id = 3", "| 1 row(s) affected"},
84
	} {
85
		out, err := runQuery(t.Context(), db, tt.query)
86
		is.Err(t, err, nil)
87
		is.Equal(t, tt.want, out)
88
	}
89
}
90
91
func TestUpdateFromBody(t *testing.T) {
92
	db := newTestDB(t)
93
	_, err := db.ExecContext(t.Context(), "insert into users values (1, 'alice', 30)")
94
	is.Err(t, err, nil)
95
96
	out, err := updateFromBody(t.Context(), db, "sqlite", "update users",
97
		[]string{"id name age", "1 alice 31", "3 carol 40", "garbage"})
98
	is.Err(t, err, nil)
99
	is.Equal(t, "| 1 updated, 1 inserted, 1 skipped (bad column count)", out)
100
101
	var count, aliceAge int
102
	err = db.QueryRowContext(t.Context(), "select count(*) from users").Scan(&count)
103
	is.Err(t, err, nil)
104
	is.Equal(t, 2, count)
105
106
	err = db.QueryRowContext(t.Context(), "select age from users where id = 1").Scan(&aliceAge)
107
	is.Err(t, err, nil)
108
	is.Equal(t, 31, aliceAge)
109
}
110
111
func TestExtractTableName(t *testing.T) {
112
	for _, tt := range []struct {
113
		name, query, want string
114
		ok                bool
115
	}{
116
		{"update", "update users set name = 'x' where id = 1", "users", true},
117
		{"update with subquery", "update users set n = (select 1 from roles) where id = 1", "users", true},
118
		{"select", "select * from users;", "users", true},
119
		{"insert", "insert into logs (a) values (1)", "logs", true},
120
	} {
121
		t.Run(tt.name, func(t *testing.T) {
122
			got, ok := extractTableName(tt.query)
123
			is.Equal(t, tt.want, got)
124
			is.Equal(t, tt.ok, ok)
125
		})
126
	}
127
}
128
129
func TestDriverForDSN(t *testing.T) {
130
	for _, tt := range []struct {
131
		dsn, want string
132
		err       any
133
	}{
134
		{"postgres://user:pass@host/db", "postgres", nil},
135
		{"host=db user=alice", "postgres", nil},
136
		{"postgres://host/name.db", "postgres", nil},
137
		{"/tmp/data.db", "sqlite", nil},
138
		{"/tmp/app.sqlite3", "sqlite", nil},
139
		{"mongodb://host/db", "", "cannot determine driver"},
140
	} {
141
		got, err := driverForDSN(tt.dsn)
142
		is.Err(t, err, tt.err)
143
		is.Equal(t, tt.want, got)
144
	}
145
}
146
147
func TestPlaceholder(t *testing.T) {
148
	is.Equal(t, "$3", placeholder("postgres", 2))
149
	for _, i := range []int{0, 3} {
150
		is.Equal(t, "?", placeholder("sqlite", i))
151
	}
152
}
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
166
func newTestDB(t *testing.T) *sql.DB {
167
	t.Helper()
168
	db, err := sql.Open("sqlite", ":memory:")
169
	is.Err(t, err, nil)
170
	db.SetMaxOpenConns(1) // keep connection shared across queries.
171
	t.Cleanup(func() { db.Close() })
172
	_, err = db.ExecContext(t.Context(), "create table users (id integer primary key, name text, age integer)")
173
	is.Err(t, err, nil)
174
	return db
175
}