all repos

viye @ master

my shot at reimplementing xiki

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

Oleksandr Smirnov Oleksandr Smirnov
olexsmir@gmail.com
refactor sql tool tests, 1 month ago
1
package sql
2
3
import (
4
	"context"
5
	"database/sql"
6
	"errors"
7
	"fmt"
8
	"strings"
9
10
	"github.com/olexsmir/viye/internal/config"
11
	"github.com/olexsmir/viye/internal/viye"
12
13
	_ "github.com/lib/pq"
14
	_ "modernc.org/sqlite"
15
)
16
17
type Tool struct{}
18
19
func (Tool) Name() string { return "sql(select, insert, update, delete, etc...)" }
20
func (Tool) Match(c *viye.Context) bool {
21
	switch strings.ToLower(c.Cmd) {
22
	case "select", "explain", "pragma":
23
		return len(c.Args) >= 1
24
	case "insert":
25
		return len(c.Args) >= 2 && strings.EqualFold(c.Args[0], "into")
26
	case "update":
27
		return len(c.Args) >= 2 && strings.EqualFold(c.Args[1], "set")
28
	case "delete":
29
		return len(c.Args) >= 2 && strings.EqualFold(c.Args[0], "from")
30
	case "create", "drop", "alter":
31
		return len(c.Args) >= 2
32
	case "analyze":
33
		return true
34
	}
35
	return false
36
}
37
38
func (Tool) Execute(c *viye.Context) (string, error) {
39
	ctx, cancel := context.WithTimeout(context.Background(), viye.Timeout)
40
	defer cancel()
41
42
	query := strings.Join(append([]string{c.Cmd}, c.Args...), " ")
43
	db, driver, err := openDB()
44
	if err != nil {
45
		return "", err
46
	}
47
	defer db.Close()
48
49
	if len(c.Body) > 0 {
50
		return updateFromBody(ctx, db, driver, query, c.Body)
51
	}
52
	return runQuery(ctx, db, query)
53
}
54
55
const maxRows = 100
56
57
func runQuery(ctx context.Context, db *sql.DB, query string) (string, error) {
58
	if !isSelect(query) {
59
		r, err := db.ExecContext(ctx, query)
60
		if err != nil {
61
			return "", fmt.Errorf("sql: %s", err)
62
		}
63
		n, _ := r.RowsAffected()
64
		return fmt.Sprintf("| %d row(s) affected", n), nil
65
	}
66
67
	rows, err := db.QueryContext(ctx, withLimit(query))
68
	if err != nil {
69
		return "", fmt.Errorf("sql: %s", err)
70
	}
71
	defer rows.Close()
72
	cols, err := rows.Columns()
73
	if err != nil {
74
		return "", fmt.Errorf("sql: %s", err)
75
	}
76
77
	colWidths := make([]int, len(cols))
78
	for i, c := range cols {
79
		colWidths[i] = len(c)
80
	}
81
82
	src := make([]sql.NullString, len(cols))
83
	ptrs := make([]any, len(cols))
84
	for i := range src {
85
		ptrs[i] = &src[i]
86
	}
87
88
	var rowsData [][]string
89
	for rows.Next() && len(rowsData) <= maxRows {
90
		if err := rows.Scan(ptrs...); err != nil {
91
			return "", fmt.Errorf("sql: scan: %w", err)
92
		}
93
		r := make([]string, len(cols))
94
		for i, v := range src {
95
			if v.Valid {
96
				r[i] = v.String
97
			} else {
98
				r[i] = "NULL"
99
			}
100
			if len(r[i]) > colWidths[i] {
101
				colWidths[i] = len(r[i])
102
			}
103
		}
104
		rowsData = append(rowsData, r)
105
	}
106
	if err := rows.Err(); err != nil {
107
		return "", fmt.Errorf("sql: rows: %w", err)
108
	}
109
110
	var buf strings.Builder
111
	writeRow := func(fields []string) {
112
		buf.WriteString(": ")
113
		for i, f := range fields {
114
			buf.WriteString(f)
115
			if i < len(fields)-1 {
116
				buf.WriteString(strings.Repeat(" ", colWidths[i]-len(f)+2))
117
			}
118
		}
119
		buf.WriteByte('\n')
120
	}
121
122
	truncated := len(rowsData) > maxRows
123
	if truncated {
124
		rowsData = rowsData[:maxRows]
125
	}
126
127
	writeRow(cols)
128
	for _, r := range rowsData {
129
		writeRow(r)
130
	}
131
	if truncated {
132
		fmt.Fprintf(&buf, "... %d+ more row(s), output truncated at %d rows\n", maxRows, maxRows)
133
	}
134
135
	return buf.String(), nil
136
}
137
138
func updateFromBody(ctx context.Context, db *sql.DB, driver, query string, body []string) (string, error) {
139
	if len(body) < 2 {
140
		return "", errors.New("sql: body must have header + at least one data row")
141
	}
142
	header := strings.Fields(body[0])
143
	if len(header) == 0 {
144
		return "", fmt.Errorf("sql: invalid header line")
145
	}
146
	pkCol := header[0]
147
	table, ok := extractTableName(query)
148
	if !ok {
149
		return "", fmt.Errorf("sql: could not determine table name")
150
	}
151
152
	var updated, inserted, skipped int
153
	for _, line := range body[1:] {
154
		vals := strings.Fields(line)
155
		if len(vals) != len(header) {
156
			skipped++
157
			continue
158
		}
159
160
		existsQuery := fmt.Sprintf("select exists(select 1 from %s where %s = %s)", table, pkCol, placeholder(driver, 0))
161
		var exists bool
162
		if err := db.QueryRowContext(ctx, existsQuery, vals[0]).Scan(&exists); err != nil {
163
			return "", fmt.Errorf("sql: check exists: %w", err)
164
		}
165
		if exists {
166
			var sets []string
167
			var args []any
168
			for i, col := range header {
169
				if col == pkCol {
170
					continue
171
				}
172
				sets = append(sets, fmt.Sprintf("%s = %s", col, placeholder(driver, len(args))))
173
				args = append(args, vals[i])
174
			}
175
			args = append(args, vals[0])
176
			updateQuery := fmt.Sprintf("update %s set %s where %s = %s", table, strings.Join(sets, ", "), pkCol, placeholder(driver, len(args)-1))
177
			if _, err := db.ExecContext(ctx, updateQuery, args...); err != nil {
178
				return "", fmt.Errorf("sql: update: %w", err)
179
			}
180
			updated++
181
		} else {
182
			phs := make([]string, len(header))
183
			args := make([]any, len(header))
184
			for i := range header {
185
				phs[i] = placeholder(driver, i)
186
				args[i] = vals[i]
187
			}
188
			insertQuery := fmt.Sprintf("INSERT INTO %s (%s) VALUES (%s)", table, strings.Join(header, ", "), strings.Join(phs, ", "))
189
			if _, err := db.ExecContext(ctx, insertQuery, args...); err != nil {
190
				return "", fmt.Errorf("sql: insert: %w", err)
191
			}
192
			inserted++
193
		}
194
	}
195
196
	msg := fmt.Sprintf("| %d updated, %d inserted", updated, inserted)
197
	if skipped > 0 {
198
		msg += fmt.Sprintf(", %d skipped (bad column count)", skipped)
199
	}
200
	return msg, nil
201
}
202
203
func placeholder(driver string, i int) string {
204
	if driver == "postgres" {
205
		return fmt.Sprintf("$%d", i+1)
206
	}
207
	return "?"
208
}
209
210
func isSelect(query string) bool {
211
	q := strings.ToLower(query)
212
	return strings.HasPrefix(q, "select") || strings.HasPrefix(q, "explain") ||
213
		strings.HasPrefix(q, "pragma")
214
}
215
216
func withLimit(query string) string {
217
	q := strings.TrimRight(strings.TrimSpace(query), "; \t\n")
218
	if !isSelect(q) || strings.HasPrefix(strings.ToLower(q), "pragma") || hasLimit(q) {
219
		return query
220
	}
221
	return q + fmt.Sprintf(" limit %d", maxRows+1)
222
}
223
224
func hasLimit(q string) bool {
225
	f := strings.Fields(q)
226
	for i := len(f) - 1; i >= 0 && i >= len(f)-4; i-- {
227
		if strings.EqualFold(f[i], "limit") {
228
			return true
229
		}
230
	}
231
	return false
232
}
233
234
func extractTableName(query string) (string, bool) {
235
	q := strings.ToUpper(strings.TrimSpace(query))
236
	for _, kw := range []string{"UPDATE ", "INTO ", "FROM "} {
237
		idx := strings.Index(q, kw)
238
		if idx >= 0 {
239
			parts := strings.Fields(query[idx+len(kw):])
240
			if len(parts) > 0 {
241
				return strings.Trim(parts[0], "\"`'[];"), true
242
			}
243
		}
244
	}
245
	return "", false
246
}
247
248
func openDB() (*sql.DB, string, error) {
249
	cfg, err := config.Load()
250
	if err != nil {
251
		return nil, "", err
252
	}
253
	dsn, ok := cfg.Get("dsn")
254
	if !ok {
255
		return nil, "", errors.New("dsn option is not set")
256
	}
257
258
	driver, err := driverForDSN(dsn)
259
	if err != nil {
260
		return nil, "", err
261
	}
262
263
	db, err := sql.Open(driver, dsn)
264
	if err != nil {
265
		return nil, "", err
266
	}
267
	if err := db.Ping(); err != nil {
268
		_ = db.Close()
269
		return nil, "", err
270
	}
271
	return db, driver, nil
272
}
273
274
func driverForDSN(dsn string) (string, error) {
275
	switch {
276
	case strings.HasPrefix(dsn, "postgres://") || strings.Contains(dsn, "user="):
277
		return "postgres", nil
278
	case strings.Contains(dsn, ".db") || strings.Contains(dsn, ".sqlite"):
279
		return "sqlite", nil
280
	default:
281
		return "", fmt.Errorf("sql: cannot determine driver from dsn %q", dsn)
282
	}
283
}