diff --git a/portalnetwork/storage/content_storage.go b/portalnetwork/storage/content_storage.go index e726d612e5..a5d0d2b26d 100644 --- a/portalnetwork/storage/content_storage.go +++ b/portalnetwork/storage/content_storage.go @@ -7,6 +7,7 @@ import ( ) var ErrContentNotFound = fmt.Errorf("content not found") +var ErrInsufficientRadius = fmt.Errorf("insufficient radius") var MaxDistance = uint256.MustFromHex("0xffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffff") diff --git a/portalnetwork/storage/ethpepple/storage.go b/portalnetwork/storage/ethpepple/storage.go index 897918bbbb..922e1a9e8e 100644 --- a/portalnetwork/storage/ethpepple/storage.go +++ b/portalnetwork/storage/ethpepple/storage.go @@ -182,6 +182,20 @@ func (c *ContentStorage) Get(contentKey []byte, contentId []byte) ([]byte, error // Put implements storage.ContentStorage. func (c *ContentStorage) Put(contentKey []byte, contentId []byte, content []byte) error { + distance := xor(contentId, c.nodeId[:]) + valid, err := c.inRadius(distance) + if err != nil { + return err + } + if !valid { + return storage.ErrInsufficientRadius + } + + err = c.db.Set(distance, content, c.writeOptions) + if err != nil { + return err + } + length := uint64(len(contentId)) + uint64(len(content)) <-c.sizeChan c.size += length @@ -193,8 +207,7 @@ func (c *ContentStorage) Put(contentKey []byte, contentId []byte, content []byte } } c.sizeChan <- struct{}{} - distance := xor(contentId, c.nodeId[:]) - return c.db.Set(distance, content, c.writeOptions) + return nil } // Radius implements storage.ContentStorage. @@ -263,6 +276,17 @@ func (c *ContentStorage) prune() error { return nil } +func (c *ContentStorage) inRadius(distance []byte) (bool, error) { + dis := uint256.NewInt(0) + err := dis.UnmarshalSSZ(distance) + if err != nil { + return false, err + } + val := c.radius.Load() + radius := val.(*uint256.Int) + return radius.Gt(dis), nil +} + func xor(contentId, nodeId []byte) []byte { // length of contentId maybe not 32bytes padding := make([]byte, 32) diff --git a/portalnetwork/storage/ethpepple/storage_test.go b/portalnetwork/storage/ethpepple/storage_test.go index 44e3cd301d..9aaca6f6d0 100644 --- a/portalnetwork/storage/ethpepple/storage_test.go +++ b/portalnetwork/storage/ethpepple/storage_test.go @@ -2,7 +2,6 @@ package ethpepple import ( "testing" - "time" "github.com/ethereum/go-ethereum/portalnetwork/storage" "github.com/holiman/uint256" @@ -102,61 +101,63 @@ func TestPrune(t *testing.T) { testcases := []struct { contentKey [32]byte content []byte - shouldPrune bool + outOfRadius bool + err error }{ { - contentKey: uint256.NewInt(1).Bytes32(), - content: genBytes(900_000), - shouldPrune: false, + contentKey: uint256.NewInt(1).Bytes32(), + content: genBytes(900_000), }, { - contentKey: uint256.NewInt(2).Bytes32(), - content: genBytes(40_000), - shouldPrune: false, + contentKey: uint256.NewInt(2).Bytes32(), + content: genBytes(40_000), }, { - contentKey: uint256.NewInt(3).Bytes32(), - content: genBytes(20_000), - shouldPrune: true, + contentKey: uint256.NewInt(3).Bytes32(), + content: genBytes(20_000), + err: storage.ErrContentNotFound, }, { - contentKey: uint256.NewInt(4).Bytes32(), - content: genBytes(20_000), - shouldPrune: true, + contentKey: uint256.NewInt(4).Bytes32(), + content: genBytes(20_000), + err: storage.ErrContentNotFound, }, { - contentKey: uint256.NewInt(5).Bytes32(), - content: genBytes(20_000), - shouldPrune: true, + contentKey: uint256.NewInt(5).Bytes32(), + content: genBytes(20_000), + err: storage.ErrContentNotFound, }, { contentKey: uint256.NewInt(6).Bytes32(), content: genBytes(20_000), - shouldPrune: false, + err: storage.ErrInsufficientRadius, + outOfRadius: true, }, { contentKey: uint256.NewInt(7).Bytes32(), content: genBytes(20_000), - shouldPrune: false, + err: storage.ErrInsufficientRadius, + outOfRadius: true, }, } for _, val := range testcases { - db.Put(val.contentKey[:], val.contentKey[:], val.content) + err := db.Put(val.contentKey[:], val.contentKey[:], val.content) + if err != nil { + assert.Equal(t, val.err, err) + } } - // // wait to prune done - time.Sleep(2 * time.Second) for _, val := range testcases { content, err := db.Get(val.contentKey[:], val.contentKey[:]) - if !val.shouldPrune { + if err == nil { assert.Equal(t, val.content, content) - } else { - assert.Error(t, err) + } else if !val.outOfRadius { + assert.Equal(t, val.err, err) } } radius := db.Radius() data, err := radius.MarshalSSZ() assert.NoError(t, err) - actual := uint256.NewInt(4).Bytes32() + actual := uint256.NewInt(2).Bytes32() assert.Equal(t, data, actual[:]) }