use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
#[derive(Clone, Default)]
pub(crate) struct CancelToken(Arc<AtomicBool>);
impl CancelToken {
pub(crate) fn new() -> Self {
Self::default()
}
pub(crate) fn cancel(&self) {
self.0.store(true, Ordering::Relaxed);
}
pub(crate) fn is_cancelled(&self) -> bool {
self.0.load(Ordering::Relaxed)
}
}
#[derive(Default)]
pub(crate) struct CancelRegistry {
live: Mutex<HashMap<String, CancelToken>>,
}
impl CancelRegistry {
pub(crate) fn new() -> Self {
Self::default()
}
pub(crate) fn register(self: &Arc<Self>, request_id: &str) -> (CancelToken, CancelGuard) {
let token = CancelToken::new();
self.live
.lock()
.unwrap_or_else(|p| p.into_inner())
.insert(request_id.to_string(), token.clone());
(
token,
CancelGuard {
registry: Arc::clone(self),
request_id: request_id.to_string(),
},
)
}
pub(crate) fn cancel(&self, request_id: &str) -> bool {
match self
.live
.lock()
.unwrap_or_else(|p| p.into_inner())
.get(request_id)
{
Some(token) => {
token.cancel();
true
}
None => false,
}
}
pub(crate) fn live_count(&self) -> usize {
self.live.lock().unwrap_or_else(|p| p.into_inner()).len()
}
fn deregister(&self, request_id: &str) {
self.live
.lock()
.unwrap_or_else(|p| p.into_inner())
.remove(request_id);
}
}
pub(crate) struct CancelGuard {
registry: Arc<CancelRegistry>,
request_id: String,
}
impl Drop for CancelGuard {
fn drop(&mut self) {
self.registry.deregister(&self.request_id);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_registered_generation_can_be_cancelled_by_id() {
let registry = Arc::new(CancelRegistry::new());
let (token, _guard) = registry.register("chatcmpl-1");
assert!(!token.is_cancelled());
assert!(registry.cancel("chatcmpl-1"));
assert!(token.is_cancelled());
}
#[test]
fn cancelling_an_unknown_id_reports_that_nothing_was_running() {
let registry = Arc::new(CancelRegistry::new());
assert!(!registry.cancel("chatcmpl-never-issued"));
}
#[test]
fn the_registration_ends_with_the_generation() {
let registry = Arc::new(CancelRegistry::new());
{
let (_token, _guard) = registry.register("chatcmpl-2");
assert_eq!(registry.live_count(), 1);
}
assert_eq!(registry.live_count(), 0);
assert!(!registry.cancel("chatcmpl-2"));
}
#[test]
fn a_panicking_generation_does_not_leak_its_id() {
let registry = Arc::new(CancelRegistry::new());
let for_thread = Arc::clone(®istry);
let handle = std::thread::spawn(move || {
let (_token, _guard) = for_thread.register("chatcmpl-3");
panic!("decode thread died");
});
assert!(handle.join().is_err());
assert_eq!(registry.live_count(), 0);
}
#[test]
fn concurrent_generations_are_cancelled_independently() {
let registry = Arc::new(CancelRegistry::new());
let (first, _g1) = registry.register("chatcmpl-a");
let (second, _g2) = registry.register("chatcmpl-b");
assert!(registry.cancel("chatcmpl-a"));
assert!(first.is_cancelled());
assert!(
!second.is_cancelled(),
"cancelling one chat must not stop another"
);
}
}