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 }