Skip to content

Commit 9ba77ed

Browse files
authored
Merge branch 'development' into ed/issue-1794
2 parents f065500 + 88c59ea commit 9ba77ed

4 files changed

Lines changed: 124 additions & 24 deletions

File tree

lib/trie/node.go

Lines changed: 13 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,7 @@ import (
4747
"sync"
4848

4949
"github.com/ChainSafe/gossamer/lib/common"
50-
"github.com/ChainSafe/gossamer/lib/scale"
50+
"github.com/ChainSafe/gossamer/pkg/scale"
5151
)
5252

5353
// node is the interface for trie methods
@@ -337,26 +337,28 @@ func (b *branch) decode(r io.Reader, header byte) (err error) {
337337
return err
338338
}
339339

340-
sd := &scale.Decoder{Reader: r}
340+
sd := scale.NewDecoder(r)
341341

342342
if nodeType == 3 {
343+
var value []byte
343344
// branch w/ value
344-
value, err := sd.Decode([]byte{})
345+
err := sd.Decode(&value)
345346
if err != nil {
346347
return err
347348
}
348-
b.value = value.([]byte)
349+
b.value = value
349350
}
350351

351352
for i := 0; i < 16; i++ {
352353
if (childrenBitmap[i/8]>>(i%8))&1 == 1 {
353-
hash, err := sd.Decode([]byte{})
354+
var hash []byte
355+
err := sd.Decode(&hash)
354356
if err != nil {
355357
return err
356358
}
357359

358360
b.children[i] = &leaf{
359-
hash: hash.([]byte),
361+
hash: hash,
360362
}
361363
}
362364
}
@@ -386,14 +388,15 @@ func (l *leaf) decode(r io.Reader, header byte) (err error) {
386388
return err
387389
}
388390

389-
sd := &scale.Decoder{Reader: r}
390-
value, err := sd.Decode([]byte{})
391+
sd := scale.NewDecoder(r)
392+
var value []byte
393+
err = sd.Decode(&value)
391394
if err != nil {
392395
return err
393396
}
394397

395-
if len(value.([]byte)) > 0 {
396-
l.value = value.([]byte)
398+
if len(value) > 0 {
399+
l.value = value
397400
}
398401

399402
l.dirty = true

lib/trie/node_test.go

Lines changed: 5 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@ import (
2222
"testing"
2323

2424
"github.com/ChainSafe/gossamer/lib/common"
25-
"github.com/ChainSafe/gossamer/lib/scale"
25+
"github.com/ChainSafe/gossamer/pkg/scale"
2626

2727
"github.com/stretchr/testify/require"
2828
)
@@ -160,14 +160,12 @@ func TestBranchEncode(t *testing.T) {
160160
expected = append(expected, nibblesToKeyLE(b.key)...)
161161
expected = append(expected, common.Uint16ToBytes(b.childrenBitmap())...)
162162

163-
buf := bytes.Buffer{}
164-
encoder := &scale.Encoder{Writer: &buf}
165-
_, err = encoder.Encode(b.value)
163+
enc, err := scale.Marshal(b.value)
166164
if err != nil {
167165
t.Fatalf("Fail when encoding value with scale: %s", err)
168166
}
169167

170-
expected = append(expected, buf.Bytes()...)
168+
expected = append(expected, enc...)
171169

172170
for _, child := range b.children {
173171
if child != nil {
@@ -207,14 +205,12 @@ func TestLeafEncode(t *testing.T) {
207205
expected = append(expected, header...)
208206
expected = append(expected, nibblesToKeyLE(l.key)...)
209207

210-
buf := bytes.Buffer{}
211-
encoder := &scale.Encoder{Writer: &buf}
212-
_, err = encoder.Encode(l.value)
208+
enc, err := scale.Marshal(l.value)
213209
if err != nil {
214210
t.Fatalf("Fail when encoding value with scale: %s", err)
215211
}
216212

217-
expected = append(expected, buf.Bytes()...)
213+
expected = append(expected, enc...)
218214

219215
hasher := newHasher(false)
220216
defer hasher.returnToPool()

pkg/scale/decode.go

Lines changed: 42 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,8 @@ import (
2121
"encoding/binary"
2222
"errors"
2323
"fmt"
24+
"io"
25+
"io/ioutil"
2426
"math/big"
2527
"reflect"
2628
)
@@ -87,7 +89,7 @@ func Unmarshal(data []byte, dst interface{}) (err error) {
8789
if err != nil {
8890
return
8991
}
90-
ds.Buffer = *buf
92+
ds.Reader = buf
9193

9294
err = ds.unmarshal(elem)
9395
if err != nil {
@@ -96,8 +98,36 @@ func Unmarshal(data []byte, dst interface{}) (err error) {
9698
return
9799
}
98100

101+
// Decoder is used to decode from an io.Reader
102+
type Decoder struct {
103+
decodeState
104+
}
105+
106+
// Decode accepts a pointer to a destination and decodes into supplied destination
107+
func (d *Decoder) Decode(dst interface{}) (err error) {
108+
dstv := reflect.ValueOf(dst)
109+
if dstv.Kind() != reflect.Ptr || dstv.IsNil() {
110+
err = fmt.Errorf("unsupported dst: %T, must be a pointer to a destination", dst)
111+
return
112+
}
113+
114+
elem := indirect(dstv)
115+
if err != nil {
116+
return
117+
}
118+
return d.unmarshal(elem)
119+
}
120+
121+
// NewDecoder is constructor for Decoder
122+
func NewDecoder(r io.Reader) (d *Decoder) {
123+
d = &Decoder{
124+
decodeState{r},
125+
}
126+
return
127+
}
128+
99129
type decodeState struct {
100-
bytes.Buffer
130+
io.Reader
101131
}
102132

103133
func (ds *decodeState) unmarshal(dstv reflect.Value) (err error) {
@@ -230,6 +260,12 @@ func (ds *decodeState) decodeCustomPrimitive(dstv reflect.Value) (err error) {
230260
return
231261
}
232262

263+
func (ds *decodeState) ReadByte() (byte, error) {
264+
b := make([]byte, 1) // make buffer
265+
_, err := ds.Reader.Read(b) // read what's in the Decoder's underlying buffer to our new buffer b
266+
return b[0], err
267+
}
268+
233269
func (ds *decodeState) decodeResult(dstv reflect.Value) (err error) {
234270
res := dstv.Interface().(Result)
235271
var rb byte
@@ -263,7 +299,8 @@ func (ds *decodeState) decodeResult(dstv reflect.Value) (err error) {
263299
}
264300
dstv.Set(reflect.ValueOf(res))
265301
default:
266-
err = fmt.Errorf("unsupported Result value: %v, bytes: %v", rb, ds.Bytes())
302+
bytes, _ := ioutil.ReadAll(ds.Reader)
303+
err = fmt.Errorf("unsupported Result value: %v, bytes: %v", rb, bytes)
267304
}
268305
return
269306
}
@@ -295,7 +332,8 @@ func (ds *decodeState) decodePointer(dstv reflect.Value) (err error) {
295332
dstv.Set(tempElem)
296333
}
297334
default:
298-
err = fmt.Errorf("unsupported Option value: %v, bytes: %v", rb, ds.Bytes())
335+
bytes, _ := ioutil.ReadAll(ds.Reader)
336+
err = fmt.Errorf("unsupported Option value: %v, bytes: %v", rb, bytes)
299337
}
300338
return
301339
}

pkg/scale/decode_test.go

Lines changed: 64 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
package scale
1818

1919
import (
20+
"bytes"
2021
"math/big"
2122
"reflect"
2223
"testing"
@@ -189,7 +190,6 @@ func Test_unmarshal_optionality(t *testing.T) {
189190
if diff != "" {
190191
t.Errorf("decodeState.unmarshal() = %s", diff)
191192
}
192-
193193
}
194194
})
195195
}
@@ -238,3 +238,66 @@ func Test_unmarshal_optionality_nil_case(t *testing.T) {
238238
})
239239
}
240240
}
241+
242+
func Test_Decoder_Decode(t *testing.T) {
243+
for _, tt := range newTests(fixedWidthIntegerTests, variableWidthIntegerTests, stringTests,
244+
boolTests, sliceTests, arrayTests,
245+
) {
246+
t.Run(tt.name, func(t *testing.T) {
247+
dst := reflect.New(reflect.TypeOf(tt.in)).Elem().Interface()
248+
wantBuf := bytes.NewBuffer(tt.want)
249+
d := NewDecoder(wantBuf)
250+
if err := d.Decode(&dst); (err != nil) != tt.wantErr {
251+
t.Errorf("Decoder.Decode() error = %v, wantErr %v", err, tt.wantErr)
252+
return
253+
}
254+
if !reflect.DeepEqual(dst, tt.in) {
255+
t.Errorf("Decoder.Decode() = %v, want %v", dst, tt.in)
256+
}
257+
})
258+
}
259+
}
260+
261+
func Test_Decoder_Decode_MultipleCalls(t *testing.T) {
262+
tests := []struct {
263+
name string
264+
ins []interface{}
265+
want []byte
266+
wantErr []bool
267+
}{
268+
{
269+
name: "int64 and []byte",
270+
ins: []interface{}{int64(9223372036854775807), []byte{0x01}},
271+
want: append([]byte{0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0x7f}, []byte{0x04, 0x01}...),
272+
},
273+
{
274+
name: "eof error",
275+
ins: []interface{}{int64(9223372036854775807), []byte{0x01}},
276+
want: []byte{0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0x7f},
277+
wantErr: []bool{false, true},
278+
},
279+
}
280+
for _, tt := range tests {
281+
t.Run(tt.name, func(t *testing.T) {
282+
buf := bytes.NewBuffer(tt.want)
283+
d := NewDecoder(buf)
284+
285+
for i := range tt.ins {
286+
in := tt.ins[i]
287+
dst := reflect.New(reflect.TypeOf(in)).Elem().Interface()
288+
var wantErr bool
289+
if len(tt.wantErr) > i {
290+
wantErr = tt.wantErr[i]
291+
}
292+
if err := d.Decode(&dst); (err != nil) != wantErr {
293+
t.Errorf("Decoder.Decode() error = %v, wantErr %v", err, tt.wantErr[i])
294+
return
295+
}
296+
if !wantErr && !reflect.DeepEqual(dst, in) {
297+
t.Errorf("Decoder.Decode() = %v, want %v", dst, in)
298+
return
299+
}
300+
}
301+
})
302+
}
303+
}

0 commit comments

Comments
 (0)