summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--sqlr.go37
-rw-r--r--sqlr_test.go18
2 files changed, 33 insertions, 22 deletions
diff --git a/sqlr.go b/sqlr.go
index 0224c7c..4ac798f 100644
--- a/sqlr.go
+++ b/sqlr.go
@@ -15,36 +15,29 @@ func Scan(rows *sql.Rows, v any) error {
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
- })
+ fields := reflect.VisibleFields(inner.Type())
cols, err := rows.Columns()
if err != nil {
return err
}
dest := make([]any, len(cols))
+Cols:
for i, c := range cols {
- field, ok := fields[c]
- if !ok {
- return fmt.Errorf("no field with tag %q", c)
+ for _, f := range fields {
+ if !f.IsExported() {
+ continue
+ }
+ tag := f.Tag.Get("sql")
+ if tag == "" {
+ tag = f.Name
+ }
+ if tag == c {
+ dest[i] = inner.FieldByIndex(f.Index).Addr().Interface()
+ continue Cols
+ }
}
- dest[i] = inner.FieldByIndex(field.Index).Addr().Interface()
+ return fmt.Errorf("no field with tag %q", c)
}
return rows.Scan(dest...)
diff --git a/sqlr_test.go b/sqlr_test.go
index 4323f11..4625fde 100644
--- a/sqlr_test.go
+++ b/sqlr_test.go
@@ -54,6 +54,17 @@ func TestScan(t *testing.T) {
ColA int `sql:"col_a"`
ColB string `sql:"col_b"`
}{100, "test"},
+ }, {
+ desc: "embedded struct",
+ db: `
+ CREATE TABLE Tbl (Col);
+ INSERT INTO Tbl VALUES (100)`,
+ query: "SELECT Col FROM Tbl",
+ want: func() any {
+ type Inner struct{ Col int }
+ type outer struct{ Inner }
+ return &outer{Inner{100}}
+ }(),
}} {
t.Run(tc.desc, func(t *testing.T) {
db, err := sql.Open("sqlite3", ":memory:")
@@ -129,6 +140,13 @@ func TestScan_Errors(t *testing.T) {
out: new(struct {
Col int `sql:"-"`
}),
+ }, {
+ desc: "unexported field",
+ db: `
+ CREATE TABLE Tbl (col);
+ INSERT INTO Tbl VALUES (100)`,
+ query: "SELECT col FROM Tbl",
+ out: new(struct{ col int }),
}} {
t.Run(tc.desc, func(t *testing.T) {
db, err := sql.Open("sqlite3", ":memory:")