all repos

viye @ f47d92e4476f8153e98a5b6fc09bf418afdb2888

my shot at reimplementing xiki
10 files changed, 577 insertions(+), 26 deletions(-)
add sql tool
Author: Oleksandr Smirnov olexsmir@gmail.com
Committed at: 2026-08-29 16:30:24 +0300
Authored at: 2026-08-29 14:14:17 +0300
Change ID: xzllrunyxlxqlwopmqvnzrolozpvvwuk
Parent: f5220eb
M go.mod
···
        2
        2
         

      
        3
        3
         go 1.27.0

      
        4
        4
         

      
        5
        
        -require olexsmir.xyz/x v0.3.1

      
        
        5
        +require (

      
        
        6
        +	github.com/lib/pq v1.12.3

      
        
        7
        +	modernc.org/sqlite v1.57.0

      
        
        8
        +	olexsmir.xyz/x v0.3.1

      
        
        9
        +)

      
        
        10
        +

      
        
        11
        +require (

      
        
        12
        +	github.com/dustin/go-humanize v1.0.1 // indirect

      
        
        13
        +	github.com/google/uuid v1.6.0 // indirect

      
        
        14
        +	github.com/mattn/go-isatty v0.0.24 // indirect

      
        
        15
        +	github.com/ncruces/go-strftime v1.0.0 // indirect

      
        
        16
        +	github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect

      
        
        17
        +	golang.org/x/sys v0.47.0 // indirect

      
        
        18
        +	modernc.org/libc v1.74.4 // indirect

      
        
        19
        +	modernc.org/mathutil v1.7.1 // indirect

      
        
        20
        +	modernc.org/memory v1.11.0 // indirect

      
        
        21
        +)

      
M go.sum
···
        
        1
        +github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=

      
        
        2
        +github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=

      
        
        3
        +github.com/google/pprof v0.0.0-20260802141513-ef3492d7dac3 h1:LMLX+LgTNWpfvCBdFebv6EsYotImrt/Ppc5cXIriCSo=

      
        
        4
        +github.com/google/pprof v0.0.0-20260802141513-ef3492d7dac3/go.mod h1:jl5iWTm0/hd5PjEYEOuwAJ57L/CibdZfrqZ5XA5GrCk=

      
        
        5
        +github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=

      
        
        6
        +github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=

      
        
        7
        +github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=

      
        
        8
        +github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=

      
        
        9
        +github.com/lib/pq v1.12.3 h1:tTWxr2YLKwIvK90ZXEw8GP7UFHtcbTtty8zsI+YjrfQ=

      
        
        10
        +github.com/lib/pq v1.12.3/go.mod h1:/p+8NSbOcwzAEI7wiMXFlgydTwcgTr3OSKMsD2BitpA=

      
        
        11
        +github.com/mattn/go-isatty v0.0.24 h1:tGZZoVgT/KiqK1c8ocVLeDS8BSWMRd47J3Lbz7vsReI=

      
        
        12
        +github.com/mattn/go-isatty v0.0.24/go.mod h1:nMCL3Zebbrt45jsMDgnfIwz6ydEQApk5oEI3HqDio6A=

      
        
        13
        +github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w=

      
        
        14
        +github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls=

      
        
        15
        +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=

      
        
        16
        +github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=

      
        
        17
        +golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ=

      
        
        18
        +golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0=

      
        
        19
        +golang.org/x/sync v0.21.0 h1:HLII4xRRTtCRkxYp4HNFF0Js/Og6q2i++KXbg0gHCwM=

      
        
        20
        +golang.org/x/sync v0.21.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=

      
        
        21
        +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=

      
        
        22
        +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=

      
        
        23
        +golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q=

      
        
        24
        +golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA=

      
        
        25
        +modernc.org/cc/v4 v4.29.1 h1:MKgdCV3WykTSPqpVrnxdEDS0HEd2FHpKZDzxzU5LyeI=

      
        
        26
        +modernc.org/cc/v4 v4.29.1/go.mod h1:OnovgIhbbMXMu1aISnJ0wvVD1KnW+cAUJkIrAWh+kVI=

      
        
        27
        +modernc.org/ccgo/v4 v4.34.6 h1:sBgfIwyN0TQ9C5hwIeuqyeAKyMWnbvj2fvpF4L11uzU=

      
        
        28
        +modernc.org/ccgo/v4 v4.34.6/go.mod h1:SZ8YcN9NG7XVsQYdm6jYBvi8PQP1qi+kqB6OhjqI3Fk=

      
        
        29
        +modernc.org/fileutil v1.4.0 h1:j6ZzNTftVS054gi281TyLjHPp6CPHr2KCxEXjEbD6SM=

      
        
        30
        +modernc.org/fileutil v1.4.0/go.mod h1:EqdKFDxiByqxLk8ozOxObDSfcVOv/54xDs/DUHdvCUU=

      
        
        31
        +modernc.org/gc/v2 v2.6.5 h1:nyqdV8q46KvTpZlsw66kWqwXRHdjIlJOhG6kxiV/9xI=

      
        
        32
        +modernc.org/gc/v2 v2.6.5/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito=

      
        
        33
        +modernc.org/gc/v3 v3.1.4 h1:2g65LGVSmFQrXeITAw97x7hCRvZFcyE1uDP+7Vng7JI=

      
        
        34
        +modernc.org/gc/v3 v3.1.4/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY=

      
        
        35
        +modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks=

      
        
        36
        +modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI=

      
        
        37
        +modernc.org/libc v1.74.4 h1:fX1Omw4o2/1C2iRkkIsrQTasJQldLhRmuPreXLoWs9k=

      
        
        38
        +modernc.org/libc v1.74.4/go.mod h1:eeQAS9W3sZeKYMFubydxJpII9ybHWshk+7or7bLG9co=

      
        
        39
        +modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=

      
        
        40
        +modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=

      
        
        41
        +modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI=

      
        
        42
        +modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw=

      
        
        43
        +modernc.org/opt v0.2.0 h1:tGyef5ApycA7FSEOMraay9SaTk5zmbx7Tu+cJs4QKZg=

      
        
        44
        +modernc.org/opt v0.2.0/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=

      
        
        45
        +modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w=

      
        
        46
        +modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE=

      
        
        47
        +modernc.org/sqlite v1.57.0 h1:qNQP6xnx5M0ISNtlnxoOX0+cD5bJ0/gr9aMmndFczzg=

      
        
        48
        +modernc.org/sqlite v1.57.0/go.mod h1:yCJ2cmAaIkHQ25oXWrF8H4O1lIfPYPR26yCEDj2P3pQ=

      
        
        49
        +modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0=

      
        
        50
        +modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A=

      
        
        51
        +modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y=

      
        
        52
        +modernc.org/token v1.1.0/go.mod h1:UGzOrNV1mAFSEB63lOFHIpNRUVMvYTc6yu1SMY/XTDM=

      
        1
        53
         olexsmir.xyz/x v0.3.1 h1:wlJjZG98duc2ZX3z1Hqx48DKUWwOFObE9dNPZqg3TK4=

      
        2
        54
         olexsmir.xyz/x v0.3.1/go.mod h1:79Z6oaW+ZCKZahm5r4Cgxbv1Ey8IQuAi1uXMhyJFJto=

      
M internal/config/config.go
···
        63
        63
         }

      
        64
        64
         

      
        65
        65
         func (c *Config) GetAll() map[string]string { return c.cfg }

      
        66
        
        -func (c *Config) Get(key string) string     { return c.cfg[key] }

      
        
        66
        +func (c *Config) Get(key string) (val string, ok bool) {

      
        
        67
        +	val, ok = c.cfg[key]

      
        
        68
        +	return val, ok

      
        
        69
        +}

      
        
        70
        +

      
        67
        71
         func (c *Config) Set(key, value string) error {

      
        68
        72
         	c.cfg[key] = value

      
        69
        73
         	return c.save()

      
M internal/config/config_test.go
···
        19
        19
         		cfg, err := Load()

      
        20
        20
         		is.Err(t, err, nil)

      
        21
        21
         		for key, val := range want.cfg {

      
        22
        
        -			is.Equal(t, val, cfg.Get(key))

      
        
        22
        +			v, _ := cfg.Get(key)

      
        
        23
        +			is.Equal(t, val, v)

      
        23
        24
         		}

      
        24
        25
         	})

      
        25
        26
         

      ···
        32
        33
         

      
        33
        34
         		cfg, err := Load()

      
        34
        35
         		is.Err(t, err, nil)

      
        35
        
        -		is.Equal(t, "value", cfg.Get("custom"))

      
        
        36
        +

      
        
        37
        +		val, ok := cfg.Get("custom")

      
        
        38
        +		is.Equal(t, "value", val)

      
        
        39
        +		is.Equal(t, true, ok)

      
        36
        40
         	})

      
        37
        41
         

      
        38
        42
         	t.Run("invalid config file returns error", func(t *testing.T) {

      
M internal/plumbing/http/http.go
···
        20
        20
         

      
        21
        21
         func (Tool) Name() string { return "http(get, post, put, patch, delete)" }

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

      
        23
        
        -	return (c.Cmd == "get" || c.Cmd == "post" || c.Cmd == "put" || c.Cmd == "patch" || c.Cmd == "delete")

      
        
        23
        +	switch c.Cmd {

      
        
        24
        +	case "get", "post", "put", "patch", "delete":

      
        
        25
        +		return len(c.Args) > 0 && isURL(c.Args[0])

      
        
        26
        +	}

      
        
        27
        +	return false

      
        24
        28
         }

      
        25
        29
         

      
        26
        30
         func (Tool) Execute(c *viye.Context) (string, error) {

      
M internal/plumbing/http/http_test.go
···
        12
        12
         )

      
        13
        13
         

      
        14
        14
         func TestMatch(t *testing.T) {

      
        15
        
        -	tests := map[string]bool{

      
        16
        
        -		"get":    true,

      
        17
        
        -		"post":   true,

      
        18
        
        -		"put":    true,

      
        19
        
        -		"patch":  true,

      
        20
        
        -		"delete": true,

      
        21
        
        -		"empty":  false,

      
        22
        
        -		"":       false,

      
        23
        
        -	}

      
        24
        
        -	for tinp, twant := range tests {

      
        25
        
        -		got := (Tool{}).Match(&viye.Context{Cmd: tinp})

      
        26
        
        -		is.Equal(t, twant, got)

      
        
        15
        +	for _, tt := range []struct {

      
        
        16
        +		cmd  string

      
        
        17
        +		args []string

      
        
        18
        +		want bool

      
        
        19
        +	}{

      
        
        20
        +		{"", nil, false},

      
        
        21
        +		{"get", []string{"http://example.com"}, true},

      
        
        22
        +		{"post", []string{"https://x.io"}, true},

      
        
        23
        +		{"put", []string{"http://x.io/a"}, true},

      
        
        24
        +		{"patch", []string{"https://x.io/a"}, true},

      
        
        25
        +		{"delete", []string{"http://example.com"}, true},

      
        
        26
        +		{"GET", []string{"http://example.com"}, false}, // case-sensitive

      
        
        27
        +		{"get", nil, false},

      
        
        28
        +		{"post", []string{"not-a-url"}, false},

      
        
        29
        +		{"delete", nil, false},

      
        
        30
        +		{"empty", nil, false},

      
        
        31
        +	} {

      
        
        32
        +		got := (Tool{}).Match(&viye.Context{Cmd: tt.cmd, Args: tt.args})

      
        
        33
        +		is.Equal(t, tt.want, got)

      
        27
        34
         	}

      
        28
        35
         }

      
        29
        36
         

      
A internal/plumbing/sql/sql.go
···
        
        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
        +}

      
A internal/plumbing/sql/sql_test.go
···
        
        1
        +package sql

      
        
        2
        +

      
        
        3
        +import (

      
        
        4
        +	"context"

      
        
        5
        +	"database/sql"

      
        
        6
        +	"strings"

      
        
        7
        +	"testing"

      
        
        8
        +

      
        
        9
        +	"github.com/olexsmir/viye/internal/viye"

      
        
        10
        +	"olexsmir.xyz/x/is"

      
        
        11
        +)

      
        
        12
        +

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

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

      
        
        15
        +		cmd  string

      
        
        16
        +		args []string

      
        
        17
        +		want bool

      
        
        18
        +	}{

      
        
        19
        +		"select":                    {"select", []string{"*", "from", "users"}, true},

      
        
        20
        +		"select with no expression": {"select", nil, false},

      
        
        21
        +		"insert":                    {"insert", []string{"into", "users"}, true},

      
        
        22
        +		"insert without into":       {"insert", []string{"users"}, false},

      
        
        23
        +		"update":                    {"update", []string{"users", "set", "x=1"}, true},

      
        
        24
        +		"update without set":        {"update", []string{"users", "x=1"}, false},

      
        
        25
        +		"delete":                    {"delete", []string{"from", "users"}, true},

      
        
        26
        +		"delete missing table":      {"delete", []string{"from"}, false},

      
        
        27
        +		"create":                    {"create", []string{"table", "t"}, true},

      
        
        28
        +		"create incomplete":         {"create", []string{"table"}, false},

      
        
        29
        +		"drop":                      {"drop", []string{"table", "t"}, true},

      
        
        30
        +		"alter":                     {"alter", []string{"table", "t"}, true},

      
        
        31
        +		"pragma":                    {"pragma", []string{"table_info(t)"}, true},

      
        
        32
        +		"explain":                   {"explain", []string{"select", "1"}, true},

      
        
        33
        +		"analyze alone":             {"analyze", nil, true},

      
        
        34
        +		"uppercase":                 {"SELECT", []string{"*", "from", "t"}, true},

      
        
        35
        +		"not sql":                   {"curl", []string{"https://x"}, false},

      
        
        36
        +	} {

      
        
        37
        +		t.Run(tname, func(t *testing.T) {

      
        
        38
        +			if got := (Tool{}).Match(&viye.Context{Cmd: tt.cmd, Args: tt.args}); got != tt.want {

      
        
        39
        +				t.Errorf("Match(%q, %v) = %v, want %v", tt.cmd, tt.args, got, tt.want)

      
        
        40
        +			}

      
        
        41
        +		})

      
        
        42
        +	}

      
        
        43
        +}

      
        
        44
        +

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

      
        
        46
        +	db := newTestDB(t)

      
        
        47
        +	ctx := context.Background()

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

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

      
        
        50
        +

      
        
        51
        +	want := `: id  name   age

      
        
        52
        +: 1   alice  30

      
        
        53
        +: 2   bob    NULL

      
        
        54
        +`

      
        
        55
        +	for _, query := range []string{

      
        
        56
        +		"select * from users order by id",

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

      
        
        58
        +	} {

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

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

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

      
        
        62
        +	}

      
        
        63
        +}

      
        
        64
        +

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

      
        
        66
        +	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")

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

      
        
        69
        +

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

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

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

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

      
        
        74
        +	}

      
        
        75
        +}

      
        
        76
        +

      
        
        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
        +func TestRunQuery_write(t *testing.T) {

      
        
        90
        +	db := newTestDB(t)

      
        
        91
        +	ctx := context.Background()

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

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

      
        
        94
        +

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

      
        
        96
        +		{"update users set age = 40 where id = 1", "| 1 row(s) affected"},

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

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

      
        
        99
        +	} {

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

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

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

      
        
        103
        +	}

      
        
        104
        +}

      
        
        105
        +

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

      
        
        107
        +	db := newTestDB(t)

      
        
        108
        +	ctx := context.Background()

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

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

      
        
        111
        +

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

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

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

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

      
        
        116
        +

      
        
        117
        +	var count, aliceAge int

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

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

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

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

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

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

      
        
        124
        +}

      
        
        125
        +

      
        
        126
        +func TestExtractTableName(t *testing.T) {

      
        
        127
        +	for _, tt := range []struct {

      
        
        128
        +		name, query, want string

      
        
        129
        +		ok                bool

      
        
        130
        +	}{

      
        
        131
        +		{"update", "update users set name = 'x' where id = 1", "users", true},

      
        
        132
        +		{"update with subquery", "update users set n = (select 1 from roles) where id = 1", "users", true},

      
        
        133
        +		{"select", "select * from users;", "users", true},

      
        
        134
        +		{"insert", "insert into logs (a) values (1)", "logs", true},

      
        
        135
        +	} {

      
        
        136
        +		t.Run(tt.name, func(t *testing.T) {

      
        
        137
        +			got, ok := extractTableName(tt.query)

      
        
        138
        +			is.Equal(t, tt.want, got)

      
        
        139
        +			is.Equal(t, tt.ok, ok)

      
        
        140
        +		})

      
        
        141
        +	}

      
        
        142
        +}

      
        
        143
        +

      
        
        144
        +func TestDriverForDSN(t *testing.T) {

      
        
        145
        +	for _, tt := range []struct {

      
        
        146
        +		dsn, want string

      
        
        147
        +		err       any

      
        
        148
        +	}{

      
        
        149
        +		{"postgres://user:pass@host/db", "postgres", nil},

      
        
        150
        +		{"host=db user=alice", "postgres", nil},

      
        
        151
        +		{"postgres://host/name.db", "postgres", nil},

      
        
        152
        +		{"/tmp/data.db", "sqlite", nil},

      
        
        153
        +		{"/tmp/app.sqlite3", "sqlite", nil},

      
        
        154
        +		{"mongodb://host/db", "", "cannot determine driver"},

      
        
        155
        +	} {

      
        
        156
        +		got, err := driverForDSN(tt.dsn)

      
        
        157
        +		is.Err(t, err, tt.err)

      
        
        158
        +		is.Equal(t, tt.want, got)

      
        
        159
        +	}

      
        
        160
        +}

      
        
        161
        +

      
        
        162
        +func TestPlaceholder(t *testing.T) {

      
        
        163
        +	is.Equal(t, "$3", placeholder("postgres", 2))

      
        
        164
        +	for _, i := range []int{0, 3} {

      
        
        165
        +		is.Equal(t, "?", placeholder("sqlite", i))

      
        
        166
        +	}

      
        
        167
        +}

      
        
        168
        +

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

      
        
        170
        +	t.Helper()

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

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

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

      
        
        174
        +	db.SetMaxOpenConns(1)

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

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

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

      
        
        178
        +	return db

      
        
        179
        +}

      
M internal/plumbing/weather/weather.go
···
        33
        33
         }

      
        34
        34
         

      
        35
        35
         func getCity(c *viye.Context) (string, error) {

      
        36
        
        -	var city string

      
        37
        36
         	if len(c.Args) == 1 {

      
        38
        
        -		city = c.Args[0]

      
        39
        
        -	} else {

      
        40
        
        -		cfg, err := config.Load()

      
        41
        
        -		if err != nil {

      
        42
        
        -			return "", err

      
        43
        
        -		}

      
        44
        
        -		city = cfg.Get("city")

      
        
        37
        +		return c.Args[0], nil

      
        45
        38
         	}

      
        46
        
        -	if city == "" {

      
        
        39
        +	cfg, err := config.Load()

      
        
        40
        +	if err != nil {

      
        
        41
        +		return "", err

      
        
        42
        +	}

      
        
        43
        +	city, ok := cfg.Get("city")

      
        
        44
        +	if !ok {

      
        47
        45
         		return "", errors.New("please provide a city option")

      
        48
        46
         	}

      
        49
        47
         	return city, nil

      
M main.go
···
        14
        14
         	"github.com/olexsmir/viye/internal/plumbing/makefile"

      
        15
        15
         	"github.com/olexsmir/viye/internal/plumbing/neofetch"

      
        16
        16
         	"github.com/olexsmir/viye/internal/plumbing/shell"

      
        
        17
        +	"github.com/olexsmir/viye/internal/plumbing/sql"

      
        17
        18
         	"github.com/olexsmir/viye/internal/plumbing/tldr"

      
        18
        19
         	"github.com/olexsmir/viye/internal/plumbing/url"

      
        19
        20
         	"github.com/olexsmir/viye/internal/plumbing/weather"

      ···
        26
        27
         	v.Register(&shell.Tool{})

      
        27
        28
         	v.Register(&http.Tool{})

      
        28
        29
         	v.Register(&url.Tool{})

      
        
        30
        +	v.Register(&sql.Tool{})

      
        29
        31
         	v.Register(&ip.Tool{})

      
        30
        32
         	v.Register(&gobin.Tool{})

      
        31
        33
         	v.Register(&json.Tool{})