diff options
| author | Rose Hogenson <rosehogenson@posteo.net> | 2024-02-06 20:27:19 -0800 |
|---|---|---|
| committer | Rose Hogenson <rosehogenson@posteo.net> | 2024-02-06 20:27:19 -0800 |
| commit | 432e3bda53ca02b9f3eb58163f8459d051f1df20 (patch) | |
| tree | 0e031c8c305f9897364094d26c308c6d5f4994d5 | |
| download | sqlr-432e3bda53ca02b9f3eb58163f8459d051f1df20.tar.zst | |
Initial commit.
| -rw-r--r-- | go.mod | 8 | ||||
| -rw-r--r-- | go.sum | 4 | ||||
| -rw-r--r-- | sqlr.go | 51 | ||||
| -rw-r--r-- | sqlr_test.go | 158 |
4 files changed, 221 insertions, 0 deletions
@@ -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 +) @@ -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= @@ -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") + } + }) + } +} |
