les: implement les/4 protocol extensions

This commit is contained in:
Zsolt Felfoldi 2020-01-03 20:51:45 +01:00
parent 55435e4ba2
commit 46fa5c031b
17 changed files with 694 additions and 200 deletions

View file

@ -101,6 +101,7 @@ var (
utils.UltraLightServersFlag, utils.UltraLightServersFlag,
utils.UltraLightFractionFlag, utils.UltraLightFractionFlag,
utils.UltraLightOnlyAnnounceFlag, utils.UltraLightOnlyAnnounceFlag,
utils.LespayTestModuleFlag,
utils.WhitelistFlag, utils.WhitelistFlag,
utils.CacheFlag, utils.CacheFlag,
utils.CacheDatabaseFlag, utils.CacheDatabaseFlag,

View file

@ -94,6 +94,7 @@ var AppHelpFlagGroups = []flagGroup{
utils.UltraLightServersFlag, utils.UltraLightServersFlag,
utils.UltraLightFractionFlag, utils.UltraLightFractionFlag,
utils.UltraLightOnlyAnnounceFlag, utils.UltraLightOnlyAnnounceFlag,
utils.LespayTestModuleFlag,
}, },
}, },
{ {

View file

@ -272,6 +272,10 @@ var (
Usage: "Maximum number of light clients to serve, or light servers to attach to", Usage: "Maximum number of light clients to serve, or light servers to attach to",
Value: eth.DefaultConfig.LightPeers, Value: eth.DefaultConfig.LightPeers,
} }
LespayTestModuleFlag = cli.BoolFlag{
Name: "lespay.testmodule",
Usage: "Enable dummy payment module (for testing only)",
}
UltraLightServersFlag = cli.StringFlag{ UltraLightServersFlag = cli.StringFlag{
Name: "ulc.servers", Name: "ulc.servers",
Usage: "List of trusted ultra-light servers", Usage: "List of trusted ultra-light servers",
@ -1009,6 +1013,9 @@ func setLes(ctx *cli.Context, cfg *eth.Config) {
if ctx.GlobalIsSet(UltraLightOnlyAnnounceFlag.Name) { if ctx.GlobalIsSet(UltraLightOnlyAnnounceFlag.Name) {
cfg.UltraLightOnlyAnnounce = ctx.GlobalBool(UltraLightOnlyAnnounceFlag.Name) cfg.UltraLightOnlyAnnounce = ctx.GlobalBool(UltraLightOnlyAnnounceFlag.Name)
} }
if ctx.GlobalIsSet(LespayTestModuleFlag.Name) {
cfg.LespayTestModule = true
}
} }
// makeDatabaseHandles raises out the number of allowed file handles per process // makeDatabaseHandles raises out the number of allowed file handles per process

View file

@ -116,6 +116,9 @@ type Config struct {
UltraLightFraction int `toml:",omitempty"` // Percentage of trusted servers to accept an announcement UltraLightFraction int `toml:",omitempty"` // Percentage of trusted servers to accept an announcement
UltraLightOnlyAnnounce bool `toml:",omitempty"` // Whether to only announce headers, or also serve them UltraLightOnlyAnnounce bool `toml:",omitempty"` // Whether to only announce headers, or also serve them
// Light client payment options
LespayTestModule bool
// Database options // Database options
SkipBcVersionCheck bool `toml:"-"` SkipBcVersionCheck bool `toml:"-"`
DatabaseHandles int `toml:"-"` DatabaseHandles int `toml:"-"`

View file

@ -32,6 +32,7 @@ func (c Config) MarshalTOML() (interface{}, error) {
UltraLightServers []string `toml:",omitempty"` UltraLightServers []string `toml:",omitempty"`
UltraLightFraction int `toml:",omitempty"` UltraLightFraction int `toml:",omitempty"`
UltraLightOnlyAnnounce bool `toml:",omitempty"` UltraLightOnlyAnnounce bool `toml:",omitempty"`
LespayTestModule bool `toml:"-"`
SkipBcVersionCheck bool `toml:"-"` SkipBcVersionCheck bool `toml:"-"`
DatabaseHandles int `toml:"-"` DatabaseHandles int `toml:"-"`
DatabaseCache int DatabaseCache int
@ -68,6 +69,7 @@ func (c Config) MarshalTOML() (interface{}, error) {
enc.UltraLightServers = c.UltraLightServers enc.UltraLightServers = c.UltraLightServers
enc.UltraLightFraction = c.UltraLightFraction enc.UltraLightFraction = c.UltraLightFraction
enc.UltraLightOnlyAnnounce = c.UltraLightOnlyAnnounce enc.UltraLightOnlyAnnounce = c.UltraLightOnlyAnnounce
enc.LespayTestModule = c.LespayTestModule
enc.SkipBcVersionCheck = c.SkipBcVersionCheck enc.SkipBcVersionCheck = c.SkipBcVersionCheck
enc.DatabaseHandles = c.DatabaseHandles enc.DatabaseHandles = c.DatabaseHandles
enc.DatabaseCache = c.DatabaseCache enc.DatabaseCache = c.DatabaseCache
@ -108,6 +110,7 @@ func (c *Config) UnmarshalTOML(unmarshal func(interface{}) error) error {
UltraLightServers []string `toml:",omitempty"` UltraLightServers []string `toml:",omitempty"`
UltraLightFraction *int `toml:",omitempty"` UltraLightFraction *int `toml:",omitempty"`
UltraLightOnlyAnnounce *bool `toml:",omitempty"` UltraLightOnlyAnnounce *bool `toml:",omitempty"`
LespayTestModule *bool `toml:"-"`
SkipBcVersionCheck *bool `toml:"-"` SkipBcVersionCheck *bool `toml:"-"`
DatabaseHandles *int `toml:"-"` DatabaseHandles *int `toml:"-"`
DatabaseCache *int DatabaseCache *int
@ -175,6 +178,9 @@ func (c *Config) UnmarshalTOML(unmarshal func(interface{}) error) error {
if dec.UltraLightOnlyAnnounce != nil { if dec.UltraLightOnlyAnnounce != nil {
c.UltraLightOnlyAnnounce = *dec.UltraLightOnlyAnnounce c.UltraLightOnlyAnnounce = *dec.UltraLightOnlyAnnounce
} }
if dec.LespayTestModule != nil {
c.LespayTestModule = *dec.LespayTestModule
}
if dec.SkipBcVersionCheck != nil { if dec.SkipBcVersionCheck != nil {
c.SkipBcVersionCheck = *dec.SkipBcVersionCheck c.SkipBcVersionCheck = *dec.SkipBcVersionCheck
} }

View file

@ -40,6 +40,9 @@ type clientHandler struct {
downloader *downloader.Downloader downloader *downloader.Downloader
backend *LightEthereum backend *LightEthereum
lespayReplyHandlers map[uint64]func([]byte, uint) bool
lespayReplyLock sync.Mutex
closeCh chan struct{} closeCh chan struct{}
wg sync.WaitGroup // WaitGroup used to track all connected peers. wg sync.WaitGroup // WaitGroup used to track all connected peers.
syncDone func() // Test hooks when syncing is done. syncDone func() // Test hooks when syncing is done.
@ -50,6 +53,7 @@ func newClientHandler(ulcServers []string, ulcFraction int, checkpoint *params.T
checkpoint: checkpoint, checkpoint: checkpoint,
backend: backend, backend: backend,
closeCh: make(chan struct{}), closeCh: make(chan struct{}),
lespayReplyHandlers: make(map[uint64]func([]byte, uint) bool),
} }
if ulcServers != nil { if ulcServers != nil {
ulc, err := newULC(ulcServers, ulcFraction) ulc, err := newULC(ulcServers, ulcFraction)
@ -112,28 +116,48 @@ func (h *clientHandler) handle(p *serverPeer) error {
p.Log().Debug("Light Ethereum handshake failed", "err", err) p.Log().Debug("Light Ethereum handshake failed", "err", err)
return err return err
} }
var (
connectedAt mclock.AbsTime
lastActive bool
)
activate := func() {
// Register the peer locally // Register the peer locally
if err := h.backend.peers.register(p); err != nil { if err := h.backend.peers.register(p); err != nil {
p.Log().Error("Light Ethereum peer registration failed", "err", err) p.Log().Error("Light Ethereum peer registration failed", "err", err)
return err return
} }
serverConnectionGauge.Update(int64(h.backend.peers.len())) serverConnectionGauge.Update(int64(h.backend.peers.len()))
connectedAt = mclock.Now()
connectedAt := mclock.Now() h.fetcher.announce(p, &announceData{Hash: p.headInfo.Hash, Number: p.headInfo.Number, Td: p.headInfo.Td})
defer func() { lastActive = true
h.backend.peers.unregister(p.id) }
deactivate := func() {
h.backend.peers.unregister(p)
connectionTimer.Update(time.Duration(mclock.Now() - connectedAt)) connectionTimer.Update(time.Duration(mclock.Now() - connectedAt))
serverConnectionGauge.Update(int64(h.backend.peers.len())) serverConnectionGauge.Update(int64(h.backend.peers.len()))
lastActive = false
}
defer func() {
if lastActive {
deactivate()
}
h.backend.peers.disconnect(p.id)
}() }()
h.fetcher.announce(p, &announceData{Hash: p.headInfo.Hash, Number: p.headInfo.Number, Td: p.headInfo.Td})
// pool entry can be nil during the unit test. // pool entry can be nil during the unit test.
if p.poolEntry != nil { if p.poolEntry != nil {
h.backend.serverPool.registered(p.poolEntry) h.backend.serverPool.registered(p.poolEntry)
} }
// Spawn a main loop to handle all incoming messages. // Spawn a main loop to handle all incoming messages.
for { for {
if p.active && !lastActive {
activate()
}
if !p.active && lastActive {
deactivate()
}
if err := h.handleMsg(p); err != nil { if err := h.handleMsg(p); err != nil {
p.Log().Debug("Light Ethereum message handling failed", "err", err) p.Log().Debug("Light Ethereum message handling failed", "err", err)
p.fcServer.DumpLogs() p.fcServer.DumpLogs()
@ -157,7 +181,10 @@ func (h *clientHandler) handleMsg(p *serverPeer) error {
} }
defer msg.Discard() defer msg.Discard()
var deliverMsg *Msg var (
deliverMsg *Msg
responseError bool
)
// Handle the message depending on its contents // Handle the message depending on its contents
switch msg.Code { switch msg.Code {
@ -193,13 +220,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) 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
SF stateFeedback
Headers []*types.Header 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 {
@ -210,13 +239,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) 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
SF stateFeedback
Data []*types.Body 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,
@ -225,13 +256,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) 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
SF stateFeedback
Data [][]byte 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,
@ -240,13 +273,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) 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
SF stateFeedback
Receipts []types.Receipts 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,
@ -255,13 +290,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) 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
SF stateFeedback
Data light.NodeList 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,
@ -270,13 +307,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) 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
SF stateFeedback
Data HelperTrieResps 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,
@ -285,13 +324,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) 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
SF stateFeedback
Status []light.TxStatus 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,
@ -302,13 +343,32 @@ func (h *clientHandler) handleMsg(p *serverPeer) 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.unfreeze() p.unfreeze()
p.Log().Debug("Service resumed") p.Log().Debug("Service resumed")
case LespayReplyMsg:
p.Log().Trace("Received tx status response")
var resp struct {
ReqID uint64
Reply lespayReply
}
if err := msg.Decode(&resp); err != nil {
return errResp(ErrDecode, "msg %v: %v", msg, err)
}
h.lespayReplyLock.Lock()
if handler := h.lespayReplyHandlers[resp.ReqID]; handler != nil {
delete(h.lespayReplyHandlers, resp.ReqID)
responseError = !handler(resp.Reply.Reply, resp.Reply.Delay)
} else {
responseError = true
}
h.lespayReplyLock.Unlock()
default: default:
p.Log().Trace("Received invalid message", "code", msg.Code) p.Log().Trace("Received invalid message", "code", msg.Code)
return errResp(ErrInvalidMsgCode, "%v", msg.Code) return errResp(ErrInvalidMsgCode, "%v", msg.Code)
@ -316,17 +376,48 @@ func (h *clientHandler) handleMsg(p *serverPeer) error {
// Deliver the received response to retriever. // Deliver the received response to retriever.
if deliverMsg != nil { if deliverMsg != nil {
if err := h.backend.retriever.deliver(p, deliverMsg); err != nil { if err := h.backend.retriever.deliver(p, deliverMsg); err != nil {
responseError = true
}
}
if responseError {
p.errCount++ p.errCount++
if p.errCount > maxResponseErrors { if p.errCount > maxResponseErrors {
return err return err
} }
} }
}
return nil return nil
} }
// makeLespayCall sends a lespay command through an LES connection and registers
// a response handler. It returns a cancel function that removes the response
// handler and calls it with a nil parameter if the response has not arrived yet.
func (h *clientHandler) makeLespayCall(p *peer, cmd []byte, handler func([]byte, uint) bool) func() bool {
reqID := genReqID()
h.lespayReplyLock.Lock()
h.lespayReplyHandlers[reqID] = handler
h.lespayReplyLock.Unlock()
if p.SendLespay(reqID, cmd) != nil {
h.lespayReplyLock.Lock()
delete(h.lespayReplyHandlers, reqID)
h.lespayReplyLock.Unlock()
return nil
}
return func() bool {
h.lespayReplyLock.Lock()
cancel := h.lespayReplyHandlers[reqID] != nil
if cancel {
delete(h.lespayReplyHandlers, reqID)
}
h.lespayReplyLock.Unlock()
if cancel {
handler(nil, 0)
}
return cancel
}
}
func (h *clientHandler) removePeer(id string) { func (h *clientHandler) removePeer(id string) {
h.backend.peers.unregister(id) h.backend.peers.disconnect(id)
} }
type peerConnection struct { type peerConnection struct {

View file

@ -38,20 +38,29 @@ 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
SF stateFeedback
Data interface{} 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, true)
defer tearDown() defer tearDown()
bc := server.handler.blockchain bc := server.handler.blockchain
@ -168,19 +177,20 @@ func testGetBlockHeaders(t *testing.T, protocol int) {
// Send the hash request and verify the response // Send the hash request and verify the response
reqID++ reqID++
cost := server.peer.peer.GetRequestCost(GetBlockHeadersMsg, int(tt.query.Amount))
sendRequest(server.peer.app, GetBlockHeadersMsg, reqID, tt.query) sendRequest(server.peer.app, GetBlockHeadersMsg, reqID, 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, true)
defer tearDown() defer tearDown()
bc := server.handler.blockchain bc := server.handler.blockchain
@ -245,20 +255,21 @@ func testGetBlockBodies(t *testing.T, protocol int) {
reqID++ reqID++
// Send the hash request and verify the response // Send the hash request and verify the response
cost := server.peer.peer.GetRequestCost(GetBlockBodiesMsg, len(hashes))
sendRequest(server.peer.app, GetBlockBodiesMsg, reqID, hashes) sendRequest(server.peer.app, GetBlockBodiesMsg, reqID, 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
server, tearDown := newServerEnv(t, 4, protocol, nil, false, true, 0) server, tearDown := newServerEnv(t, 4, protocol, nil, false, true, 0, true)
defer tearDown() defer tearDown()
bc := server.handler.blockchain bc := server.handler.blockchain
@ -276,18 +287,19 @@ func testGetCode(t *testing.T, protocol int) {
} }
} }
cost := server.peer.peer.GetRequestCost(GetCodeMsg, len(codereqs))
sendRequest(server.peer.app, GetCodeMsg, 42, codereqs) sendRequest(server.peer.app, GetCodeMsg, 42, 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, true)
defer tearDown() defer tearDown()
bc := server.handler.blockchain bc := server.handler.blockchain
@ -296,8 +308,9 @@ func testGetStaleCode(t *testing.T, protocol int) {
BHash: bc.GetHeaderByNumber(number).Hash(), BHash: bc.GetHeaderByNumber(number).Hash(),
AccKey: crypto.Keccak256(testContractAddr[:]), AccKey: crypto.Keccak256(testContractAddr[:]),
} }
cost := server.peer.peer.GetRequestCost(GetCodeMsg, 1)
sendRequest(server.peer.app, GetCodeMsg, 42, []*CodeReq{req}) sendRequest(server.peer.app, GetCodeMsg, 42, []*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)
} }
} }
@ -307,12 +320,12 @@ 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
server, tearDown := newServerEnv(t, 4, protocol, nil, false, true, 0) server, tearDown := newServerEnv(t, 4, protocol, nil, false, true, 0, true)
defer tearDown() defer tearDown()
bc := server.handler.blockchain bc := server.handler.blockchain
@ -327,19 +340,20 @@ func testGetReceipt(t *testing.T, protocol int) {
receipts = append(receipts, rawdb.ReadRawReceipts(server.db, block.Hash(), block.NumberU64())) receipts = append(receipts, rawdb.ReadRawReceipts(server.db, block.Hash(), block.NumberU64()))
} }
// Send the hash request and verify the response // Send the hash request and verify the response
cost := server.peer.peer.GetRequestCost(GetReceiptsMsg, len(hashes))
sendRequest(server.peer.app, GetReceiptsMsg, 42, hashes) sendRequest(server.peer.app, GetReceiptsMsg, 42, 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
server, tearDown := newServerEnv(t, 4, protocol, nil, false, true, 0) server, tearDown := newServerEnv(t, 4, protocol, nil, false, true, 0, true)
defer tearDown() defer tearDown()
bc := server.handler.blockchain bc := server.handler.blockchain
@ -362,18 +376,19 @@ 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))
sendRequest(server.peer.app, GetProofsV2Msg, 42, proofreqs) sendRequest(server.peer.app, GetProofsV2Msg, 42, 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, true)
defer tearDown() defer tearDown()
bc := server.handler.blockchain bc := server.handler.blockchain
@ -395,7 +410,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)
} }
} }
@ -405,8 +420,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
@ -420,7 +435,7 @@ func testGetCHTProofs(t *testing.T, protocol int) {
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
} }
} }
server, tearDown := newServerEnv(t, int(config.ChtSize+config.ChtConfirms), protocol, waitIndexers, false, true, 0) server, tearDown := newServerEnv(t, int(config.ChtSize+config.ChtConfirms), protocol, waitIndexers, false, true, 0, true)
defer tearDown() defer tearDown()
bc := server.handler.blockchain bc := server.handler.blockchain
@ -446,14 +461,15 @@ func testGetCHTProofs(t *testing.T, protocol int) {
AuxReq: auxHeader, AuxReq: auxHeader,
}} }}
// Send the proof request and verify the response // Send the proof request and verify the response
cost := server.peer.peer.GetRequestCost(GetHelperTrieProofsMsg, len(requestsV2))
sendRequest(server.peer.app, GetHelperTrieProofsMsg, 42, requestsV2) sendRequest(server.peer.app, GetHelperTrieProofsMsg, 42, 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) {
@ -468,7 +484,7 @@ func testGetBloombitsProofs(t *testing.T, protocol int) {
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
} }
} }
server, tearDown := newServerEnv(t, int(config.BloomTrieSize+config.BloomTrieConfirms), protocol, waitIndexers, false, true, 0) server, tearDown := newServerEnv(t, int(config.BloomTrieSize+config.BloomTrieConfirms), protocol, waitIndexers, false, true, 0, true)
defer tearDown() defer tearDown()
bc := server.handler.blockchain bc := server.handler.blockchain
@ -494,18 +510,19 @@ func testGetBloombitsProofs(t *testing.T, protocol int) {
trie.Prove(key, 0, &proofs.Proofs) trie.Prove(key, 0, &proofs.Proofs)
// Send the proof request and verify the response // Send the proof request and verify the response
cost := server.peer.peer.GetRequestCost(GetHelperTrieProofsMsg, len(requests))
sendRequest(server.peer.app, GetHelperTrieProofsMsg, 42, requests) sendRequest(server.peer.app, GetHelperTrieProofsMsg, 42, 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, true)
defer tearDown() defer tearDown()
server.handler.addTxsSync = true server.handler.addTxsSync = true
@ -515,12 +532,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)
sendRequest(server.peer.app, SendTxV2Msg, reqID, types.Transactions{tx}) sendRequest(server.peer.app, SendTxV2Msg, reqID, types.Transactions{tx})
} else { } else {
cost = server.peer.peer.GetRequestCost(GetTxStatusMsg, 1)
sendRequest(server.peer.app, GetTxStatusMsg, reqID, []common.Hash{tx.Hash()}) sendRequest(server.peer.app, GetTxStatusMsg, reqID, []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")
} }
} }
@ -595,8 +615,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, true)
defer tearDown() defer tearDown()
server.handler.server.costTracker.testing = true server.handler.server.costTracker.testing = true
@ -616,7 +639,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)
} }
} }
@ -635,7 +658,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

@ -40,6 +40,8 @@ var (
miscInTxsTrafficMeter = metrics.NewRegisteredMeter("les/misc/in/traffic/txs", nil) miscInTxsTrafficMeter = metrics.NewRegisteredMeter("les/misc/in/traffic/txs", nil)
miscInTxStatusPacketsMeter = metrics.NewRegisteredMeter("les/misc/in/packets/txStatus", nil) miscInTxStatusPacketsMeter = metrics.NewRegisteredMeter("les/misc/in/packets/txStatus", nil)
miscInTxStatusTrafficMeter = metrics.NewRegisteredMeter("les/misc/in/traffic/txStatus", nil) miscInTxStatusTrafficMeter = metrics.NewRegisteredMeter("les/misc/in/traffic/txStatus", nil)
miscInLespayPacketsMeter = metrics.NewRegisteredMeter("les/misc/in/packets/lespay", nil)
miscInLespayTrafficMeter = metrics.NewRegisteredMeter("les/misc/in/traffic/lespay", nil)
miscOutPacketsMeter = metrics.NewRegisteredMeter("les/misc/out/packets/total", nil) miscOutPacketsMeter = metrics.NewRegisteredMeter("les/misc/out/packets/total", nil)
miscOutTrafficMeter = metrics.NewRegisteredMeter("les/misc/out/traffic/total", nil) miscOutTrafficMeter = metrics.NewRegisteredMeter("les/misc/out/traffic/total", nil)
@ -59,6 +61,8 @@ var (
miscOutTxsTrafficMeter = metrics.NewRegisteredMeter("les/misc/out/traffic/txs", nil) miscOutTxsTrafficMeter = metrics.NewRegisteredMeter("les/misc/out/traffic/txs", nil)
miscOutTxStatusPacketsMeter = metrics.NewRegisteredMeter("les/misc/out/packets/txStatus", nil) miscOutTxStatusPacketsMeter = metrics.NewRegisteredMeter("les/misc/out/packets/txStatus", nil)
miscOutTxStatusTrafficMeter = metrics.NewRegisteredMeter("les/misc/out/traffic/txStatus", nil) miscOutTxStatusTrafficMeter = metrics.NewRegisteredMeter("les/misc/out/traffic/txStatus", nil)
miscOutLespayPacketsMeter = metrics.NewRegisteredMeter("les/misc/out/packets/lespay", nil)
miscOutLespayTrafficMeter = metrics.NewRegisteredMeter("les/misc/out/traffic/lespay", nil)
miscServingTimeHeaderTimer = metrics.NewRegisteredTimer("les/misc/serve/header", nil) miscServingTimeHeaderTimer = metrics.NewRegisteredTimer("les/misc/serve/header", nil)
miscServingTimeBodyTimer = metrics.NewRegisteredTimer("les/misc/serve/body", nil) miscServingTimeBodyTimer = metrics.NewRegisteredTimer("les/misc/serve/body", nil)
@ -68,6 +72,7 @@ var (
miscServingTimeHelperTrieTimer = metrics.NewRegisteredTimer("les/misc/serve/helperTrie", nil) miscServingTimeHelperTrieTimer = metrics.NewRegisteredTimer("les/misc/serve/helperTrie", nil)
miscServingTimeTxTimer = metrics.NewRegisteredTimer("les/misc/serve/txs", nil) miscServingTimeTxTimer = metrics.NewRegisteredTimer("les/misc/serve/txs", nil)
miscServingTimeTxStatusTimer = metrics.NewRegisteredTimer("les/misc/serve/txStatus", nil) miscServingTimeTxStatusTimer = metrics.NewRegisteredTimer("les/misc/serve/txStatus", nil)
miscServingTimeLespayTimer = metrics.NewRegisteredTimer("les/misc/serve/lespay", nil)
connectionTimer = metrics.NewRegisteredTimer("les/connection/duration", nil) connectionTimer = metrics.NewRegisteredTimer("les/connection/duration", nil)
serverConnectionGauge = metrics.NewRegisteredGauge("les/connection/server", nil) serverConnectionGauge = metrics.NewRegisteredGauge("les/connection/server", nil)

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
@ -236,7 +236,7 @@ func testOdr(t *testing.T, protocol int, expFail uint64, checkCached bool, fn od
// still expect all retrievals to pass, now data should be cached locally // still expect all retrievals to pass, now data should be cached locally
if checkCached { if checkCached {
client.handler.backend.peers.unregister(client.peer.speer.id) client.handler.backend.peers.disconnect(client.peer.speer.id)
time.Sleep(time.Millisecond * 10) // ensure that all peerSetNotify callbacks are executed time.Sleep(time.Millisecond * 10) // ensure that all peerSetNotify callbacks are executed
test(5) test(5)
} }

View file

@ -132,6 +132,7 @@ type peerCommons struct {
frozen uint32 // Flag whether the peer is frozen. frozen uint32 // Flag whether the peer is frozen.
announceType uint64 // New block announcement type. announceType uint64 // New block announcement type.
headInfo blockInfo // Latest block information. headInfo blockInfo // Latest block information.
active bool
// Background task queue for caching peer tasks and executing in order. // Background task queue for caching peer tasks and executing in order.
sendQueue *execQueue sendQueue *execQueue
@ -478,12 +479,18 @@ func (p *serverPeer) requestTxStatus(reqID uint64, txHashes []common.Hash) error
return sendRequest(p.rw, GetTxStatusMsg, reqID, txHashes) return sendRequest(p.rw, GetTxStatusMsg, reqID, txHashes)
} }
// SendTxStatus creates a reply with a batch of transactions to be added to the remote transaction pool. // sendTxs creates a reply with a batch of transactions to be added to the remote transaction pool.
func (p *serverPeer) sendTxs(reqID uint64, txs rlp.RawValue) error { func (p *serverPeer) sendTxs(reqID uint64, txs rlp.RawValue) error {
p.Log().Debug("Sending batch of transactions", "size", len(txs)) p.Log().Debug("Sending batch of transactions", "size", len(txs))
return sendRequest(p.rw, SendTxV2Msg, reqID, txs) return sendRequest(p.rw, SendTxV2Msg, reqID, txs)
} }
// sendLespay sends a set of commands to the service token sale module
func (p *serverPeer) sendLespay(reqID uint64, cmd []byte) error {
p.Log().Debug("Sending batch of lespay commands", "size", len(cmd))
return sendRequest(p.rw, LespayMsg, reqID, cmd)
}
// waitBefore implements distPeer interface // waitBefore implements distPeer interface
func (p *serverPeer) waitBefore(maxCost uint64) (time.Duration, float64) { func (p *serverPeer) waitBefore(maxCost uint64) (time.Duration, float64) {
return p.fcServer.CanSend(maxCost) return p.fcServer.CanSend(maxCost)
@ -554,10 +561,12 @@ func (p *serverPeer) updateFlowControl(update keyValueMap) {
// If any of the flow control params is nil, refuse to update. // If any of the flow control params is nil, refuse to update.
var params flowcontrol.ServerParams var params flowcontrol.ServerParams
updated := false
if update.get("flowControl/BL", &params.BufLimit) == nil && update.get("flowControl/MRR", &params.MinRecharge) == nil { if update.get("flowControl/BL", &params.BufLimit) == nil && update.get("flowControl/MRR", &params.MinRecharge) == nil {
// todo can light client set a minimal acceptable flow control params? // todo can light client set a minimal acceptable flow control params?
p.fcParams = params p.fcParams = params
p.fcServer.UpdateParams(params) p.fcServer.UpdateParams(params)
updated = true
} }
var MRC RequestCostList var MRC RequestCostList
if update.get("flowControl/MRC", &MRC) == nil { if update.get("flowControl/MRC", &MRC) == nil {
@ -565,7 +574,18 @@ func (p *serverPeer) updateFlowControl(update keyValueMap) {
for code, cost := range costUpdate { for code, cost := range costUpdate {
p.fcCosts[code] = cost p.fcCosts[code] = cost
} }
updated = true
} }
if updated {
p.active = p.paramsUseful()
}
}
// paramsUseful returns true if the server parameters ensure the minimum required
// buffer limit and recharge
func (p *serverPeer) paramsUseful() bool {
reqRecharge, reqBufLimit := p.fcCosts.reqParams()
return p.fcParams.MinRecharge >= reqRecharge && p.fcParams.BufLimit >= reqBufLimit
} }
// Handshake executes the les protocol handshake, negotiating version number, // Handshake executes the les protocol handshake, negotiating version number,
@ -634,6 +654,9 @@ func (p *serverPeer) Handshake(td *big.Int, head common.Hash, headNum uint64, ge
type clientPeer struct { type clientPeer struct {
peerCommons peerCommons
activate, deactivate func()
getBalance func() uint64
// responseLock ensures that responses are queued in the same order as // responseLock ensures that responses are queued in the same order as
// RequestProcessed is called // RequestProcessed is called
responseLock sync.Mutex responseLock sync.Mutex
@ -681,8 +704,8 @@ func (p *clientPeer) sendStop() error {
} }
// sendResume notifies the client about getting out of frozen state // sendResume notifies the client about getting out of frozen state
func (p *clientPeer) sendResume(bv uint64) error { func (p *clientPeer) sendResume(sf stateFeedback) error {
return p2p.Send(p.rw, ResumeMsg, bv) return p2p.Send(p.rw, ResumeMsg, sf)
} }
// freeze temporarily puts the client in a frozen state which means all unprocessed // freeze temporarily puts the client in a frozen state which means all unprocessed
@ -711,7 +734,19 @@ func (p *clientPeer) freeze() {
continue continue
} }
atomic.StoreUint32(&p.frozen, 0) atomic.StoreUint32(&p.frozen, 0)
p.sendResume(bufValue) var balance uint64
if p.getBalance != nil {
balance = p.getBalance()
}
sf := stateFeedback{
protocolVersion: p.version,
stateFeedbackV4: stateFeedbackV4{
BV: bufValue,
RealCost: 0,
TokenBalance: balance,
},
}
p.sendResume(sf)
return return
} }
}() }()
@ -728,12 +763,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
SF stateFeedback
Data rlp.RawValue 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
@ -786,6 +822,12 @@ func (p *clientPeer) replyTxStatus(reqID uint64, stats []light.TxStatus) *reply
return &reply{p.rw, TxStatusMsg, reqID, data} return &reply{p.rw, TxStatusMsg, reqID, data}
} }
// replyLespay sends a set of replies to lespay commands
func (p *clientPeer) replyLespay(reqID uint64, reply []byte, delay uint) error {
p.Log().Debug("Sending batch of lespay replies", "size", len(reply))
return sendRequest(p.rw, LespayReplyMsg, reqID, lespayReply{reply, delay})
}
// sendAnnounce announces the availability of a number of blocks through // sendAnnounce announces the availability of a number of blocks through
// a hash notification. // a hash notification.
func (p *clientPeer) sendAnnounce(request announceData) error { func (p *clientPeer) sendAnnounce(request announceData) error {
@ -798,12 +840,21 @@ func (p *clientPeer) updateCapacity(cap uint64) {
p.lock.Lock() p.lock.Lock()
defer p.lock.Unlock() defer p.lock.Unlock()
if !p.active && cap != 0 && p.activate != nil {
p.activate()
}
if cap != 0 || p.version >= lpv4 {
p.fcParams = flowcontrol.ServerParams{MinRecharge: cap, BufLimit: cap * bufLimitRatio} p.fcParams = flowcontrol.ServerParams{MinRecharge: cap, BufLimit: cap * bufLimitRatio}
p.fcClient.UpdateParams(p.fcParams) p.fcClient.UpdateParams(p.fcParams)
var kvList keyValueList var kvList keyValueList
kvList = kvList.add("flowControl/MRR", cap)
kvList = kvList.add("flowControl/BL", cap*bufLimitRatio) kvList = kvList.add("flowControl/BL", cap*bufLimitRatio)
kvList = kvList.add("flowControl/MRR", cap)
p.mustQueueSend(func() { p.sendAnnounce(announceData{Update: kvList}) }) p.mustQueueSend(func() { p.sendAnnounce(announceData{Update: kvList}) })
}
if p.active && cap == 0 && p.deactivate != nil {
p.deactivate()
}
} }
// freezeClient temporarily puts the client in a frozen state which means all // freezeClient temporarily puts the client in a frozen state which means all
@ -859,8 +910,14 @@ func (p *clientPeer) Handshake(td *big.Int, head common.Hash, headNum uint64, ge
*lists = (*lists).add("serveRecentState", stateRecent) *lists = (*lists).add("serveRecentState", stateRecent)
*lists = (*lists).add("txRelay", nil) *lists = (*lists).add("txRelay", nil)
} }
*lists = (*lists).add("flowControl/BL", server.defParams.BufLimit) p.active = p.version < lpv4
*lists = (*lists).add("flowControl/MRR", server.defParams.MinRecharge) if p.active {
p.fcParams = server.defParams
} else {
p.fcParams = flowcontrol.ServerParams{}
}
*lists = (*lists).add("flowControl/BL", p.fcParams.BufLimit)
*lists = (*lists).add("flowControl/MRR", p.fcParams.MinRecharge)
var costList RequestCostList var costList RequestCostList
if server.costTracker.testCostList != nil { if server.costTracker.testCostList != nil {
@ -870,7 +927,6 @@ func (p *clientPeer) Handshake(td *big.Int, head common.Hash, headNum uint64, ge
} }
*lists = (*lists).add("flowControl/MRC", costList) *lists = (*lists).add("flowControl/MRC", costList)
p.fcCosts = costList.decode(ProtocolLengths[uint(p.version)]) p.fcCosts = costList.decode(ProtocolLengths[uint(p.version)])
p.fcParams = server.defParams
// Add advertised checkpoint and register block height which // Add advertised checkpoint and register block height which
// client can verify the checkpoint validity. // client can verify the checkpoint validity.
@ -890,7 +946,7 @@ func (p *clientPeer) Handshake(td *big.Int, head common.Hash, headNum uint64, ge
// set default announceType on server side // set default announceType on server side
p.announceType = announceTypeSimple p.announceType = announceTypeSimple
} }
p.fcClient = flowcontrol.NewClientNode(server.fcManager, server.defParams) p.fcClient = flowcontrol.NewClientNode(server.fcManager, p.fcParams)
} }
return nil return nil
}) })
@ -913,7 +969,7 @@ type clientPeerSubscriber interface {
// clientPeerSet represents the set of active client peers currently // clientPeerSet represents the set of active client peers currently
// participating in the Light Ethereum sub-protocol. // participating in the Light Ethereum sub-protocol.
type clientPeerSet struct { type clientPeerSet struct {
peers map[string]*clientPeer active, inactive map[string]*clientPeer
// subscribers is a batch of subscribers and peerset will notify // subscribers is a batch of subscribers and peerset will notify
// these subscribers when the peerset changes(new client peer is // these subscribers when the peerset changes(new client peer is
// added or removed) // added or removed)
@ -924,17 +980,23 @@ type clientPeerSet struct {
// newClientPeerSet creates a new peer set to track the client peers. // newClientPeerSet creates a new peer set to track the client peers.
func newClientPeerSet() *clientPeerSet { func newClientPeerSet() *clientPeerSet {
return &clientPeerSet{peers: make(map[string]*clientPeer)} return &clientPeerSet{
active: make(map[string]*clientPeer),
inactive: make(map[string]*clientPeer),
}
} }
// subscribe adds a service to be notified about added or removed // subscribe adds a service to be notified about added or removed
// peers and also register all active peers into the given service. // peers and also register all active peers into the given service.
func (ps *clientPeerSet) subscribe(sub clientPeerSubscriber) { func (ps *clientPeerSet) subscribe(sub clientPeerSubscriber) {
ps.lock.Lock() ps.lock.Lock()
defer ps.lock.Unlock()
ps.subscribers = append(ps.subscribers, sub) ps.subscribers = append(ps.subscribers, sub)
for _, p := range ps.peers { notify := make([]*clientPeer, 0, len(ps.active))
for _, p := range ps.active {
notify = append(notify, p)
}
ps.lock.Unlock()
for _, p := range notify {
sub.registerPeer(p) sub.registerPeer(p)
} }
} }
@ -956,17 +1018,23 @@ func (ps *clientPeerSet) unSubscribe(sub clientPeerSubscriber) {
// peer is already known. // peer is already known.
func (ps *clientPeerSet) register(peer *clientPeer) error { func (ps *clientPeerSet) register(peer *clientPeer) error {
ps.lock.Lock() ps.lock.Lock()
defer ps.lock.Unlock()
if ps.closed { if ps.closed {
ps.lock.Unlock()
return errClosed return errClosed
} }
if _, exist := ps.peers[peer.id]; exist { if _, ok := ps.active[p.id]; ok {
ps.lock.Unlock()
return errAlreadyRegistered return errAlreadyRegistered
} }
ps.peers[peer.id] = peer delete(ps.inactive, p.id)
for _, sub := range ps.subscribers { ps.active[p.id] = p
sub.registerPeer(peer)
peers := make([]clientPeerSubscriber, len(ps.subscribers))
copy(peers, ps.subscribers)
ps.lock.Unlock()
for _, n := range peers {
n.registerPeer(p)
} }
return nil return nil
} }
@ -976,27 +1044,58 @@ func (ps *clientPeerSet) register(peer *clientPeer) error {
// at the networking layer. // at the networking layer.
func (ps *clientPeerSet) unregister(id string) error { func (ps *clientPeerSet) unregister(id string) error {
ps.lock.Lock() ps.lock.Lock()
defer ps.lock.Unlock() if _, ok := ps.active[p.id]; !ok {
ps.lock.Unlock()
return errNotRegistered
} else {
delete(ps.active, p.id)
ps.inactive[p.id] = p
peers := make([]clientPeerSubscriber, len(ps.subscribers))
copy(peers, ps.subscribers)
ps.lock.Unlock()
p, ok := ps.peers[id] for _, n := range peers {
if !ok { n.unregisterPeer(p)
}
return nil
}
}
// disconnect removes a remote peer from either the active or inactive set and
// initiates disconnection at the networking layer.
func (ps *clientPeerSet) disconnect(id string) error {
ps.lock.Lock()
var (
peers []clientPeerSubscriber
p *clientPeer
ok bool
)
if p, ok = ps.active[id]; ok {
delete(ps.active, p.id)
peers = make([]clientPeerSubscriber, len(ps.subscribers))
copy(peers, ps.subscribers)
} else if p, ok = ps.inactive[id]; ok {
delete(ps.inactive, id)
} else {
ps.lock.Unlock()
return errNotRegistered return errNotRegistered
} }
delete(ps.peers, id) ps.lock.Unlock()
for _, sub := range ps.subscribers { for _, n := range peers {
sub.unregisterPeer(p) n.unregisterPeer(p)
} }
p.Peer.Disconnect(p2p.DiscRequested) p.Peer.Disconnect(p2p.DiscUselessPeer)
return nil return nil
} }
// ids returns a list of all registered peer IDs // ids returns a list of all active peer IDs
func (ps *clientPeerSet) ids() []string { func (ps *clientPeerSet) ids() []string {
ps.lock.RLock() ps.lock.RLock()
defer ps.lock.RUnlock() defer ps.lock.RUnlock()
var ids []string var ids []string
for id := range ps.peers { for id := range ps.active {
ids = append(ids, id) ids = append(ids, id)
} }
return ids return ids
@ -1007,24 +1106,27 @@ func (ps *clientPeerSet) peer(id string) *clientPeer {
ps.lock.RLock() ps.lock.RLock()
defer ps.lock.RUnlock() defer ps.lock.RUnlock()
return ps.peers[id] if p, ok := ps.active[id]; ok {
return p
}
return ps.inactive[id]
} }
// len returns if the current number of peers in the set. // len returns if the current number of peers in the active set.
func (ps *clientPeerSet) len() int { func (ps *clientPeerSet) len() int {
ps.lock.RLock() ps.lock.RLock()
defer ps.lock.RUnlock() defer ps.lock.RUnlock()
return len(ps.peers) return len(ps.active)
} }
// allClientPeers returns all client peers in a list. // allClientPeers returns all active client peers in a list.
func (ps *clientPeerSet) allPeers() []*clientPeer { func (ps *clientPeerSet) allPeers() []*clientPeer {
ps.lock.RLock() ps.lock.RLock()
defer ps.lock.RUnlock() defer ps.lock.RUnlock()
list := make([]*clientPeer, 0, len(ps.peers)) list := make([]*clientPeer, 0, len(ps.peers))
for _, p := range ps.peers { for _, p := range ps.active {
list = append(list, p) list = append(list, p)
} }
return list return list
@ -1036,7 +1138,10 @@ func (ps *clientPeerSet) close() {
ps.lock.Lock() ps.lock.Lock()
defer ps.lock.Unlock() defer ps.lock.Unlock()
for _, p := range ps.peers { for _, p := range ps.active {
p.Disconnect(p2p.DiscQuitting)
}
for _, p := range ps.inactive {
p.Disconnect(p2p.DiscQuitting) p.Disconnect(p2p.DiscQuitting)
} }
ps.closed = true ps.closed = true
@ -1045,7 +1150,7 @@ func (ps *clientPeerSet) close() {
// serverPeerSet represents the set of active server peers currently // serverPeerSet represents the set of active server peers currently
// participating in the Light Ethereum sub-protocol. // participating in the Light Ethereum sub-protocol.
type serverPeerSet struct { type serverPeerSet struct {
peers map[string]*serverPeer active, inactive map[string]*serverPeer
// subscribers is a batch of subscribers and peerset will notify // subscribers is a batch of subscribers and peerset will notify
// these subscribers when the peerset changes(new server peer is // these subscribers when the peerset changes(new server peer is
// added or removed) // added or removed)
@ -1056,17 +1161,23 @@ type serverPeerSet struct {
// newServerPeerSet creates a new peer set to track the active server peers. // newServerPeerSet creates a new peer set to track the active server peers.
func newServerPeerSet() *serverPeerSet { func newServerPeerSet() *serverPeerSet {
return &serverPeerSet{peers: make(map[string]*serverPeer)} return &serverPeerSet{
active: make(map[string]*serverPeer),
inactive: make(map[string]*serverPeer),
}
} }
// subscribe adds a service to be notified about added or removed // subscribe adds a service to be notified about added or removed
// peers and also register all active peers into the given service. // peers and also register all active peers into the given service.
func (ps *serverPeerSet) subscribe(sub serverPeerSubscriber) { func (ps *serverPeerSet) subscribe(sub serverPeerSubscriber) {
ps.lock.Lock() ps.lock.Lock()
defer ps.lock.Unlock()
ps.subscribers = append(ps.subscribers, sub) ps.subscribers = append(ps.subscribers, sub)
for _, p := range ps.peers { notify := make([]*serverPeer, 0, len(ps.active))
for _, p := range ps.active {
notify = append(notify, p)
}
ps.lock.Unlock()
for _, p := range notify {
sub.registerPeer(p) sub.registerPeer(p)
} }
} }
@ -1088,17 +1199,23 @@ func (ps *serverPeerSet) unSubscribe(sub serverPeerSubscriber) {
// peer is already known. // peer is already known.
func (ps *serverPeerSet) register(peer *serverPeer) error { func (ps *serverPeerSet) register(peer *serverPeer) error {
ps.lock.Lock() ps.lock.Lock()
defer ps.lock.Unlock()
if ps.closed { if ps.closed {
ps.lock.Unlock()
return errClosed return errClosed
} }
if _, exist := ps.peers[peer.id]; exist { if _, ok := ps.active[p.id]; ok {
ps.lock.Unlock()
return errAlreadyRegistered return errAlreadyRegistered
} }
ps.peers[peer.id] = peer delete(ps.inactive, p.id)
for _, sub := range ps.subscribers { ps.active[p.id] = p
sub.registerPeer(peer)
peers := make([]serverPeerSubscriber, len(ps.subscribers))
copy(peers, ps.subscribers)
ps.lock.Unlock()
for _, n := range peers {
n.registerPeer(p)
} }
return nil return nil
} }
@ -1108,27 +1225,58 @@ func (ps *serverPeerSet) register(peer *serverPeer) error {
// the networking layer. // the networking layer.
func (ps *serverPeerSet) unregister(id string) error { func (ps *serverPeerSet) unregister(id string) error {
ps.lock.Lock() ps.lock.Lock()
defer ps.lock.Unlock() if _, ok := ps.active[p.id]; !ok {
ps.lock.Unlock()
return errNotRegistered
} else {
delete(ps.active, p.id)
ps.inactive[p.id] = p
peers := make([]serverPeerSubscriber, len(ps.subscribers))
copy(peers, ps.subscribers)
ps.lock.Unlock()
p, ok := ps.peers[id] for _, n := range peers {
if !ok { n.unregisterPeer(p)
}
return nil
}
}
// disconnect removes a remote peer from either the active or inactive set and
// initiates disconnection at the networking layer.
func (ps *serverPeerSet) disconnect(id string) error {
ps.lock.Lock()
var (
peers []serverPeerSubscriber
p *serverPeer
ok bool
)
if p, ok = ps.active[id]; ok {
delete(ps.active, p.id)
peers = make([]serverPeerSubscriber, len(ps.subscribers))
copy(peers, ps.subscribers)
} else if p, ok = ps.inactive[id]; ok {
delete(ps.inactive, id)
} else {
ps.lock.Unlock()
return errNotRegistered return errNotRegistered
} }
delete(ps.peers, id) ps.lock.Unlock()
for _, sub := range ps.subscribers { for _, n := range peers {
sub.unregisterPeer(p) n.unregisterPeer(p)
} }
p.Peer.Disconnect(p2p.DiscRequested) p.Peer.Disconnect(p2p.DiscUselessPeer)
return nil return nil
} }
// ids returns a list of all registered peer IDs // ids returns a list of all active peer IDs
func (ps *serverPeerSet) ids() []string { func (ps *serverPeerSet) ids() []string {
ps.lock.RLock() ps.lock.RLock()
defer ps.lock.RUnlock() defer ps.lock.RUnlock()
var ids []string var ids []string
for id := range ps.peers { for id := range ps.active {
ids = append(ids, id) ids = append(ids, id)
} }
return ids return ids
@ -1139,15 +1287,18 @@ func (ps *serverPeerSet) peer(id string) *serverPeer {
ps.lock.RLock() ps.lock.RLock()
defer ps.lock.RUnlock() defer ps.lock.RUnlock()
return ps.peers[id] if p, ok := ps.active[id]; ok {
return p
}
return ps.inactive[id]
} }
// len returns if the current number of peers in the set. // len returns if the current number of peers in the active set.
func (ps *serverPeerSet) len() int { func (ps *serverPeerSet) len() int {
ps.lock.RLock() ps.lock.RLock()
defer ps.lock.RUnlock() defer ps.lock.RUnlock()
return len(ps.peers) return len(ps.active)
} }
// bestPeer retrieves the known peer with the currently highest total difficulty. // bestPeer retrieves the known peer with the currently highest total difficulty.
@ -1161,7 +1312,7 @@ func (ps *serverPeerSet) bestPeer() *serverPeer {
bestPeer *serverPeer bestPeer *serverPeer
bestTd *big.Int bestTd *big.Int
) )
for _, p := range ps.peers { for _, p := range ps.active {
if td := p.Td(); bestTd == nil || td.Cmp(bestTd) > 0 { if td := p.Td(); bestTd == nil || td.Cmp(bestTd) > 0 {
bestPeer, bestTd = p, td bestPeer, bestTd = p, td
} }
@ -1169,12 +1320,12 @@ func (ps *serverPeerSet) bestPeer() *serverPeer {
return bestPeer return bestPeer
} }
// allServerPeers returns all server peers in a list. // allPeers returns all active server peers in a list.
func (ps *serverPeerSet) allPeers() []*serverPeer { func (ps *serverPeerSet) allPeers() []*serverPeer {
ps.lock.RLock() ps.lock.RLock()
defer ps.lock.RUnlock() defer ps.lock.RUnlock()
list := make([]*serverPeer, 0, len(ps.peers)) list := make([]*serverPeer, 0, len(ps.active))
for _, p := range ps.peers { for _, p := range ps.peers {
list = append(list, p) list = append(list, p)
} }
@ -1187,7 +1338,10 @@ func (ps *serverPeerSet) close() {
ps.lock.Lock() ps.lock.Lock()
defer ps.lock.Unlock() defer ps.lock.Unlock()
for _, p := range ps.peers { for _, p := range ps.active {
p.Disconnect(p2p.DiscQuitting)
}
for _, p := range ps.inactive {
p.Disconnect(p2p.DiscQuitting) p.Disconnect(p2p.DiscQuitting)
} }
ps.closed = true ps.closed = true

View file

@ -33,17 +33,18 @@ import (
const ( const (
lpv2 = 2 lpv2 = 2
lpv3 = 3 lpv3 = 3
lpv4 = 4
) )
// Supported versions of the les protocol (first is primary) // Supported versions of the les protocol (first is primary)
var ( var (
ClientProtocolVersions = []uint{lpv2, lpv3} ClientProtocolVersions = []uint{lpv2, lpv3, lpv4}
ServerProtocolVersions = []uint{lpv2, lpv3} ServerProtocolVersions = []uint{lpv2, lpv3, lpv4}
AdvertiseProtocolVersions = []uint{lpv2} // clients are searching for the first advertised protocol in the list AdvertiseProtocolVersions = []uint{lpv2} // clients are searching for the first advertised protocol in the list
) )
// Number of implemented message corresponding to different protocol versions. // Number of implemented message corresponding to different protocol versions.
var ProtocolLengths = map[uint]uint64{lpv2: 22, lpv3: 24} var ProtocolLengths = map[uint]uint64{lpv2: 22, lpv3: 24, lpv4: 26}
const ( const (
NetworkId = 1 NetworkId = 1
@ -74,6 +75,9 @@ const (
// Protocol messages introduced in LPV3 // Protocol messages introduced in LPV3
StopMsg = 0x16 StopMsg = 0x16
ResumeMsg = 0x17 ResumeMsg = 0x17
// Protocol messages introduced in LPV4
LespayMsg = 0x18
LespayReplyMsg = 0x19
) )
type requestInfo struct { type requestInfo struct {
@ -201,6 +205,11 @@ type hashOrNumber struct {
Number uint64 // Block hash from which to retrieve headers (excludes Hash) Number uint64 // Block hash from which to retrieve headers (excludes Hash)
} }
type lespayReply struct {
Reply []byte
Delay uint
}
// EncodeRLP is a specialized encoder for hashOrNumber to encode only one of the // EncodeRLP is a specialized encoder for hashOrNumber to encode only one of the
// two contained union fields. // two contained union fields.
func (hn *hashOrNumber) EncodeRLP(w io.Writer) error { func (hn *hashOrNumber) EncodeRLP(w io.Writer) error {
@ -235,3 +244,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

@ -345,7 +345,7 @@ func (r *sentReq) tryRequest() {
if hrto { if hrto {
pp.Log().Debug("Request timed out hard") pp.Log().Debug("Request timed out hard")
if r.rm.peers != nil { if r.rm.peers != nil {
r.rm.peers.unregister(pp.id) r.rm.peers.disconnect(pp.id)
} }
} }

View file

@ -44,6 +44,7 @@ type LesServer struct {
handler *serverHandler handler *serverHandler
lesTopics []discv5.Topic lesTopics []discv5.Topic
privateKey *ecdsa.PrivateKey privateKey *ecdsa.PrivateKey
srvr *p2p.Server
// Flow control and capacity management // Flow control and capacity management
fcManager *flowcontrol.ClientManager fcManager *flowcontrol.ClientManager
@ -51,6 +52,7 @@ type LesServer struct {
defParams flowcontrol.ServerParams defParams flowcontrol.ServerParams
servingQueue *servingQueue servingQueue *servingQueue
clientPool *clientPool clientPool *clientPool
tokenSale *tokenSale
minCapacity, maxCapacity, freeCapacity uint64 minCapacity, maxCapacity, freeCapacity uint64
threadsIdle int // Request serving threads count when system is idle. threadsIdle int // Request serving threads count when system is idle.
@ -116,8 +118,13 @@ func NewLesServer(e *eth.Ethereum, config *eth.Config) (*LesServer, error) {
srv.maxCapacity = totalRecharge srv.maxCapacity = totalRecharge
} }
srv.fcManager.SetCapacityLimits(srv.freeCapacity, srv.maxCapacity, srv.freeCapacity*2) srv.fcManager.SetCapacityLimits(srv.freeCapacity, srv.maxCapacity, srv.freeCapacity*2)
srv.clientPool = newClientPool(srv.chainDb, srv.freeCapacity, mclock.System{}, func(id enode.ID) { go srv.peers.unregister(peerIdToString(id)) }) srv.clientPool = newClientPool(srv.chainDb, srv.minCapacity, srv.freeCapacity, mclock.System{}, func(id enode.ID) { go srv.peers.disconnect(peerIdToString(id)) })
srv.clientPool.setDefaultFactors(priceFactors{0, 1, 1}, priceFactors{0, 1, 1}) srv.clientPool.setDefaultFactors(priceFactors{0, 1, 1}, priceFactors{0, 1, 1})
srv.tokenSale = newTokenSale(srv.clientPool, 0.1, 100)
if config.LespayTestModule {
srv.tokenSale.addReceiver("test", testReceiver{})
srv.clientPool.setExpirationTCs(defaultPosExpTC, defaultNegExpTC)
}
checkpoint := srv.latestLocalCheckpoint() checkpoint := srv.latestLocalCheckpoint()
if !checkpoint.Empty() { if !checkpoint.Empty() {
@ -148,6 +155,12 @@ func (s *LesServer) APIs() []rpc.API {
Service: NewPrivateDebugAPI(s), Service: NewPrivateDebugAPI(s),
Public: false, Public: false,
}, },
{
Namespace: "lespay",
Version: "1.0",
Service: NewPrivateLespayAPI(s.lesCommons.peers, nil, s.srvr.DiscV5, s.tokenSale),
Public: false,
},
} }
} }
@ -167,6 +180,7 @@ func (s *LesServer) Protocols() []p2p.Protocol {
// Start starts the LES server // Start starts the LES server
func (s *LesServer) Start(srvr *p2p.Server) { func (s *LesServer) Start(srvr *p2p.Server) {
s.srvr = srvr
s.privateKey = srvr.PrivateKey s.privateKey = srvr.PrivateKey
s.handler.start() s.handler.start()
@ -174,6 +188,7 @@ func (s *LesServer) Start(srvr *p2p.Server) {
go s.capacityManagement() go s.capacityManagement()
if srvr.DiscV5 != nil { if srvr.DiscV5 != nil {
srvr.DiscV5.RegisterTalkHandler("lespay", s.handler.talkRequestHandler)
for _, topic := range s.lesTopics { for _, topic := range s.lesTopics {
topic := topic topic := topic
go func() { go func() {
@ -191,6 +206,11 @@ func (s *LesServer) Start(srvr *p2p.Server) {
func (s *LesServer) Stop() { func (s *LesServer) Stop() {
close(s.closeCh) close(s.closeCh)
if s.srvr.DiscV5 != nil {
s.srvr.DiscV5.RemoveTalkHandler("lespay")
}
s.tokenSale.stop()
// Disconnect existing sessions. // Disconnect existing sessions.
// This also closes the gate for any new registrations on the peer set. // This also closes the gate for any new registrations on the peer set.
// sessions which are already established but not added to pm.peers yet // sessions which are already established but not added to pm.peers yet

View file

@ -20,6 +20,7 @@ import (
"encoding/binary" "encoding/binary"
"encoding/json" "encoding/json"
"errors" "errors"
"net"
"sync" "sync"
"sync/atomic" "sync/atomic"
"time" "time"
@ -35,6 +36,7 @@ import (
"github.com/ethereum/go-ethereum/log" "github.com/ethereum/go-ethereum/log"
"github.com/ethereum/go-ethereum/metrics" "github.com/ethereum/go-ethereum/metrics"
"github.com/ethereum/go-ethereum/p2p" "github.com/ethereum/go-ethereum/p2p"
"github.com/ethereum/go-ethereum/p2p/enode"
"github.com/ethereum/go-ethereum/rlp" "github.com/ethereum/go-ethereum/rlp"
"github.com/ethereum/go-ethereum/trie" "github.com/ethereum/go-ethereum/trie"
) )
@ -54,10 +56,7 @@ const (
MaxTxStatus = 256 // Amount of transactions to queried per request MaxTxStatus = 256 // Amount of transactions to queried per request
) )
var ( var errTooManyInvalidRequest = errors.New("too many invalid requests made")
errTooManyInvalidRequest = errors.New("too many invalid requests made")
errFullClientPool = errors.New("client pool is full")
)
// serverHandler is responsible for serving light client and process // serverHandler is responsible for serving light client and process
// all incoming light requests. // all incoming light requests.
@ -103,6 +102,9 @@ func (h *serverHandler) stop() {
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 := newClientPeer(int(version), h.server.config.NetworkId, p, newMeteredMsgWriter(rw, int(version))) peer := newClientPeer(int(version), h.server.config.NetworkId, p, newMeteredMsgWriter(rw, int(version)))
defer peer.close() defer peer.close()
peer.getBalance = func() uint64 {
return h.server.clientPool.getPosBalance(p.ID()).value.value(h.server.clientPool.posExpiration(mclock.Now()))
}
h.wg.Add(1) h.wg.Add(1)
defer h.wg.Done() defer h.wg.Done()
return h.handle(peer) return h.handle(peer)
@ -134,28 +136,58 @@ func (h *serverHandler) handle(p *clientPeer) error {
} }
defer p.fcClient.Disconnect() defer p.fcClient.Disconnect()
// Disconnect the inbound peer if it's rejected by clientPool var (
if !h.server.clientPool.connect(p, 0) { connectedAt mclock.AbsTime
p.Log().Debug("Light Ethereum peer registration failed", "err", errFullClientPool) wg *sync.WaitGroup // Wait group used to track all in-flight task routines.
return errFullClientPool )
} p.activate = func() {
// Register the peer locally // Register the peer locally
if err := h.server.peers.register(p); err != nil { if err := h.server.peers.Register(p); err != nil {
h.server.clientPool.disconnect(p) h.server.clientPool.disconnect(p)
p.Log().Error("Light Ethereum peer registration failed", "err", err) p.Log().Error("Light Ethereum peer registration failed", "err", err)
return err return
}
clientConnectionGauge.Update(int64(h.server.peers.Len()))
connectedAt = mclock.Now()
wg = new(sync.WaitGroup)
p.active = true
}
p.deactivate = func() {
h.server.peers.unregister(p)
if p.version < lpv4 {
h.server.peers.disconnect(p.id)
}
clientConnectionGauge.Update(int64(h.server.peers.Len()))
connectionTimer.Update(time.Duration(mclock.Now() - connectedAt))
p.active = false
}
if p.active {
p.activate()
} }
clientConnectionGauge.Update(int64(h.server.peers.len()))
var wg sync.WaitGroup // Wait group used to track all in-flight task routines. if capacity, err := h.server.clientPool.connect(p, 0); err != nil {
// Disconnect the inbound peer if it's rejected by clientPool
p.Log().Debug("Light Ethereum peer registration failed", "err", err)
return err
} else if capacity != p.fcParams.MinRecharge {
if p.version < lpv4 {
h.server.peers.Disconnect(p.id)
} else {
p.updateCapacity(capacity)
}
}
connectedAt := mclock.Now()
defer func() { defer func() {
wg.Wait() // Ensure all background task routines have exited. wg.Wait() // Ensure all background task routines have exited.
h.server.peers.unregister(p.id)
h.server.clientPool.disconnect(p) h.server.clientPool.disconnect(p)
clientConnectionGauge.Update(int64(h.server.peers.len())) p.responseLock.Lock()
connectionTimer.Update(time.Duration(mclock.Now() - connectedAt)) if p.active {
p.deactivate()
}
p.activate = nil
p.deactivate = nil
p.responseLock.Unlock()
h.server.peers.disconnect(p.id)
}() }()
// Spawn a main loop to handle all incoming messages. // Spawn a main loop to handle all incoming messages.
@ -166,7 +198,7 @@ func (h *serverHandler) handle(p *clientPeer) error {
return err return err
default: default:
} }
if err := h.handleMsg(p, &wg); err != nil { if err := h.handleMsg(p, wg); err != nil {
p.Log().Debug("Light Ethereum message handling failed", "err", err) p.Log().Debug("Light Ethereum message handling failed", "err", err)
return err return err
} }
@ -245,22 +277,33 @@ func (h *serverHandler) handleMsg(p *clientPeer, 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.mustQueueSend(func() { p.mustQueueSend(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:
@ -375,6 +418,8 @@ func (h *serverHandler) handleMsg(p *clientPeer, wg *sync.WaitGroup) error {
} }
reply := p.replyBlockHeaders(req.ReqID, headers) reply := p.replyBlockHeaders(req.ReqID, headers)
sendResponse(req.ReqID, query.Amount, p.replyBlockHeaders(req.ReqID, headers), task.done()) sendResponse(req.ReqID, query.Amount, p.replyBlockHeaders(req.ReqID, headers), task.done())
reply := p.ReplyBlockHeaders(req.ReqID, headers)
sendResponse(req.ReqID, query.Amount, reply, task.done())
if metrics.EnabledExpensive { if metrics.EnabledExpensive {
miscOutHeaderPacketsMeter.Mark(1) miscOutHeaderPacketsMeter.Mark(1)
miscOutHeaderTrafficMeter.Mark(int64(reply.size())) miscOutHeaderTrafficMeter.Mark(int64(reply.size()))
@ -824,6 +869,38 @@ func (h *serverHandler) handleMsg(p *clientPeer, wg *sync.WaitGroup) error {
} }
}() }()
} }
case LespayMsg:
p.Log().Trace("Received transaction status query request")
if metrics.EnabledExpensive {
miscInLespayPacketsMeter.Mark(1)
miscInLespayTrafficMeter.Mark(int64(msg.Size))
defer func(start time.Time) { miscServingTimeLespayTimer.UpdateSince(start) }(time.Now())
}
var req struct {
ReqID uint64
Cmd []byte
}
if err := msg.Decode(&req); err != nil {
clientErrorMeter.Mark(1)
return errResp(ErrDecode, "msg %v: %v", msg, err)
}
if !h.server.tokenSale.queueCommand(p.id, lespayCmd{
cmd: req.Cmd,
id: p.ID(),
freeID: p.freeClientId(),
send: func(reply []byte, delay uint) {
if metrics.EnabledExpensive {
miscOutLespayPacketsMeter.Mark(1)
miscOutLespayTrafficMeter.Mark(int64(len(reply)))
}
p.queueSend(func() {
p.ReplyLespay(req.ReqID, reply, delay)
})
},
}) {
clientErrorMeter.Mark(1)
return errResp(ErrRequestRejected, "")
}
default: default:
p.Log().Trace("Received invalid message", "code", msg.Code) p.Log().Trace("Received invalid message", "code", msg.Code)
@ -959,3 +1036,49 @@ func (h *serverHandler) broadcastHeaders() {
} }
} }
} }
// talkRequestHandler implements discv5.TalkRequestHandler. It processes a list of
// lespay token sale commands and returns the results and the recommended delay.
//
// Note: the UDP talk format for lespay commands allows multiple commands in a single
// packet because UDP does not guarantee the correct order of messages which might be
// important in some cases (like deposit followed by buyTokens).
func (h *serverHandler) talkRequestHandler(id enode.ID, addr *net.UDPAddr, payload interface{}) (interface{}, uint, bool) {
c, ok := payload.([]interface{})
if !ok {
return nil, 0, false
}
type result struct {
data []byte
delay uint
}
resultCh := make(chan result, len(c))
results := make([][]byte, len(c))
for _, c := range c {
cmd, ok := c.([]byte)
if !ok {
return nil, 0, false
}
if !h.server.tokenSale.queueCommand(id.String(), lespayCmd{
cmd: cmd,
id: id,
freeID: addr.IP.String(),
send: func(reply []byte, delay uint) {
resultCh <- result{reply, delay}
},
}) {
return nil, 0, false
}
}
var lastDelay uint
for i := range results {
select {
case r := <-resultCh:
results[i], lastDelay = r.data, r.delay
case <-h.closeCh:
return nil, 0, false
}
}
return results, lastDelay, true
}

View file

@ -78,10 +78,10 @@ var (
processConfirms = big.NewInt(1) processConfirms = big.NewInt(1)
// The token bucket buffer limit for testing purpose. // The token bucket buffer limit for testing purpose.
testBufLimit = uint64(1000000) testBufLimit = uint64(6000)
// The buffer recharging speed for testing purpose. // The buffer recharging speed for testing purpose.
testBufRecharge = uint64(1000) testBufRecharge = uint64(1)
) )
/* /*
@ -281,7 +281,7 @@ func newTestServerHandler(blocks int, indexers []*core.ChainIndexer, db ethdb.Da
} }
server.costTracker, server.freeCapacity = newCostTracker(db, server.config) server.costTracker, server.freeCapacity = newCostTracker(db, server.config)
server.costTracker.testCostList = testCostList(0) // Disable flow control mechanism. server.costTracker.testCostList = testCostList(0) // Disable flow control mechanism.
server.clientPool = newClientPool(db, 1, clock, nil) server.clientPool = newClientPool(db, 1, 1, clock, nil)
server.clientPool.setLimits(10000, 10000) // Assign enough capacity for clientpool server.clientPool.setLimits(10000, 10000) // Assign enough capacity for clientpool
server.handler = newServerHandler(server, simulation.Blockchain(), db, txpool, func() bool { return true }) server.handler = newServerHandler(server, simulation.Blockchain(), db, txpool, func() bool { return true })
if server.oracle != nil { if server.oracle != nil {
@ -395,8 +395,13 @@ func (p *testPeer) handshake(t *testing.T, td *big.Int, head common.Hash, headNu
expList = expList.add("serveStateSince", uint64(0)) expList = expList.add("serveStateSince", uint64(0))
expList = expList.add("serveRecentState", uint64(core.TriesInMemory-4)) expList = expList.add("serveRecentState", uint64(core.TriesInMemory-4))
expList = expList.add("txRelay", nil) expList = expList.add("txRelay", nil)
if p.peer.version >= lpv4 {
expList = expList.add("flowControl/BL", uint64(0))
expList = expList.add("flowControl/MRR", uint64(0))
} else {
expList = expList.add("flowControl/BL", testBufLimit) expList = expList.add("flowControl/BL", testBufLimit)
expList = expList.add("flowControl/MRR", testBufRecharge) expList = expList.add("flowControl/MRR", testBufRecharge)
}
expList = expList.add("flowControl/MRC", costList) expList = expList.add("flowControl/MRC", costList)
if err := p2p.ExpectMsg(p.app, StatusMsg, expList); err != nil { if err := p2p.ExpectMsg(p.app, StatusMsg, expList); err != nil {
@ -408,6 +413,16 @@ func (p *testPeer) handshake(t *testing.T, td *big.Int, head common.Hash, headNu
p.cpeer.fcParams = flowcontrol.ServerParams{ p.cpeer.fcParams = flowcontrol.ServerParams{
BufLimit: testBufLimit, BufLimit: testBufLimit,
MinRecharge: testBufRecharge, MinRecharge: testBufRecharge,
}
func (p *testPeer) expectCapUpdate(t *testing.T) {
if p.peer.version >= lpv4 {
var expList keyValueList
expList = expList.add("flowControl/BL", testBufLimit)
expList = expList.add("flowControl/MRR", testBufRecharge)
if err := p2p.ExpectMsg(p.app, AnnounceMsg, announceData{Update: expList}); err != nil {
t.Fatalf("status recv: %v", err)
}
} }
} }
@ -438,7 +453,7 @@ type testServer struct {
bloomTrieIndexer *core.ChainIndexer bloomTrieIndexer *core.ChainIndexer
} }
func newServerEnv(t *testing.T, blocks int, protocol int, callback indexerCallback, simClock bool, newPeer bool, testCost uint64) (*testServer, func()) { func newServerEnv(t *testing.T, blocks int, protocol int, callback indexerCallback, simClock bool, newPeer bool, testCost uint64, expectCapUpdate bool) (*testServer, func()) {
db := rawdb.NewMemoryDatabase() db := rawdb.NewMemoryDatabase()
indexers := testIndexers(db, nil, light.TestServerIndexerConfig) indexers := testIndexers(db, nil, light.TestServerIndexerConfig)
@ -480,6 +495,9 @@ func newServerEnv(t *testing.T, blocks int, protocol int, callback indexerCallba
cIndexer.Close() cIndexer.Close()
bIndexer.Close() bIndexer.Close()
} }
if expectCapUpdate {
server.peer.expectCapUpdate(t)
}
return server, teardown return server, teardown
} }

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.
@ -124,9 +124,9 @@ func connect(server *serverHandler, serverId enode.ID, client *clientHandler, pr
return peer1, peer2, nil return peer1, peer2, nil
} }
// newTestServerPeer creates server peer. // newServerPeer creates server peer.
func newTestServerPeer(t *testing.T, blocks int, protocol int) (*testServer, *enode.Node, func()) { func newTestServerPeer(t *testing.T, blocks int, protocol int) (*testServer, *enode.Node, func()) {
s, teardown := newServerEnv(t, blocks, protocol, nil, false, false, 0) s, teardown := newServerEnv(t, blocks, protocol, nil, false, false, 0, false)
key, err := crypto.GenerateKey() key, err := crypto.GenerateKey()
if err != nil { if err != nil {
t.Fatal("generate key err:", err) t.Fatal("generate key err:", err)