diff options
| -rw-r--r-- | go.mod | 2 | ||||
| -rw-r--r-- | sqlr.go | 35 | ||||
| -rw-r--r-- | sqlr_test.go | 75 |
3 files changed, 106 insertions, 6 deletions
@@ -1,6 +1,6 @@ module gitlab.com/rhogenson/sqlr -go 1.21.6 +go 1.23 require ( github.com/google/go-cmp v0.6.0 @@ -3,8 +3,10 @@ package sqlr import ( + "context" "database/sql" "fmt" + "iter" "reflect" ) @@ -42,3 +44,36 @@ Cols: return rows.Scan(dest...) } + +// Querier is the underlying type used to query the database, +// usually *sql.DB or *sql.Tx. +type Querier interface { + QueryContext(context.Context, string, ...any) (*sql.Rows, error) +} + +// Query returns an iterator over the rows matched by a query. +func Query[Row any](ctx context.Context, querier Querier, query string, args ...any) iter.Seq2[Row, error] { + return func(yield func(Row, error) bool) { + rows, err := querier.QueryContext(ctx, query, args...) + if err != nil { + var zero Row + yield(zero, err) + return + } + defer rows.Close() + for rows.Next() { + var row Row + if err := Scan(rows, &row); err != nil { + yield(row, err) + return + } + if !yield(row, nil) { + return + } + } + if err := rows.Err(); err != nil { + var zero Row + yield(zero, err) + } + } +} diff --git a/sqlr_test.go b/sqlr_test.go index 4625fde..f2ebfbf 100644 --- a/sqlr_test.go +++ b/sqlr_test.go @@ -1,12 +1,15 @@ -package sqlr +package sqlr_test import ( "context" "database/sql" - "github.com/google/go-cmp/cmp" - _ "github.com/mattn/go-sqlite3" + "fmt" "reflect" "testing" + + "github.com/google/go-cmp/cmp" + _ "github.com/mattn/go-sqlite3" + "gitlab.com/rhogenson/sqlr" ) func TestScan(t *testing.T) { @@ -85,7 +88,7 @@ func TestScan(t *testing.T) { } got := reflect.New(reflect.TypeOf(tc.want).Elem()).Interface() - if err := Scan(rows, got); err != nil { + if err := sqlr.Scan(rows, got); err != nil { t.Fatalf("Scan failed: %s", err) } if err := rows.Close(); err != nil { @@ -167,10 +170,72 @@ func TestScan_Errors(t *testing.T) { t.Fatalf("No data") } - err = Scan(rows, tc.out) + err = sqlr.Scan(rows, tc.out) if err == nil { t.Errorf("Scan returned nil, want error") } }) } } + +func TestQuery(t *testing.T) { + ctx := context.Background() + 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, ` + CREATE TABLE Tbl (Name, HighScore); + INSERT INTO Tbl VALUES ("rose", 100) + `); err != nil { + t.Fatalf("Failed to initialize test database: %s", err) + } + + type row struct { + Name string + HighScore int + } + const query = "SELECT Name, HighScore FROM Tbl" + var got []row + for row, err := range sqlr.Query[row](ctx, db, query) { + if err != nil { + t.Fatalf("Query(%q) failed: %s", query, err) + } + got = append(got, row) + } + want := []row{{ + Name: "rose", + HighScore: 100, + }} + if diff := cmp.Diff(want, got); diff != "" { + t.Errorf("Query(%q) returned unexpected diff (-want +got):\n%s", query, diff) + } +} + +func ExampleQuery() { + ctx := context.Background() + db, err := sql.Open("sqlite3", ":memory:") + if err != nil { + panic(err) + } + defer db.Close() + if _, err := db.ExecContext(ctx, ` + CREATE TABLE Tbl (Name, HighScore); + INSERT INTO Tbl VALUES ("rose", 100) + `); err != nil { + panic(err) + } + + type row struct { + Name string + HighScore int + } + for row, err := range sqlr.Query[row](ctx, db, "SELECT Name, HighScore FROM Tbl") { + if err != nil { + panic(err) + } + fmt.Printf("%s %d\n", row.Name, row.HighScore) + } + // Output: rose 100 +} |
