diff --git a/go.mod b/go.mod index 41b06af..54427af 100644 --- a/go.mod +++ b/go.mod @@ -2,7 +2,10 @@ module code.aneur.in/go/sqimple go 1.25.6 -require github.com/alecthomas/assert/v2 v2.11.0 +require ( + code.aneur.in/go/record v0.0.0-20260704022026-db3f74ccdeb4 + github.com/alecthomas/assert/v2 v2.11.0 +) require ( github.com/alecthomas/repr v0.4.0 // indirect diff --git a/go.sum b/go.sum index f571a34..87d97f4 100644 --- a/go.sum +++ b/go.sum @@ -1,6 +1,16 @@ +code.aneur.in/go/record v0.0.0-20260704022026-db3f74ccdeb4 h1:cE8tCjhHePkEMBEQbkuTy/NQ5DlAoDwRPQERqcZK5DM= +code.aneur.in/go/record v0.0.0-20260704022026-db3f74ccdeb4/go.mod h1:nIqf1d1YWOPSVrbBFJY2ZBEMMFSfQCgXL7nFVXOzVqo= github.com/alecthomas/assert/v2 v2.11.0 h1:2Q9r3ki8+JYXvGsDyBXwH3LcJ+WK5D0gc5E8vS6K3D0= github.com/alecthomas/assert/v2 v2.11.0/go.mod h1:Bze95FyfUr7x34QZrjL+XP+0qgp/zg8yS+TtBj1WA3k= github.com/alecthomas/repr v0.4.0 h1:GhI2A8MACjfegCPVq9f1FLvIBS+DrQ2KQBFZP1iFzXc= github.com/alecthomas/repr v0.4.0/go.mod h1:Fr0507jx4eOXV7AlPV6AVZLYrLIuIeSOWtW57eE/O/4= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/hexops/gotextdiff v1.0.3 h1:gitA9+qJrrTCsiCl7+kh75nPqQt1cx4ZkudSTLoUqJM= github.com/hexops/gotextdiff v1.0.3/go.mod h1:pSWU5MAI3yDq+fZBTazCSJysOMbxWL1BSow5/V2vxeg= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/record.go b/record.go new file mode 100644 index 0000000..7e2a4de --- /dev/null +++ b/record.go @@ -0,0 +1,167 @@ +package sqimple + +import ( + "database/sql" + "errors" + "fmt" + + "code.aneur.in/go/record" +) + +func ScanOneRecord(rows *sql.Rows) (record.Record, error) { + records, err := ScanRecords(rows) + if err != nil { + return nil, err + } + + if len(records) == 0 { + return nil, errors.New("no rows") + } + + return records[0], nil +} + +// ScanRecords transforms SQL row data into record.Record objects, making them easier to handle with type safety. +// +// https://go.dev/ref/spec#Types +func ScanRecords(rows *sql.Rows) ([]record.Record, error) { + columns, err := rows.ColumnTypes() + if err != nil { + return nil, err + } + + data := []record.Record{} + + for rows.Next() { + values := []any{} + for _, column := range columns { + switch column.ScanType().Name() { + case "bool": + var value bool + values = append(values, &value) + case "byte": + var value byte + values = append(values, &value) + case "complex64": + var value complex64 + values = append(values, &value) + case "complex128": + var value complex128 + values = append(values, &value) + case "float32": + var value float32 + values = append(values, &value) + case "float64": + var value float64 + values = append(values, &value) + case "int": + var value int + values = append(values, &value) + case "int8": + var value int8 + values = append(values, &value) + case "int16": + var value int16 + values = append(values, &value) + case "int32": + var value int32 + values = append(values, &value) + case "int64": + var value int64 + values = append(values, &value) + case "rune": + var value rune + values = append(values, &value) + case "string": + var value string + values = append(values, &value) + case "uint": + var value uint + values = append(values, &value) + case "uint8": + var value uint8 + values = append(values, &value) + case "uint16": + var value uint16 + values = append(values, &value) + case "uint32": + var value uint32 + values = append(values, &value) + case "uint64": + var value uint64 + values = append(values, &value) + default: + return nil, fmt.Errorf("unknown scan type %q", column.ScanType().Name()) + } + } + + err := rows.Scan(values...) + if err != nil { + return nil, err + } + + row := record.Record{} + for i, column := range columns { + switch column.ScanType().Name() { + case "bool": + value := values[i].(*bool) + row[column.Name()] = *value + case "byte": + value := values[i].(*byte) + row[column.Name()] = *value + case "complex64": + value := values[i].(*complex64) + row[column.Name()] = *value + case "complex128": + value := values[i].(*complex128) + row[column.Name()] = *value + case "float32": + value := values[i].(*float32) + row[column.Name()] = *value + case "float64": + value := values[i].(*float64) + row[column.Name()] = *value + case "int": + value := values[i].(*int) + row[column.Name()] = *value + case "int8": + value := values[i].(*int8) + row[column.Name()] = *value + case "int16": + value := values[i].(*int16) + row[column.Name()] = *value + case "int32": + value := values[i].(*int32) + row[column.Name()] = *value + case "int64": + value := values[i].(*int64) + row[column.Name()] = *value + case "rune": + value := values[i].(*rune) + row[column.Name()] = *value + case "string": + value := values[i].(*string) + row[column.Name()] = *value + case "uint": + value := values[i].(*uint) + row[column.Name()] = *value + case "uint8": + value := values[i].(*uint8) + row[column.Name()] = *value + case "uint16": + value := values[i].(*uint16) + row[column.Name()] = *value + case "uint32": + value := values[i].(*uint32) + row[column.Name()] = *value + case "uint64": + value := values[i].(*uint64) + row[column.Name()] = *value + } + } + + data = append(data, row) + } + + return data, nil +} diff --git a/select.go b/select.go index e1c1b0e..ff3b3a6 100644 --- a/select.go +++ b/select.go @@ -5,6 +5,8 @@ import ( "database/sql" "fmt" "strings" + + "code.aneur.in/go/record" ) // https://www.sqlite.org/lang_select.html @@ -123,6 +125,24 @@ func (q *SelectQuery) Query() (*sql.Rows, error) { return nil, ErrNoDatabase } +func (q *SelectQuery) QueryRecord() (record.Record, error) { + rows, err := q.Query() + if err != nil { + return nil, err + } + + return ScanOneRecord(rows) +} + +func (q *SelectQuery) QueryRecords() ([]record.Record, error) { + rows, err := q.Query() + if err != nil { + return nil, err + } + + return ScanRecords(rows) +} + func (q *SelectQuery) QueryRow() *sql.Row { if q.DB != nil { if q.Ctx != nil {