7 changed files with 164 additions and 12 deletions
@ -0,0 +1,99 @@
|
||||
package group |
||||
|
||||
import ( |
||||
"hash/fnv" |
||||
"sync" |
||||
) |
||||
|
||||
// ConcurrentMap 分段锁 map, 提升并发性
|
||||
type ConcurrentMap[K comparable, V any] struct { |
||||
hashKeyFunc func(K) string |
||||
equalsFunc func(v1, v2 V) bool |
||||
counter int64 |
||||
segments int |
||||
segmentsMap []map[K]V |
||||
segmentsLock []*sync.RWMutex |
||||
} |
||||
|
||||
// NewConcurrentMap 分段锁并发 map
|
||||
// segments: 分段数
|
||||
// hashKeyFunc: key转string函数
|
||||
// equalsFunc: value比较函数, Put时新旧值相同则不返回旧值
|
||||
func NewConcurrentMap[K comparable, V any](segments int, hashKeyFunc func(K) string) *ConcurrentMap[K, V] { |
||||
m := &ConcurrentMap[K, V]{ |
||||
hashKeyFunc: hashKeyFunc, |
||||
segments: segments, |
||||
segmentsMap: make([]map[K]V, segments), |
||||
segmentsLock: make([]*sync.RWMutex, segments), |
||||
} |
||||
for i := 0; i < segments; i++ { |
||||
m.segmentsMap[i] = make(map[K]V, 16) |
||||
m.segmentsLock[i] = &sync.RWMutex{} |
||||
} |
||||
return m |
||||
} |
||||
|
||||
// segment 根据 key 确定分段
|
||||
func (m *ConcurrentMap[K, V]) segment(k K) int { |
||||
hashK := m.hashKeyFunc(k) |
||||
hash := fnv32Hash(hashK) |
||||
return int(hash) % m.segments |
||||
} |
||||
|
||||
// Store 放置新值
|
||||
func (m *ConcurrentMap[K, V]) Store(k K, v V) { // (old V, hasOld bool) // 返回旧值
|
||||
segment := m.segment(k) |
||||
lock := m.segmentsLock[segment] |
||||
lock.Lock() |
||||
defer lock.Unlock() |
||||
//if prev, ok := m.segmentsMap[segment][k]; ok {
|
||||
// if !m.equalsFunc(v, prev) { // 两值不同返回旧值
|
||||
// old, hasOld = prev, true
|
||||
// }
|
||||
//}
|
||||
m.segmentsMap[segment][k] = v |
||||
return |
||||
} |
||||
|
||||
func (m *ConcurrentMap[K, V]) Load(k K) (v V, ok bool) { |
||||
segment := m.segment(k) |
||||
lock := m.segmentsLock[segment] |
||||
lock.RLock() |
||||
defer lock.RUnlock() |
||||
v, ok = m.segmentsMap[segment][k] |
||||
return |
||||
} |
||||
|
||||
func (m *ConcurrentMap[K, V]) Delete(k K) { |
||||
segment := m.segment(k) |
||||
lock := m.segmentsLock[segment] |
||||
lock.Lock() |
||||
defer lock.Unlock() |
||||
delete(m.segmentsMap[segment], k) |
||||
} |
||||
|
||||
func (m *ConcurrentMap[K, V]) Range(f func(key, value any) bool) { |
||||
for i := 0; i < m.segments; i++ { |
||||
lock := m.segmentsLock[i] |
||||
func() { |
||||
lock.RLock() |
||||
defer lock.RUnlock() |
||||
|
||||
for k, v := range m.segmentsMap[i] { |
||||
if !f(k, v) { |
||||
i = m.segments // stop range
|
||||
return |
||||
} |
||||
} |
||||
}() |
||||
} |
||||
} |
||||
|
||||
func fnv32Hash(k string) uint32 { |
||||
f := fnv.New32() |
||||
_, err := f.Write([]byte(k)) |
||||
if err != nil { |
||||
panic(err) |
||||
} |
||||
return f.Sum32() |
||||
} |
||||
@ -0,0 +1,46 @@
|
||||
package group |
||||
|
||||
import ( |
||||
"fmt" |
||||
"math/rand" |
||||
"sync" |
||||
"testing" |
||||
"time" |
||||
) |
||||
|
||||
func TestConcurrentMap(t *testing.T) { |
||||
cm := NewConcurrentMap[string, string](16, func(k string) string { return k }) |
||||
concurrent := 1000 |
||||
|
||||
wg := sync.WaitGroup{} |
||||
wg.Add(concurrent) |
||||
|
||||
for i := 0; i < concurrent; i++ { |
||||
go func(loop int) { |
||||
for j := 0; j < concurrent; j++ { |
||||
k := fmt.Sprintf("%d-%d", loop, j) |
||||
cm.Store(k, k) |
||||
} |
||||
wg.Done() |
||||
}(i) |
||||
} |
||||
wg.Wait() |
||||
|
||||
wg.Add(concurrent) |
||||
|
||||
for i := 0; i < concurrent; i++ { |
||||
go func() { |
||||
r := rand.New(rand.NewSource(time.Now().UnixMilli())) |
||||
for i := 0; i < concurrent; i++ { |
||||
k := fmt.Sprintf("%d-%d", r.Intn(concurrent), r.Intn(concurrent)) |
||||
v, ok := cm.Load(k) |
||||
if !ok || v != k { |
||||
t.Errorf("value load error: %s", k) |
||||
} |
||||
} |
||||
wg.Done() |
||||
}() |
||||
} |
||||
|
||||
wg.Wait() |
||||
} |
||||
Loading…
Reference in new issue