Skip to content

Commit 0eac140

Browse files
authored
Implement nested structs for columns and values (#60)
1 parent da49f61 commit 0eac140

5 files changed

Lines changed: 210 additions & 32 deletions

File tree

columns.go

Lines changed: 33 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -74,46 +74,60 @@ func columns(v interface{}, strict bool, excluded ...string) ([]string, error) {
7474
return res, nil
7575
}
7676

77+
names := columnNames(model, strict, excluded...)
78+
toCache := append(names, excluded...)
79+
columnsCache.Store(model, toCache)
80+
return names, nil
81+
}
82+
83+
func columnNames(model reflect.Value, strict bool, excluded ...string) []string {
7784
numfield := model.NumField()
7885
names := make([]string, 0, numfield)
7986

80-
isExcluded := func(name string) bool {
81-
for _, ex := range excluded {
82-
if ex == name {
83-
return true
84-
}
85-
}
86-
return false
87-
}
88-
8987
for i := 0; i < numfield; i++ {
9088
valField := model.Field(i)
9189
if !valField.IsValid() || !valField.CanSet() {
9290
continue
9391
}
9492

9593
typeField := model.Type().Field(i)
96-
if tag, ok := typeField.Tag.Lookup(dbTag); ok {
97-
if tag != "-" && !isExcluded(tag) {
98-
names = append(names, tag)
99-
}
94+
95+
if typeField.Type.Kind() == reflect.Struct {
96+
embeddedNames := columnNames(valField, strict, excluded...)
97+
names = append(names, embeddedNames...)
10098
continue
10199
}
102100

103-
if strict {
101+
fieldName := typeField.Name
102+
if tag, hasTag := typeField.Tag.Lookup(dbTag); hasTag {
103+
if tag == "-" {
104+
continue
105+
}
106+
fieldName = tag
107+
} else if strict {
108+
// there's no tag name and we're in strict mode so move on
104109
continue
105110
}
106111

107-
if isExcluded(typeField.Name) || !supportedColumnType(valField.Kind()) {
112+
if isExcluded(fieldName, excluded...) {
108113
continue
109114
}
110115

111-
names = append(names, typeField.Name)
116+
if supportedColumnType(valField.Kind()) {
117+
names = append(names, fieldName)
118+
}
112119
}
113120

114-
toCache := append(names, excluded...)
115-
columnsCache.Store(model, toCache)
116-
return names, nil
121+
return names
122+
}
123+
124+
func isExcluded(name string, excluded ...string) bool {
125+
for _, ex := range excluded {
126+
if ex == name {
127+
return true
128+
}
129+
}
130+
return false
117131
}
118132

119133
func reflectValue(v interface{}) (reflect.Value, error) {

columns_test.go

Lines changed: 63 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -63,16 +63,76 @@ func TestColumnsIgnoresPrivateFields(t *testing.T) {
6363
assert.EqualValues(t, []string{"Age"}, cols)
6464
}
6565

66-
func TestColumnsAddsComplexTypesWhenStructTag(t *testing.T) {
66+
func TestColumnsAddsComplexTypesWhenNoStructTag(t *testing.T) {
6767
type person struct {
6868
Address struct {
6969
Street string
70+
}
71+
}
72+
73+
cols, err := Columns(&person{})
74+
assert.NoError(t, err)
75+
assert.EqualValues(t, []string{"Street"}, cols)
76+
}
77+
78+
func TestColumnsAddsComplexTypesWhenStructTag(t *testing.T) {
79+
type person struct {
80+
Address struct {
81+
Street string `db:"address.street"`
82+
}
83+
}
84+
85+
cols, err := Columns(&person{})
86+
assert.NoError(t, err)
87+
assert.EqualValues(t, []string{"address.street"}, cols)
88+
}
89+
90+
func TestColumnsDoesNotAddStructTag(t *testing.T) {
91+
type person struct {
92+
Address struct {
93+
Street string `db:"address.street"`
7094
} `db:"address"`
7195
}
7296

7397
cols, err := Columns(&person{})
7498
assert.NoError(t, err)
75-
assert.EqualValues(t, []string{"address"}, cols)
99+
assert.EqualValues(t, []string{"address.street"}, cols)
100+
}
101+
102+
func TestColumnsStrictAddsComplexTypesWhenStructTag(t *testing.T) {
103+
type person struct {
104+
Address struct {
105+
Street string `db:"address.street"`
106+
}
107+
}
108+
109+
cols, err := ColumnsStrict(&person{})
110+
assert.NoError(t, err)
111+
assert.EqualValues(t, []string{"address.street"}, cols)
112+
}
113+
114+
func TestColumnsStrictDoesNotAddComplexTypesWhenNoStructTag(t *testing.T) {
115+
type person struct {
116+
Address struct {
117+
Street string
118+
}
119+
}
120+
121+
cols, err := ColumnsStrict(&person{})
122+
assert.NoError(t, err)
123+
assert.EqualValues(t, []string{}, cols)
124+
}
125+
126+
func TestColumnsStrictAddsComplexTypesRegardlessOfStructTag(t *testing.T) {
127+
type person struct {
128+
Address struct {
129+
Street string `db:"address.street"`
130+
} `db:"-"`
131+
}
132+
133+
cols, err := ColumnsStrict(&person{})
134+
assert.NoError(t, err)
135+
assert.EqualValues(t, []string{"address.street"}, cols)
76136
}
77137

78138
func TestColumnsIgnoresComplexTypesWhenNoStructTag(t *testing.T) {
@@ -84,7 +144,7 @@ func TestColumnsIgnoresComplexTypesWhenNoStructTag(t *testing.T) {
84144

85145
cols, err := Columns(&person{})
86146
assert.NoError(t, err)
87-
assert.EqualValues(t, []string{}, cols)
147+
assert.EqualValues(t, []string{"Street"}, cols)
88148
}
89149

90150
func TestColumnsExcludesFields(t *testing.T) {

examples_values_columns_test.go

Lines changed: 73 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,31 @@ func ExampleValues() {
2222
// [1 Brett]
2323
}
2424

25+
func ExampleValues_nested() {
26+
type Address struct {
27+
Street string
28+
City string
29+
}
30+
31+
person := struct {
32+
ID int
33+
Name string
34+
Address
35+
}{
36+
Name: "Brett",
37+
ID: 1,
38+
Address: Address{
39+
City: "San Francisco",
40+
},
41+
}
42+
43+
cols := []string{"Name", "City"}
44+
vals, _ := scan.Values(cols, &person)
45+
fmt.Printf("%+v", vals)
46+
// Output:
47+
// [Brett San Francisco]
48+
}
49+
2550
func ExampleColumns() {
2651
var person struct {
2752
ID int `db:"person_id"`
@@ -59,3 +84,51 @@ func ExampleColumnsStrict() {
5984
// Output:
6085
// [id age]
6186
}
87+
88+
func ExampleColumnsNested() {
89+
var person struct {
90+
ID int `db:"person.id"`
91+
Name string `db:"person.name"`
92+
Company struct {
93+
ID int `db:"company.id"`
94+
Name string
95+
}
96+
}
97+
98+
cols, _ := scan.Columns(&person)
99+
fmt.Printf("%+v", cols)
100+
// Output:
101+
// [person.id person.name company.id Name]
102+
}
103+
104+
func ExampleColumnsNestedStrict() {
105+
var person struct {
106+
ID int `db:"person.id"`
107+
Name string `db:"person.name"`
108+
Company struct {
109+
ID int `db:"company.id"`
110+
Name string
111+
}
112+
}
113+
114+
cols, _ := scan.ColumnsStrict(&person)
115+
fmt.Printf("%+v", cols)
116+
// Output:
117+
// [person.id person.name company.id]
118+
}
119+
120+
func ExampleColumnsNested_exclude() {
121+
var person struct {
122+
ID int `db:"person.id"`
123+
Name string `db:"person.name"`
124+
Company struct {
125+
ID int `db:"-"`
126+
Name string `db:"company.name"`
127+
}
128+
}
129+
130+
cols, _ := scan.Columns(&person)
131+
fmt.Printf("%+v", cols)
132+
// Output:
133+
// [person.id person.name company.name]
134+
}

values.go

Lines changed: 20 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -26,34 +26,45 @@ func Values(cols []string, v interface{}) ([]interface{}, error) {
2626
return nil, fmt.Errorf("field %T.%q either does not exist or is unexported: %w", v, col, ErrStructFieldMissing)
2727
}
2828

29-
vals[i] = model.Field(j).Interface()
29+
vals[i] = model.FieldByIndex(j).Interface()
3030
}
3131
return vals, nil
3232
}
3333

34-
func loadFields(val reflect.Value) map[string]int {
34+
func loadFields(val reflect.Value) map[string][]int {
3535
if cache, cached := valuesCache.Load(val); cached {
36-
return cache.(map[string]int)
36+
return cache.(map[string][]int)
3737
}
3838
return writeFieldsCache(val)
3939
}
4040

41-
func writeFieldsCache(val reflect.Value) map[string]int {
41+
func writeFieldsCache(val reflect.Value) map[string][]int {
42+
m := map[string][]int{}
43+
writeFields(val, m, []int{})
44+
valuesCache.Store(val, m)
45+
return m
46+
}
47+
48+
func writeFields(val reflect.Value, m map[string][]int, index []int) {
4249
typ := val.Type()
4350
numfield := val.NumField()
44-
m := map[string]int{}
4551

4652
for i := 0; i < numfield; i++ {
4753
if !val.Field(i).CanSet() {
4854
continue
4955
}
5056

5157
field := typ.Field(i)
52-
m[field.Name] = i
58+
fieldIndex := append(index, field.Index...)
59+
60+
if field.Type.Kind() == reflect.Struct {
61+
writeFields(val.Field(i), m, fieldIndex)
62+
continue
63+
}
64+
65+
m[field.Name] = fieldIndex
5366
if tag, ok := field.Tag.Lookup(dbTag); ok {
54-
m[tag] = i
67+
m[tag] = fieldIndex
5568
}
5669
}
57-
valuesCache.Store(val, m)
58-
return m
5970
}

values_test.go

Lines changed: 21 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,26 @@ func TestValuesScansDBTags(t *testing.T) {
3333
assert.EqualValues(t, []interface{}{"Brett"}, vals)
3434
}
3535

36+
func TestValuesScansNestedFields(t *testing.T) {
37+
type Address struct {
38+
Street string
39+
City string
40+
}
41+
42+
type Person struct {
43+
Name string
44+
Age int
45+
Address
46+
}
47+
48+
p := &Person{Name: "Brett", Address: Address{Street: "123 Main St", City: "San Francisco"}}
49+
50+
vals, err := Values([]string{"Name", "Street", "City"}, p)
51+
require.NoError(t, err)
52+
53+
assert.EqualValues(t, []interface{}{"Brett", "123 Main St", "San Francisco"}, vals)
54+
}
55+
3656
func TestValuesReturnsErrorWhenPassingNonPointer(t *testing.T) {
3757
_, err := Values([]string{"Name"}, "")
3858
require.Error(t, err)
@@ -80,7 +100,7 @@ func TestValuesReadsFromCacheFirst(t *testing.T) {
80100
}
81101

82102
v := reflect.Indirect(reflect.ValueOf(&person))
83-
valuesCache.Store(v, map[string]int{"Name": 0})
103+
valuesCache.Store(v, map[string][]int{"Name": {0}})
84104

85105
vals, err := Values([]string{"Name"}, &person)
86106
require.NoError(t, err)

0 commit comments

Comments
 (0)