swarm/network/bitvector: add tests, fix Get method and add Length

This commit is contained in:
Janos Guljas 2018-01-02 16:37:04 +01:00 committed by Balint Gabor
parent d6caf09d65
commit ca2882fc81
2 changed files with 107 additions and 5 deletions

View file

@ -1,34 +1,48 @@
package bitvector package bitvector
import "errors"
var errInvalidLength = errors.New("invalid length")
type BitVector struct { type BitVector struct {
len int len int
b []byte b []byte
} }
func New(l int) *BitVector { func New(l int) (bv *BitVector, err error) {
return NewFromBytes(make([]byte, l/8+1), l) 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{ return &BitVector{
len: l, len: l,
b: b, b: b,
} }, nil
} }
func (bv *BitVector) Get(i int) bool { func (bv *BitVector) Get(i int) bool {
bi := i / 8 bi := i / 8
return uint8(bv.b[bi])&0x1>>uint(i%8) != 0 return uint8(bv.b[bi])&(0x1<<uint(i%8)) != 0
} }
func (bv *BitVector) Set(i int, v bool) { func (bv *BitVector) Set(i int, v bool) {
bi := i / 8 bi := i / 8
cv := bv.Get(i) cv := bv.Get(i)
if cv != v { if cv != v {
bv.b[bi] ^= 0x1 >> uint8(i%8) bv.b[bi] ^= 0x1 << uint8(i%8)
} }
} }
func (bv *BitVector) Bytes() []byte { func (bv *BitVector) Bytes() []byte {
return bv.b return bv.b
} }
func (bv *BitVector) Length() int {
return bv.len
}

View file

@ -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])
}
}