diff --git a/exec/query_executor_internal_test.go b/exec/query_executor_internal_test.go index ad5babd..d1ab333 100644 --- a/exec/query_executor_internal_test.go +++ b/exec/query_executor_internal_test.go @@ -1426,6 +1426,93 @@ func (qes *queryExecutorSuite) TestScanVal_withStruct_sqlScanner() { qes.Equal(sql.NullString{String: "hello", Valid: true}, ns) } +// TestScanStructs_timeField_fromTime 验证 time.Time 和 *time.Time 字段直接从 time.Time 驱动值扫描, +// 纳秒精度和时区 Location 必须原样保留。 +func (qes *queryExecutorSuite) TestScanStructs_timeField_fromTime() { + type Row struct { + T time.Time `db:"t"` + PT *time.Time `db:"pt"` + } + db, mock, err := sqlmock.New() + qes.NoError(err) + now := time.Now() // 含纳秒、Local 时区 + mock.ExpectQuery(`SELECT \* FROM "items"`). + WillReturnRows(sqlmock.NewRows([]string{"t", "pt"}). + AddRow(now, now). + AddRow(now, nil), + ) + + e := newQueryExecutor(db, nil, `SELECT * FROM "items"`) + var items []Row + qes.NoError(e.ScanStructs(&items)) + qes.Len(items, 2) + + // 时间相等(Equal 只比较时刻,不比较 Location) + qes.True(items[0].T.Equal(now)) + qes.NotNil(items[0].PT) + qes.True(items[0].PT.Equal(now)) + // 纳秒精度保留 + qes.Equal(now.Nanosecond(), items[0].T.Nanosecond()) + qes.Equal(now.Nanosecond(), items[0].PT.Nanosecond()) + // NULL → nil 指针 + qes.Nil(items[1].PT) +} + +// TestScanStructs_timeField_fromBytes 验证 time.Time 和 *time.Time 字段从 []byte 字符串扫描, +// 模拟 MySQL 驱动不配置 parseTime=true 时返回 []byte 的场景(原 bug:scan error on *time.Time)。 +func (qes *queryExecutorSuite) TestScanStructs_timeField_fromBytes() { + type Row struct { + CreatedAt time.Time `db:"created_at"` + RetryAt *time.Time `db:"retry_at"` + } + db, mock, err := sqlmock.New(sqlmock.ValueConverterOption(anyValueConverter{})) + qes.NoError(err) + + datetimeBytes := []byte("2024-03-15 10:30:45") + mock.ExpectQuery(`SELECT \* FROM "items"`). + WillReturnRows(mock.NewRows([]string{"created_at", "retry_at"}). + AddRow(datetimeBytes, datetimeBytes). + AddRow(datetimeBytes, nil), + ) + + e := newQueryExecutor(db, nil, `SELECT * FROM "items"`) + var items []Row + qes.NoError(e.ScanStructs(&items)) + qes.Len(items, 2) + + want := time.Date(2024, 3, 15, 10, 30, 45, 0, time.UTC) + qes.True(items[0].CreatedAt.Equal(want)) + qes.NotNil(items[0].RetryAt) + qes.True(items[0].RetryAt.Equal(want)) + qes.Nil(items[1].RetryAt) +} + +// TestScanStructs_timeField_fromDateOnlyBytes 验证仅含日期的 []byte(无时分秒)也能正确解析。 +func (qes *queryExecutorSuite) TestScanStructs_timeField_fromDateOnlyBytes() { + type Row struct { + Birthday time.Time `db:"birthday"` + ExpiredAt *time.Time `db:"expired_at"` + } + db, mock, err := sqlmock.New(sqlmock.ValueConverterOption(anyValueConverter{})) + qes.NoError(err) + + dateBytes := []byte("2024-03-15") + mock.ExpectQuery(`SELECT \* FROM "items"`). + WillReturnRows(mock.NewRows([]string{"birthday", "expired_at"}). + AddRow(dateBytes, dateBytes), + ) + + e := newQueryExecutor(db, nil, `SELECT * FROM "items"`) + var items []Row + qes.NoError(e.ScanStructs(&items)) + qes.Len(items, 1) + + want := time.Date(2024, 3, 15, 0, 0, 0, 0, time.UTC) + qes.True(items[0].Birthday.Equal(want)) + qes.NotNil(items[0].ExpiredAt) + qes.True(items[0].ExpiredAt.Equal(want)) +} + func TestQueryExecutorSuite(t *testing.T) { suite.Run(t, new(queryExecutorSuite)) } diff --git a/exec/scanner.go b/exec/scanner.go index ca4aa55..749235a 100644 --- a/exec/scanner.go +++ b/exec/scanner.go @@ -4,6 +4,7 @@ import ( "database/sql" "encoding/json" "reflect" + "time" "git.fsdpf.net/go/db/internal/errors" "git.fsdpf.net/go/db/internal/util" @@ -65,72 +66,92 @@ func (s *scanner) GetRecords() ([]map[string]any, error) { } func (s *scanner) GetRecord() (record map[string]any, err error) { - // Setup columns, but only once. - if s.columns == nil || s.columnMap == nil { + if s.columns == nil { colsType, err := s.rows.ColumnTypes() if err != nil { return nil, err } - s.columns = make([]string, len(colsType)) - s.columnMap = util.ColumnMap{} - for i := range colsType { - s.columns[i] = colsType[i].Name() - - typ := colsType[i].ScanType() - - // 强制日期时间为字符串 - if typ == reflect.TypeOf(sql.NullTime{}) { - typ = reflect.TypeOf(sql.Null[[]uint8]{}) - } - - s.columnMap[colsType[i].Name()] = util.ColumnData{ - ColumnName: colsType[i].Name(), - GoType: typ, + 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([]interface{}, len(s.columns)) - + scans := make([]any, len(s.columns)) for i, col := range s.columns { - scans[i] = reflect.New(s.columnMap[col].GoType).Interface() + 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]interface{}, len(s.columns)) + record = make(map[string]any, len(s.columns)) for i, col := range s.columns { - var vv any - switch v := scans[i].(type) { - case *sql.Null[[]uint8]: // 强制日期时间为字符串 - if vv, err = v.Value(); err == nil && vv != nil { - vv = string(vv.([]uint8)) - } - case *sql.RawBytes: - if len(*v) > 1 { - if rune((*v)[0]) == rune('[') || rune((*v)[0]) == rune('{') { - err = json.Unmarshal(*v, &vv) - } else { - vv = string(*v) - } + if isNumericKind(s.columnMap[col].GoType.Kind()) { + inner := reflect.ValueOf(scans[i]).Elem() // *T + if inner.IsNil() { + record[col] = nil } else { - vv = string(*v) + 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: - vv = reflect.Indirect(reflect.ValueOf(v)).Interface() + record[col] = v } - if err != nil { - return - } - record[col] = vv } 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. @@ -177,8 +198,14 @@ func (s *scanner) ScanStruct(i interface{}) error { record := map[string]interface{}{} for index, col := range s.columns { if pi, ok := scans[index].(*interface{}); ok { - raw := toJSONRawMessage(*pi) - record[col] = &raw + 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] } @@ -217,21 +244,8 @@ func (s *scanner) ScanVal(i interface{}) error { if err := s.rows.Scan(&raw); err != nil { return err } - switch rv := raw.(type) { - case []byte: - *v = append([]byte(nil), rv...) - case sql.RawBytes: - *v = append([]byte(nil), []byte(rv)...) - case string: - *v = []byte(rv) - default: - if raw != nil { - var err error - *v, err = json.Marshal(raw) - if err != nil { - return err - } - } + if msg := toJSONRawMessage(raw); msg != nil { + *v = []byte(*msg) } default: // 指针-结构体且未实现 sql.Scanner:通过 JSON 中间层转换 @@ -321,8 +335,16 @@ func toJSONRawMessage(v interface{}) *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 @@ -342,14 +364,24 @@ func createColumnScans(cols []string, cm util.ColumnMap) (scans []interface{}, e // 处理 converting NULL to string is unsupported // 前面将 string 和 int 类型转为了 *string 和 *int switch data.GoType.Kind() { - case reflect.String, - reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64, + 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.Map, reflect.Slice, reflect.Struct: - // 使用 *interface{} 接受任意驱动值(兼容 DuckDB 返回 map[string]interface{}) - scans = append(scans, new(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()) } diff --git a/internal/util/reflect.go b/internal/util/reflect.go index b44e74a..0747ebb 100644 --- a/internal/util/reflect.go +++ b/internal/util/reflect.go @@ -6,6 +6,7 @@ import ( "reflect" "strings" "sync" + "time" "git.fsdpf.net/go/db/internal/errors" ) @@ -20,6 +21,27 @@ const ( var scannerType = reflect.TypeOf((*sql.Scanner)(nil)).Elem() +var timeType = reflect.TypeOf(time.Time{}) + +var timeParseFmts = []string{ + time.DateTime, // "2006-01-02 15:04:05" + time.DateOnly, // "2006-01-02" + "2006-01-02 15:04:05.999999999", // MySQL 带微秒 + "2006-01-02T15:04:05", // ISO8601 无时区 + time.RFC3339Nano, + time.RFC3339, +} + +func parseTimeFromRaw(u *json.RawMessage) (time.Time, bool) { + s := strings.Trim(string(*u), `"`) + for _, f := range timeParseFmts { + if t, err := time.Parse(f, s); err == nil { + return t, true + } + } + return time.Time{}, false +} + func IsUint(k reflect.Kind) bool { return (k == reflect.Uint) || (k == reflect.Uint8) || @@ -207,8 +229,15 @@ func SafeSetVarValue(v reflect.Value, src interface{}) error { // src 可能是 **T(createColumnScans 的扫描目标)或裸值(测试/直接调用) if srcReflect.Kind() != reflect.Ptr { - if srcReflect.IsValid() && v.Type().ConvertibleTo(srcReflect.Type()) { - v.Set(srcReflect.Convert(v.Type())) + if srcReflect.IsValid() { + if v.Type().ConvertibleTo(srcReflect.Type()) { + v.Set(srcReflect.Convert(v.Type())) + } else if v.Kind() == reflect.Ptr && srcReflect.Type().ConvertibleTo(v.Type().Elem()) { + // src = T,v = *T:分配新指针并赋值(如 time.Time → *time.Time) + p := reflect.New(v.Type().Elem()) + p.Elem().Set(srcReflect.Convert(v.Type().Elem())) + v.Set(p) + } } return nil } @@ -221,11 +250,24 @@ func SafeSetVarValue(v reflect.Value, src interface{}) error { } if v.Kind() == reflect.Ptr { - // v 是指针字段(如 *sql.NullString) + // v 是指针字段(如 *sql.NullString、*time.Time) if srcVal.Kind() == reflect.Ptr { // src = **T, srcVal = *T → v = *T if v.Type().ConvertibleTo(srcVal.Type()) { v.Set(srcVal.Convert(v.Type())) + } else if u, ok := srcVal.Interface().(*json.RawMessage); ok && len(*u) >= 2 { + if v.Type().Elem() == timeType { + if t, ok := parseTimeFromRaw(u); ok { + p := reflect.New(timeType) + p.Elem().Set(reflect.ValueOf(t)) + v.Set(p) + } + } else { + p := reflect.New(v.Type().Elem()) + if err := json.Unmarshal(*u, p.Interface()); err == nil { + v.Set(p) + } + } } } else { // src = *T, srcVal = T → allocate new *T and set @@ -244,6 +286,12 @@ func SafeSetVarValue(v reflect.Value, src interface{}) error { if v.Type().ConvertibleTo(srcVal.Type().Elem()) { v.Set(srcVal.Elem().Convert(v.Type())) } else if u, ok := srcVal.Interface().(*json.RawMessage); ok && len(*u) >= 2 { + if v.Type() == timeType { + if t, ok := parseTimeFromRaw(u); ok { + v.Set(reflect.ValueOf(t)) + return nil + } + } if err := json.Unmarshal(*u, v.Addr().Interface()); err != nil { return err }