master
go 116 lines 3.2 KB
Raw
1 // SPDX-License-Identifier: GPL-3.0-or-later
2
3 package sqlquery
4
5 import (
6 "strings"
7 "testing"
8
9 "github.com/DATA-DOG/go-sqlmock"
10 "github.com/stretchr/testify/assert"
11 "github.com/stretchr/testify/require"
12 )
13
14 func TestScanTypedRows(t *testing.T) {
15 t.Run("all supported types", func(t *testing.T) {
16 db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherEqual))
17 require.NoError(t, err)
18 defer func() { _ = db.Close() }()
19
20 mock.ExpectQuery("SELECT test").
21 WillReturnRows(sqlmock.NewRows([]string{"s", "i", "f"}).
22 AddRow("abc", int64(7), float64(1.5)))
23
24 rows, err := db.Query("SELECT test")
25 require.NoError(t, err)
26 defer func() { _ = rows.Close() }()
27
28 data, err := ScanTypedRows(rows, []ScanColumnSpec{
29 {Type: ScanValueString},
30 {Type: ScanValueInteger},
31 {Type: ScanValueFloat},
32 })
33 require.NoError(t, err)
34 require.Len(t, data, 1)
35 assert.Equal(t, "abc", data[0][0])
36 assert.EqualValues(t, 7, data[0][1])
37 assert.EqualValues(t, 1.5, data[0][2])
38 assert.NoError(t, mock.ExpectationsWereMet())
39 })
40
41 t.Run("null defaults", func(t *testing.T) {
42 db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherEqual))
43 require.NoError(t, err)
44 defer func() { _ = db.Close() }()
45
46 mock.ExpectQuery("SELECT test").
47 WillReturnRows(sqlmock.NewRows([]string{"s", "i", "f"}).
48 AddRow(nil, nil, nil))
49
50 rows, err := db.Query("SELECT test")
51 require.NoError(t, err)
52 defer func() { _ = rows.Close() }()
53
54 data, err := ScanTypedRows(rows, []ScanColumnSpec{
55 {Type: ScanValueString},
56 {Type: ScanValueInteger},
57 {Type: ScanValueFloat},
58 })
59 require.NoError(t, err)
60 require.Len(t, data, 1)
61 assert.Equal(t, "", data[0][0])
62 assert.EqualValues(t, 0, data[0][1])
63 assert.EqualValues(t, 0.0, data[0][2])
64 assert.NoError(t, mock.ExpectationsWereMet())
65 })
66
67 t.Run("transform on non-null only", func(t *testing.T) {
68 db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherEqual))
69 require.NoError(t, err)
70 defer func() { _ = db.Close() }()
71
72 mock.ExpectQuery("SELECT test").
73 WillReturnRows(sqlmock.NewRows([]string{"s"}).
74 AddRow(" query ").
75 AddRow(nil))
76
77 rows, err := db.Query("SELECT test")
78 require.NoError(t, err)
79 defer func() { _ = rows.Close() }()
80
81 data, err := ScanTypedRows(rows, []ScanColumnSpec{
82 {
83 Type: ScanValueString,
84 Transform: func(v any) any {
85 return strings.TrimSpace(v.(string))
86 },
87 },
88 })
89 require.NoError(t, err)
90 require.Len(t, data, 2)
91 assert.Equal(t, "query", data[0][0])
92 assert.Equal(t, "", data[1][0])
93 assert.NoError(t, mock.ExpectationsWereMet())
94 })
95
96 t.Run("scan error propagation", func(t *testing.T) {
97 db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherEqual))
98 require.NoError(t, err)
99 defer func() { _ = db.Close() }()
100
101 mock.ExpectQuery("SELECT test").
102 WillReturnRows(sqlmock.NewRows([]string{"only_one"}).AddRow("abc"))
103
104 rows, err := db.Query("SELECT test")
105 require.NoError(t, err)
106 defer func() { _ = rows.Close() }()
107
108 _, err = ScanTypedRows(rows, []ScanColumnSpec{
109 {Type: ScanValueString},
110 {Type: ScanValueString},
111 })
112 require.Error(t, err)
113 assert.Contains(t, err.Error(), "scan row")
114 assert.NoError(t, mock.ExpectationsWereMet())
115 })
116 }