master
go 101 lines 2.21 KB
Raw
1 // SPDX-License-Identifier: GPL-3.0-or-later
2
3 package sqlquery
4
5 import (
6 "database/sql"
7 "fmt"
8 )
9
10 type ScanValueType string
11
12 const (
13 ScanValueString ScanValueType = "string"
14 ScanValueInteger ScanValueType = "integer"
15 ScanValueFloat ScanValueType = "float"
16 ScanValueDiscard ScanValueType = "discard"
17 )
18
19 // ScanColumnSpec describes how a query result column should be scanned and normalized.
20 type ScanColumnSpec struct {
21 Type ScanValueType
22 Transform func(any) any
23 }
24
25 // RowScanner is the minimal interface required for typed row scanning.
26 type RowScanner interface {
27 Next() bool
28 Scan(dest ...any) error
29 }
30
31 // ScanTypedRows scans all rows according to specs and returns normalized [][]any rows.
32 // The function intentionally does not call rows.Err(); callers keep control over that.
33 func ScanTypedRows(rows RowScanner, specs []ScanColumnSpec) ([][]any, error) {
34 data := make([][]any, 0, 500)
35 holders := makeScanHolders(specs)
36
37 for rows.Next() {
38 if err := rows.Scan(holders.ptrs...); err != nil {
39 return nil, fmt.Errorf("scan row: %w", err)
40 }
41
42 row := make([]any, len(specs))
43 for i := range holders.ptrs {
44 value, ok := holders.value(i)
45 if ok && specs[i].Transform != nil {
46 value = specs[i].Transform(value)
47 }
48 row[i] = value
49 }
50 data = append(data, row)
51 }
52
53 return data, nil
54 }
55
56 type scanHolders struct {
57 ptrs []any
58 }
59
60 func makeScanHolders(specs []ScanColumnSpec) scanHolders {
61 holders := scanHolders{ptrs: make([]any, len(specs))}
62 for i, spec := range specs {
63 switch spec.Type {
64 case ScanValueString:
65 holders.ptrs[i] = &sql.NullString{}
66 case ScanValueInteger:
67 holders.ptrs[i] = &sql.NullInt64{}
68 case ScanValueFloat:
69 holders.ptrs[i] = &sql.NullFloat64{}
70 case ScanValueDiscard:
71 holders.ptrs[i] = new(any)
72 default:
73 holders.ptrs[i] = &sql.NullString{}
74 }
75 }
76 return holders
77 }
78
79 func (h scanHolders) value(i int) (any, bool) {
80 switch v := h.ptrs[i].(type) {
81 case *sql.NullString:
82 if v.Valid {
83 return v.String, true
84 }
85 return "", false
86 case *sql.NullInt64:
87 if v.Valid {
88 return v.Int64, true
89 }
90 return int64(0), false
91 case *sql.NullFloat64:
92 if v.Valid {
93 return v.Float64, true
94 }
95 return float64(0), false
96 case *any:
97 return nil, false
98 default:
99 return nil, false
100 }
101 }