les: new serving queue

This commit is contained in:
Zsolt Felfoldi 2019-01-28 22:01:37 +01:00
parent b3f5a40502
commit a67f0ae4ad
2 changed files with 283 additions and 340 deletions

View file

@ -335,16 +335,16 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
p.responseCount++ p.responseCount++
responseCount := p.responseCount responseCount := p.responseCount
var ( var (
maxCost uint64 maxCost uint64
priority int64 task *servingTask
) )
reject := func(reqID, reqCnt, maxCnt uint64) bool { accept := func(reqID, reqCnt, maxCnt uint64) bool {
if reqCnt == 0 { if reqCnt == 0 {
return true return false
} }
if p.fcClient == nil || reqCnt > maxCnt { if p.fcClient == nil || reqCnt > maxCnt {
return true return false
} }
maxCost = p.fcCosts.getCost(msg.Code, reqCnt) maxCost = p.fcCosts.getCost(msg.Code, reqCnt)
@ -352,11 +352,11 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
if bufShort > 0 { if bufShort > 0 {
p.Log().Error("Request came too early", "remaining", common.PrettyDuration(time.Duration(bufShort*1000000/p.fcParams.MinRecharge))) p.Log().Error("Request came too early", "remaining", common.PrettyDuration(time.Duration(bufShort*1000000/p.fcParams.MinRecharge)))
} }
return true return false
} else { } else {
priority = servingPriority task = pm.servingQueue.newTask(servingPriority)
} }
return false return task.start()
} }
if msg.Size > ProtocolMaxMsgSize { if msg.Size > ProtocolMaxMsgSize {
@ -366,12 +366,6 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
var deliverMsg *Msg var deliverMsg *Msg
errorFn := func(err error) {
if err != nil {
p.errCh <- err
}
}
sendReq := func(reqID, amount uint64, reply *reply, servingTime uint64) { sendReq := func(reqID, amount uint64, reply *reply, servingTime uint64) {
p.responseLock.Lock() p.responseLock.Lock()
defer p.responseLock.Unlock() defer p.responseLock.Unlock()
@ -448,24 +442,21 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
} }
query := req.Query query := req.Query
if reject(req.ReqID, query.Amount, MaxHeaderFetch) { if !accept(req.ReqID, query.Amount, MaxHeaderFetch) {
return errResp(ErrRequestRejected, "") return errResp(ErrRequestRejected, "")
} }
go func() {
hashMode := query.Origin.Hash != (common.Hash{})
first := true
maxNonCanonical := uint64(100)
hashMode := query.Origin.Hash != (common.Hash{}) // Gather headers until the fetch or network limits is reached
first := true var (
maxNonCanonical := uint64(100) bytes common.StorageSize
headers []*types.Header
// Gather headers until the fetch or network limits is reached unknown bool
var ( )
bytes common.StorageSize for !unknown && len(headers) < int(query.Amount) && bytes < softResponseLimit {
headers []*types.Header
unknown bool
)
pm.servingQueue.addTask(&servingTask{
priority: priority,
run: func() (bool, error) {
// Retrieve the next header satisfying the query // Retrieve the next header satisfying the query
var origin *types.Header var origin *types.Header
if hashMode { if hashMode {
@ -482,7 +473,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
origin = pm.blockchain.GetHeaderByNumber(query.Origin.Number) origin = pm.blockchain.GetHeaderByNumber(query.Origin.Number)
} }
if origin == nil { if origin == nil {
return true, nil break
} }
headers = append(headers, origin) headers = append(headers, origin)
bytes += estHeaderRlpSize bytes += estHeaderRlpSize
@ -533,14 +524,9 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
// Number based traversal towards the leaf block // Number based traversal towards the leaf block
query.Origin.Number += query.Skip + 1 query.Origin.Number += query.Skip + 1
} }
return unknown || len(headers) >= int(query.Amount) || bytes >= softResponseLimit, nil }
}, sendReq(req.ReqID, query.Amount, p.ReplyBlockHeaders(req.ReqID, headers), task.done())
// after: sendFunc(query.Amount, func(bv uint64) error { return p.SendBlockHeaders(req.ReqID, bv, headers) }), }()
send: func(servingTime uint64) {
sendReq(req.ReqID, query.Amount, p.ReplyBlockHeaders(req.ReqID, headers), servingTime)
},
fail: errorFn,
})
case BlockHeadersMsg: case BlockHeadersMsg:
if pm.downloader == nil { if pm.downloader == nil {
@ -582,18 +568,13 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
bodies []rlp.RawValue bodies []rlp.RawValue
) )
reqCnt := len(req.Hashes) reqCnt := len(req.Hashes)
if reject(req.ReqID, uint64(reqCnt), MaxBodyFetch) { if !accept(req.ReqID, uint64(reqCnt), MaxBodyFetch) {
return errResp(ErrRequestRejected, "") return errResp(ErrRequestRejected, "")
} }
go func() {
index := 0 for _, hash := range req.Hashes {
pm.servingQueue.addTask(&servingTask{
priority: priority,
run: func() (bool, error) {
hash := req.Hashes[index]
index++
if bytes >= softResponseLimit { if bytes >= softResponseLimit {
return true, nil break
} }
// Retrieve the requested block body, stopping if enough was found // Retrieve the requested block body, stopping if enough was found
if number := rawdb.ReadHeaderNumber(pm.chainDb, hash); number != nil { if number := rawdb.ReadHeaderNumber(pm.chainDb, hash); number != nil {
@ -602,13 +583,9 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
bytes += len(data) bytes += len(data)
} }
} }
return index == reqCnt, nil }
}, sendReq(req.ReqID, uint64(reqCnt), p.ReplyBlockBodiesRLP(req.ReqID, bodies), task.done())
send: func(servingTime uint64) { }()
sendReq(req.ReqID, uint64(reqCnt), p.ReplyBlockBodiesRLP(req.ReqID, bodies), servingTime)
},
fail: errorFn,
})
case BlockBodiesMsg: case BlockBodiesMsg:
if pm.odr == nil { if pm.odr == nil {
@ -647,42 +624,33 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
data [][]byte data [][]byte
) )
reqCnt := len(req.Reqs) reqCnt := len(req.Reqs)
if reject(req.ReqID, uint64(reqCnt), MaxCodeFetch) { if !accept(req.ReqID, uint64(reqCnt), MaxCodeFetch) {
return errResp(ErrRequestRejected, "") return errResp(ErrRequestRejected, "")
} }
index := 0 go func() {
pm.servingQueue.addTask(&servingTask{ for _, req := range req.Reqs {
priority: priority,
run: func() (bool, error) {
req := req.Reqs[index]
index++
done := index == reqCnt
// Retrieve the requested state entry, stopping if enough was found // Retrieve the requested state entry, stopping if enough was found
if number := rawdb.ReadHeaderNumber(pm.chainDb, req.BHash); number != nil { if number := rawdb.ReadHeaderNumber(pm.chainDb, req.BHash); number != nil {
if header := rawdb.ReadHeader(pm.chainDb, req.BHash, *number); header != nil { if header := rawdb.ReadHeader(pm.chainDb, req.BHash, *number); header != nil {
statedb, err := pm.blockchain.State() statedb, err := pm.blockchain.State()
if err != nil { if err != nil {
return done, nil continue
} }
account, err := pm.getAccount(statedb, header.Root, common.BytesToHash(req.AccKey)) account, err := pm.getAccount(statedb, header.Root, common.BytesToHash(req.AccKey))
if err != nil { if err != nil {
return done, nil continue
} }
code, _ := statedb.Database().TrieDB().Node(common.BytesToHash(account.CodeHash)) code, _ := statedb.Database().TrieDB().Node(common.BytesToHash(account.CodeHash))
data = append(data, code) data = append(data, code)
if bytes += len(code); bytes >= softResponseLimit { if bytes += len(code); bytes >= softResponseLimit {
return true, nil break
} }
} }
} }
return done, nil }
}, sendReq(req.ReqID, uint64(reqCnt), p.ReplyCode(req.ReqID, data), task.done())
send: func(servingTime uint64) { }()
sendReq(req.ReqID, uint64(reqCnt), p.ReplyCode(req.ReqID, data), servingTime)
},
fail: errorFn,
})
case CodeMsg: case CodeMsg:
if pm.odr == nil { if pm.odr == nil {
@ -721,19 +689,13 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
receipts []rlp.RawValue receipts []rlp.RawValue
) )
reqCnt := len(req.Hashes) reqCnt := len(req.Hashes)
if reject(req.ReqID, uint64(reqCnt), MaxReceiptFetch) { if !accept(req.ReqID, uint64(reqCnt), MaxReceiptFetch) {
return errResp(ErrRequestRejected, "") return errResp(ErrRequestRejected, "")
} }
go func() {
index := 0 for _, hash := range req.Hashes {
pm.servingQueue.addTask(&servingTask{
priority: priority,
run: func() (bool, error) {
hash := req.Hashes[index]
index++
done := index == reqCnt
if bytes >= softResponseLimit { if bytes >= softResponseLimit {
return true, nil break
} }
// Retrieve the requested block's receipts, skipping if unknown to us // Retrieve the requested block's receipts, skipping if unknown to us
var results types.Receipts var results types.Receipts
@ -742,7 +704,7 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
} }
if results == nil { if results == nil {
if header := pm.blockchain.GetHeaderByHash(hash); header == nil || header.ReceiptHash != types.EmptyRootHash { if header := pm.blockchain.GetHeaderByHash(hash); header == nil || header.ReceiptHash != types.EmptyRootHash {
return done, nil continue
} }
} }
// If known, encode and queue for response packet // If known, encode and queue for response packet
@ -752,13 +714,9 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
receipts = append(receipts, encoded) receipts = append(receipts, encoded)
bytes += len(encoded) bytes += len(encoded)
} }
return done, nil }
}, sendReq(req.ReqID, uint64(reqCnt), p.ReplyReceiptsRLP(req.ReqID, receipts), task.done())
send: func(servingTime uint64) { }()
sendReq(req.ReqID, uint64(reqCnt), p.ReplyReceiptsRLP(req.ReqID, receipts), servingTime)
},
fail: errorFn,
})
case ReceiptsMsg: case ReceiptsMsg:
if pm.odr == nil { if pm.odr == nil {
@ -797,29 +755,23 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
proofs proofsData proofs proofsData
) )
reqCnt := len(req.Reqs) reqCnt := len(req.Reqs)
if reject(req.ReqID, uint64(reqCnt), MaxProofsFetch) { if !accept(req.ReqID, uint64(reqCnt), MaxProofsFetch) {
return errResp(ErrRequestRejected, "") return errResp(ErrRequestRejected, "")
} }
go func() {
index := 0 for _, req := range req.Reqs {
pm.servingQueue.addTask(&servingTask{
priority: priority,
run: func() (bool, error) {
req := req.Reqs[index]
index++
done := index == reqCnt
// Retrieve the requested state entry, stopping if enough was found // Retrieve the requested state entry, stopping if enough was found
if number := rawdb.ReadHeaderNumber(pm.chainDb, req.BHash); number != nil { if number := rawdb.ReadHeaderNumber(pm.chainDb, req.BHash); number != nil {
if header := rawdb.ReadHeader(pm.chainDb, req.BHash, *number); header != nil { if header := rawdb.ReadHeader(pm.chainDb, req.BHash, *number); header != nil {
statedb, err := pm.blockchain.State() statedb, err := pm.blockchain.State()
if err != nil { if err != nil {
return done, nil continue
} }
var trie state.Trie var trie state.Trie
if len(req.AccKey) > 0 { if len(req.AccKey) > 0 {
account, err := pm.getAccount(statedb, header.Root, common.BytesToHash(req.AccKey)) account, err := pm.getAccount(statedb, header.Root, common.BytesToHash(req.AccKey))
if err != nil { if err != nil {
return done, nil continue
} }
trie, _ = statedb.Database().OpenStorageTrie(common.BytesToHash(req.AccKey), account.Root) trie, _ = statedb.Database().OpenStorageTrie(common.BytesToHash(req.AccKey), account.Root)
} else { } else {
@ -831,18 +783,14 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
proofs = append(proofs, proof) proofs = append(proofs, proof)
if bytes += proof.DataSize(); bytes >= softResponseLimit { if bytes += proof.DataSize(); bytes >= softResponseLimit {
return true, nil break
} }
} }
} }
} }
return done, nil }
}, sendReq(req.ReqID, uint64(reqCnt), p.ReplyProofs(req.ReqID, proofs), task.done())
send: func(servingTime uint64) { }()
sendReq(req.ReqID, uint64(reqCnt), p.ReplyProofs(req.ReqID, proofs), servingTime)
},
fail: errorFn,
})
case GetProofsV2Msg: case GetProofsV2Msg:
p.Log().Trace("Received les/2 proofs request") p.Log().Trace("Received les/2 proofs request")
@ -861,19 +809,14 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
root common.Hash root common.Hash
) )
reqCnt := len(req.Reqs) reqCnt := len(req.Reqs)
if reject(req.ReqID, uint64(reqCnt), MaxProofsFetch) { if !accept(req.ReqID, uint64(reqCnt), MaxProofsFetch) {
return errResp(ErrRequestRejected, "") return errResp(ErrRequestRejected, "")
} }
go func() {
nodes := light.NewNodeSet() nodes := light.NewNodeSet()
index := 0 for _, req := range req.Reqs {
pm.servingQueue.addTask(&servingTask{
priority: priority,
run: func() (bool, error) {
req := req.Reqs[index]
index++
done := index == reqCnt
// Look up the state belonging to the request // Look up the state belonging to the request
if statedb == nil || req.BHash != lastBHash { if statedb == nil || req.BHash != lastBHash {
statedb, root, lastBHash = nil, common.Hash{}, req.BHash statedb, root, lastBHash = nil, common.Hash{}, req.BHash
@ -886,34 +829,30 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
} }
} }
if statedb == nil { if statedb == nil {
return done, nil continue
} }
// Pull the account or storage trie of the request // Pull the account or storage trie of the request
var trie state.Trie var trie state.Trie
if len(req.AccKey) > 0 { if len(req.AccKey) > 0 {
account, err := pm.getAccount(statedb, root, common.BytesToHash(req.AccKey)) account, err := pm.getAccount(statedb, root, common.BytesToHash(req.AccKey))
if err != nil { if err != nil {
return done, nil continue
} }
trie, _ = statedb.Database().OpenStorageTrie(common.BytesToHash(req.AccKey), account.Root) trie, _ = statedb.Database().OpenStorageTrie(common.BytesToHash(req.AccKey), account.Root)
} else { } else {
trie, _ = statedb.Database().OpenTrie(root) trie, _ = statedb.Database().OpenTrie(root)
} }
if trie == nil { if trie == nil {
return done, nil continue
} }
// Prove the user's request from the account or stroage trie // Prove the user's request from the account or stroage trie
trie.Prove(req.Key, req.FromLevel, nodes) trie.Prove(req.Key, req.FromLevel, nodes)
if nodes.DataSize() >= softResponseLimit { if nodes.DataSize() >= softResponseLimit {
return true, nil break
} }
return done, nil }
}, sendReq(req.ReqID, uint64(reqCnt), p.ReplyProofsV2(req.ReqID, nodes.NodeList()), task.done())
send: func(servingTime uint64) { }()
sendReq(req.ReqID, uint64(reqCnt), p.ReplyProofsV2(req.ReqID, nodes.NodeList()), servingTime)
},
fail: errorFn,
})
case ProofsV1Msg: case ProofsV1Msg:
if pm.odr == nil { if pm.odr == nil {
@ -973,24 +912,18 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
proofs []ChtResp proofs []ChtResp
) )
reqCnt := len(req.Reqs) reqCnt := len(req.Reqs)
if reject(req.ReqID, uint64(reqCnt), MaxHelperTrieProofsFetch) { if !accept(req.ReqID, uint64(reqCnt), MaxHelperTrieProofsFetch) {
return errResp(ErrRequestRejected, "") return errResp(ErrRequestRejected, "")
} }
trieDb := trie.NewDatabase(ethdb.NewTable(pm.chainDb, light.ChtTablePrefix)) go func() {
trieDb := trie.NewDatabase(ethdb.NewTable(pm.chainDb, light.ChtTablePrefix))
index := 0 for _, req := range req.Reqs {
pm.servingQueue.addTask(&servingTask{
priority: priority,
run: func() (bool, error) {
req := req.Reqs[index]
index++
done := index == reqCnt
if header := pm.blockchain.GetHeaderByNumber(req.BlockNum); header != nil { if header := pm.blockchain.GetHeaderByNumber(req.BlockNum); header != nil {
sectionHead := rawdb.ReadCanonicalHash(pm.chainDb, req.ChtNum*pm.iConfig.ChtSize-1) sectionHead := rawdb.ReadCanonicalHash(pm.chainDb, req.ChtNum*pm.iConfig.ChtSize-1)
if root := light.GetChtRoot(pm.chainDb, req.ChtNum-1, sectionHead); root != (common.Hash{}) { if root := light.GetChtRoot(pm.chainDb, req.ChtNum-1, sectionHead); root != (common.Hash{}) {
trie, err := trie.New(root, trieDb) trie, err := trie.New(root, trieDb)
if err != nil { if err != nil {
return done, nil continue
} }
var encNumber [8]byte var encNumber [8]byte
binary.BigEndian.PutUint64(encNumber[:], req.BlockNum) binary.BigEndian.PutUint64(encNumber[:], req.BlockNum)
@ -1000,17 +933,13 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
proofs = append(proofs, ChtResp{Header: header, Proof: proof}) proofs = append(proofs, ChtResp{Header: header, Proof: proof})
if bytes += proof.DataSize() + estHeaderRlpSize; bytes >= softResponseLimit { if bytes += proof.DataSize() + estHeaderRlpSize; bytes >= softResponseLimit {
return true, nil break
} }
} }
} }
return done, nil }
}, sendReq(req.ReqID, uint64(reqCnt), p.ReplyHeaderProofs(req.ReqID, proofs), task.done())
send: func(servingTime uint64) { }()
sendReq(req.ReqID, uint64(reqCnt), p.ReplyHeaderProofs(req.ReqID, proofs), servingTime)
},
fail: errorFn,
})
case GetHelperTrieProofsMsg: case GetHelperTrieProofsMsg:
p.Log().Trace("Received helper trie proof request") p.Log().Trace("Received helper trie proof request")
@ -1028,24 +957,19 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
auxData [][]byte auxData [][]byte
) )
reqCnt := len(req.Reqs) reqCnt := len(req.Reqs)
if reject(req.ReqID, uint64(reqCnt), MaxHelperTrieProofsFetch) { if !accept(req.ReqID, uint64(reqCnt), MaxHelperTrieProofsFetch) {
return errResp(ErrRequestRejected, "") return errResp(ErrRequestRejected, "")
} }
go func() {
var ( var (
lastIdx uint64 lastIdx uint64
lastType uint lastType uint
root common.Hash root common.Hash
auxTrie *trie.Trie auxTrie *trie.Trie
) )
nodes := light.NewNodeSet() nodes := light.NewNodeSet()
for _, req := range req.Reqs {
index := 0
pm.servingQueue.addTask(&servingTask{
priority: priority,
run: func() (bool, error) {
req := req.Reqs[index]
index++
if auxTrie == nil || req.Type != lastType || req.TrieIdx != lastIdx { if auxTrie == nil || req.Type != lastType || req.TrieIdx != lastIdx {
auxTrie, lastType, lastIdx = nil, req.Type, req.TrieIdx auxTrie, lastType, lastIdx = nil, req.Type, req.TrieIdx
@ -1072,15 +996,11 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
} }
} }
if nodes.DataSize()+auxBytes >= softResponseLimit { if nodes.DataSize()+auxBytes >= softResponseLimit {
return true, nil break
} }
return index == reqCnt, nil }
}, sendReq(req.ReqID, uint64(reqCnt), p.ReplyHelperTrieProofs(req.ReqID, HelperTrieResps{Proofs: nodes.NodeList(), AuxData: auxData}), task.done())
send: func(servingTime uint64) { }()
sendReq(req.ReqID, uint64(reqCnt), p.ReplyHelperTrieProofs(req.ReqID, HelperTrieResps{Proofs: nodes.NodeList(), AuxData: auxData}), servingTime)
},
fail: errorFn,
})
case HeaderProofsMsg: case HeaderProofsMsg:
if pm.odr == nil { if pm.odr == nil {
@ -1133,21 +1053,13 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
return errResp(ErrDecode, "msg %v: %v", msg, err) return errResp(ErrDecode, "msg %v: %v", msg, err)
} }
reqCnt := len(txs) reqCnt := len(txs)
if reject(0, uint64(reqCnt), MaxTxSend) { if !accept(0, uint64(reqCnt), MaxTxSend) {
return errResp(ErrRequestRejected, "") return errResp(ErrRequestRejected, "")
} }
go func() {
pm.servingQueue.addTask(&servingTask{ pm.txpool.AddRemotes(txs)
priority: priority, sendReq(0, uint64(reqCnt), nil, task.done())
run: func() (bool, error) { }()
pm.txpool.AddRemotes(txs)
return true, nil
},
send: func(servingTime uint64) {
sendReq(0, uint64(reqCnt), nil, servingTime)
},
fail: errorFn,
})
case SendTxV2Msg: case SendTxV2Msg:
if pm.txpool == nil { if pm.txpool == nil {
@ -1162,35 +1074,26 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
return errResp(ErrDecode, "msg %v: %v", msg, err) return errResp(ErrDecode, "msg %v: %v", msg, err)
} }
reqCnt := len(req.Txs) reqCnt := len(req.Txs)
if reject(req.ReqID, uint64(reqCnt), MaxTxSend) { if !accept(req.ReqID, uint64(reqCnt), MaxTxSend) {
return errResp(ErrRequestRejected, "") return errResp(ErrRequestRejected, "")
} }
go func() {
var stats []txStatus hashes := make([]common.Hash, len(req.Txs))
pm.servingQueue.addTask(&servingTask{ for i, tx := range req.Txs {
priority: priority, hashes[i] = tx.Hash()
run: func() (bool, error) { }
hashes := make([]common.Hash, len(req.Txs)) stats := pm.txStatus(hashes)
for i, tx := range req.Txs { for i, stat := range stats {
hashes[i] = tx.Hash() if stat.Status == core.TxStatusUnknown {
} if errs := pm.txpool.AddRemotes([]*types.Transaction{req.Txs[i]}); errs[0] != nil {
stats = pm.txStatus(hashes) stats[i].Error = errs[0].Error()
for i, stat := range stats { continue
if stat.Status == core.TxStatusUnknown {
if errs := pm.txpool.AddRemotes([]*types.Transaction{req.Txs[i]}); errs[0] != nil {
stats[i].Error = errs[0].Error()
continue
}
stats[i] = pm.txStatus([]common.Hash{hashes[i]})[0]
} }
stats[i] = pm.txStatus([]common.Hash{hashes[i]})[0]
} }
return true, nil }
}, sendReq(req.ReqID, uint64(reqCnt), p.ReplyTxStatus(req.ReqID, stats), task.done())
send: func(servingTime uint64) { }()
sendReq(req.ReqID, uint64(reqCnt), p.ReplyTxStatus(req.ReqID, stats), servingTime)
},
fail: errorFn,
})
case GetTxStatusMsg: case GetTxStatusMsg:
if pm.txpool == nil { if pm.txpool == nil {
@ -1205,22 +1108,13 @@ func (pm *ProtocolManager) handleMsg(p *peer) error {
return errResp(ErrDecode, "msg %v: %v", msg, err) return errResp(ErrDecode, "msg %v: %v", msg, err)
} }
reqCnt := len(req.Hashes) reqCnt := len(req.Hashes)
if reject(req.ReqID, uint64(reqCnt), MaxTxStatus) { if !accept(req.ReqID, uint64(reqCnt), MaxTxStatus) {
return errResp(ErrRequestRejected, "") return errResp(ErrRequestRejected, "")
} }
go func() {
var stats []txStatus stats := pm.txStatus(req.Hashes)
pm.servingQueue.addTask(&servingTask{ sendReq(req.ReqID, uint64(reqCnt), p.ReplyTxStatus(req.ReqID, stats), task.done())
priority: priority, }()
run: func() (bool, error) {
stats = pm.txStatus(req.Hashes)
return true, nil
},
send: func(servingTime uint64) {
sendReq(req.ReqID, uint64(reqCnt), p.ReplyTxStatus(req.ReqID, stats), servingTime)
},
fail: errorFn,
})
case TxStatusMsg: case TxStatusMsg:
if pm.odr == nil { if pm.odr == nil {

View file

@ -26,13 +26,16 @@ import (
// servingQueue runs serving tasks in a limited number of threads and puts the // servingQueue runs serving tasks in a limited number of threads and puts the
// waiting tasks in a priority queue // waiting tasks in a priority queue
type servingQueue struct { type servingQueue struct {
lock sync.Mutex tokenCh chan runToken
threadCount int // number of currently running threads queueAddCh, queueBestCh chan *servingTask
stopCount int // number of threads to be stopped after they finish their current task stopThreadCh, quit chan struct{}
queue *prque.Prque // priority queue for waiting or suspended tasks setThreadsCh chan int
best *servingTask // either best == nil (queue empty) or waitingForTask is empty
waiting []chan *servingTask // threads waiting for a task wg sync.WaitGroup
suspendBias int64 // priority bias against suspending an already running task 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
suspendBias int64 // priority bias against suspending an already running task
} }
// servingTask represents a request serving task. Tasks can be implemented to // servingTask represents a request serving task. Tasks can be implemented to
@ -44,145 +47,191 @@ type servingQueue struct {
// - run: execute a single step; return true if finished // - run: execute a single step; return true if finished
// - after: executed after run finishes or returns an error, receives the total serving time // - after: executed after run finishes or returns an error, receives the total serving time
type servingTask struct { type servingTask struct {
sq *servingQueue
servingTime uint64 servingTime uint64
done bool
err error
priority int64 priority int64
run func() (finished bool, err error) biasAdded bool
send func(servingTime uint64) token runToken
fail func(err error) tokenCh chan runToken
}
type runToken chan struct{}
func (t *servingTask) start() bool {
select {
case t.token = <-t.sq.tokenCh:
default:
t.tokenCh = make(chan runToken, 1)
select {
case t.sq.queueAddCh <- t:
case <-t.sq.quit:
return false
}
select {
case t.token = <-t.tokenCh:
case <-t.sq.quit:
return false
}
}
if t.token == nil {
return false
}
t.servingTime -= uint64(mclock.Now())
return true
}
func (t *servingTask) done() uint64 {
t.servingTime += uint64(mclock.Now())
close(t.token)
return t.servingTime
}
func (t *servingTask) waitOrStop() bool {
t.done()
if !t.biasAdded {
t.priority += t.sq.suspendBias
t.biasAdded = true
}
return t.start()
} }
// newServingQueue returns a new servingQueue // newServingQueue returns a new servingQueue
func newServingQueue(_suspendBias int64) *servingQueue { func newServingQueue(suspendBias int64) *servingQueue {
return &servingQueue{ sq := &servingQueue{
queue: prque.New(nil), queue: prque.New(nil),
suspendBias: _suspendBias, suspendBias: suspendBias,
tokenCh: make(chan runToken),
queueAddCh: make(chan *servingTask, 100),
queueBestCh: make(chan *servingTask),
stopThreadCh: make(chan struct{}),
quit: make(chan struct{}),
setThreadsCh: make(chan int, 10),
}
sq.wg.Add(2)
go sq.queueLoop()
go sq.threadCountLoop()
return sq
}
func (sq *servingQueue) newTask(priority int64) *servingTask {
return &servingTask{
sq: sq,
priority: priority,
} }
} }
// addTask adds a new task, either starting it immediately or queueing it func (sq *servingQueue) threadController() {
func (sq *servingQueue) addTask(task *servingTask) { for {
sq.lock.Lock() token := make(runToken)
defer sq.lock.Unlock() select {
case best := <-sq.queueBestCh:
if l := len(sq.waiting); l != 0 { best.tokenCh <- token
l-- default:
sq.waiting[l] <- task select {
sq.waiting = sq.waiting[:l] case best := <-sq.queueBestCh:
return best.tokenCh <- token
case sq.tokenCh <- token:
case <-sq.stopThreadCh:
sq.wg.Done()
return
case <-sq.quit:
sq.wg.Done()
return
}
}
<-token
select {
case <-sq.stopThreadCh:
sq.wg.Done()
return
case <-sq.quit:
sq.wg.Done()
return
default:
}
} }
}
func (sq *servingQueue) addTask(task *servingTask) {
if sq.best == nil { if sq.best == nil {
sq.best = task sq.best = task
return } else if task.priority > sq.best.priority {
}
if task.priority < sq.best.priority {
sq.queue.Push(sq.best, sq.best.priority) sq.queue.Push(sq.best, sq.best.priority)
sq.best = task sq.best = task
return return
} else {
sq.queue.Push(task, task.priority)
} }
sq.queue.Push(task, task.priority)
} }
// getNewTask selects a new task to be processed. If blocking == true then it waits func (sq *servingQueue) queueLoop() {
// until a runnable task arrives or returns nil if the thread should be stopped. for {
// if currentTask != nil then it returns immediately and only returns a new task if sq.best != nil {
// if the current one should be suspended. select {
// Note: either blocking should be false or currentTask should be nil. case task := <-sq.queueAddCh:
func (sq *servingQueue) getNewTask(currentTask *servingTask, blocking bool) *servingTask { sq.addTask(task)
sq.lock.Lock() case sq.queueBestCh <- sq.best:
if sq.stopCount == 0 { if sq.queue.Size() == 0 {
if sq.best != nil && (currentTask == nil || sq.best.priority <= currentTask.priority-sq.suspendBias) { sq.best = nil
best := sq.best } else {
if sq.queue.Size() == 0 { sq.best, _ = sq.queue.PopItem().(*servingTask)
sq.best = nil }
} else { case <-sq.quit:
sq.best, _ = sq.queue.PopItem().(*servingTask) sq.wg.Done()
return
}
} else {
select {
case task := <-sq.queueAddCh:
sq.addTask(task)
case <-sq.quit:
sq.wg.Done()
return
}
}
}
}
func (sq *servingQueue) threadCountLoop() {
var threadCountTarget int
for {
for threadCountTarget > sq.threadCount {
sq.wg.Add(1)
go sq.threadController()
sq.threadCount++
}
if threadCountTarget < sq.threadCount {
select {
case threadCountTarget = <-sq.setThreadsCh:
case sq.stopThreadCh <- struct{}{}:
sq.threadCount--
case <-sq.quit:
sq.wg.Done()
return
}
} else {
select {
case threadCountTarget = <-sq.setThreadsCh:
case <-sq.quit:
sq.wg.Done()
return
} }
sq.lock.Unlock()
return best
} }
if blocking {
ch := make(chan *servingTask)
sq.waiting = append(sq.waiting, ch)
sq.lock.Unlock()
return <-ch
}
} else {
sq.stopCount--
sq.threadCount--
} }
sq.lock.Unlock()
return nil
} }
// setThreads sets the processing thread count, suspending tasks as soon as // setThreads sets the processing thread count, suspending tasks as soon as
// possible if necessary. // possible if necessary.
func (sq *servingQueue) setThreads(threadCount int) { func (sq *servingQueue) setThreads(threadCount int) {
sq.lock.Lock() select {
defer sq.lock.Unlock() case sq.setThreadsCh <- threadCount:
case <-sq.quit:
diff := threadCount - sq.threadCount + sq.stopCount return
if diff > 0 {
// start more threads
if sq.stopCount >= diff {
sq.stopCount -= diff
} else {
diff -= sq.stopCount
sq.stopCount = 0
sq.threadCount += diff
for ; diff > 0; diff-- {
go sq.servingThread()
}
}
}
if diff < 0 {
// stop some threads
lw := len(sq.waiting)
sq.stopCount -= diff
for diff < 0 && lw > 0 {
diff++
lw--
sq.waiting[lw] <- nil
sq.stopCount--
sq.threadCount--
}
sq.waiting = sq.waiting[:lw]
} }
} }
// stop stops task processing as soon as possible // stop stops task processing as soon as possible
func (sq *servingQueue) stop() { func (sq *servingQueue) stop() {
sq.setThreads(0) close(sq.quit)
} sq.wg.Wait()
// servingThread implements a single serving thread
func (sq *servingQueue) servingThread() {
for {
task := sq.getNewTask(nil, true)
if task == nil {
return
}
task.servingTime -= uint64(mclock.Now())
for {
task.done, task.err = task.run()
if task.done || task.err != nil {
task.servingTime += uint64(mclock.Now())
if task.err == nil {
task.send(task.servingTime)
} else {
task.fail(task.err)
}
break
}
if newTask := sq.getNewTask(task, false); newTask != nil {
now := uint64(mclock.Now())
task.servingTime += now
sq.addTask(task)
task = newTask
task.servingTime -= now
}
}
}
} }