use batch instead of channel

Signed-off-by: jsvisa <delweng@gmail.com>
This commit is contained in:
jsvisa 2025-08-18 12:01:53 +08:00
parent 25d3b5d88b
commit b45b0419d7
2 changed files with 74 additions and 39 deletions

View file

@ -38,9 +38,21 @@ func storageKey(accountHash common.Hash, slotHash common.Hash) [64]byte {
return key
}
// trienodeKey returns a key for uniquely identifying the trie node.
func trienodeKey(accountHash common.Hash, path string) string {
return accountHash.Hex() + path
// trienodeKey uses a fixed-size byte array instead of string to avoid string allocations.
type trienodeKey [96]byte // 32 bytes for hash + up to 64 bytes for path
// makeTrienodeKey returns a key for uniquely identifying the trie node.
func makeTrienodeKey(accountHash common.Hash, path string) trienodeKey {
var key trienodeKey
copy(key[:32], accountHash[:])
copy(key[32:], path)
return key
}
// shardTask used to batch task by shard to minimize lock contention
type shardTask struct {
accountHash common.Hash
path string
}
// lookup is an internal structure used to efficiently determine the layer in
@ -69,7 +81,7 @@ type lookup struct {
// The key is the account address hash and the trie path of the node,
// the value is a slice of **diff layer** IDs indicating where the
// slot was modified, with the order from oldest to newest.
storageNodes [storageNodesShardCount]map[string][]common.Hash
storageNodes [storageNodesShardCount]map[trienodeKey][]common.Hash
// descendant is the callback indicating whether the layer with
// given root is a descendant of the one specified by `ancestor`.
@ -103,7 +115,7 @@ func newLookup(head layer, descendant func(state common.Hash, ancestor common.Ha
}
// Initialize all 16 storage node shards
for i := 0; i < storageNodesShardCount; i++ {
l.storageNodes[i] = make(map[string][]common.Hash)
l.storageNodes[i] = make(map[trienodeKey][]common.Hash)
}
// Apply the diff layers from bottom to top
@ -216,7 +228,7 @@ func (l *lookup) nodeTip(accountHash common.Hash, path string, stateID common.Ha
list = l.accountNodes[path]
} else {
shardIndex := getStorageShardIndex(path) // Use only path for sharding
list = l.storageNodes[shardIndex][trienodeKey(accountHash, path)]
list = l.storageNodes[shardIndex][makeTrienodeKey(accountHash, path)]
}
for i := len(list) - 1; i >= 0; i-- {
// If the current state matches the stateID, or the requested state is a
@ -323,19 +335,38 @@ func (l *lookup) addStorageNodes(state common.Hash, nodes map[common.Hash]map[st
var (
wg sync.WaitGroup
tasks = make([]chan string, storageNodesShardCount)
locks [storageNodesShardCount]sync.Mutex
tasks = make([][]shardTask, storageNodesShardCount)
)
wg.Add(storageNodesShardCount)
for i := 0; i < storageNodesShardCount; i++ {
tasks[i] = make(chan string, 10) // Buffer to avoid blocking
// Pre-allocate work lists
for accountHash, slots := range nodes {
for path := range slots {
shardIndex := getStorageShardIndex(path)
tasks[shardIndex] = append(tasks[shardIndex], shardTask{
accountHash: accountHash,
path: path,
})
}
}
// Start all workers, each handling its own shard
wg.Add(storageNodesShardCount)
for shardIndex := 0; shardIndex < storageNodesShardCount; shardIndex++ {
go func(shardIdx int) {
defer wg.Done()
taskList := tasks[shardIdx]
if len(taskList) == 0 {
return
}
locks[shardIdx].Lock()
defer locks[shardIdx].Unlock()
shard := l.storageNodes[shardIdx]
for key := range tasks[shardIdx] {
for _, task := range taskList {
key := makeTrienodeKey(task.accountHash, task.path)
list, exists := shard[key]
if !exists {
list = make([]common.Hash, 0, 16) // TODO(rjl493456442) use sync pool
@ -343,18 +374,8 @@ func (l *lookup) addStorageNodes(state common.Hash, nodes map[common.Hash]map[st
list = append(list, state)
shard[key] = list
}
}(shardIndex)
}
for accountHash, slots := range nodes {
for path := range slots {
shardIndex := getStorageShardIndex(path)
tasks[shardIndex] <- trienodeKey(accountHash, path)
}
}
// Close all channels to signal workers to finish
for i := 0; i < storageNodesShardCount; i++ {
close(tasks[i])
}(shardIndex)
}
wg.Wait()
}
@ -456,20 +477,39 @@ func (l *lookup) removeStorageNodes(state common.Hash, nodes map[common.Hash]map
var (
eg errgroup.Group
tasks = make([]chan string, storageNodesShardCount)
locks [storageNodesShardCount]sync.Mutex
tasks = make([][]shardTask, storageNodesShardCount)
)
for i := 0; i < storageNodesShardCount; i++ {
tasks[i] = make(chan string, 10) // Buffer to avoid blocking
// Pre-allocate work lists
for accountHash, slots := range nodes {
for path := range slots {
shardIndex := getStorageShardIndex(path)
tasks[shardIndex] = append(tasks[shardIndex], shardTask{
accountHash: accountHash,
path: path,
})
}
}
// Start all workers, each handling its own shard
for shardIndex := 0; shardIndex < storageNodesShardCount; shardIndex++ {
shardIdx := shardIndex // Capture the variable
eg.Go(func() error {
taskList := tasks[shardIdx]
if len(taskList) == 0 {
return nil
}
locks[shardIdx].Lock()
defer locks[shardIdx].Unlock()
shard := l.storageNodes[shardIdx]
for key := range tasks[shardIdx] {
for _, task := range taskList {
key := makeTrienodeKey(task.accountHash, task.path)
found, list := removeFromList(shard[key], state)
if !found {
return fmt.Errorf("storage lookup is not found, key: %s, state: %x", key, state)
return fmt.Errorf("storage lookup is not found, key: %x, state: %x", key, state)
}
if len(list) != 0 {
shard[key] = list
@ -480,15 +520,5 @@ func (l *lookup) removeStorageNodes(state common.Hash, nodes map[common.Hash]map
return nil
})
}
for accountHash, slots := range nodes {
for path := range slots {
shardIndex := getStorageShardIndex(path)
tasks[shardIndex] <- trienodeKey(accountHash, path)
}
}
for i := 0; i < storageNodesShardCount; i++ {
close(tasks[i])
}
return eg.Wait()
}

View file

@ -63,6 +63,11 @@ func BenchmarkAddNodes(b *testing.B) {
accountNodeCount: 2000,
nodesPerAccount: 40,
},
{
name: "XLarge-5000-accounts-50-nodes",
accountNodeCount: 5000,
nodesPerAccount: 50,
},
}
for _, tc := range tests {
@ -75,7 +80,7 @@ func BenchmarkAddNodes(b *testing.B) {
// Initialize all 16 storage node shards
for i := 0; i < storageNodesShardCount; i++ {
lookup.storageNodes[i] = make(map[string][]common.Hash)
lookup.storageNodes[i] = make(map[trienodeKey][]common.Hash)
}
var state common.Hash
@ -87,7 +92,7 @@ func BenchmarkAddNodes(b *testing.B) {
for i := 0; i < b.N; i++ {
// Reset the lookup instance for each benchmark iteration
for j := 0; j < storageNodesShardCount; j++ {
lookup.storageNodes[j] = make(map[string][]common.Hash)
lookup.storageNodes[j] = make(map[trienodeKey][]common.Hash)
}
lookup.addStorageNodes(state, storageNodes)