use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use tokio::sync::watch;
pub type SfShard = Mutex<HashMap<String, Arc<watch::Sender<()>>>>;
pub const SF_SHARDS: usize = 64;
pub fn shard_index(key: &str) -> usize {
use std::hash::{Hash, Hasher};
let mut hasher = std::collections::hash_map::DefaultHasher::new();
key.hash(&mut hasher);
(hasher.finish() as usize) & (SF_SHARDS - 1)
}
pub struct AsyncSfGuard {
shards: &'static [SfShard; SF_SHARDS],
idx: usize,
key: String,
signal: Arc<watch::Sender<()>>,
finished: bool,
}
impl AsyncSfGuard {
pub fn new(
shards: &'static [SfShard; SF_SHARDS],
idx: usize,
key: String,
signal: Arc<watch::Sender<()>>,
) -> Self {
Self {
shards,
idx,
key,
signal,
finished: false,
}
}
pub fn finish(&mut self) {
if self.finished {
return;
}
self.finished = true;
if let Ok(mut map) = self.shards[self.idx].lock() {
map.remove(&self.key);
}
let _ = self.signal.send(());
}
}
impl Drop for AsyncSfGuard {
fn drop(&mut self) {
if !self.finished {
self.finish();
}
}
}
pub async fn wait_flight(mut rx: watch::Receiver<()>) {
let _ = rx.changed().await;
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
static SHARDS: std::sync::LazyLock<[SfShard; SF_SHARDS]> =
std::sync::LazyLock::new(|| std::array::from_fn(|_| Mutex::new(HashMap::new())));
fn register(key: &str) -> (Arc<watch::Sender<()>>, watch::Receiver<()>) {
let idx = shard_index(key);
let mut map = SHARDS[idx].lock().unwrap();
match map.entry(key.to_string()) {
std::collections::hash_map::Entry::Occupied(e) => {
let tx = e.get().clone();
let rx = tx.subscribe();
(tx, rx)
}
std::collections::hash_map::Entry::Vacant(e) => {
let (tx, _rx) = watch::channel(());
let tx = Arc::new(tx);
let rx = tx.subscribe();
e.insert(tx.clone());
(tx, rx)
}
}
}
#[test]
fn test_shard_index_in_range_and_stable() {
for key in [
"",
"a",
"user:123",
"很长很长的中文key🎯",
&"x".repeat(1024),
] {
let idx = shard_index(key);
assert!(idx < SF_SHARDS);
assert_eq!(idx, shard_index(key));
}
}
#[tokio::test]
async fn test_finish_removes_entry_and_signals() {
let (tx, rx) = register("k-finish");
let mut guard = AsyncSfGuard::new(&SHARDS, shard_index("k-finish"), "k-finish".into(), tx);
guard.finish();
guard.finish();
assert!(
SHARDS[shard_index("k-finish")]
.lock()
.unwrap()
.get("k-finish")
.is_none()
);
let waited = tokio::time::timeout(Duration::from_secs(1), wait_flight(rx)).await;
assert!(waited.is_ok(), "finish 后等待者必须被放行");
}
#[tokio::test]
async fn test_drop_without_finish_cleans_and_releases() {
let (tx, rx) = register("k-drop");
{
let _guard = AsyncSfGuard::new(&SHARDS, shard_index("k-drop"), "k-drop".into(), tx);
}
assert!(
SHARDS[shard_index("k-drop")]
.lock()
.unwrap()
.get("k-drop")
.is_none()
);
let waited = tokio::time::timeout(Duration::from_secs(1), wait_flight(rx)).await;
assert!(waited.is_ok(), "Drop 兜底必须放行等待者");
}
#[tokio::test]
async fn test_wait_flight_survives_sender_drop() {
let (tx, _rx) = watch::channel(());
let rx2 = tx.subscribe();
drop(tx);
let waited = tokio::time::timeout(Duration::from_secs(1), wait_flight(rx2)).await;
assert!(waited.is_ok(), "通道关闭必须放行等待者");
}
}