common: implement sql scanner/valuer on Address

This commit is contained in:
Vincent Serpoul 2018-06-08 12:15:28 +08:00
parent aaca11da13
commit 25556b2aa7
2 changed files with 140 additions and 2 deletions

View file

@ -255,6 +255,30 @@ func (a *Address) UnmarshalJSON(input []byte) error {
return hexutil.UnmarshalFixedJSON(addressT, input, a[:]) return hexutil.UnmarshalFixedJSON(addressT, input, a[:])
} }
// Scan implements Scanner for database/sql
func (a *Address) Scan(src interface{}) error {
srcB, ok := src.([]byte)
if !ok {
return fmt.Errorf("Address Scan: couldn't scan %v into Address", src)
}
if len(srcB) != AddressLength {
return fmt.Errorf(
"Address Scan: len %d instead of expected %d",
len(srcB),
AddressLength,
)
}
*a = BytesToAddress(srcB)
return nil
}
// Value implements valuer for database/sql
func (a Address) Value() (driver.Value, error) {
return a.Bytes(), nil
}
// UnprefixedAddress allows marshaling an Address without 0x prefix. // UnprefixedAddress allows marshaling an Address without 0x prefix.
type UnprefixedAddress Address type UnprefixedAddress Address

View file

@ -260,7 +260,7 @@ func BenchmarkHash_Scan(b *testing.B) {
h := &Hash{} h := &Hash{}
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
if err := h.Scan(tst); err != nil { if err := h.Scan(tst); err != nil {
b.Errorf("BenchmarkHash_Scan: error Scan on hash %v", err) b.Errorf("BenchmarkHash_Scan: error Scan on Hash %v", err)
} }
} }
} }
@ -312,7 +312,121 @@ func BenchmarkHash_Value(b *testing.B) {
usedH.SetBytes(tst) usedH.SetBytes(tst)
for i := 0; i < b.N; i++ { for i := 0; i < b.N; i++ {
if _, err := usedH.Value(); err != nil { if _, err := usedH.Value(); err != nil {
b.Errorf("BenchmarkHash_Value: error Value on hash %v", err) b.Errorf("BenchmarkHash_Value: error Value on Hash %v", err)
return
}
}
}
func TestAddress_Scan(t *testing.T) {
type args struct {
src interface{}
}
tests := []struct {
name string
args args
wantErr bool
}{
{
name: "working scan",
args: args{src: []byte{
0xb2, 0x6f, 0x2b, 0x34, 0x2a, 0xab, 0x24, 0xbc, 0xf6, 0x3e,
0xa2, 0x18, 0xc6, 0xa9, 0x27, 0x4d, 0x30, 0xab, 0x9a, 0x15,
}},
wantErr: false,
},
{
name: "non working scan",
args: args{src: int64(1234567890)},
wantErr: true,
},
{
name: "invalid length scan",
args: args{src: []byte{
0xb2, 0x6f, 0x2b, 0x34, 0x2a, 0xab, 0x24, 0xbc, 0xf6, 0x3e,
0xa2, 0x18, 0xc6, 0xa9, 0x27, 0x4d, 0x30, 0xab, 0x9a,
}},
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
a := &Address{}
if err := a.Scan(tt.args.src); (err != nil) != tt.wantErr {
t.Errorf("Address.Scan() error = %v, wantErr %v", err, tt.wantErr)
}
if !tt.wantErr {
for i := range a {
if a[i] != tt.args.src.([]byte)[i] {
t.Errorf(
"Address.Scan() didn't scan the %d src correctly (have %X, want %X)",
i, a[i], tt.args.src.([]byte)[i],
)
}
}
}
})
}
}
func BenchmarkAddress_Scan(b *testing.B) {
tst := []byte{
0xb2, 0x6f, 0x2b, 0x34, 0x2a, 0xab, 0x24, 0xbc, 0xf6, 0x3e,
0xa2, 0x18, 0xc6, 0xa9, 0x27, 0x4d, 0x30, 0xab, 0x9a, 0x15,
}
a := &Address{}
for i := 0; i < b.N; i++ {
if err := a.Scan(tst); err != nil {
b.Errorf("BenchmarkAddress_Scan: error Scan on Address %v", err)
}
}
}
func TestAddress_Value(t *testing.T) {
b := []byte{
0xb2, 0x6f, 0x2b, 0x34, 0x2a, 0xab, 0x24, 0xbc, 0xf6, 0x3e,
0xa2, 0x18, 0xc6, 0xa9, 0x27, 0x4d, 0x30, 0xab, 0x9a, 0x15,
}
var usedA Address
usedA.SetBytes(b)
tests := []struct {
name string
a Address
want driver.Value
wantErr bool
}{
{
name: "Working value",
a: usedA,
want: b,
wantErr: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := tt.a.Value()
if (err != nil) != tt.wantErr {
t.Errorf("Address.Value() error = %v, wantErr %v", err, tt.wantErr)
return
}
if !reflect.DeepEqual(got, tt.want) {
t.Errorf("Address.Value() = %v, want %v", got, tt.want)
}
})
}
}
func BenchmarkAddress_Value(b *testing.B) {
tst := []byte{
0xb2, 0x6f, 0x2b, 0x34, 0x2a, 0xab, 0x24, 0xbc, 0xf6, 0x3e,
0xa2, 0x18, 0xc6, 0xa9, 0x27, 0x4d, 0x30, 0xab, 0x9a, 0x15,
}
var usedA Address
usedA.SetBytes(tst)
for i := 0; i < b.N; i++ {
if _, err := usedA.Value(); err != nil {
b.Errorf("BenchmarkAddress_Value: error Value on Address %v", err)
return return
} }
} }