diff --git a/swarm/network/bitvector/bitvector.go b/swarm/network/bitvector/bitvector.go index 1769fa4f1d..93e55a9d09 100644 --- a/swarm/network/bitvector/bitvector.go +++ b/swarm/network/bitvector/bitvector.go @@ -1,34 +1,48 @@ package bitvector +import "errors" + +var errInvalidLength = errors.New("invalid length") + type BitVector struct { len int b []byte } -func New(l int) *BitVector { +func New(l int) (bv *BitVector, err error) { return NewFromBytes(make([]byte, l/8+1), l) } -func NewFromBytes(b []byte, l int) *BitVector { +func NewFromBytes(b []byte, l int) (bv *BitVector, err error) { + if l <= 0 { + return nil, errInvalidLength + } + if len(b)*8 < l { + return nil, errInvalidLength + } return &BitVector{ len: l, b: b, - } + }, nil } func (bv *BitVector) Get(i int) bool { bi := i / 8 - return uint8(bv.b[bi])&0x1>>uint(i%8) != 0 + return uint8(bv.b[bi])&(0x1<> uint8(i%8) + bv.b[bi] ^= 0x1 << uint8(i%8) } } func (bv *BitVector) Bytes() []byte { return bv.b } + +func (bv *BitVector) Length() int { + return bv.len +} diff --git a/swarm/network/bitvector/bitvector_test.go b/swarm/network/bitvector/bitvector_test.go new file mode 100644 index 0000000000..ae759404d1 --- /dev/null +++ b/swarm/network/bitvector/bitvector_test.go @@ -0,0 +1,88 @@ +package bitvector + +import "testing" + +func TestBitvectorNew(t *testing.T) { + _, err := New(0) + if err != errInvalidLength { + t.Errorf("expected err %v, got %v", errInvalidLength, err) + } + + _, err = NewFromBytes(nil, 0) + if err != errInvalidLength { + t.Errorf("expected err %v, got %v", errInvalidLength, err) + } + + _, err = NewFromBytes([]byte{0}, 9) + if err != errInvalidLength { + t.Errorf("expected err %v, got %v", errInvalidLength, err) + } + + _, err = NewFromBytes(make([]byte, 8), 8) + if err != nil { + t.Error(err) + } +} + +func TestBitvectorGetSet(t *testing.T) { + for _, length := range []int{ + 1, + 2, + 4, + 8, + 9, + 15, + 16, + } { + bv, err := New(length) + if err != nil { + t.Errorf("error for length %v: %v", length, err) + } + + for i := 0; i < length; i++ { + if bv.Get(i) { + t.Errorf("expected false for element on index %v", i) + } + } + + func() { + defer func() { + if err := recover(); err == nil { + t.Errorf("expecting panic") + } + }() + bv.Get(length + 8) + }() + + for i := 0; i < length; i++ { + bv.Set(i, true) + for j := 0; j < length; j++ { + if j == i { + if bv.Get(j) != true { + t.Errorf("element on index %v is not set to true", i) + } + } else { + if bv.Get(j) != false { + t.Errorf("element on index %v is not false", i) + } + } + } + + bv.Set(i, false) + + if bv.Get(i) != false { + t.Errorf("element on index %v is not set to false", i) + } + } + } +} + +func TestBitvectorNewFromBytesGet(t *testing.T) { + bv, err := NewFromBytes([]byte{8}, 8) + if err != nil { + t.Error(err) + } + if bv.Get(3) != true { + t.Fatalf("element 3 is not set to true: state %08b", bv.b[0]) + } +}