all repos

viye @ 66ef839e7fe6e1655d8036d32fce619faa7b1be0

my shot at reimplementing xiki

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

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