viye/internal/plumbing/sql/sql.go (view raw)
| 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 | } |