diff --git a/les/handler.go b/les/handler.go index 801853636f..9854d6b494 100644 --- a/les/handler.go +++ b/les/handler.go @@ -366,7 +366,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { var deliverMsg *Msg - sendReq := func(reqID, amount uint64, reply *reply, servingTime uint64) { + sendResponse := func(reqID, amount uint64, reply *reply, servingTime uint64) { p.responseLock.Lock() defer p.responseLock.Unlock() @@ -457,11 +457,13 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { unknown bool ) for !unknown && len(headers) < int(query.Amount) && bytes < softResponseLimit { + if !first && !task.waitOrStop() { + return + } // Retrieve the next header satisfying the query var origin *types.Header if hashMode { if first { - first = false origin = pm.blockchain.GetHeaderByHash(query.Origin.Hash) if origin != nil { query.Origin.Number = origin.Number.Uint64() @@ -524,8 +526,9 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { // Number based traversal towards the leaf block query.Origin.Number += query.Skip + 1 } + first = false } - sendReq(req.ReqID, query.Amount, p.ReplyBlockHeaders(req.ReqID, headers), task.done()) + sendResponse(req.ReqID, query.Amount, p.ReplyBlockHeaders(req.ReqID, headers), task.done()) }() case BlockHeadersMsg: @@ -572,7 +575,10 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { return errResp(ErrRequestRejected, "") } go func() { - for _, hash := range req.Hashes { + for i, hash := range req.Hashes { + if i != 0 && !task.waitOrStop() { + return + } if bytes >= softResponseLimit { break } @@ -584,7 +590,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { } } } - sendReq(req.ReqID, uint64(reqCnt), p.ReplyBlockBodiesRLP(req.ReqID, bodies), task.done()) + sendResponse(req.ReqID, uint64(reqCnt), p.ReplyBlockBodiesRLP(req.ReqID, bodies), task.done()) }() case BlockBodiesMsg: @@ -628,7 +634,10 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { return errResp(ErrRequestRejected, "") } go func() { - for _, req := range req.Reqs { + for i, req := range req.Reqs { + if i != 0 && !task.waitOrStop() { + return + } // Retrieve the requested state entry, stopping if enough was found if number := rawdb.ReadHeaderNumber(pm.chainDb, req.BHash); number != nil { if header := rawdb.ReadHeader(pm.chainDb, req.BHash, *number); header != nil { @@ -649,7 +658,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { } } } - sendReq(req.ReqID, uint64(reqCnt), p.ReplyCode(req.ReqID, data), task.done()) + sendResponse(req.ReqID, uint64(reqCnt), p.ReplyCode(req.ReqID, data), task.done()) }() case CodeMsg: @@ -693,7 +702,10 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { return errResp(ErrRequestRejected, "") } go func() { - for _, hash := range req.Hashes { + for i, hash := range req.Hashes { + if i != 0 && !task.waitOrStop() { + return + } if bytes >= softResponseLimit { break } @@ -715,7 +727,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { bytes += len(encoded) } } - sendReq(req.ReqID, uint64(reqCnt), p.ReplyReceiptsRLP(req.ReqID, receipts), task.done()) + sendResponse(req.ReqID, uint64(reqCnt), p.ReplyReceiptsRLP(req.ReqID, receipts), task.done()) }() case ReceiptsMsg: @@ -759,7 +771,10 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { return errResp(ErrRequestRejected, "") } go func() { - for _, req := range req.Reqs { + for i, req := range req.Reqs { + if i != 0 && !task.waitOrStop() { + return + } // Retrieve the requested state entry, stopping if enough was found if number := rawdb.ReadHeaderNumber(pm.chainDb, req.BHash); number != nil { if header := rawdb.ReadHeader(pm.chainDb, req.BHash, *number); header != nil { @@ -789,7 +804,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { } } } - sendReq(req.ReqID, uint64(reqCnt), p.ReplyProofs(req.ReqID, proofs), task.done()) + sendResponse(req.ReqID, uint64(reqCnt), p.ReplyProofs(req.ReqID, proofs), task.done()) }() case GetProofsV2Msg: @@ -816,7 +831,10 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { nodes := light.NewNodeSet() - for _, req := range req.Reqs { + for i, req := range req.Reqs { + if i != 0 && !task.waitOrStop() { + return + } // Look up the state belonging to the request if statedb == nil || req.BHash != lastBHash { statedb, root, lastBHash = nil, common.Hash{}, req.BHash @@ -851,7 +869,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { break } } - sendReq(req.ReqID, uint64(reqCnt), p.ReplyProofsV2(req.ReqID, nodes.NodeList()), task.done()) + sendResponse(req.ReqID, uint64(reqCnt), p.ReplyProofsV2(req.ReqID, nodes.NodeList()), task.done()) }() case ProofsV1Msg: @@ -917,7 +935,10 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { } go func() { trieDb := trie.NewDatabase(ethdb.NewTable(pm.chainDb, light.ChtTablePrefix)) - for _, req := range req.Reqs { + for i, req := range req.Reqs { + if i != 0 && !task.waitOrStop() { + return + } if header := pm.blockchain.GetHeaderByNumber(req.BlockNum); header != nil { sectionHead := rawdb.ReadCanonicalHash(pm.chainDb, req.ChtNum*pm.iConfig.ChtSize-1) if root := light.GetChtRoot(pm.chainDb, req.ChtNum-1, sectionHead); root != (common.Hash{}) { @@ -938,7 +959,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { } } } - sendReq(req.ReqID, uint64(reqCnt), p.ReplyHeaderProofs(req.ReqID, proofs), task.done()) + sendResponse(req.ReqID, uint64(reqCnt), p.ReplyHeaderProofs(req.ReqID, proofs), task.done()) }() case GetHelperTrieProofsMsg: @@ -969,7 +990,10 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { auxTrie *trie.Trie ) nodes := light.NewNodeSet() - for _, req := range req.Reqs { + for i, req := range req.Reqs { + if i != 0 && !task.waitOrStop() { + return + } if auxTrie == nil || req.Type != lastType || req.TrieIdx != lastIdx { auxTrie, lastType, lastIdx = nil, req.Type, req.TrieIdx @@ -999,7 +1023,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { break } } - sendReq(req.ReqID, uint64(reqCnt), p.ReplyHelperTrieProofs(req.ReqID, HelperTrieResps{Proofs: nodes.NodeList(), AuxData: auxData}), task.done()) + sendResponse(req.ReqID, uint64(reqCnt), p.ReplyHelperTrieProofs(req.ReqID, HelperTrieResps{Proofs: nodes.NodeList(), AuxData: auxData}), task.done()) }() case HeaderProofsMsg: @@ -1057,8 +1081,13 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { return errResp(ErrRequestRejected, "") } go func() { - pm.txpool.AddRemotes(txs) - sendReq(0, uint64(reqCnt), nil, task.done()) + for i, tx := range txs { + if i != 0 && !task.waitOrStop() { + return + } + pm.txpool.AddRemotes([]*types.Transaction{tx}) + } + sendResponse(0, uint64(reqCnt), nil, task.done()) }() case SendTxV2Msg: @@ -1078,21 +1107,22 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { return errResp(ErrRequestRejected, "") } go func() { - hashes := make([]common.Hash, len(req.Txs)) + stats := make([]txStatus, len(req.Txs)) for i, tx := range req.Txs { - hashes[i] = tx.Hash() - } - stats := pm.txStatus(hashes) - for i, stat := range stats { - if stat.Status == core.TxStatusUnknown { - if errs := pm.txpool.AddRemotes([]*types.Transaction{req.Txs[i]}); errs[0] != nil { + if i != 0 && !task.waitOrStop() { + return + } + hash := tx.Hash() + stats[i] = pm.txStatus(hash) + if stats[i].Status == core.TxStatusUnknown { + if errs := pm.txpool.AddRemotes([]*types.Transaction{tx}); errs[0] != nil { stats[i].Error = errs[0].Error() continue } - stats[i] = pm.txStatus([]common.Hash{hashes[i]})[0] + stats[i] = pm.txStatus(hash) } } - sendReq(req.ReqID, uint64(reqCnt), p.ReplyTxStatus(req.ReqID, stats), task.done()) + sendResponse(req.ReqID, uint64(reqCnt), p.ReplyTxStatus(req.ReqID, stats), task.done()) }() case GetTxStatusMsg: @@ -1112,8 +1142,14 @@ func (pm *ProtocolManager) handleMsg(p *peer) error { return errResp(ErrRequestRejected, "") } go func() { - stats := pm.txStatus(req.Hashes) - sendReq(req.ReqID, uint64(reqCnt), p.ReplyTxStatus(req.ReqID, stats), task.done()) + stats := make([]txStatus, len(req.Hashes)) + for i, hash := range req.Hashes { + if i != 0 && !task.waitOrStop() { + return + } + stats[i] = pm.txStatus(hash) + } + sendResponse(req.ReqID, uint64(reqCnt), p.ReplyTxStatus(req.ReqID, stats), task.done()) }() case TxStatusMsg: @@ -1190,21 +1226,17 @@ func (pm *ProtocolManager) getHelperTrieAuxData(req HelperTrieReq) []byte { return nil } -func (pm *ProtocolManager) txStatus(hashes []common.Hash) []txStatus { - stats := make([]txStatus, len(hashes)) - for i, stat := range pm.txpool.Status(hashes) { - // Save the status we've got from the transaction pool - stats[i].Status = stat - - // If the transaction is unknown to the pool, try looking it up locally - if stat == core.TxStatusUnknown { - if tx, blockHash, blockNumber, txIndex := rawdb.ReadTransaction(pm.chainDb, hashes[i]); tx != nil { - stats[i].Status = core.TxStatusIncluded - stats[i].Lookup = &rawdb.LegacyTxLookupEntry{BlockHash: blockHash, BlockIndex: blockNumber, Index: txIndex} - } +func (pm *ProtocolManager) txStatus(hash common.Hash) txStatus { + var stat txStatus + stat.Status = pm.txpool.Status([]common.Hash{hash})[0] + // If the transaction is unknown to the pool, try looking it up locally + if stat.Status == core.TxStatusUnknown { + if tx, blockHash, blockNumber, txIndex := rawdb.ReadTransaction(pm.chainDb, hash); tx != nil { + stat.Status = core.TxStatusIncluded + stat.Lookup = &rawdb.LegacyTxLookupEntry{BlockHash: blockHash, BlockIndex: blockNumber, Index: txIndex} } } - return stats + return stat } // isULCEnabled returns true if we can use ULC diff --git a/les/server.go b/les/server.go index 87126c779a..270640f025 100644 --- a/les/server.go +++ b/les/server.go @@ -154,18 +154,20 @@ func (s *LesServer) startEventLoop() { totalRechargeCh := make(chan uint64, 100) totalRecharge := s.costTracker.subscribeTotalRecharge(totalRechargeCh) totalCapacityCh := make(chan uint64, 100) + updateRecharge := func() { + if processing { + s.protocolManager.servingQueue.setThreads(s.thcBlockProcessing) + s.fcManager.SetRechargeCurve(flowcontrol.PieceWiseLinear{{0, 0}, {totalRecharge, totalRecharge}}) + } else { + s.protocolManager.servingQueue.setThreads(s.thcNormal) + s.fcManager.SetRechargeCurve(flowcontrol.PieceWiseLinear{{0, 0}, {totalRecharge / 10, totalRecharge}, {totalRecharge, totalRecharge}}) + } + } + updateRecharge() totalCapacity := s.fcManager.SubscribeTotalCapacity(totalCapacityCh) + s.priorityClientPool.setLimits(s.maxPeers, totalCapacity) go func() { - updateRecharge := func() { - if processing { - s.protocolManager.servingQueue.setThreads(s.thcBlockProcessing) - s.fcManager.SetRechargeCurve(flowcontrol.PieceWiseLinear{{0, 0}, {totalRecharge, totalRecharge}}) - } else { - s.protocolManager.servingQueue.setThreads(s.thcNormal) - s.fcManager.SetRechargeCurve(flowcontrol.PieceWiseLinear{{0, 0}, {totalRecharge / 10, totalRecharge}, {totalRecharge, totalRecharge}}) - } - } for { select { case processing = <-blockProcFeed: diff --git a/les/servingqueue.go b/les/servingqueue.go index 6b4b277aa9..2438fdfe3c 100644 --- a/les/servingqueue.go +++ b/les/servingqueue.go @@ -23,7 +23,7 @@ import ( "github.com/ethereum/go-ethereum/common/prque" ) -// servingQueue runs serving tasks in a limited number of threads and puts the +// servingQueue allows running tasks in a limited number of threads and puts the // waiting tasks in a priority queue type servingQueue struct { tokenCh chan runToken @@ -34,7 +34,7 @@ type servingQueue struct { wg sync.WaitGroup threadCount int // number of currently running threads queue *prque.Prque // priority queue for waiting or suspended tasks - best *servingTask // either best == nil (queue empty) or waitingForTask is empty + best *servingTask // the highest priority task (not included in the queue) suspendBias int64 // priority bias against suspending an already running task } @@ -55,8 +55,13 @@ type servingTask struct { tokenCh chan runToken } +// runToken received by servingTask.start allows the task to run. Closing the +// channel by servingTask.stop signals the thread controller to allow a new task +// to start running. type runToken chan struct{} +// start blocks until the task can start and returns true if it is allowed to run. +// Returning false means that the task should be cancelled. func (t *servingTask) start() bool { select { case t.token = <-t.sq.tokenCh: @@ -80,12 +85,18 @@ func (t *servingTask) start() bool { return true } +// done signals the thread controller about the task being finished and returns +// the total serving time of the task in nanoseconds. func (t *servingTask) done() uint64 { t.servingTime += uint64(mclock.Now()) close(t.token) return t.servingTime } +// waitOrStop can be called during the execution of the task. It blocks if there +// is a higher priority task waiting (a bias is applied in favor of the currently +// running task). Returning true means that the execution can be resumed. False +// means the task should be cancelled. func (t *servingTask) waitOrStop() bool { t.done() if !t.biasAdded { @@ -113,6 +124,7 @@ func newServingQueue(suspendBias int64) *servingQueue { return sq } +// newTask creates a new task with the given priority func (sq *servingQueue) newTask(priority int64) *servingTask { return &servingTask{ sq: sq, @@ -120,6 +132,12 @@ func (sq *servingQueue) newTask(priority int64) *servingTask { } } +// threadController is started in multiple goroutines and controls the execution +// of tasks. The number of active thread controllers equals the allowed number of +// concurrently running threads. It tries to fetch the highest priority queued +// task first. If there are no queued tasks waiting then it can directly catch +// run tokens from the token channel and allow the corresponding tasks to run +// without entering the priority queue. func (sq *servingQueue) threadController() { for { token := make(runToken) @@ -152,6 +170,7 @@ func (sq *servingQueue) threadController() { } } +// addTask inserts a task into the priority queue func (sq *servingQueue) addTask(task *servingTask) { if sq.best == nil { sq.best = task @@ -164,6 +183,9 @@ func (sq *servingQueue) addTask(task *servingTask) { } } +// queueLoop is an event loop running in a goroutine. It receives tasks from queueAddCh +// and always tries to send the highest priority task to queueBestCh. Successfully sent +// tasks are removed from the queue. func (sq *servingQueue) queueLoop() { for { if sq.best != nil { @@ -192,6 +214,8 @@ func (sq *servingQueue) queueLoop() { } } +// threadCountLoop is an event loop running in a goroutine. It adjusts the number +// of active thread controller goroutines. func (sq *servingQueue) threadCountLoop() { var threadCountTarget int for { @@ -220,7 +244,7 @@ func (sq *servingQueue) threadCountLoop() { } } -// setThreads sets the processing thread count, suspending tasks as soon as +// setThreads sets the allowed processing thread count, suspending tasks as soon as // possible if necessary. func (sq *servingQueue) setThreads(threadCount int) { select { @@ -230,7 +254,7 @@ func (sq *servingQueue) setThreads(threadCount int) { } } -// stop stops task processing as soon as possible +// stop stops task processing as soon as possible and shuts down the serving queue. func (sq *servingQueue) stop() { close(sq.quit) sq.wg.Wait()