- createColumnScans: *time.Time 等指针到结构体类型改用 *interface{} 扫描目标,
避免 database/sql 因无法将 []byte 转换为 *time.Time 而报 scan error
- ScanStruct: time.Time 值直接存入 record,跳过 toJSONRawMessage 格式化,
保留纳秒精度和时区 Location
- SafeSetVarValue: 新增 T→*T 赋值路径(适用于 time.Time→*time.Time 等场景),
以及从 *json.RawMessage 解析 *time.Time 的路径(处理 []byte datetime 字符串)
- 新增 parseTimeFromRaw 支持 MySQL datetime、date-only、ISO8601、RFC3339 等格式
- 新增测试:覆盖 time.Time/[]byte 驱动值、date-only 格式、NULL 指针等场景
392 lines
9.6 KiB
Go
392 lines
9.6 KiB
Go
package exec
|
||
|
||
import (
|
||
"database/sql"
|
||
"encoding/json"
|
||
"reflect"
|
||
"time"
|
||
|
||
"git.fsdpf.net/go/db/internal/errors"
|
||
"git.fsdpf.net/go/db/internal/util"
|
||
)
|
||
|
||
type (
|
||
// Scanner knows how to scan sql.Rows into structs.
|
||
Scanner interface {
|
||
Next() bool
|
||
ScanStruct(i interface{}) error
|
||
ScanStructs(i interface{}) error
|
||
ScanVal(i interface{}) error
|
||
ScanVals(i interface{}) error
|
||
GetRecord() (map[string]any, error)
|
||
GetRecords() ([]map[string]any, error)
|
||
Close() error
|
||
Err() error
|
||
}
|
||
|
||
scanner struct {
|
||
rows *sql.Rows
|
||
columnMap util.ColumnMap
|
||
columns []string
|
||
}
|
||
)
|
||
|
||
func unableToFindFieldError(col string) error {
|
||
return errors.New(`unable to find corresponding field to column "%s" returned by query`, col)
|
||
}
|
||
|
||
// NewScanner returns a scanner that can be used for scanning rows into structs.
|
||
func NewScanner(rows *sql.Rows) Scanner {
|
||
return &scanner{rows: rows}
|
||
}
|
||
|
||
// Next prepares the next row for Scanning. See sql.Rows#Next for more
|
||
// information.
|
||
func (s *scanner) Next() bool {
|
||
return s.rows.Next()
|
||
}
|
||
|
||
// Err returns the error, if any that was encountered during iteration. See
|
||
// sql.Rows#Err for more information.
|
||
func (s *scanner) Err() error {
|
||
return s.rows.Err()
|
||
}
|
||
|
||
func (s *scanner) GetRecords() ([]map[string]any, error) {
|
||
records := []map[string]any{}
|
||
for s.Next() {
|
||
row, err := s.GetRecord()
|
||
if err != nil {
|
||
return records, err
|
||
}
|
||
records = append(records, row)
|
||
}
|
||
|
||
return records, s.Err()
|
||
}
|
||
|
||
func (s *scanner) GetRecord() (record map[string]any, err error) {
|
||
if s.columns == nil {
|
||
colsType, err := s.rows.ColumnTypes()
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
s.columns = make([]string, len(colsType))
|
||
s.columnMap = make(util.ColumnMap, len(colsType))
|
||
for i, ct := range colsType {
|
||
s.columns[i] = ct.Name()
|
||
s.columnMap[ct.Name()] = util.ColumnData{
|
||
ColumnName: ct.Name(),
|
||
GoType: ct.ScanType(),
|
||
}
|
||
}
|
||
}
|
||
|
||
scans := make([]any, len(s.columns))
|
||
for i, col := range s.columns {
|
||
if isNumericKind(s.columnMap[col].GoType.Kind()) {
|
||
// **T:NULL 时内层指针为 nil,否则持有精确数值类型
|
||
scans[i] = reflect.New(reflect.PointerTo(s.columnMap[col].GoType)).Interface()
|
||
} else {
|
||
scans[i] = new(any)
|
||
}
|
||
}
|
||
|
||
if err := s.rows.Scan(scans...); err != nil {
|
||
return nil, err
|
||
}
|
||
|
||
record = make(map[string]any, len(s.columns))
|
||
for i, col := range s.columns {
|
||
if isNumericKind(s.columnMap[col].GoType.Kind()) {
|
||
inner := reflect.ValueOf(scans[i]).Elem() // *T
|
||
if inner.IsNil() {
|
||
record[col] = nil
|
||
} else {
|
||
record[col] = inner.Elem().Interface()
|
||
}
|
||
continue
|
||
}
|
||
|
||
v := *(scans[i].(*any))
|
||
switch val := v.(type) {
|
||
case time.Time:
|
||
layout := time.DateTime
|
||
if val.Hour() == 0 && val.Minute() == 0 && val.Second() == 0 && val.Nanosecond() == 0 {
|
||
layout = time.DateOnly
|
||
}
|
||
record[col] = val.Format(layout)
|
||
case []byte:
|
||
if len(val) > 0 && (val[0] == '[' || val[0] == '{') {
|
||
var parsed any
|
||
if json.Unmarshal(val, &parsed) == nil {
|
||
record[col] = parsed
|
||
continue
|
||
}
|
||
}
|
||
record[col] = string(val)
|
||
case sql.RawBytes:
|
||
if len(val) > 0 && (val[0] == '[' || val[0] == '{') {
|
||
var parsed any
|
||
if json.Unmarshal(val, &parsed) == nil {
|
||
record[col] = parsed
|
||
continue
|
||
}
|
||
}
|
||
record[col] = string(val)
|
||
default:
|
||
record[col] = v
|
||
}
|
||
}
|
||
|
||
return record, s.Err()
|
||
}
|
||
|
||
func isNumericKind(k reflect.Kind) bool {
|
||
switch k {
|
||
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64,
|
||
reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64,
|
||
reflect.Float32, reflect.Float64:
|
||
return true
|
||
}
|
||
return false
|
||
}
|
||
|
||
// ScanStruct will scan the current row into i.
|
||
func (s *scanner) ScanStruct(i interface{}) error {
|
||
// Setup columnMap and columns, but only once.
|
||
if s.columnMap == nil || s.columns == nil {
|
||
cm, err := util.GetColumnMap(i)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
|
||
cols, err := s.rows.Columns()
|
||
if err != nil {
|
||
return err
|
||
}
|
||
|
||
// 补全未知字段类型
|
||
if len(cols) != len(cm) {
|
||
colTypes, ctErr := s.rows.ColumnTypes()
|
||
if ctErr != nil {
|
||
return ctErr
|
||
}
|
||
for _, t := range colTypes {
|
||
if _, ok := cm[t.Name()]; !ok {
|
||
cm[t.Name()] = util.ColumnData{
|
||
ColumnName: t.Name(),
|
||
GoType: t.ScanType(),
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
s.columnMap = cm
|
||
s.columns = cols
|
||
}
|
||
|
||
scans, err := createColumnScans(s.columns, s.columnMap)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
|
||
if err := s.rows.Scan(scans...); err != nil {
|
||
return err
|
||
}
|
||
|
||
record := map[string]interface{}{}
|
||
for index, col := range s.columns {
|
||
if pi, ok := scans[index].(*interface{}); ok {
|
||
v := *pi
|
||
if t, isTime := v.(time.Time); isTime {
|
||
// time.Time 直接存储,避免格式化字符串丢失纳秒和时区
|
||
record[col] = t
|
||
} else {
|
||
raw := toJSONRawMessage(v)
|
||
record[col] = &raw
|
||
}
|
||
} else {
|
||
record[col] = scans[index]
|
||
}
|
||
}
|
||
|
||
util.AssignStructVals(i, record, s.columnMap)
|
||
|
||
return s.Err()
|
||
}
|
||
|
||
// ScanStructs scans results in slice of structs
|
||
func (s *scanner) ScanStructs(i interface{}) error {
|
||
val, err := checkScanStructsTarget(i)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
return s.scanIntoSlice(val, func(i interface{}) error {
|
||
return s.ScanStruct(i)
|
||
})
|
||
}
|
||
|
||
// ScanVal will scan the current row and column into i.
|
||
func (s *scanner) ScanVal(i interface{}) error {
|
||
switch v := i.(type) {
|
||
case *sql.RawBytes:
|
||
// 零拷贝扫描,rows.Close 前立即拷贝防止驱动回收缓冲区
|
||
if err := s.rows.Scan(v); err != nil {
|
||
return err
|
||
}
|
||
buf := make(sql.RawBytes, len(*v))
|
||
copy(buf, *v)
|
||
*v = buf
|
||
case *[]byte:
|
||
// 先扫描到 interface{},驱动可能返回 []byte/string/map 等任意类型
|
||
var raw interface{}
|
||
if err := s.rows.Scan(&raw); err != nil {
|
||
return err
|
||
}
|
||
if msg := toJSONRawMessage(raw); msg != nil {
|
||
*v = []byte(*msg)
|
||
}
|
||
default:
|
||
// 指针-结构体且未实现 sql.Scanner:通过 JSON 中间层转换
|
||
if rv := reflect.ValueOf(i); rv.Kind() == reflect.Ptr && rv.Elem().Kind() == reflect.Struct {
|
||
if _, ok := i.(sql.Scanner); !ok {
|
||
var raw interface{}
|
||
if err := s.rows.Scan(&raw); err != nil {
|
||
return err
|
||
}
|
||
if raw == nil {
|
||
return s.Err()
|
||
}
|
||
data, err := json.Marshal(raw)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
return json.Unmarshal(data, i)
|
||
}
|
||
}
|
||
if err := s.rows.Scan(i); err != nil {
|
||
return err
|
||
}
|
||
}
|
||
return s.Err()
|
||
}
|
||
|
||
// ScanStructs scans results in slice of values
|
||
func (s *scanner) ScanVals(i interface{}) error {
|
||
val, err := checkScanValsTarget(i)
|
||
if err != nil {
|
||
return err
|
||
}
|
||
return s.scanIntoSlice(val, func(i interface{}) error {
|
||
return s.ScanVal(i)
|
||
})
|
||
}
|
||
|
||
// Close closes the Rows, preventing further enumeration. See sql.Rows#Close
|
||
// for more info.
|
||
func (s *scanner) Close() error {
|
||
return s.rows.Close()
|
||
}
|
||
|
||
func (s *scanner) scanIntoSlice(val reflect.Value, it func(i interface{}) error) error {
|
||
elemType := util.GetSliceElementType(val)
|
||
|
||
for s.Next() {
|
||
row := reflect.New(elemType)
|
||
if rowErr := it(row.Interface()); rowErr != nil {
|
||
return rowErr
|
||
}
|
||
util.AppendSliceElement(val, row)
|
||
}
|
||
|
||
return s.Err()
|
||
}
|
||
|
||
func checkScanStructsTarget(i interface{}) (reflect.Value, error) {
|
||
val := reflect.ValueOf(i)
|
||
if !util.IsPointer(val.Kind()) {
|
||
return val, errUnsupportedScanStructsType
|
||
}
|
||
val = reflect.Indirect(val)
|
||
if !util.IsSlice(val.Kind()) {
|
||
return val, errUnsupportedScanStructsType
|
||
}
|
||
return val, nil
|
||
}
|
||
|
||
func checkScanValsTarget(i interface{}) (reflect.Value, error) {
|
||
val := reflect.ValueOf(i)
|
||
if !util.IsPointer(val.Kind()) {
|
||
return val, errUnsupportedScanValsType
|
||
}
|
||
val = reflect.Indirect(val)
|
||
if !util.IsSlice(val.Kind()) {
|
||
return val, errUnsupportedScanValsType
|
||
}
|
||
return val, nil
|
||
}
|
||
|
||
func toJSONRawMessage(v interface{}) *json.RawMessage {
|
||
if v == nil {
|
||
return nil
|
||
}
|
||
var raw json.RawMessage
|
||
switch s := v.(type) {
|
||
case []byte:
|
||
raw = append(json.RawMessage(nil), s...)
|
||
case sql.RawBytes:
|
||
raw = append(json.RawMessage(nil), s...)
|
||
case string:
|
||
raw = json.RawMessage(s)
|
||
case time.Time:
|
||
layout := time.DateTime
|
||
if s.Hour() == 0 && s.Minute() == 0 && s.Second() == 0 && s.Nanosecond() == 0 {
|
||
layout = time.DateOnly
|
||
}
|
||
raw = json.RawMessage(s.Format(layout))
|
||
default:
|
||
b, _ := json.Marshal(v)
|
||
raw = b
|
||
}
|
||
return &raw
|
||
}
|
||
|
||
func createColumnScans(cols []string, cm util.ColumnMap) (scans []interface{}, err error) {
|
||
scans = make([]interface{}, 0, len(cols))
|
||
|
||
for _, col := range cols {
|
||
data, ok := cm[col]
|
||
|
||
if !ok {
|
||
return scans, unableToFindFieldError(col)
|
||
}
|
||
// 处理 converting NULL to string is unsupported
|
||
// 前面将 string 和 int 类型转为了 *string 和 *int
|
||
switch data.GoType.Kind() {
|
||
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64,
|
||
reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64,
|
||
reflect.Float32, reflect.Float64,
|
||
reflect.Bool:
|
||
scans = append(scans, reflect.New(reflect.PointerTo(data.GoType)).Interface())
|
||
case reflect.String, reflect.Map, reflect.Slice, reflect.Struct:
|
||
// string 也用 *interface{}:DuckDB 的 MAP/LIST 列 ScanType() 可能报告为 string,
|
||
// 但驱动实际返回 map[string]interface{} 或 []interface{},必须用 *interface{} 接收。
|
||
scans = append(scans, new(any))
|
||
case reflect.Ptr:
|
||
// 指针到结构体(如 *time.Time):驱动可能返回 []byte,
|
||
// 用 **struct 作扫描目标时 database/sql 无法完成转换,改用 *interface{} 接收。
|
||
// 指针到基础类型(*string、*int 等)仍用 **T,database/sql 可自行处理。
|
||
if data.GoType.Elem().Kind() == reflect.Struct {
|
||
scans = append(scans, new(any))
|
||
} else {
|
||
scans = append(scans, reflect.New(data.GoType).Interface())
|
||
}
|
||
default:
|
||
scans = append(scans, reflect.New(data.GoType).Interface())
|
||
}
|
||
}
|
||
|
||
return scans, nil
|
||
}
|