package collect 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函数 func NewConcurrentMap[K comparable, V any](concurrencyLevel int, hashKeyFunc func(K) string) *ConcurrentMap[K, V] { segments := 1 // segments = 2^n for segments < concurrencyLevel { segments <<= 1 } 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 - 1) } func (m *ConcurrentMap[K, V]) update(k K, update func(map[K]V)) { segment := m.segment(k) lock := m.segmentsLock[segment] lock.Lock() defer lock.Unlock() update(m.segmentsMap[segment]) } // Store 放置新值 func (m *ConcurrentMap[K, V]) Store(k K, v V) { m.update(k, func(segment map[K]V) { 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) { m.update(k, func(segment map[K]V) { delete(segment, k) }) } func (m *ConcurrentMap[K, V]) Range(f func(key K, value V) 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 (m *ConcurrentMap[K, V]) LoadAndDelete(k K) (prev V, loaded bool) { m.update(k, func(segment map[K]V) { prev, loaded = segment[k] delete(segment, k) }) return } func (m *ConcurrentMap[K, V]) Swap(k K, v V) (prev V, loaded bool) { m.update(k, func(segment map[K]V) { prev, loaded = segment[k] segment[k] = v }) return } // ComputeIfAbsent 加载, 如果值不存在使用 mapping(k) 填充并返回 // mapped 值是否是 mapping(k) 填充的 func (m *ConcurrentMap[K, V]) ComputeIfAbsent(k K, mapping func(k K) V) (res V, mapped bool) { segment := m.segment(k) lock := m.segmentsLock[segment] lock.RLock() v, ok := m.segmentsMap[segment][k] lock.RUnlock() if ok { res = v return } lock.Lock() defer lock.Unlock() // double check if v, ok = m.segmentsMap[segment][k]; ok { res = v return } // write mapping value res = mapping(k) m.segmentsMap[segment][k] = res mapped = true return } func fnv32Hash(k string) uint32 { f := fnv.New32() _, err := f.Write([]byte(k)) if err != nil { panic(err) } return f.Sum32() }