les: using stateFeedback instead of BV

This commit is contained in:
Zsolt Felfoldi 2019-12-17 18:46:42 +01:00
parent c52605cac9
commit 21f760b68e
10 changed files with 178 additions and 81 deletions

View file

@ -290,12 +290,12 @@ func (bt *balanceTracker) updateAfter(dt time.Duration) {
} }
// requestCost should be called after serving a request for the given peer // requestCost should be called after serving a request for the given peer
func (bt *balanceTracker) requestCost(cost uint64) { func (bt *balanceTracker) requestCost(cost uint64) uint64 {
bt.lock.Lock() bt.lock.Lock()
defer bt.lock.Unlock() defer bt.lock.Unlock()
if bt.stopped { if bt.stopped {
return return 0
} }
now := bt.clock.Now() now := bt.clock.Now()
bt.addBalance(now) bt.addBalance(now)
@ -323,6 +323,7 @@ func (bt *balanceTracker) requestCost(cost uint64) {
} }
} }
bt.sumReqCost += cost bt.sumReqCost += cost
return bt.balance.pos
} }
// getBalance returns the current positive and negative balance // getBalance returns the current positive and negative balance

View file

@ -200,13 +200,15 @@ func (h *clientHandler) handleMsg(p *peer) error {
case BlockHeadersMsg: case BlockHeadersMsg:
p.Log().Trace("Received block header response message") p.Log().Trace("Received block header response message")
var resp struct { var resp struct {
ReqID, BV uint64 ReqID uint64
Headers []*types.Header SF stateFeedback
Headers []*types.Header
} }
resp.SF.protocolVersion = p.version
if err := msg.Decode(&resp); err != nil { if err := msg.Decode(&resp); err != nil {
return errResp(ErrDecode, "msg %v: %v", msg, err) return errResp(ErrDecode, "msg %v: %v", msg, err)
} }
p.fcServer.ReceivedReply(resp.ReqID, resp.BV) p.fcServer.ReceivedReply(resp.ReqID, resp.SF.BV)
if h.fetcher.requestedID(resp.ReqID) { if h.fetcher.requestedID(resp.ReqID) {
h.fetcher.deliverHeaders(p, resp.ReqID, resp.Headers) h.fetcher.deliverHeaders(p, resp.ReqID, resp.Headers)
} else { } else {
@ -217,13 +219,15 @@ func (h *clientHandler) handleMsg(p *peer) error {
case BlockBodiesMsg: case BlockBodiesMsg:
p.Log().Trace("Received block bodies response") p.Log().Trace("Received block bodies response")
var resp struct { var resp struct {
ReqID, BV uint64 ReqID uint64
Data []*types.Body SF stateFeedback
Data []*types.Body
} }
resp.SF.protocolVersion = p.version
if err := msg.Decode(&resp); err != nil { if err := msg.Decode(&resp); err != nil {
return errResp(ErrDecode, "msg %v: %v", msg, err) return errResp(ErrDecode, "msg %v: %v", msg, err)
} }
p.fcServer.ReceivedReply(resp.ReqID, resp.BV) p.fcServer.ReceivedReply(resp.ReqID, resp.SF.BV)
deliverMsg = &Msg{ deliverMsg = &Msg{
MsgType: MsgBlockBodies, MsgType: MsgBlockBodies,
ReqID: resp.ReqID, ReqID: resp.ReqID,
@ -232,13 +236,15 @@ func (h *clientHandler) handleMsg(p *peer) error {
case CodeMsg: case CodeMsg:
p.Log().Trace("Received code response") p.Log().Trace("Received code response")
var resp struct { var resp struct {
ReqID, BV uint64 ReqID uint64
Data [][]byte SF stateFeedback
Data [][]byte
} }
resp.SF.protocolVersion = p.version
if err := msg.Decode(&resp); err != nil { if err := msg.Decode(&resp); err != nil {
return errResp(ErrDecode, "msg %v: %v", msg, err) return errResp(ErrDecode, "msg %v: %v", msg, err)
} }
p.fcServer.ReceivedReply(resp.ReqID, resp.BV) p.fcServer.ReceivedReply(resp.ReqID, resp.SF.BV)
deliverMsg = &Msg{ deliverMsg = &Msg{
MsgType: MsgCode, MsgType: MsgCode,
ReqID: resp.ReqID, ReqID: resp.ReqID,
@ -247,13 +253,15 @@ func (h *clientHandler) handleMsg(p *peer) error {
case ReceiptsMsg: case ReceiptsMsg:
p.Log().Trace("Received receipts response") p.Log().Trace("Received receipts response")
var resp struct { var resp struct {
ReqID, BV uint64 ReqID uint64
Receipts []types.Receipts SF stateFeedback
Receipts []types.Receipts
} }
resp.SF.protocolVersion = p.version
if err := msg.Decode(&resp); err != nil { if err := msg.Decode(&resp); err != nil {
return errResp(ErrDecode, "msg %v: %v", msg, err) return errResp(ErrDecode, "msg %v: %v", msg, err)
} }
p.fcServer.ReceivedReply(resp.ReqID, resp.BV) p.fcServer.ReceivedReply(resp.ReqID, resp.SF.BV)
deliverMsg = &Msg{ deliverMsg = &Msg{
MsgType: MsgReceipts, MsgType: MsgReceipts,
ReqID: resp.ReqID, ReqID: resp.ReqID,
@ -262,13 +270,15 @@ func (h *clientHandler) handleMsg(p *peer) error {
case ProofsV2Msg: case ProofsV2Msg:
p.Log().Trace("Received les/2 proofs response") p.Log().Trace("Received les/2 proofs response")
var resp struct { var resp struct {
ReqID, BV uint64 ReqID uint64
Data light.NodeList SF stateFeedback
Data light.NodeList
} }
resp.SF.protocolVersion = p.version
if err := msg.Decode(&resp); err != nil { if err := msg.Decode(&resp); err != nil {
return errResp(ErrDecode, "msg %v: %v", msg, err) return errResp(ErrDecode, "msg %v: %v", msg, err)
} }
p.fcServer.ReceivedReply(resp.ReqID, resp.BV) p.fcServer.ReceivedReply(resp.ReqID, resp.SF.BV)
deliverMsg = &Msg{ deliverMsg = &Msg{
MsgType: MsgProofsV2, MsgType: MsgProofsV2,
ReqID: resp.ReqID, ReqID: resp.ReqID,
@ -277,13 +287,15 @@ func (h *clientHandler) handleMsg(p *peer) error {
case HelperTrieProofsMsg: case HelperTrieProofsMsg:
p.Log().Trace("Received helper trie proof response") p.Log().Trace("Received helper trie proof response")
var resp struct { var resp struct {
ReqID, BV uint64 ReqID uint64
Data HelperTrieResps SF stateFeedback
Data HelperTrieResps
} }
resp.SF.protocolVersion = p.version
if err := msg.Decode(&resp); err != nil { if err := msg.Decode(&resp); err != nil {
return errResp(ErrDecode, "msg %v: %v", msg, err) return errResp(ErrDecode, "msg %v: %v", msg, err)
} }
p.fcServer.ReceivedReply(resp.ReqID, resp.BV) p.fcServer.ReceivedReply(resp.ReqID, resp.SF.BV)
deliverMsg = &Msg{ deliverMsg = &Msg{
MsgType: MsgHelperTrieProofs, MsgType: MsgHelperTrieProofs,
ReqID: resp.ReqID, ReqID: resp.ReqID,
@ -292,13 +304,15 @@ func (h *clientHandler) handleMsg(p *peer) error {
case TxStatusMsg: case TxStatusMsg:
p.Log().Trace("Received tx status response") p.Log().Trace("Received tx status response")
var resp struct { var resp struct {
ReqID, BV uint64 ReqID uint64
Status []light.TxStatus SF stateFeedback
Status []light.TxStatus
} }
resp.SF.protocolVersion = p.version
if err := msg.Decode(&resp); err != nil { if err := msg.Decode(&resp); err != nil {
return errResp(ErrDecode, "msg %v: %v", msg, err) return errResp(ErrDecode, "msg %v: %v", msg, err)
} }
p.fcServer.ReceivedReply(resp.ReqID, resp.BV) p.fcServer.ReceivedReply(resp.ReqID, resp.SF.BV)
deliverMsg = &Msg{ deliverMsg = &Msg{
MsgType: MsgTxStatus, MsgType: MsgTxStatus,
ReqID: resp.ReqID, ReqID: resp.ReqID,
@ -309,11 +323,12 @@ func (h *clientHandler) handleMsg(p *peer) error {
h.backend.retriever.frozen(p) h.backend.retriever.frozen(p)
p.Log().Debug("Service stopped") p.Log().Debug("Service stopped")
case ResumeMsg: case ResumeMsg:
var bv uint64 var sf stateFeedback
if err := msg.Decode(&bv); err != nil { sf.protocolVersion = p.version
if err := msg.Decode(&sf); err != nil {
return errResp(ErrDecode, "msg %v: %v", msg, err) return errResp(ErrDecode, "msg %v: %v", msg, err)
} }
p.fcServer.ResumeFreeze(bv) p.fcServer.ResumeFreeze(sf.BV)
p.freezeServer(false) p.freezeServer(false)
p.Log().Debug("Service resumed") p.Log().Debug("Service resumed")
case LespayReplyMsg: case LespayReplyMsg:

View file

@ -619,16 +619,17 @@ func (f *clientPool) setCapacityLocked(id enode.ID, freeID string, capacity uint
return f.setCapacity(id, freeID, capacity, minConnTime, setCap) return f.setCapacity(id, freeID, capacity, minConnTime, setCap)
} }
// requestCost feeds request cost after serving a request from the given peer. // requestCost feeds request cost after serving a request from the given peer and
func (f *clientPool) requestCost(p *peer, cost uint64) { // returns the remaining token balance
func (f *clientPool) requestCost(p *peer, cost uint64) uint64 {
f.lock.Lock() f.lock.Lock()
defer f.lock.Unlock() defer f.lock.Unlock()
info, exist := f.connectedMap[p.ID()] info, exist := f.connectedMap[p.ID()]
if !exist || f.closed { if !exist || f.closed {
return return 0
} }
info.balanceTracker.requestCost(cost) return info.balanceTracker.requestCost(cost)
} }
// logOffset calculates the time-dependent offset for the logarithmic // logOffset calculates the time-dependent offset for the logarithmic
@ -654,7 +655,12 @@ func (f *clientPool) getPosBalance(id enode.ID) posBalance {
f.lock.Lock() f.lock.Lock()
defer f.lock.Unlock() defer f.lock.Unlock()
return f.ndb.getOrNewPB(id) if c := f.connectedMap[id]; c != nil {
pb, _ := c.balanceTracker.getBalance(mclock.Now())
return posBalance{value: pb, meta: c.balanceMetaInfo}
} else {
return f.ndb.getOrNewPB(id)
}
} }
// addBalance updates the balance of a client (either overwrites it or adds to it). // addBalance updates the balance of a client (either overwrites it or adds to it).

View file

@ -38,17 +38,26 @@ import (
"github.com/ethereum/go-ethereum/trie" "github.com/ethereum/go-ethereum/trie"
) )
func expectResponse(r p2p.MsgReader, msgcode, reqID, bv uint64, data interface{}) error { func expectResponse(r p2p.MsgReader, protocol int, msgcode, reqID, bv, cost uint64, data interface{}) error {
type resp struct { type resp struct {
ReqID, BV uint64 ReqID uint64
Data interface{} SF stateFeedback
Data interface{}
} }
return p2p.ExpectMsg(r, msgcode, resp{reqID, bv, data}) sf := stateFeedback{
protocolVersion: protocol,
stateFeedbackV4: stateFeedbackV4{
BV: bv,
RealCost: cost,
TokenBalance: 0,
},
}
return p2p.ExpectMsg(r, msgcode, resp{reqID, sf, data})
} }
// Tests that block headers can be retrieved from a remote chain based on user queries. // Tests that block headers can be retrieved from a remote chain based on user queries.
func TestGetBlockHeadersLes2(t *testing.T) { testGetBlockHeaders(t, 2) }
func TestGetBlockHeadersLes3(t *testing.T) { testGetBlockHeaders(t, 3) } func TestGetBlockHeadersLes3(t *testing.T) { testGetBlockHeaders(t, 3) }
func TestGetBlockHeadersLes4(t *testing.T) { testGetBlockHeaders(t, 4) }
func testGetBlockHeaders(t *testing.T, protocol int) { func testGetBlockHeaders(t *testing.T, protocol int) {
server, tearDown := newServerEnv(t, downloader.MaxHashFetch+15, protocol, nil, false, true, 0) server, tearDown := newServerEnv(t, downloader.MaxHashFetch+15, protocol, nil, false, true, 0)
@ -170,15 +179,15 @@ func testGetBlockHeaders(t *testing.T, protocol int) {
cost := server.peer.peer.GetRequestCost(GetBlockHeadersMsg, int(tt.query.Amount)) cost := server.peer.peer.GetRequestCost(GetBlockHeadersMsg, int(tt.query.Amount))
sendRequest(server.peer.app, GetBlockHeadersMsg, reqID, cost, tt.query) sendRequest(server.peer.app, GetBlockHeadersMsg, reqID, cost, tt.query)
if err := expectResponse(server.peer.app, BlockHeadersMsg, reqID, testBufLimit, headers); err != nil { if err := expectResponse(server.peer.app, protocol, BlockHeadersMsg, reqID, testBufLimit, cost, headers); err != nil {
t.Errorf("test %d: headers mismatch: %v", i, err) t.Errorf("test %d: headers mismatch: %v", i, err)
} }
} }
} }
// Tests that block contents can be retrieved from a remote chain based on their hashes. // Tests that block contents can be retrieved from a remote chain based on their hashes.
func TestGetBlockBodiesLes2(t *testing.T) { testGetBlockBodies(t, 2) }
func TestGetBlockBodiesLes3(t *testing.T) { testGetBlockBodies(t, 3) } func TestGetBlockBodiesLes3(t *testing.T) { testGetBlockBodies(t, 3) }
func TestGetBlockBodiesLes4(t *testing.T) { testGetBlockBodies(t, 4) }
func testGetBlockBodies(t *testing.T, protocol int) { func testGetBlockBodies(t *testing.T, protocol int) {
server, tearDown := newServerEnv(t, downloader.MaxBlockFetch+15, protocol, nil, false, true, 0) server, tearDown := newServerEnv(t, downloader.MaxBlockFetch+15, protocol, nil, false, true, 0)
@ -248,15 +257,15 @@ func testGetBlockBodies(t *testing.T, protocol int) {
// Send the hash request and verify the response // Send the hash request and verify the response
cost := server.peer.peer.GetRequestCost(GetBlockBodiesMsg, len(hashes)) cost := server.peer.peer.GetRequestCost(GetBlockBodiesMsg, len(hashes))
sendRequest(server.peer.app, GetBlockBodiesMsg, reqID, cost, hashes) sendRequest(server.peer.app, GetBlockBodiesMsg, reqID, cost, hashes)
if err := expectResponse(server.peer.app, BlockBodiesMsg, reqID, testBufLimit, bodies); err != nil { if err := expectResponse(server.peer.app, protocol, BlockBodiesMsg, reqID, testBufLimit, cost, bodies); err != nil {
t.Errorf("test %d: bodies mismatch: %v", i, err) t.Errorf("test %d: bodies mismatch: %v", i, err)
} }
} }
} }
// Tests that the contract codes can be retrieved based on account addresses. // Tests that the contract codes can be retrieved based on account addresses.
func TestGetCodeLes2(t *testing.T) { testGetCode(t, 2) }
func TestGetCodeLes3(t *testing.T) { testGetCode(t, 3) } func TestGetCodeLes3(t *testing.T) { testGetCode(t, 3) }
func TestGetCodeLes4(t *testing.T) { testGetCode(t, 4) }
func testGetCode(t *testing.T, protocol int) { func testGetCode(t *testing.T, protocol int) {
// Assemble the test environment // Assemble the test environment
@ -280,14 +289,14 @@ func testGetCode(t *testing.T, protocol int) {
cost := server.peer.peer.GetRequestCost(GetCodeMsg, len(codereqs)) cost := server.peer.peer.GetRequestCost(GetCodeMsg, len(codereqs))
sendRequest(server.peer.app, GetCodeMsg, 42, cost, codereqs) sendRequest(server.peer.app, GetCodeMsg, 42, cost, codereqs)
if err := expectResponse(server.peer.app, CodeMsg, 42, testBufLimit, codes); err != nil { if err := expectResponse(server.peer.app, protocol, CodeMsg, 42, testBufLimit, cost, codes); err != nil {
t.Errorf("codes mismatch: %v", err) t.Errorf("codes mismatch: %v", err)
} }
} }
// Tests that the stale contract codes can't be retrieved based on account addresses. // Tests that the stale contract codes can't be retrieved based on account addresses.
func TestGetStaleCodeLes2(t *testing.T) { testGetStaleCode(t, 2) }
func TestGetStaleCodeLes3(t *testing.T) { testGetStaleCode(t, 3) } func TestGetStaleCodeLes3(t *testing.T) { testGetStaleCode(t, 3) }
func TestGetStaleCodeLes4(t *testing.T) { testGetStaleCode(t, 4) }
func testGetStaleCode(t *testing.T, protocol int) { func testGetStaleCode(t *testing.T, protocol int) {
server, tearDown := newServerEnv(t, core.TriesInMemory+4, protocol, nil, false, true, 0) server, tearDown := newServerEnv(t, core.TriesInMemory+4, protocol, nil, false, true, 0)
@ -301,7 +310,7 @@ func testGetStaleCode(t *testing.T, protocol int) {
} }
cost := server.peer.peer.GetRequestCost(GetCodeMsg, 1) cost := server.peer.peer.GetRequestCost(GetCodeMsg, 1)
sendRequest(server.peer.app, GetCodeMsg, 42, cost, []*CodeReq{req}) sendRequest(server.peer.app, GetCodeMsg, 42, cost, []*CodeReq{req})
if err := expectResponse(server.peer.app, CodeMsg, 42, testBufLimit, expected); err != nil { if err := expectResponse(server.peer.app, protocol, CodeMsg, 42, testBufLimit, cost, expected); err != nil {
t.Errorf("codes mismatch: %v", err) t.Errorf("codes mismatch: %v", err)
} }
} }
@ -311,8 +320,8 @@ func testGetStaleCode(t *testing.T, protocol int) {
} }
// Tests that the transaction receipts can be retrieved based on hashes. // Tests that the transaction receipts can be retrieved based on hashes.
func TestGetReceiptLes2(t *testing.T) { testGetReceipt(t, 2) }
func TestGetReceiptLes3(t *testing.T) { testGetReceipt(t, 3) } func TestGetReceiptLes3(t *testing.T) { testGetReceipt(t, 3) }
func TestGetReceiptLes4(t *testing.T) { testGetReceipt(t, 4) }
func testGetReceipt(t *testing.T, protocol int) { func testGetReceipt(t *testing.T, protocol int) {
// Assemble the test environment // Assemble the test environment
@ -333,14 +342,14 @@ func testGetReceipt(t *testing.T, protocol int) {
// Send the hash request and verify the response // Send the hash request and verify the response
cost := server.peer.peer.GetRequestCost(GetReceiptsMsg, len(hashes)) cost := server.peer.peer.GetRequestCost(GetReceiptsMsg, len(hashes))
sendRequest(server.peer.app, GetReceiptsMsg, 42, cost, hashes) sendRequest(server.peer.app, GetReceiptsMsg, 42, cost, hashes)
if err := expectResponse(server.peer.app, ReceiptsMsg, 42, testBufLimit, receipts); err != nil { if err := expectResponse(server.peer.app, protocol, ReceiptsMsg, 42, testBufLimit, cost, receipts); err != nil {
t.Errorf("receipts mismatch: %v", err) t.Errorf("receipts mismatch: %v", err)
} }
} }
// Tests that trie merkle proofs can be retrieved // Tests that trie merkle proofs can be retrieved
func TestGetProofsLes2(t *testing.T) { testGetProofs(t, 2) }
func TestGetProofsLes3(t *testing.T) { testGetProofs(t, 3) } func TestGetProofsLes3(t *testing.T) { testGetProofs(t, 3) }
func TestGetProofsLes4(t *testing.T) { testGetProofs(t, 4) }
func testGetProofs(t *testing.T, protocol int) { func testGetProofs(t *testing.T, protocol int) {
// Assemble the test environment // Assemble the test environment
@ -369,14 +378,14 @@ func testGetProofs(t *testing.T, protocol int) {
// Send the proof request and verify the response // Send the proof request and verify the response
cost := server.peer.peer.GetRequestCost(GetProofsV2Msg, len(proofreqs)) cost := server.peer.peer.GetRequestCost(GetProofsV2Msg, len(proofreqs))
sendRequest(server.peer.app, GetProofsV2Msg, 42, cost, proofreqs) sendRequest(server.peer.app, GetProofsV2Msg, 42, cost, proofreqs)
if err := expectResponse(server.peer.app, ProofsV2Msg, 42, testBufLimit, proofsV2.NodeList()); err != nil { if err := expectResponse(server.peer.app, protocol, ProofsV2Msg, 42, testBufLimit, cost, proofsV2.NodeList()); err != nil {
t.Errorf("proofs mismatch: %v", err) t.Errorf("proofs mismatch: %v", err)
} }
} }
// Tests that the stale contract codes can't be retrieved based on account addresses. // Tests that the stale contract codes can't be retrieved based on account addresses.
func TestGetStaleProofLes2(t *testing.T) { testGetStaleProof(t, 2) }
func TestGetStaleProofLes3(t *testing.T) { testGetStaleProof(t, 3) } func TestGetStaleProofLes3(t *testing.T) { testGetStaleProof(t, 3) }
func TestGetStaleProofLes4(t *testing.T) { testGetStaleProof(t, 4) }
func testGetStaleProof(t *testing.T, protocol int) { func testGetStaleProof(t *testing.T, protocol int) {
server, tearDown := newServerEnv(t, core.TriesInMemory+4, protocol, nil, false, true, 0) server, tearDown := newServerEnv(t, core.TriesInMemory+4, protocol, nil, false, true, 0)
@ -402,7 +411,7 @@ func testGetStaleProof(t *testing.T, protocol int) {
t.Prove(account, 0, proofsV2) t.Prove(account, 0, proofsV2)
expected = proofsV2.NodeList() expected = proofsV2.NodeList()
} }
if err := expectResponse(server.peer.app, ProofsV2Msg, 42, testBufLimit, expected); err != nil { if err := expectResponse(server.peer.app, protocol, ProofsV2Msg, 42, testBufLimit, cost, expected); err != nil {
t.Errorf("codes mismatch: %v", err) t.Errorf("codes mismatch: %v", err)
} }
} }
@ -412,8 +421,8 @@ func testGetStaleProof(t *testing.T, protocol int) {
} }
// Tests that CHT proofs can be correctly retrieved. // Tests that CHT proofs can be correctly retrieved.
func TestGetCHTProofsLes2(t *testing.T) { testGetCHTProofs(t, 2) }
func TestGetCHTProofsLes3(t *testing.T) { testGetCHTProofs(t, 3) } func TestGetCHTProofsLes3(t *testing.T) { testGetCHTProofs(t, 3) }
func TestGetCHTProofsLes4(t *testing.T) { testGetCHTProofs(t, 4) }
func testGetCHTProofs(t *testing.T, protocol int) { func testGetCHTProofs(t *testing.T, protocol int) {
config := light.TestServerIndexerConfig config := light.TestServerIndexerConfig
@ -455,13 +464,13 @@ func testGetCHTProofs(t *testing.T, protocol int) {
// Send the proof request and verify the response // Send the proof request and verify the response
cost := server.peer.peer.GetRequestCost(GetHelperTrieProofsMsg, len(requestsV2)) cost := server.peer.peer.GetRequestCost(GetHelperTrieProofsMsg, len(requestsV2))
sendRequest(server.peer.app, GetHelperTrieProofsMsg, 42, cost, requestsV2) sendRequest(server.peer.app, GetHelperTrieProofsMsg, 42, cost, requestsV2)
if err := expectResponse(server.peer.app, HelperTrieProofsMsg, 42, testBufLimit, proofsV2); err != nil { if err := expectResponse(server.peer.app, protocol, HelperTrieProofsMsg, 42, testBufLimit, cost, proofsV2); err != nil {
t.Errorf("proofs mismatch: %v", err) t.Errorf("proofs mismatch: %v", err)
} }
} }
func TestGetBloombitsProofsLes2(t *testing.T) { testGetBloombitsProofs(t, 2) }
func TestGetBloombitsProofsLes3(t *testing.T) { testGetBloombitsProofs(t, 3) } func TestGetBloombitsProofsLes3(t *testing.T) { testGetBloombitsProofs(t, 3) }
func TestGetBloombitsProofsLes4(t *testing.T) { testGetBloombitsProofs(t, 4) }
// Tests that bloombits proofs can be correctly retrieved. // Tests that bloombits proofs can be correctly retrieved.
func testGetBloombitsProofs(t *testing.T, protocol int) { func testGetBloombitsProofs(t *testing.T, protocol int) {
@ -504,14 +513,14 @@ func testGetBloombitsProofs(t *testing.T, protocol int) {
// Send the proof request and verify the response // Send the proof request and verify the response
cost := server.peer.peer.GetRequestCost(GetHelperTrieProofsMsg, len(requests)) cost := server.peer.peer.GetRequestCost(GetHelperTrieProofsMsg, len(requests))
sendRequest(server.peer.app, GetHelperTrieProofsMsg, 42, cost, requests) sendRequest(server.peer.app, GetHelperTrieProofsMsg, 42, cost, requests)
if err := expectResponse(server.peer.app, HelperTrieProofsMsg, 42, testBufLimit, proofs); err != nil { if err := expectResponse(server.peer.app, protocol, HelperTrieProofsMsg, 42, testBufLimit, cost, proofs); err != nil {
t.Errorf("bit %d: proofs mismatch: %v", bit, err) t.Errorf("bit %d: proofs mismatch: %v", bit, err)
} }
} }
} }
func TestTransactionStatusLes2(t *testing.T) { testTransactionStatus(t, 2) }
func TestTransactionStatusLes3(t *testing.T) { testTransactionStatus(t, 3) } func TestTransactionStatusLes3(t *testing.T) { testTransactionStatus(t, 3) }
func TestTransactionStatusLes4(t *testing.T) { testTransactionStatus(t, 4) }
func testTransactionStatus(t *testing.T, protocol int) { func testTransactionStatus(t *testing.T, protocol int) {
server, tearDown := newServerEnv(t, 0, protocol, nil, false, true, 0) server, tearDown := newServerEnv(t, 0, protocol, nil, false, true, 0)
@ -524,14 +533,15 @@ func testTransactionStatus(t *testing.T, protocol int) {
test := func(tx *types.Transaction, send bool, expStatus light.TxStatus) { test := func(tx *types.Transaction, send bool, expStatus light.TxStatus) {
reqID++ reqID++
var cost uint64
if send { if send {
cost := server.peer.peer.GetRequestCost(SendTxV2Msg, 1) cost = server.peer.peer.GetRequestCost(SendTxV2Msg, 1)
sendRequest(server.peer.app, SendTxV2Msg, reqID, cost, types.Transactions{tx}) sendRequest(server.peer.app, SendTxV2Msg, reqID, cost, types.Transactions{tx})
} else { } else {
cost := server.peer.peer.GetRequestCost(GetTxStatusMsg, 1) cost = server.peer.peer.GetRequestCost(GetTxStatusMsg, 1)
sendRequest(server.peer.app, GetTxStatusMsg, reqID, cost, []common.Hash{tx.Hash()}) sendRequest(server.peer.app, GetTxStatusMsg, reqID, cost, []common.Hash{tx.Hash()})
} }
if err := expectResponse(server.peer.app, TxStatusMsg, reqID, testBufLimit, []light.TxStatus{expStatus}); err != nil { if err := expectResponse(server.peer.app, protocol, TxStatusMsg, reqID, testBufLimit, cost, []light.TxStatus{expStatus}); err != nil {
t.Errorf("transaction status mismatch") t.Errorf("transaction status mismatch")
} }
} }
@ -606,8 +616,11 @@ func testTransactionStatus(t *testing.T, protocol int) {
test(tx2, false, light.TxStatus{Status: core.TxStatusPending}) test(tx2, false, light.TxStatus{Status: core.TxStatusPending})
} }
func TestStopResumeLes3(t *testing.T) { func TestStopResumeLes3(t *testing.T) { testStopResume(t, 3) }
server, tearDown := newServerEnv(t, 0, 3, nil, true, true, testBufLimit/10) func TestStopResumeLes4(t *testing.T) { testStopResume(t, 4) }
func testStopResume(t *testing.T, protocol int) {
server, tearDown := newServerEnv(t, 0, protocol, nil, true, true, testBufLimit/10)
defer tearDown() defer tearDown()
server.handler.server.costTracker.testing = true server.handler.server.costTracker.testing = true
@ -627,7 +640,7 @@ func TestStopResumeLes3(t *testing.T) {
for expBuf >= testCost { for expBuf >= testCost {
req() req()
expBuf -= testCost expBuf -= testCost
if err := expectResponse(server.peer.app, BlockHeadersMsg, reqID, expBuf, []*types.Header{header}); err != nil { if err := expectResponse(server.peer.app, protocol, BlockHeadersMsg, reqID, expBuf, testCost, []*types.Header{header}); err != nil {
t.Errorf("expected response and failed: %v", err) t.Errorf("expected response and failed: %v", err)
} }
} }
@ -646,7 +659,15 @@ func TestStopResumeLes3(t *testing.T) {
// expect a ResumeMsg with the partially recharged buffer value // expect a ResumeMsg with the partially recharged buffer value
expBuf += testBufRecharge * wait expBuf += testBufRecharge * wait
if err := p2p.ExpectMsg(server.peer.app, ResumeMsg, expBuf); err != nil { sf := stateFeedback{
protocolVersion: protocol,
stateFeedbackV4: stateFeedbackV4{
BV: expBuf,
RealCost: 0,
TokenBalance: 0,
},
}
if err := p2p.ExpectMsg(server.peer.app, ResumeMsg, sf); err != nil {
t.Errorf("expected ResumeMsg and failed: %v", err) t.Errorf("expected ResumeMsg and failed: %v", err)
} }
} }

View file

@ -38,8 +38,8 @@ import (
type odrTestFn func(ctx context.Context, db ethdb.Database, config *params.ChainConfig, bc *core.BlockChain, lc *light.LightChain, bhash common.Hash) []byte type odrTestFn func(ctx context.Context, db ethdb.Database, config *params.ChainConfig, bc *core.BlockChain, lc *light.LightChain, bhash common.Hash) []byte
func TestOdrGetBlockLes2(t *testing.T) { testOdr(t, 2, 1, true, odrGetBlock) }
func TestOdrGetBlockLes3(t *testing.T) { testOdr(t, 3, 1, true, odrGetBlock) } func TestOdrGetBlockLes3(t *testing.T) { testOdr(t, 3, 1, true, odrGetBlock) }
func TestOdrGetBlockLes4(t *testing.T) { testOdr(t, 4, 1, true, odrGetBlock) }
func odrGetBlock(ctx context.Context, db ethdb.Database, config *params.ChainConfig, bc *core.BlockChain, lc *light.LightChain, bhash common.Hash) []byte { func odrGetBlock(ctx context.Context, db ethdb.Database, config *params.ChainConfig, bc *core.BlockChain, lc *light.LightChain, bhash common.Hash) []byte {
var block *types.Block var block *types.Block
@ -55,8 +55,8 @@ func odrGetBlock(ctx context.Context, db ethdb.Database, config *params.ChainCon
return rlp return rlp
} }
func TestOdrGetReceiptsLes2(t *testing.T) { testOdr(t, 2, 1, true, odrGetReceipts) }
func TestOdrGetReceiptsLes3(t *testing.T) { testOdr(t, 3, 1, true, odrGetReceipts) } func TestOdrGetReceiptsLes3(t *testing.T) { testOdr(t, 3, 1, true, odrGetReceipts) }
func TestOdrGetReceiptsLes4(t *testing.T) { testOdr(t, 4, 1, true, odrGetReceipts) }
func odrGetReceipts(ctx context.Context, db ethdb.Database, config *params.ChainConfig, bc *core.BlockChain, lc *light.LightChain, bhash common.Hash) []byte { func odrGetReceipts(ctx context.Context, db ethdb.Database, config *params.ChainConfig, bc *core.BlockChain, lc *light.LightChain, bhash common.Hash) []byte {
var receipts types.Receipts var receipts types.Receipts
@ -76,8 +76,8 @@ func odrGetReceipts(ctx context.Context, db ethdb.Database, config *params.Chain
return rlp return rlp
} }
func TestOdrAccountsLes2(t *testing.T) { testOdr(t, 2, 1, true, odrAccounts) }
func TestOdrAccountsLes3(t *testing.T) { testOdr(t, 3, 1, true, odrAccounts) } func TestOdrAccountsLes3(t *testing.T) { testOdr(t, 3, 1, true, odrAccounts) }
func TestOdrAccountsLes4(t *testing.T) { testOdr(t, 4, 1, true, odrAccounts) }
func odrAccounts(ctx context.Context, db ethdb.Database, config *params.ChainConfig, bc *core.BlockChain, lc *light.LightChain, bhash common.Hash) []byte { func odrAccounts(ctx context.Context, db ethdb.Database, config *params.ChainConfig, bc *core.BlockChain, lc *light.LightChain, bhash common.Hash) []byte {
dummyAddr := common.HexToAddress("1234567812345678123456781234567812345678") dummyAddr := common.HexToAddress("1234567812345678123456781234567812345678")
@ -105,8 +105,8 @@ func odrAccounts(ctx context.Context, db ethdb.Database, config *params.ChainCon
return res return res
} }
func TestOdrContractCallLes2(t *testing.T) { testOdr(t, 2, 2, true, odrContractCall) }
func TestOdrContractCallLes3(t *testing.T) { testOdr(t, 3, 2, true, odrContractCall) } func TestOdrContractCallLes3(t *testing.T) { testOdr(t, 3, 2, true, odrContractCall) }
func TestOdrContractCallLes4(t *testing.T) { testOdr(t, 4, 2, true, odrContractCall) }
type callmsg struct { type callmsg struct {
types.Message types.Message
@ -155,8 +155,8 @@ func odrContractCall(ctx context.Context, db ethdb.Database, config *params.Chai
return res return res
} }
func TestOdrTxStatusLes2(t *testing.T) { testOdr(t, 2, 1, false, odrTxStatus) }
func TestOdrTxStatusLes3(t *testing.T) { testOdr(t, 3, 1, false, odrTxStatus) } func TestOdrTxStatusLes3(t *testing.T) { testOdr(t, 3, 1, false, odrTxStatus) }
func TestOdrTxStatusLes4(t *testing.T) { testOdr(t, 4, 1, false, odrTxStatus) }
func odrTxStatus(ctx context.Context, db ethdb.Database, config *params.ChainConfig, bc *core.BlockChain, lc *light.LightChain, bhash common.Hash) []byte { func odrTxStatus(ctx context.Context, db ethdb.Database, config *params.ChainConfig, bc *core.BlockChain, lc *light.LightChain, bhash common.Hash) []byte {
var txs types.Transactions var txs types.Transactions

View file

@ -113,6 +113,8 @@ type peer struct {
fcParams flowcontrol.ServerParams fcParams flowcontrol.ServerParams
fcCosts requestCostTable fcCosts requestCostTable
getBalance func() posBalance
trusted, server bool trusted, server bool
onlyAnnounce bool onlyAnnounce bool
chainSince, chainRecent uint64 chainSince, chainRecent uint64
@ -198,7 +200,19 @@ func (p *peer) freezeClient() {
time.Sleep(freezeCheckPeriod) time.Sleep(freezeCheckPeriod)
} else { } else {
atomic.StoreUint32(&p.frozen, 0) atomic.StoreUint32(&p.frozen, 0)
p.SendResume(bufValue) var balance uint64
if p.getBalance != nil {
balance = p.getBalance().value
}
sf := stateFeedback{
protocolVersion: p.version,
stateFeedbackV4: stateFeedbackV4{
BV: bufValue,
RealCost: 0,
TokenBalance: balance,
},
}
p.SendResume(sf)
break break
} }
} }
@ -314,12 +328,13 @@ type reply struct {
} }
// send sends the reply with the calculated buffer value // send sends the reply with the calculated buffer value
func (r *reply) send(bv uint64) error { func (r *reply) send(sf stateFeedback) error {
type resp struct { type resp struct {
ReqID, BV uint64 ReqID uint64
Data rlp.RawValue SF stateFeedback
Data rlp.RawValue
} }
return p2p.Send(r.w, r.msgcode, resp{r.reqID, bv, r.data}) return p2p.Send(r.w, r.msgcode, resp{r.reqID, sf, r.data})
} }
// size returns the RLP encoded size of the message data // size returns the RLP encoded size of the message data
@ -394,8 +409,8 @@ func (p *peer) SendStop() error {
} }
// SendResume notifies the client about getting out of frozen state // SendResume notifies the client about getting out of frozen state
func (p *peer) SendResume(bv uint64) error { func (p *peer) SendResume(sf stateFeedback) error {
return p2p.Send(p.rw, ResumeMsg, bv) return p2p.Send(p.rw, ResumeMsg, sf)
} }
// ReplyBlockHeaders creates a reply with a batch of block headers // ReplyBlockHeaders creates a reply with a batch of block headers

View file

@ -239,3 +239,28 @@ func (hn *hashOrNumber) DecodeRLP(s *rlp.Stream) error {
type CodeData []struct { type CodeData []struct {
Value []byte Value []byte
} }
type stateFeedbackV4 struct {
BV, RealCost, TokenBalance uint64
}
type stateFeedback struct {
protocolVersion int
stateFeedbackV4
}
func (sf stateFeedback) EncodeRLP(w io.Writer) error {
if sf.protocolVersion >= lpv4 {
return rlp.Encode(w, sf.stateFeedbackV4)
} else {
return rlp.Encode(w, sf.BV)
}
}
func (sf *stateFeedback) DecodeRLP(s *rlp.Stream) error {
if sf.protocolVersion >= lpv4 {
return s.Decode(&sf.stateFeedbackV4)
} else {
return s.Decode(&sf.BV)
}
}

View file

@ -36,22 +36,22 @@ func secAddr(addr common.Address) []byte {
type accessTestFn func(db ethdb.Database, bhash common.Hash, number uint64) light.OdrRequest type accessTestFn func(db ethdb.Database, bhash common.Hash, number uint64) light.OdrRequest
func TestBlockAccessLes2(t *testing.T) { testAccess(t, 2, tfBlockAccess) }
func TestBlockAccessLes3(t *testing.T) { testAccess(t, 3, tfBlockAccess) } func TestBlockAccessLes3(t *testing.T) { testAccess(t, 3, tfBlockAccess) }
func TestBlockAccessLes4(t *testing.T) { testAccess(t, 4, tfBlockAccess) }
func tfBlockAccess(db ethdb.Database, bhash common.Hash, number uint64) light.OdrRequest { func tfBlockAccess(db ethdb.Database, bhash common.Hash, number uint64) light.OdrRequest {
return &light.BlockRequest{Hash: bhash, Number: number} return &light.BlockRequest{Hash: bhash, Number: number}
} }
func TestReceiptsAccessLes2(t *testing.T) { testAccess(t, 2, tfReceiptsAccess) }
func TestReceiptsAccessLes3(t *testing.T) { testAccess(t, 3, tfReceiptsAccess) } func TestReceiptsAccessLes3(t *testing.T) { testAccess(t, 3, tfReceiptsAccess) }
func TestReceiptsAccessLes4(t *testing.T) { testAccess(t, 4, tfReceiptsAccess) }
func tfReceiptsAccess(db ethdb.Database, bhash common.Hash, number uint64) light.OdrRequest { func tfReceiptsAccess(db ethdb.Database, bhash common.Hash, number uint64) light.OdrRequest {
return &light.ReceiptsRequest{Hash: bhash, Number: number} return &light.ReceiptsRequest{Hash: bhash, Number: number}
} }
func TestTrieEntryAccessLes2(t *testing.T) { testAccess(t, 2, tfTrieEntryAccess) }
func TestTrieEntryAccessLes3(t *testing.T) { testAccess(t, 3, tfTrieEntryAccess) } func TestTrieEntryAccessLes3(t *testing.T) { testAccess(t, 3, tfTrieEntryAccess) }
func TestTrieEntryAccessLes4(t *testing.T) { testAccess(t, 4, tfTrieEntryAccess) }
func tfTrieEntryAccess(db ethdb.Database, bhash common.Hash, number uint64) light.OdrRequest { func tfTrieEntryAccess(db ethdb.Database, bhash common.Hash, number uint64) light.OdrRequest {
if number := rawdb.ReadHeaderNumber(db, bhash); number != nil { if number := rawdb.ReadHeaderNumber(db, bhash); number != nil {
@ -60,8 +60,8 @@ func tfTrieEntryAccess(db ethdb.Database, bhash common.Hash, number uint64) ligh
return nil return nil
} }
func TestCodeAccessLes2(t *testing.T) { testAccess(t, 2, tfCodeAccess) }
func TestCodeAccessLes3(t *testing.T) { testAccess(t, 3, tfCodeAccess) } func TestCodeAccessLes3(t *testing.T) { testAccess(t, 3, tfCodeAccess) }
func TestCodeAccessLes4(t *testing.T) { testAccess(t, 4, tfCodeAccess) }
func tfCodeAccess(db ethdb.Database, bhash common.Hash, num uint64) light.OdrRequest { func tfCodeAccess(db ethdb.Database, bhash common.Hash, num uint64) light.OdrRequest {
number := rawdb.ReadHeaderNumber(db, bhash) number := rawdb.ReadHeaderNumber(db, bhash)

View file

@ -106,6 +106,9 @@ func (h *serverHandler) stop() {
// runPeer is the p2p protocol run function for the given version. // runPeer is the p2p protocol run function for the given version.
func (h *serverHandler) runPeer(version uint, p *p2p.Peer, rw p2p.MsgReadWriter) error { func (h *serverHandler) runPeer(version uint, p *p2p.Peer, rw p2p.MsgReadWriter) error {
peer := newPeer(int(version), h.server.config.NetworkId, false, p, newMeteredMsgWriter(rw, int(version))) peer := newPeer(int(version), h.server.config.NetworkId, false, p, newMeteredMsgWriter(rw, int(version)))
peer.getBalance = func() posBalance {
return h.server.clientPool.getPosBalance(p.ID())
}
h.wg.Add(1) h.wg.Add(1)
defer h.wg.Done() defer h.wg.Done()
return h.handle(peer) return h.handle(peer)
@ -247,22 +250,33 @@ func (h *serverHandler) handleMsg(p *peer, wg *sync.WaitGroup) error {
if reply != nil { if reply != nil {
replySize = reply.size() replySize = reply.size()
} }
var realCost uint64 var realCost, balance uint64
if h.server.costTracker.testing { if h.server.costTracker.testing {
realCost = maxCost // Assign a fake cost for testing purpose realCost = maxCost // Assign a fake cost for testing purpose
} else { } else {
realCost = h.server.costTracker.realCost(servingTime, msg.Size, replySize) realCost = h.server.costTracker.realCost(servingTime, msg.Size, replySize)
if realCost > maxCost {
realCost = maxCost
}
} }
bv := p.fcClient.RequestProcessed(reqID, responseCount, maxCost, realCost) bv := p.fcClient.RequestProcessed(reqID, responseCount, maxCost, realCost)
if amount != 0 { if amount != 0 {
// Feed cost tracker request serving statistic. // Feed cost tracker request serving statistic.
h.server.costTracker.updateStats(msg.Code, amount, servingTime, realCost) h.server.costTracker.updateStats(msg.Code, amount, servingTime, realCost)
// Reduce priority "balance" for the specific peer. // Reduce priority "balance" for the specific peer.
h.server.clientPool.requestCost(p, realCost) balance = h.server.clientPool.requestCost(p, realCost)
}
sf := stateFeedback{
protocolVersion: p.version,
stateFeedbackV4: stateFeedbackV4{
BV: bv,
RealCost: realCost,
TokenBalance: balance,
},
} }
if reply != nil { if reply != nil {
p.queueSend(func() { p.queueSend(func() {
if err := reply.send(bv); err != nil { if err := reply.send(sf); err != nil {
select { select {
case p.errCh <- err: case p.errCh <- err:
default: default:

View file

@ -28,8 +28,8 @@ import (
"github.com/ethereum/go-ethereum/p2p/enode" "github.com/ethereum/go-ethereum/p2p/enode"
) )
func TestULCAnnounceThresholdLes2(t *testing.T) { testULCAnnounceThreshold(t, 2) }
func TestULCAnnounceThresholdLes3(t *testing.T) { testULCAnnounceThreshold(t, 3) } func TestULCAnnounceThresholdLes3(t *testing.T) { testULCAnnounceThreshold(t, 3) }
func TestULCAnnounceThresholdLes4(t *testing.T) { testULCAnnounceThreshold(t, 4) }
func testULCAnnounceThreshold(t *testing.T, protocol int) { func testULCAnnounceThreshold(t *testing.T, protocol int) {
// todo figure out why it takes fetcher so longer to fetcher the announced header. // todo figure out why it takes fetcher so longer to fetcher the announced header.