package db import ( "fmt" "reflect" "time" "github.com/shopspring/decimal" "go.mongodb.org/mongo-driver/bson" "go.mongodb.org/mongo-driver/bson/bsoncodec" "go.mongodb.org/mongo-driver/bson/bsonrw" "go.mongodb.org/mongo-driver/bson/bsontype" "go.mongodb.org/mongo-driver/bson/primitive" ) var registry = func() *bsoncodec.Registry { builder := bson.NewRegistryBuilder() builder.RegisterTypeDecoder(reflect.TypeOf(time.Time{}), &localTimeDecoder{}) builder.RegisterTypeDecoder(reflect.TypeOf(decimal.Decimal{}), &decimalDecoder{}) builder.RegisterTypeEncoder(reflect.TypeOf(decimal.Decimal{}), &decimalEncoder{}) builder.RegisterDefaultDecoder(reflect.Float32, &floatDecoder{}) builder.RegisterDefaultDecoder(reflect.Float64, &floatDecoder{}) return builder.Build() }() type floatDecoder struct { } func (dvd *floatDecoder) DecodeValue(dc bsoncodec.DecodeContext, vr bsonrw.ValueReader, val reflect.Value) error { var f float64 var err error switch vr.Type() { case bsontype.Int32: i32, err := vr.ReadInt32() if err != nil { return err } f = float64(i32) case bsontype.Int64: i64, err := vr.ReadInt64() if err != nil { return err } f = float64(i64) case bsontype.Double: f, err = vr.ReadDouble() if err != nil { return err } default: return fmt.Errorf("cannot decode %v into a float32 or float64 type", vr.Type()) } val.SetFloat(f) return nil } type localTimeDecoder struct{} func (*localTimeDecoder) DecodeValue(dc bsoncodec.DecodeContext, vr bsonrw.ValueReader, val reflect.Value) error { if err := (&bsoncodec.TimeCodec{}).DecodeValue(dc, vr, val); err != nil { return err } t := val.Interface().(time.Time) val.Set(reflect.ValueOf(t.Local())) return nil } type decimalDecoder struct{} func (*decimalDecoder) DecodeValue(dc bsoncodec.DecodeContext, vr bsonrw.ValueReader, val reflect.Value) error { if !val.IsValid() || val.Type() != reflect.TypeOf(decimal.Decimal{}) { return bsoncodec.ValueDecoderError{Name: "DecimalDecodeValue", Types: []reflect.Type{reflect.TypeOf(decimal.Decimal{})}, Received: val} } if vr.Type() == bson.TypeInt32 || vr.Type() == bson.TypeInt64 { _, _ = vr.ReadInt32() _, _ = vr.ReadInt64() d := decimal.NewFromFloat(0.0) val.Set(reflect.ValueOf(d)) return nil } mongodecimal, err := vr.ReadDecimal128() if err != nil { return err } d, err := decimal.NewFromString(mongodecimal.String()) if err != nil { return err } val.Set(reflect.ValueOf(d)) return nil } type decimalEncoder struct{} func (*decimalEncoder) EncodeValue(ctx bsoncodec.EncodeContext, vw bsonrw.ValueWriter, val reflect.Value) error { if !val.IsValid() || val.Type() != reflect.TypeOf(decimal.Decimal{}) { return bsoncodec.ValueDecoderError{Name: "DecimalEncodeValue", Types: []reflect.Type{reflect.TypeOf(decimal.Decimal{})}, Received: val} } if d, ok := val.Interface().(decimal.Decimal); ok { mongodecimal, err := primitive.ParseDecimal128(d.StringFixed(2)) if err != nil { return err } val = reflect.ValueOf(mongodecimal) } dve := bsoncodec.DefaultValueEncoders{} return dve.Decimal128EncodeValue(ctx, vw, val) }