- exec/scanner: 用 *interface{} 替换 **json.RawMessage 扫描目标,兼容 DuckDB 返回 map[string]interface{} 的场景;新增 toJSONRawMessage 转换函数
- exec/scanner: ScanVal 支持结构体指针,通过 JSON 中间层转换(DuckDB STRUCT 列)
- exec/scanner: 将 *sql.RawBytes 和 *[]byte 的处理从 ScanValContext 移入 scanner.ScanVal
- exec/query_executor: 简化 ScanValContext,移除私有 scan 方法
- exec: 补充 scanner 级别 ScanVal 测试用例
- internal/util/reflect: 重写 SafeSetVarValue,修复非指针 src 及 nil 指针字段的 panic
- internal/util/column_map: 恢复非匿名带标签结构体字段的展开逻辑
- schema: 新增 vector 列类型支持
- engine: 补充 DuckDB 相关配置
- dialect/sqlite3/vtab: 完善虚拟表适配器
- 各方言测试改用 sqlmock 虚拟连接
360 lines
8.1 KiB
Go
360 lines
8.1 KiB
Go
package exec
|
||
|
||
import (
|
||
"database/sql"
|
||
"encoding/json"
|
||
"reflect"
|
||
|
||
"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) {
|
||
// Setup columns, but only once.
|
||
if s.columns == nil || s.columnMap == 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,
|
||
}
|
||
}
|
||
}
|
||
|
||
scans := make([]interface{}, len(s.columns))
|
||
|
||
for i, col := range s.columns {
|
||
scans[i] = reflect.New(s.columnMap[col].GoType).Interface()
|
||
}
|
||
|
||
if err := s.rows.Scan(scans...); err != nil {
|
||
return nil, err
|
||
}
|
||
|
||
record = make(map[string]interface{}, 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)
|
||
}
|
||
} else {
|
||
vv = string(*v)
|
||
}
|
||
default:
|
||
vv = reflect.Indirect(reflect.ValueOf(v)).Interface()
|
||
}
|
||
if err != nil {
|
||
return
|
||
}
|
||
record[col] = vv
|
||
}
|
||
|
||
return record, s.Err()
|
||
}
|
||
|
||
// 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 {
|
||
raw := toJSONRawMessage(*pi)
|
||
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
|
||
}
|
||
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
|
||
}
|
||
}
|
||
}
|
||
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 string:
|
||
raw = json.RawMessage(s)
|
||
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.String,
|
||
reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64,
|
||
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{}))
|
||
default:
|
||
scans = append(scans, reflect.New(data.GoType).Interface())
|
||
}
|
||
}
|
||
|
||
return scans, nil
|
||
}
|