summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorRose Hogenson <rosehogenson@posteo.net>2024-02-06 20:27:19 -0800
committerRose Hogenson <rosehogenson@posteo.net>2024-02-06 20:27:19 -0800
commit432e3bda53ca02b9f3eb58163f8459d051f1df20 (patch)
tree0e031c8c305f9897364094d26c308c6d5f4994d5
downloadsqlr-432e3bda53ca02b9f3eb58163f8459d051f1df20.tar.zst
Initial commit.
-rw-r--r--go.mod8
-rw-r--r--go.sum4
-rw-r--r--sqlr.go51
-rw-r--r--sqlr_test.go158
4 files changed, 221 insertions, 0 deletions
diff --git a/go.mod b/go.mod
new file mode 100644
index 0000000..a696c94
--- /dev/null
+++ b/go.mod
@@ -0,0 +1,8 @@
+module gitlab.com/rhogenson/sqlr
+
+go 1.21.6
+
+require (
+ github.com/google/go-cmp v0.6.0
+ github.com/mattn/go-sqlite3 v1.14.22
+)
diff --git a/go.sum b/go.sum
new file mode 100644
index 0000000..4917b35
--- /dev/null
+++ b/go.sum
@@ -0,0 +1,4 @@
+github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
+github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
+github.com/mattn/go-sqlite3 v1.14.22 h1:2gZY6PC6kBnID23Tichd1K+Z0oS6nE/XwU+Vz/5o4kU=
+github.com/mattn/go-sqlite3 v1.14.22/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y=
diff --git a/sqlr.go b/sqlr.go
new file mode 100644
index 0000000..0224c7c
--- /dev/null
+++ b/sqlr.go
@@ -0,0 +1,51 @@
+// Package sqlr ("squealer") provides some convenience wrappers around
+// database/sql in the spirit of encoding/json.
+package sqlr
+
+import (
+ "database/sql"
+ "fmt"
+ "reflect"
+)
+
+// Scan calls rows.Scan to unpack a single row into the fields of v.
+func Scan(rows *sql.Rows, v any) error {
+ rv := reflect.ValueOf(v)
+ if rv.Kind() != reflect.Pointer || rv.Elem().Kind() != reflect.Struct {
+ return fmt.Errorf("Scan needs a pointer to a struct")
+ }
+ inner := rv.Elem()
+ innerType := inner.Type()
+
+ fields := make(map[string]reflect.StructField)
+ inner.FieldByNameFunc(func(fieldName string) bool {
+ field, ok := innerType.FieldByName(fieldName)
+ if !ok {
+ return false
+ }
+ tag, ok := field.Tag.Lookup("sql")
+ if !ok {
+ fields[fieldName] = field
+ return false
+ }
+ if tag != "-" {
+ fields[tag] = field
+ }
+ return false
+ })
+
+ cols, err := rows.Columns()
+ if err != nil {
+ return err
+ }
+ dest := make([]any, len(cols))
+ for i, c := range cols {
+ field, ok := fields[c]
+ if !ok {
+ return fmt.Errorf("no field with tag %q", c)
+ }
+ dest[i] = inner.FieldByIndex(field.Index).Addr().Interface()
+ }
+
+ return rows.Scan(dest...)
+}
diff --git a/sqlr_test.go b/sqlr_test.go
new file mode 100644
index 0000000..4323f11
--- /dev/null
+++ b/sqlr_test.go
@@ -0,0 +1,158 @@
+package sqlr
+
+import (
+ "context"
+ "database/sql"
+ "github.com/google/go-cmp/cmp"
+ _ "github.com/mattn/go-sqlite3"
+ "reflect"
+ "testing"
+)
+
+func TestScan(t *testing.T) {
+ ctx := context.Background()
+
+ for _, tc := range []struct {
+ desc string
+ db string
+ query string
+ want any
+ }{{
+ desc: "match field names",
+ db: `
+ CREATE TABLE Tbl (ColA, ColB);
+ INSERT INTO Tbl VALUES (100, "test")`,
+ query: "SELECT ColA, ColB FROM Tbl",
+ want: &struct {
+ ColA int
+ ColB string
+ }{100, "test"},
+ }, {
+ desc: "select star",
+ db: `
+ CREATE TABLE Tbl (ColA, ColB);
+ INSERT INTO Tbl VALUES (100, "test")`,
+ query: "SELECT * FROM Tbl",
+ want: &struct {
+ ColA int
+ ColB string
+ }{100, "test"},
+ }, {
+ desc: "ignored field",
+ db: `
+ CREATE TABLE Tbl (Col);
+ INSERT INTO Tbl VALUES (100)`,
+ query: "SELECT Col FROM Tbl",
+ want: &struct{ Col, OtherField int }{Col: 100},
+ }, {
+ desc: "tag",
+ db: `
+ CREATE TABLE Tbl (col_a, col_b);
+ INSERT INTO Tbl VALUES (100, "test")`,
+ query: "SELECT col_a, col_b FROM Tbl",
+ want: &struct {
+ ColA int `sql:"col_a"`
+ ColB string `sql:"col_b"`
+ }{100, "test"},
+ }} {
+ t.Run(tc.desc, func(t *testing.T) {
+ db, err := sql.Open("sqlite3", ":memory:")
+ if err != nil {
+ t.Fatalf("Failed to create in-memory database: %s", err)
+ }
+ defer db.Close()
+ if _, err := db.ExecContext(ctx, tc.db); err != nil {
+ t.Fatalf("Failed to initialize test database: %s", err)
+ }
+
+ rows, err := db.QueryContext(ctx, tc.query)
+ if err != nil {
+ t.Fatalf("Failed to query test database: %s", err)
+ }
+ if !rows.Next() {
+ t.Fatalf("No data")
+ }
+
+ got := reflect.New(reflect.TypeOf(tc.want).Elem()).Interface()
+ if err := Scan(rows, got); err != nil {
+ t.Fatalf("Scan failed: %s", err)
+ }
+ if err := rows.Close(); err != nil {
+ t.Fatalf("Close rows: %s", err)
+ }
+ if err := rows.Err(); err != nil {
+ t.Fatalf("Database read error: %s", err)
+ }
+
+ if diff := cmp.Diff(tc.want, got); diff != "" {
+ t.Errorf("Scan returned unexpected diff (-want +got):\n%s", diff)
+ }
+ })
+ }
+}
+
+func TestScan_Errors(t *testing.T) {
+ ctx := context.Background()
+
+ for _, tc := range []struct {
+ desc string
+ db string
+ query string
+ out any
+ }{{
+ desc: "wrong type",
+ db: `
+ CREATE TABLE Tbl (Col);
+ INSERT INTO Tbl VALUES (100)`,
+ query: "SELECT Col FROM Tbl",
+ out: 5,
+ }, {
+ desc: "nil pointer",
+ db: `
+ CREATE TABLE Tbl (Col);
+ INSERT INTO Tbl VALUES (100)`,
+ query: "SELECT Col FROM Tbl",
+ out: nil,
+ }, {
+ desc: "no field",
+ db: `
+ CREATE TABLE Tbl (ColA, ColB);
+ INSERT INTO Tbl VALUES (100, "test")`,
+ query: "SELECT ColA, ColB FROM Tbl",
+ out: new(struct{ ColA int }),
+ }, {
+ desc: "ignored field",
+ db: `
+ CREATE TABLE Tbl (Col);
+ INSERT INTO Tbl VALUES (100)`,
+ query: "SELECT Col FROM Tbl",
+ out: new(struct {
+ Col int `sql:"-"`
+ }),
+ }} {
+ t.Run(tc.desc, func(t *testing.T) {
+ db, err := sql.Open("sqlite3", ":memory:")
+ if err != nil {
+ t.Fatalf("Failed to create in-memory database: %s", err)
+ }
+ defer db.Close()
+ if _, err := db.ExecContext(ctx, tc.db); err != nil {
+ t.Fatalf("Failed to initialize test database: %s", err)
+ }
+
+ rows, err := db.QueryContext(ctx, tc.query)
+ if err != nil {
+ t.Fatalf("Failed to query test database: %s", err)
+ }
+ defer rows.Close()
+ if !rows.Next() {
+ t.Fatalf("No data")
+ }
+
+ err = Scan(rows, tc.out)
+ if err == nil {
+ t.Errorf("Scan returned nil, want error")
+ }
+ })
+ }
+}