You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
219 lines
4.6 KiB
219 lines
4.6 KiB
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() |
|
}
|
|
|