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 & uint32(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 }) } 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) }) } // Range 遍历时勿修改值造成死锁 func (m *ConcurrentMap[K, V]) Range(f func(key K, value V) (next 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]) RangeUpdate(f func(key K, value V) (next, remove bool, newV V)) { for i := 0; i < m.segments; i++ { lock := m.segmentsLock[i] func() { lock.Lock() defer lock.Unlock() for k, v := range m.segmentsMap[i] { next, remove, newV := f(k, v) if remove { delete(m.segmentsMap[i], k) } else { m.segmentsMap[i][k] = newV } if !next { i = m.segments // stop range return } } }() } } func (m *ConcurrentMap[K, V]) LoadRLock(k K, f func(v V, ok bool)) { segment := m.segment(k) lock := m.segmentsLock[segment] lock.RLock() defer lock.RUnlock() v, ok := m.segmentsMap[segment][k] f(v, ok) } // LoadAndUpdate 更新新值 func (m *ConcurrentMap[K, V]) LoadAndUpdate(k K, f func(v V) (remove bool, nextV V)) (value V) { m.update(k, func(segment map[K]V) { remove, nextV := f(segment[k]) if remove { delete(segment, k) } else { segment[k] = nextV value = nextV } }) return } func (m *ConcurrentMap[K, V]) LoadAndDelete(k K) (value V, loaded bool) { m.update(k, func(segment map[K]V) { value, loaded = segment[k] if loaded { 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 } func (m *ConcurrentMap[K, V]) Size() (size int) { for i := 0; i < m.segments; i++ { lock := m.segmentsLock[i] func() { lock.RLock() defer lock.RUnlock() size += len(m.segmentsMap[i]) }() } return } // ComputeIfAbsent 加载, 如果值不存在使用 mapping(k) 填充并返回 // mapped 值是否是 mapping(k) 填充的 func (m *ConcurrentMap[K, V]) ComputeIfAbsent(k K, mapping func(k K) V) (res V, mapped bool) { res, _, mapped = m.ComputeIfAbsentE(k, func(k K) (V, error) { return mapping(k), nil }) return } // ComputeIfAbsentE 加载, 如果值不存在使用 mapping(k) 填充并返回 // mapped 值是否是 mapping(k) 填充的 func (m *ConcurrentMap[K, V]) ComputeIfAbsentE(k K, mapping func(k K) (V, error)) (res V, err error, 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, err = mapping(k) if err != nil { return } 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() }