package main import ( "hash/fnv" "sort" "sync" ) type Shard struct{ mu sync.RWMutex id int data map[string]Item } func (db *RedisDB) getShard(key string)*Shard{ hash := fnv.New64a() hash.Write([]byte(key)) val := hash.Sum64() return db.shards[val%NumShards] } func (db *RedisDB) Execute(key string, fn func(*Shard) interface{}) interface{} { shard := db.getShard(key) shard.mu.Lock() defer shard.mu.Unlock() return fn(shard) } func (db *RedisDB) ExecuteRead(key string, fn func(*Shard) interface{}) interface{} { shard := db.getShard(key) shard.mu.RLock() defer shard.mu.RUnlock() return fn(shard) } func ( db *RedisDB) ExecuteMulti(keys []string, fn func([]*Shard) interface{})interface{}{ if len(keys) == 1{ shard := db.getShard(keys[0]) shard.mu.Lock() defer shard.mu.Unlock() shards := []*Shard{shard} return fn(shards) }else{ shardMap := make(map[int]*Shard) for _, k := range keys { shard := db.getShard(k) shardMap[shard.id] = shard } var sortedIDs []int for id := range shardMap { sortedIDs = append(sortedIDs, id) } sort.Ints(sortedIDs) for _, id := range sortedIDs{ shardMap[id].mu.Lock() } defer func(){ for i := len(sortedIDs)-1; i >= 0; i --{ shardMap[sortedIDs[i]].mu.Unlock() } }() shards := make([]*Shard, 0, len(shardMap)) for _, id := range sortedIDs { shards = append(shards, shardMap[id]) } return fn(shards) } }