From 46fa5c031b4bfe82fb5b933d8fd2c0d9880da248 Mon Sep 17 00:00:00 2001 From: Zsolt Felfoldi Date: Fri, 3 Jan 2020 20:51:45 +0100 Subject: [PATCH] les: implement les/4 protocol extensions --- cmd/geth/main.go | 1 + cmd/geth/usage.go | 1 + cmd/utils/flags.go | 7 + eth/config.go | 3 + eth/gen_config.go | 6 + les/client_handler.go | 179 ++++++++++++++++++------- les/handler_test.go | 107 +++++++++------ les/metrics.go | 5 + les/odr_test.go | 12 +- les/peer.go | 296 ++++++++++++++++++++++++++++++++---------- les/protocol.go | 40 +++++- les/request_test.go | 8 +- les/retrieve.go | 2 +- les/server.go | 22 +++- les/server_handler.go | 169 ++++++++++++++++++++---- les/test_helper.go | 30 ++++- les/ulc_test.go | 6 +- 17 files changed, 694 insertions(+), 200 deletions(-) diff --git a/cmd/geth/main.go b/cmd/geth/main.go index 99ef78238f..70bcf7a962 100644 --- a/cmd/geth/main.go +++ b/cmd/geth/main.go @@ -101,6 +101,7 @@ var ( utils.UltraLightServersFlag, utils.UltraLightFractionFlag, utils.UltraLightOnlyAnnounceFlag, + utils.LespayTestModuleFlag, utils.WhitelistFlag, utils.CacheFlag, utils.CacheDatabaseFlag, diff --git a/cmd/geth/usage.go b/cmd/geth/usage.go index 6f3197b9c6..4d98a8e372 100644 --- a/cmd/geth/usage.go +++ b/cmd/geth/usage.go @@ -94,6 +94,7 @@ var AppHelpFlagGroups = []flagGroup{ utils.UltraLightServersFlag, utils.UltraLightFractionFlag, utils.UltraLightOnlyAnnounceFlag, + utils.LespayTestModuleFlag, }, }, { diff --git a/cmd/utils/flags.go b/cmd/utils/flags.go index bdadebd852..1002024792 100644 --- a/cmd/utils/flags.go +++ b/cmd/utils/flags.go @@ -272,6 +272,10 @@ var ( Usage: "Maximum number of light clients to serve, or light servers to attach to", Value: eth.DefaultConfig.LightPeers, } + LespayTestModuleFlag = cli.BoolFlag{ + Name: "lespay.testmodule", + Usage: "Enable dummy payment module (for testing only)", + } UltraLightServersFlag = cli.StringFlag{ Name: "ulc.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) { 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 diff --git a/eth/config.go b/eth/config.go index 2eaf21fbc3..4a2ea18ad1 100644 --- a/eth/config.go +++ b/eth/config.go @@ -116,6 +116,9 @@ type Config struct { 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 + // Light client payment options + LespayTestModule bool + // Database options SkipBcVersionCheck bool `toml:"-"` DatabaseHandles int `toml:"-"` diff --git a/eth/gen_config.go b/eth/gen_config.go index 1c659c393c..eb889e3aa8 100644 --- a/eth/gen_config.go +++ b/eth/gen_config.go @@ -32,6 +32,7 @@ func (c Config) MarshalTOML() (interface{}, error) { UltraLightServers []string `toml:",omitempty"` UltraLightFraction int `toml:",omitempty"` UltraLightOnlyAnnounce bool `toml:",omitempty"` + LespayTestModule bool `toml:"-"` SkipBcVersionCheck bool `toml:"-"` DatabaseHandles int `toml:"-"` DatabaseCache int @@ -68,6 +69,7 @@ func (c Config) MarshalTOML() (interface{}, error) { enc.UltraLightServers = c.UltraLightServers enc.UltraLightFraction = c.UltraLightFraction enc.UltraLightOnlyAnnounce = c.UltraLightOnlyAnnounce + enc.LespayTestModule = c.LespayTestModule enc.SkipBcVersionCheck = c.SkipBcVersionCheck enc.DatabaseHandles = c.DatabaseHandles enc.DatabaseCache = c.DatabaseCache @@ -108,6 +110,7 @@ func (c *Config) UnmarshalTOML(unmarshal func(interface{}) error) error { UltraLightServers []string `toml:",omitempty"` UltraLightFraction *int `toml:",omitempty"` UltraLightOnlyAnnounce *bool `toml:",omitempty"` + LespayTestModule *bool `toml:"-"` SkipBcVersionCheck *bool `toml:"-"` DatabaseHandles *int `toml:"-"` DatabaseCache *int @@ -175,6 +178,9 @@ func (c *Config) UnmarshalTOML(unmarshal func(interface{}) error) error { if dec.UltraLightOnlyAnnounce != nil { c.UltraLightOnlyAnnounce = *dec.UltraLightOnlyAnnounce } + if dec.LespayTestModule != nil { + c.LespayTestModule = *dec.LespayTestModule + } if dec.SkipBcVersionCheck != nil { c.SkipBcVersionCheck = *dec.SkipBcVersionCheck } diff --git a/les/client_handler.go b/les/client_handler.go index d04574c8c7..a54187cb8e 100644 --- a/les/client_handler.go +++ b/les/client_handler.go @@ -40,6 +40,9 @@ type clientHandler struct { downloader *downloader.Downloader backend *LightEthereum + lespayReplyHandlers map[uint64]func([]byte, uint) bool + lespayReplyLock sync.Mutex + closeCh chan struct{} wg sync.WaitGroup // WaitGroup used to track all connected peers. syncDone func() // Test hooks when syncing is done. @@ -47,9 +50,10 @@ type clientHandler struct { func newClientHandler(ulcServers []string, ulcFraction int, checkpoint *params.TrustedCheckpoint, backend *LightEthereum) *clientHandler { handler := &clientHandler{ - checkpoint: checkpoint, - backend: backend, - closeCh: make(chan struct{}), + checkpoint: checkpoint, + backend: backend, + closeCh: make(chan struct{}), + lespayReplyHandlers: make(map[uint64]func([]byte, uint) bool), } if ulcServers != nil { 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) return err } - // Register the peer locally - if err := h.backend.peers.register(p); err != nil { - p.Log().Error("Light Ethereum peer registration failed", "err", err) - return err - } - serverConnectionGauge.Update(int64(h.backend.peers.len())) - connectedAt := mclock.Now() - defer func() { - h.backend.peers.unregister(p.id) + var ( + connectedAt mclock.AbsTime + lastActive bool + ) + activate := func() { + // Register the peer locally + if err := h.backend.peers.register(p); err != nil { + p.Log().Error("Light Ethereum peer registration failed", "err", err) + return + } + serverConnectionGauge.Update(int64(h.backend.peers.len())) + connectedAt = mclock.Now() + h.fetcher.announce(p, &announceData{Hash: p.headInfo.Hash, Number: p.headInfo.Number, Td: p.headInfo.Td}) + lastActive = true + } + deactivate := func() { + h.backend.peers.unregister(p) connectionTimer.Update(time.Duration(mclock.Now() - connectedAt)) 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. if p.poolEntry != nil { h.backend.serverPool.registered(p.poolEntry) } + // Spawn a main loop to handle all incoming messages. for { + if p.active && !lastActive { + activate() + } + if !p.active && lastActive { + deactivate() + } if err := h.handleMsg(p); err != nil { p.Log().Debug("Light Ethereum message handling failed", "err", err) p.fcServer.DumpLogs() @@ -157,7 +181,10 @@ func (h *clientHandler) handleMsg(p *serverPeer) error { } defer msg.Discard() - var deliverMsg *Msg + var ( + deliverMsg *Msg + responseError bool + ) // Handle the message depending on its contents switch msg.Code { @@ -193,13 +220,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) error { case BlockHeadersMsg: p.Log().Trace("Received block header response message") var resp struct { - ReqID, BV uint64 - Headers []*types.Header + ReqID uint64 + SF stateFeedback + Headers []*types.Header } + resp.SF.protocolVersion = p.version if err := msg.Decode(&resp); err != nil { 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) { h.fetcher.deliverHeaders(p, resp.ReqID, resp.Headers) } else { @@ -210,13 +239,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) error { case BlockBodiesMsg: p.Log().Trace("Received block bodies response") var resp struct { - ReqID, BV uint64 - Data []*types.Body + ReqID uint64 + SF stateFeedback + Data []*types.Body } + resp.SF.protocolVersion = p.version if err := msg.Decode(&resp); err != nil { 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{ MsgType: MsgBlockBodies, ReqID: resp.ReqID, @@ -225,13 +256,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) error { case CodeMsg: p.Log().Trace("Received code response") var resp struct { - ReqID, BV uint64 - Data [][]byte + ReqID uint64 + SF stateFeedback + Data [][]byte } + resp.SF.protocolVersion = p.version if err := msg.Decode(&resp); err != nil { 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{ MsgType: MsgCode, ReqID: resp.ReqID, @@ -240,13 +273,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) error { case ReceiptsMsg: p.Log().Trace("Received receipts response") var resp struct { - ReqID, BV uint64 - Receipts []types.Receipts + ReqID uint64 + SF stateFeedback + Receipts []types.Receipts } + resp.SF.protocolVersion = p.version if err := msg.Decode(&resp); err != nil { 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{ MsgType: MsgReceipts, ReqID: resp.ReqID, @@ -255,13 +290,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) error { case ProofsV2Msg: p.Log().Trace("Received les/2 proofs response") var resp struct { - ReqID, BV uint64 - Data light.NodeList + ReqID uint64 + SF stateFeedback + Data light.NodeList } + resp.SF.protocolVersion = p.version if err := msg.Decode(&resp); err != nil { 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{ MsgType: MsgProofsV2, ReqID: resp.ReqID, @@ -270,13 +307,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) error { case HelperTrieProofsMsg: p.Log().Trace("Received helper trie proof response") var resp struct { - ReqID, BV uint64 - Data HelperTrieResps + ReqID uint64 + SF stateFeedback + Data HelperTrieResps } + resp.SF.protocolVersion = p.version if err := msg.Decode(&resp); err != nil { 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{ MsgType: MsgHelperTrieProofs, ReqID: resp.ReqID, @@ -285,13 +324,15 @@ func (h *clientHandler) handleMsg(p *serverPeer) error { case TxStatusMsg: p.Log().Trace("Received tx status response") var resp struct { - ReqID, BV uint64 - Status []light.TxStatus + ReqID uint64 + SF stateFeedback + Status []light.TxStatus } + resp.SF.protocolVersion = p.version if err := msg.Decode(&resp); err != nil { 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{ MsgType: MsgTxStatus, ReqID: resp.ReqID, @@ -302,13 +343,32 @@ func (h *clientHandler) handleMsg(p *serverPeer) error { h.backend.retriever.frozen(p) p.Log().Debug("Service stopped") case ResumeMsg: - var bv uint64 - if err := msg.Decode(&bv); err != nil { + var sf stateFeedback + sf.protocolVersion = p.version + if err := msg.Decode(&sf); err != nil { return errResp(ErrDecode, "msg %v: %v", msg, err) } - p.fcServer.ResumeFreeze(bv) + p.fcServer.ResumeFreeze(sf.BV) p.unfreeze() 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: p.Log().Trace("Received invalid message", "code", 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. if deliverMsg != nil { if err := h.backend.retriever.deliver(p, deliverMsg); err != nil { - p.errCount++ - if p.errCount > maxResponseErrors { - return err - } + responseError = true + } + } + if responseError { + p.errCount++ + if p.errCount > maxResponseErrors { + return err } } 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) { - h.backend.peers.unregister(id) + h.backend.peers.disconnect(id) } type peerConnection struct { diff --git a/les/handler_test.go b/les/handler_test.go index 1612caf427..8107270354 100644 --- a/les/handler_test.go +++ b/les/handler_test.go @@ -38,20 +38,29 @@ import ( "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 { - ReqID, BV uint64 - Data interface{} + ReqID uint64 + 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. -func TestGetBlockHeadersLes2(t *testing.T) { testGetBlockHeaders(t, 2) } func TestGetBlockHeadersLes3(t *testing.T) { testGetBlockHeaders(t, 3) } +func TestGetBlockHeadersLes4(t *testing.T) { testGetBlockHeaders(t, 4) } 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() bc := server.handler.blockchain @@ -168,19 +177,20 @@ func testGetBlockHeaders(t *testing.T, protocol int) { // Send the hash request and verify the response reqID++ + cost := server.peer.peer.GetRequestCost(GetBlockHeadersMsg, int(tt.query.Amount)) 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) } } } // 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 TestGetBlockBodiesLes4(t *testing.T) { testGetBlockBodies(t, 4) } 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() bc := server.handler.blockchain @@ -245,20 +255,21 @@ func testGetBlockBodies(t *testing.T, protocol int) { reqID++ // Send the hash request and verify the response + cost := server.peer.peer.GetRequestCost(GetBlockBodiesMsg, len(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) } } } // 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 TestGetCodeLes4(t *testing.T) { testGetCode(t, 4) } func testGetCode(t *testing.T, protocol int) { // 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() 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) - 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) } } // 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 TestGetStaleCodeLes4(t *testing.T) { testGetStaleCode(t, 4) } 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() bc := server.handler.blockchain @@ -296,8 +308,9 @@ func testGetStaleCode(t *testing.T, protocol int) { BHash: bc.GetHeaderByNumber(number).Hash(), AccKey: crypto.Keccak256(testContractAddr[:]), } + cost := server.peer.peer.GetRequestCost(GetCodeMsg, 1) 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) } } @@ -307,12 +320,12 @@ func testGetStaleCode(t *testing.T, protocol int) { } // 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 TestGetReceiptLes4(t *testing.T) { testGetReceipt(t, 4) } func testGetReceipt(t *testing.T, protocol int) { // 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() 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())) } // Send the hash request and verify the response + cost := server.peer.peer.GetRequestCost(GetReceiptsMsg, len(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) } } // 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 TestGetProofsLes4(t *testing.T) { testGetProofs(t, 4) } func testGetProofs(t *testing.T, protocol int) { // 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() bc := server.handler.blockchain @@ -362,18 +376,19 @@ func testGetProofs(t *testing.T, protocol int) { } } // Send the proof request and verify the response + cost := server.peer.peer.GetRequestCost(GetProofsV2Msg, len(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) } } // 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 TestGetStaleProofLes4(t *testing.T) { testGetStaleProof(t, 4) } 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() bc := server.handler.blockchain @@ -395,7 +410,7 @@ func testGetStaleProof(t *testing.T, protocol int) { t.Prove(account, 0, proofsV2) 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) } } @@ -405,8 +420,8 @@ func testGetStaleProof(t *testing.T, protocol int) { } // 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 TestGetCHTProofsLes4(t *testing.T) { testGetCHTProofs(t, 4) } func testGetCHTProofs(t *testing.T, protocol int) { config := light.TestServerIndexerConfig @@ -420,7 +435,7 @@ func testGetCHTProofs(t *testing.T, protocol int) { 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() bc := server.handler.blockchain @@ -446,14 +461,15 @@ func testGetCHTProofs(t *testing.T, protocol int) { AuxReq: auxHeader, }} // Send the proof request and verify the response + cost := server.peer.peer.GetRequestCost(GetHelperTrieProofsMsg, len(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) } } -func TestGetBloombitsProofsLes2(t *testing.T) { testGetBloombitsProofs(t, 2) } func TestGetBloombitsProofsLes3(t *testing.T) { testGetBloombitsProofs(t, 3) } +func TestGetBloombitsProofsLes4(t *testing.T) { testGetBloombitsProofs(t, 4) } // Tests that bloombits proofs can be correctly retrieved. func testGetBloombitsProofs(t *testing.T, protocol int) { @@ -468,7 +484,7 @@ func testGetBloombitsProofs(t *testing.T, protocol int) { 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() bc := server.handler.blockchain @@ -494,18 +510,19 @@ func testGetBloombitsProofs(t *testing.T, protocol int) { trie.Prove(key, 0, &proofs.Proofs) // Send the proof request and verify the response + cost := server.peer.peer.GetRequestCost(GetHelperTrieProofsMsg, len(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) } } } -func TestTransactionStatusLes2(t *testing.T) { testTransactionStatus(t, 2) } func TestTransactionStatusLes3(t *testing.T) { testTransactionStatus(t, 3) } +func TestTransactionStatusLes4(t *testing.T) { testTransactionStatus(t, 4) } 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() 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) { reqID++ + var cost uint64 if send { + cost = server.peer.peer.GetRequestCost(SendTxV2Msg, 1) sendRequest(server.peer.app, SendTxV2Msg, reqID, types.Transactions{tx}) } else { + cost = server.peer.peer.GetRequestCost(GetTxStatusMsg, 1) 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") } } @@ -595,8 +615,11 @@ func testTransactionStatus(t *testing.T, protocol int) { test(tx2, false, light.TxStatus{Status: core.TxStatusPending}) } -func TestStopResumeLes3(t *testing.T) { - server, tearDown := newServerEnv(t, 0, 3, nil, true, true, testBufLimit/10) +func TestStopResumeLes3(t *testing.T) { testStopResume(t, 3) } +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() server.handler.server.costTracker.testing = true @@ -616,7 +639,7 @@ func TestStopResumeLes3(t *testing.T) { for expBuf >= testCost { req() 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) } } @@ -635,7 +658,15 @@ func TestStopResumeLes3(t *testing.T) { // expect a ResumeMsg with the partially recharged buffer value 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) } } diff --git a/les/metrics.go b/les/metrics.go index 9ef8c36518..12780346b6 100644 --- a/les/metrics.go +++ b/les/metrics.go @@ -40,6 +40,8 @@ var ( miscInTxsTrafficMeter = metrics.NewRegisteredMeter("les/misc/in/traffic/txs", nil) miscInTxStatusPacketsMeter = metrics.NewRegisteredMeter("les/misc/in/packets/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) miscOutTrafficMeter = metrics.NewRegisteredMeter("les/misc/out/traffic/total", nil) @@ -59,6 +61,8 @@ var ( miscOutTxsTrafficMeter = metrics.NewRegisteredMeter("les/misc/out/traffic/txs", nil) miscOutTxStatusPacketsMeter = metrics.NewRegisteredMeter("les/misc/out/packets/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) miscServingTimeBodyTimer = metrics.NewRegisteredTimer("les/misc/serve/body", nil) @@ -68,6 +72,7 @@ var ( miscServingTimeHelperTrieTimer = metrics.NewRegisteredTimer("les/misc/serve/helperTrie", nil) miscServingTimeTxTimer = metrics.NewRegisteredTimer("les/misc/serve/txs", nil) miscServingTimeTxStatusTimer = metrics.NewRegisteredTimer("les/misc/serve/txStatus", nil) + miscServingTimeLespayTimer = metrics.NewRegisteredTimer("les/misc/serve/lespay", nil) connectionTimer = metrics.NewRegisteredTimer("les/connection/duration", nil) serverConnectionGauge = metrics.NewRegisteredGauge("les/connection/server", nil) diff --git a/les/odr_test.go b/les/odr_test.go index bbe439dfec..a56de31769 100644 --- a/les/odr_test.go +++ b/les/odr_test.go @@ -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 -func TestOdrGetBlockLes2(t *testing.T) { testOdr(t, 2, 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 { var block *types.Block @@ -55,8 +55,8 @@ func odrGetBlock(ctx context.Context, db ethdb.Database, config *params.ChainCon 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 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 { var receipts types.Receipts @@ -76,8 +76,8 @@ func odrGetReceipts(ctx context.Context, db ethdb.Database, config *params.Chain 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 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 { dummyAddr := common.HexToAddress("1234567812345678123456781234567812345678") @@ -105,8 +105,8 @@ func odrAccounts(ctx context.Context, db ethdb.Database, config *params.ChainCon 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 TestOdrContractCallLes4(t *testing.T) { testOdr(t, 4, 2, true, odrContractCall) } type callmsg struct { types.Message @@ -155,8 +155,8 @@ func odrContractCall(ctx context.Context, db ethdb.Database, config *params.Chai 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 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 { 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 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 test(5) } diff --git a/les/peer.go b/les/peer.go index 28ec201bc9..41ea7194bc 100644 --- a/les/peer.go +++ b/les/peer.go @@ -132,6 +132,7 @@ type peerCommons struct { frozen uint32 // Flag whether the peer is frozen. announceType uint64 // New block announcement type. headInfo blockInfo // Latest block information. + active bool // Background task queue for caching peer tasks and executing in order. sendQueue *execQueue @@ -478,12 +479,18 @@ func (p *serverPeer) requestTxStatus(reqID uint64, txHashes []common.Hash) error 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 { p.Log().Debug("Sending batch of transactions", "size", len(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 func (p *serverPeer) waitBefore(maxCost uint64) (time.Duration, float64) { 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. var params flowcontrol.ServerParams + updated := false if update.get("flowControl/BL", ¶ms.BufLimit) == nil && update.get("flowControl/MRR", ¶ms.MinRecharge) == nil { // todo can light client set a minimal acceptable flow control params? p.fcParams = params p.fcServer.UpdateParams(params) + updated = true } var MRC RequestCostList if update.get("flowControl/MRC", &MRC) == nil { @@ -565,7 +574,18 @@ func (p *serverPeer) updateFlowControl(update keyValueMap) { for code, cost := range costUpdate { 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, @@ -634,6 +654,9 @@ func (p *serverPeer) Handshake(td *big.Int, head common.Hash, headNum uint64, ge type clientPeer struct { peerCommons + activate, deactivate func() + getBalance func() uint64 + // responseLock ensures that responses are queued in the same order as // RequestProcessed is called responseLock sync.Mutex @@ -681,8 +704,8 @@ func (p *clientPeer) sendStop() error { } // sendResume notifies the client about getting out of frozen state -func (p *clientPeer) sendResume(bv uint64) error { - return p2p.Send(p.rw, ResumeMsg, bv) +func (p *clientPeer) sendResume(sf stateFeedback) error { + return p2p.Send(p.rw, ResumeMsg, sf) } // freeze temporarily puts the client in a frozen state which means all unprocessed @@ -711,7 +734,19 @@ func (p *clientPeer) freeze() { continue } 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 } }() @@ -728,12 +763,13 @@ type reply struct { } // 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 { - ReqID, BV uint64 - Data rlp.RawValue + ReqID uint64 + 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 @@ -786,6 +822,12 @@ func (p *clientPeer) replyTxStatus(reqID uint64, stats []light.TxStatus) *reply 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 // a hash notification. func (p *clientPeer) sendAnnounce(request announceData) error { @@ -798,12 +840,21 @@ func (p *clientPeer) updateCapacity(cap uint64) { p.lock.Lock() defer p.lock.Unlock() - p.fcParams = flowcontrol.ServerParams{MinRecharge: cap, BufLimit: cap * bufLimitRatio} - p.fcClient.UpdateParams(p.fcParams) - var kvList keyValueList - kvList = kvList.add("flowControl/MRR", cap) - kvList = kvList.add("flowControl/BL", cap*bufLimitRatio) - p.mustQueueSend(func() { p.sendAnnounce(announceData{Update: kvList}) }) + 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.fcClient.UpdateParams(p.fcParams) + var kvList keyValueList + kvList = kvList.add("flowControl/BL", cap*bufLimitRatio) + kvList = kvList.add("flowControl/MRR", cap) + 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 @@ -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("txRelay", nil) } - *lists = (*lists).add("flowControl/BL", server.defParams.BufLimit) - *lists = (*lists).add("flowControl/MRR", server.defParams.MinRecharge) + p.active = p.version < lpv4 + 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 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) p.fcCosts = costList.decode(ProtocolLengths[uint(p.version)]) - p.fcParams = server.defParams // Add advertised checkpoint and register block height which // 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 p.announceType = announceTypeSimple } - p.fcClient = flowcontrol.NewClientNode(server.fcManager, server.defParams) + p.fcClient = flowcontrol.NewClientNode(server.fcManager, p.fcParams) } return nil }) @@ -913,7 +969,7 @@ type clientPeerSubscriber interface { // clientPeerSet represents the set of active client peers currently // participating in the Light Ethereum sub-protocol. type clientPeerSet struct { - peers map[string]*clientPeer + active, inactive map[string]*clientPeer // subscribers is a batch of subscribers and peerset will notify // these subscribers when the peerset changes(new client peer is // added or removed) @@ -924,17 +980,23 @@ type clientPeerSet struct { // newClientPeerSet creates a new peer set to track the client peers. 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 // peers and also register all active peers into the given service. func (ps *clientPeerSet) subscribe(sub clientPeerSubscriber) { ps.lock.Lock() - defer ps.lock.Unlock() - 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) } } @@ -956,17 +1018,23 @@ func (ps *clientPeerSet) unSubscribe(sub clientPeerSubscriber) { // peer is already known. func (ps *clientPeerSet) register(peer *clientPeer) error { ps.lock.Lock() - defer ps.lock.Unlock() - if ps.closed { + ps.lock.Unlock() return errClosed } - if _, exist := ps.peers[peer.id]; exist { + if _, ok := ps.active[p.id]; ok { + ps.lock.Unlock() return errAlreadyRegistered } - ps.peers[peer.id] = peer - for _, sub := range ps.subscribers { - sub.registerPeer(peer) + delete(ps.inactive, p.id) + ps.active[p.id] = p + + peers := make([]clientPeerSubscriber, len(ps.subscribers)) + copy(peers, ps.subscribers) + ps.lock.Unlock() + + for _, n := range peers { + n.registerPeer(p) } return nil } @@ -976,27 +1044,58 @@ func (ps *clientPeerSet) register(peer *clientPeer) error { // at the networking layer. func (ps *clientPeerSet) unregister(id string) error { 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] - if !ok { + for _, n := range peers { + 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 } - delete(ps.peers, id) - for _, sub := range ps.subscribers { - sub.unregisterPeer(p) + ps.lock.Unlock() + for _, n := range peers { + n.unregisterPeer(p) } - p.Peer.Disconnect(p2p.DiscRequested) + p.Peer.Disconnect(p2p.DiscUselessPeer) 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 { ps.lock.RLock() defer ps.lock.RUnlock() var ids []string - for id := range ps.peers { + for id := range ps.active { ids = append(ids, id) } return ids @@ -1007,24 +1106,27 @@ func (ps *clientPeerSet) peer(id string) *clientPeer { ps.lock.RLock() 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 { ps.lock.RLock() 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 { ps.lock.RLock() defer ps.lock.RUnlock() list := make([]*clientPeer, 0, len(ps.peers)) - for _, p := range ps.peers { + for _, p := range ps.active { list = append(list, p) } return list @@ -1036,7 +1138,10 @@ func (ps *clientPeerSet) close() { ps.lock.Lock() 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) } ps.closed = true @@ -1045,7 +1150,7 @@ func (ps *clientPeerSet) close() { // serverPeerSet represents the set of active server peers currently // participating in the Light Ethereum sub-protocol. type serverPeerSet struct { - peers map[string]*serverPeer + active, inactive map[string]*serverPeer // subscribers is a batch of subscribers and peerset will notify // these subscribers when the peerset changes(new server peer is // added or removed) @@ -1056,17 +1161,23 @@ type serverPeerSet struct { // newServerPeerSet creates a new peer set to track the active server peers. 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 // peers and also register all active peers into the given service. func (ps *serverPeerSet) subscribe(sub serverPeerSubscriber) { ps.lock.Lock() - defer ps.lock.Unlock() - 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) } } @@ -1088,17 +1199,23 @@ func (ps *serverPeerSet) unSubscribe(sub serverPeerSubscriber) { // peer is already known. func (ps *serverPeerSet) register(peer *serverPeer) error { ps.lock.Lock() - defer ps.lock.Unlock() - if ps.closed { + ps.lock.Unlock() return errClosed } - if _, exist := ps.peers[peer.id]; exist { + if _, ok := ps.active[p.id]; ok { + ps.lock.Unlock() return errAlreadyRegistered } - ps.peers[peer.id] = peer - for _, sub := range ps.subscribers { - sub.registerPeer(peer) + delete(ps.inactive, p.id) + ps.active[p.id] = p + + peers := make([]serverPeerSubscriber, len(ps.subscribers)) + copy(peers, ps.subscribers) + ps.lock.Unlock() + + for _, n := range peers { + n.registerPeer(p) } return nil } @@ -1108,27 +1225,58 @@ func (ps *serverPeerSet) register(peer *serverPeer) error { // the networking layer. func (ps *serverPeerSet) unregister(id string) error { 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] - if !ok { + for _, n := range peers { + 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 } - delete(ps.peers, id) - for _, sub := range ps.subscribers { - sub.unregisterPeer(p) + ps.lock.Unlock() + for _, n := range peers { + n.unregisterPeer(p) } - p.Peer.Disconnect(p2p.DiscRequested) + p.Peer.Disconnect(p2p.DiscUselessPeer) 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 { ps.lock.RLock() defer ps.lock.RUnlock() var ids []string - for id := range ps.peers { + for id := range ps.active { ids = append(ids, id) } return ids @@ -1139,15 +1287,18 @@ func (ps *serverPeerSet) peer(id string) *serverPeer { ps.lock.RLock() 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 { ps.lock.RLock() defer ps.lock.RUnlock() - return len(ps.peers) + return len(ps.active) } // bestPeer retrieves the known peer with the currently highest total difficulty. @@ -1161,7 +1312,7 @@ func (ps *serverPeerSet) bestPeer() *serverPeer { bestPeer *serverPeer bestTd *big.Int ) - for _, p := range ps.peers { + for _, p := range ps.active { if td := p.Td(); bestTd == nil || td.Cmp(bestTd) > 0 { bestPeer, bestTd = p, td } @@ -1169,12 +1320,12 @@ func (ps *serverPeerSet) bestPeer() *serverPeer { return bestPeer } -// allServerPeers returns all server peers in a list. +// allPeers returns all active server peers in a list. func (ps *serverPeerSet) allPeers() []*serverPeer { ps.lock.RLock() defer ps.lock.RUnlock() - list := make([]*serverPeer, 0, len(ps.peers)) + list := make([]*serverPeer, 0, len(ps.active)) for _, p := range ps.peers { list = append(list, p) } @@ -1187,7 +1338,10 @@ func (ps *serverPeerSet) close() { ps.lock.Lock() 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) } ps.closed = true diff --git a/les/protocol.go b/les/protocol.go index 36af88aea6..5140795ecc 100644 --- a/les/protocol.go +++ b/les/protocol.go @@ -33,17 +33,18 @@ import ( const ( lpv2 = 2 lpv3 = 3 + lpv4 = 4 ) // Supported versions of the les protocol (first is primary) var ( - ClientProtocolVersions = []uint{lpv2, lpv3} - ServerProtocolVersions = []uint{lpv2, lpv3} + ClientProtocolVersions = []uint{lpv2, lpv3, lpv4} + ServerProtocolVersions = []uint{lpv2, lpv3, lpv4} AdvertiseProtocolVersions = []uint{lpv2} // clients are searching for the first advertised protocol in the list ) // 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 ( NetworkId = 1 @@ -74,6 +75,9 @@ const ( // Protocol messages introduced in LPV3 StopMsg = 0x16 ResumeMsg = 0x17 + // Protocol messages introduced in LPV4 + LespayMsg = 0x18 + LespayReplyMsg = 0x19 ) type requestInfo struct { @@ -201,6 +205,11 @@ type hashOrNumber struct { 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 // two contained union fields. func (hn *hashOrNumber) EncodeRLP(w io.Writer) error { @@ -235,3 +244,28 @@ func (hn *hashOrNumber) DecodeRLP(s *rlp.Stream) error { type CodeData []struct { 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) + } +} diff --git a/les/request_test.go b/les/request_test.go index f58ebca9c1..f5e370a6ba 100644 --- a/les/request_test.go +++ b/les/request_test.go @@ -36,22 +36,22 @@ func secAddr(addr common.Address) []byte { 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 TestBlockAccessLes4(t *testing.T) { testAccess(t, 4, tfBlockAccess) } func tfBlockAccess(db ethdb.Database, bhash common.Hash, number uint64) light.OdrRequest { 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 TestReceiptsAccessLes4(t *testing.T) { testAccess(t, 4, tfReceiptsAccess) } func tfReceiptsAccess(db ethdb.Database, bhash common.Hash, number uint64) light.OdrRequest { 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 TestTrieEntryAccessLes4(t *testing.T) { testAccess(t, 4, tfTrieEntryAccess) } func tfTrieEntryAccess(db ethdb.Database, bhash common.Hash, number uint64) light.OdrRequest { 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 } -func TestCodeAccessLes2(t *testing.T) { testAccess(t, 2, 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 { number := rawdb.ReadHeaderNumber(db, bhash) diff --git a/les/retrieve.go b/les/retrieve.go index 5fa68b7456..45976bc9ec 100644 --- a/les/retrieve.go +++ b/les/retrieve.go @@ -345,7 +345,7 @@ func (r *sentReq) tryRequest() { if hrto { pp.Log().Debug("Request timed out hard") if r.rm.peers != nil { - r.rm.peers.unregister(pp.id) + r.rm.peers.disconnect(pp.id) } } diff --git a/les/server.go b/les/server.go index f72f31321a..d6225a9c05 100644 --- a/les/server.go +++ b/les/server.go @@ -44,6 +44,7 @@ type LesServer struct { handler *serverHandler lesTopics []discv5.Topic privateKey *ecdsa.PrivateKey + srvr *p2p.Server // Flow control and capacity management fcManager *flowcontrol.ClientManager @@ -51,6 +52,7 @@ type LesServer struct { defParams flowcontrol.ServerParams servingQueue *servingQueue clientPool *clientPool + tokenSale *tokenSale minCapacity, maxCapacity, freeCapacity uint64 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.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.tokenSale = newTokenSale(srv.clientPool, 0.1, 100) + if config.LespayTestModule { + srv.tokenSale.addReceiver("test", testReceiver{}) + srv.clientPool.setExpirationTCs(defaultPosExpTC, defaultNegExpTC) + } checkpoint := srv.latestLocalCheckpoint() if !checkpoint.Empty() { @@ -148,6 +155,12 @@ func (s *LesServer) APIs() []rpc.API { Service: NewPrivateDebugAPI(s), 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 func (s *LesServer) Start(srvr *p2p.Server) { + s.srvr = srvr s.privateKey = srvr.PrivateKey s.handler.start() @@ -174,6 +188,7 @@ func (s *LesServer) Start(srvr *p2p.Server) { go s.capacityManagement() if srvr.DiscV5 != nil { + srvr.DiscV5.RegisterTalkHandler("lespay", s.handler.talkRequestHandler) for _, topic := range s.lesTopics { topic := topic go func() { @@ -191,6 +206,11 @@ func (s *LesServer) Start(srvr *p2p.Server) { func (s *LesServer) Stop() { close(s.closeCh) + if s.srvr.DiscV5 != nil { + s.srvr.DiscV5.RemoveTalkHandler("lespay") + } + s.tokenSale.stop() + // Disconnect existing sessions. // 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 diff --git a/les/server_handler.go b/les/server_handler.go index 186bdcbb03..07bf821102 100644 --- a/les/server_handler.go +++ b/les/server_handler.go @@ -20,6 +20,7 @@ import ( "encoding/binary" "encoding/json" "errors" + "net" "sync" "sync/atomic" "time" @@ -35,6 +36,7 @@ import ( "github.com/ethereum/go-ethereum/log" "github.com/ethereum/go-ethereum/metrics" "github.com/ethereum/go-ethereum/p2p" + "github.com/ethereum/go-ethereum/p2p/enode" "github.com/ethereum/go-ethereum/rlp" "github.com/ethereum/go-ethereum/trie" ) @@ -54,10 +56,7 @@ const ( MaxTxStatus = 256 // Amount of transactions to queried per request ) -var ( - errTooManyInvalidRequest = errors.New("too many invalid requests made") - errFullClientPool = errors.New("client pool is full") -) +var errTooManyInvalidRequest = errors.New("too many invalid requests made") // serverHandler is responsible for serving light client and process // 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 { peer := newClientPeer(int(version), h.server.config.NetworkId, p, newMeteredMsgWriter(rw, int(version))) 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) defer h.wg.Done() return h.handle(peer) @@ -134,28 +136,58 @@ func (h *serverHandler) handle(p *clientPeer) error { } defer p.fcClient.Disconnect() - // Disconnect the inbound peer if it's rejected by clientPool - if !h.server.clientPool.connect(p, 0) { - p.Log().Debug("Light Ethereum peer registration failed", "err", errFullClientPool) - return errFullClientPool + var ( + connectedAt mclock.AbsTime + wg *sync.WaitGroup // Wait group used to track all in-flight task routines. + ) + p.activate = func() { + // Register the peer locally + if err := h.server.peers.Register(p); err != nil { + h.server.clientPool.disconnect(p) + p.Log().Error("Light Ethereum peer registration failed", "err", err) + return + } + clientConnectionGauge.Update(int64(h.server.peers.Len())) + connectedAt = mclock.Now() + wg = new(sync.WaitGroup) + p.active = true } - // Register the peer locally - if err := h.server.peers.register(p); err != nil { - h.server.clientPool.disconnect(p) - p.Log().Error("Light Ethereum peer registration failed", "err", err) + 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() + } + + 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) + } } - clientConnectionGauge.Update(int64(h.server.peers.len())) - var wg sync.WaitGroup // Wait group used to track all in-flight task routines. - - connectedAt := mclock.Now() defer func() { wg.Wait() // Ensure all background task routines have exited. - h.server.peers.unregister(p.id) h.server.clientPool.disconnect(p) - clientConnectionGauge.Update(int64(h.server.peers.len())) - connectionTimer.Update(time.Duration(mclock.Now() - connectedAt)) + p.responseLock.Lock() + 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. @@ -166,7 +198,7 @@ func (h *serverHandler) handle(p *clientPeer) error { return err 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) return err } @@ -245,22 +277,33 @@ func (h *serverHandler) handleMsg(p *clientPeer, wg *sync.WaitGroup) error { if reply != nil { replySize = reply.size() } - var realCost uint64 + var realCost, balance uint64 if h.server.costTracker.testing { realCost = maxCost // Assign a fake cost for testing purpose } else { realCost = h.server.costTracker.realCost(servingTime, msg.Size, replySize) + if realCost > maxCost { + realCost = maxCost + } } bv := p.fcClient.RequestProcessed(reqID, responseCount, maxCost, realCost) if amount != 0 { // Feed cost tracker request serving statistic. h.server.costTracker.updateStats(msg.Code, amount, servingTime, realCost) // 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 { p.mustQueueSend(func() { - if err := reply.send(bv); err != nil { + if err := reply.send(sf); err != nil { select { case p.errCh <- err: default: @@ -375,6 +418,8 @@ func (h *serverHandler) handleMsg(p *clientPeer, wg *sync.WaitGroup) error { } reply := p.replyBlockHeaders(req.ReqID, headers) 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 { miscOutHeaderPacketsMeter.Mark(1) 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: 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 +} diff --git a/les/test_helper.go b/les/test_helper.go index d9ffe32db2..b53c28f2e5 100644 --- a/les/test_helper.go +++ b/les/test_helper.go @@ -78,10 +78,10 @@ var ( processConfirms = big.NewInt(1) // The token bucket buffer limit for testing purpose. - testBufLimit = uint64(1000000) + testBufLimit = uint64(6000) // 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.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.handler = newServerHandler(server, simulation.Blockchain(), db, txpool, func() bool { return true }) 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("serveRecentState", uint64(core.TriesInMemory-4)) expList = expList.add("txRelay", nil) - expList = expList.add("flowControl/BL", testBufLimit) - expList = expList.add("flowControl/MRR", testBufRecharge) + 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/MRR", testBufRecharge) + } expList = expList.add("flowControl/MRC", costList) 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{ BufLimit: testBufLimit, 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 } -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() indexers := testIndexers(db, nil, light.TestServerIndexerConfig) @@ -480,6 +495,9 @@ func newServerEnv(t *testing.T, blocks int, protocol int, callback indexerCallba cIndexer.Close() bIndexer.Close() } + if expectCapUpdate { + server.peer.expectCapUpdate(t) + } return server, teardown } diff --git a/les/ulc_test.go b/les/ulc_test.go index 273c63e4bd..d26a782e23 100644 --- a/les/ulc_test.go +++ b/les/ulc_test.go @@ -28,8 +28,8 @@ import ( "github.com/ethereum/go-ethereum/p2p/enode" ) -func TestULCAnnounceThresholdLes2(t *testing.T) { testULCAnnounceThreshold(t, 2) } func TestULCAnnounceThresholdLes3(t *testing.T) { testULCAnnounceThreshold(t, 3) } +func TestULCAnnounceThresholdLes4(t *testing.T) { testULCAnnounceThreshold(t, 4) } func testULCAnnounceThreshold(t *testing.T, protocol int) { // 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 } -// newTestServerPeer creates server peer. +// newServerPeer creates server peer. 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() if err != nil { t.Fatal("generate key err:", err)