From 25556b2aa7bac40eceaf772f898345652ccc4549 Mon Sep 17 00:00:00 2001 From: Vincent Serpoul Date: Fri, 8 Jun 2018 12:15:28 +0800 Subject: [PATCH] common: implement sql scanner/valuer on Address --- common/types.go | 24 +++++++++ common/types_test.go | 118 ++++++++++++++++++++++++++++++++++++++++++- 2 files changed, 140 insertions(+), 2 deletions(-) diff --git a/common/types.go b/common/types.go index 1e70b97a44..9787aaf04f 100644 --- a/common/types.go +++ b/common/types.go @@ -255,6 +255,30 @@ func (a *Address) UnmarshalJSON(input []byte) error { 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. type UnprefixedAddress Address diff --git a/common/types_test.go b/common/types_test.go index 4fad8f4581..443bf285a7 100644 --- a/common/types_test.go +++ b/common/types_test.go @@ -260,7 +260,7 @@ func BenchmarkHash_Scan(b *testing.B) { h := &Hash{} for i := 0; i < b.N; i++ { 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) for i := 0; i < b.N; i++ { 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 } }