summaryrefslogtreecommitdiff
path: root/src/context.rs
diff options
context:
space:
mode:
authorMica White <botahamec@outlook.com>2026-08-26 20:46:31 -0400
committerMica White <botahamec@outlook.com>2026-08-26 20:46:31 -0400
commit55b3a2425b242fbc5c6e471f220eeb8b949e8751 (patch)
tree5ceb7910b60cc7331024b7ce1d49ec259e0302a8 /src/context.rs
parent6f6e030ea7edb9d155ebf21b5d42936c20801b50 (diff)
Add tests
Diffstat (limited to 'src/context.rs')
-rw-r--r--src/context.rs617
1 files changed, 617 insertions, 0 deletions
diff --git a/src/context.rs b/src/context.rs
new file mode 100644
index 0000000..2f9ca64
--- /dev/null
+++ b/src/context.rs
@@ -0,0 +1,617 @@
+use std::marker::PhantomData;
+
+use crate::ThreadKey;
+
+mod context;
+mod guard;
+pub mod iterator;
+pub mod tuple;
+
+/// Allows iterating over a lock collection, without locking every element at
+/// once.
+///
+/// Sometimes, partial allocation of locks is useful. For example, you may
+/// want to acquire a lock on the first element of a list before deciding if
+/// the second element should be locked. This function creates a
+/// [`LockContext`] which is capable of doing exactly that.
+///
+/// Upon using this context, the [`ThreadKey`] is stored inside this context.
+/// This ensures that nothing else, besides the types exposed by the context,
+/// can be locked until this context is dropped. To re-acquire the `ThreadKey`,
+/// call [`LockContext::unlock`].
+///
+/// A [`LockContext`] can be created by calling [`OwnedLockCollection::context`].
+///
+/// # Examples
+///
+/// Iterating through a tuple.
+///
+/// ```
+/// use happylock::{Mutex, ThreadKey};
+/// use happylock::collection::OwnedLockCollection;
+///
+/// let key = ThreadKey::get().unwrap();
+/// let data = (Mutex::new(true), Mutex::new(42), Mutex::new(67));
+/// let locks = OwnedLockCollection::new(data);
+/// let mut ctx = locks.context();
+/// let tuple = ctx.tuple(key);
+///
+/// let (use_other, tuple) = tuple.lock_0();
+/// let number = if **use_other {
+/// tuple.lock_2().0
+/// } else {
+/// tuple.lock_1().0
+/// };
+/// assert_eq!(**number, 67);
+/// ```
+///
+/// Iterating through a list
+///
+/// ```
+/// use happylock::{Mutex, ThreadKey};
+/// use happylock::collection::OwnedLockCollection;
+///
+/// let key = ThreadKey::get().unwrap();
+/// let data = [Mutex::new(1), Mutex::new(3), Mutex::new(8)];
+/// let locks = OwnedLockCollection::new(data);
+/// let mut ctx = locks.context();
+/// let mut iter = ctx.iter(key);
+///
+/// let mut sum = 0;
+/// while let Some(item) = iter.lock_next() {
+/// sum += **item;
+/// }
+///
+/// assert_eq!(sum, 12);
+/// ```
+///
+/// [`OwnedLockCollection::context`]: crate::collection::OwnedLockCollection::context
+pub struct LockContext<'l, L> {
+ key: Option<ThreadKey>,
+ lockable: &'l L,
+}
+
+/// Iterates through a collection of locks, allowing for partial allocation of
+/// locks, or for some locks to be skipped.
+///
+/// Sometimes, partial allocation of locks is useful. For example, you may
+/// want to acquire a lock on the first element of a list before deciding if
+/// the second element should be locked. If the list is iterable, then a
+/// [`LockingIterator`] is capable of doing exactly that.
+///
+/// A [`LockingIterator`] can be created by calling the [`LockContext::iter`]
+/// method.
+///
+/// # Example
+///
+/// ```
+/// use happylock::{Mutex, ThreadKey};
+/// use happylock::collection::OwnedLockCollection;
+///
+/// let key = ThreadKey::get().unwrap();
+/// let data = [Mutex::new(1), Mutex::new(3), Mutex::new(8)];
+/// let locks = OwnedLockCollection::new(data);
+/// let mut ctx = locks.context();
+/// let mut iter = ctx.iter(key);
+///
+/// let mut sum = 0;
+/// while let Some(item) = iter.lock_next() {
+/// sum += **item;
+/// }
+///
+/// assert_eq!(sum, 12);
+/// ```
+pub struct LockingIterator<'context, I, Outer = ()> {
+ key: &'context ThreadKey,
+ iterator: I,
+ outer: Outer,
+}
+
+/// Iterates through a tuple of locks, requiring that elements are only locked
+/// before any successive elements are locked.
+///
+/// Sometimes, partial allocation of locks is useful. For example, you may want
+/// to acquire a lock on one item before deciding if the second item should be
+/// locked. If the locks can be organized into a tuple, [`LockingTuple`] is
+/// capable of doing exactly that.
+///
+/// A [`LockingTuple`] can be created by calling the [`LockContext::iter`]
+/// method.
+///
+/// # Example
+///
+/// ```
+/// use happylock::{Mutex, ThreadKey};
+/// use happylock::collection::OwnedLockCollection;
+///
+/// let key = ThreadKey::get().unwrap();
+/// let data = (Mutex::new(true), Mutex::new(42), Mutex::new(67));
+/// let locks = OwnedLockCollection::new(data);
+/// let mut ctx = locks.context();
+/// let tuple = ctx.tuple(key);
+///
+/// let (use_other, tuple) = tuple.lock_0();
+/// let number = if **use_other {
+/// tuple.lock_2().0
+/// } else {
+/// tuple.lock_1().0
+/// };
+/// assert_eq!(**number, 67);
+/// ```
+pub struct LockingTuple<'context, L, C, Outer = ()> {
+ _lockable: PhantomData<L>,
+ key: &'context ThreadKey,
+ tuple: &'context C,
+ outer: Outer,
+}
+
+/// An RAII implementation of a “scoped lock”. When this structure
+/// is dropped (falls out of scope), the lock will be unlocked.
+///
+/// The data protected by the mutex can be accessed through this guard via its
+/// [`Deref`] and [`DerefMut`] implementations.
+///
+/// This is created by calling the [`lock`] and [`try_lock`] methods on [`Mutex`]
+///
+/// Unlike other guards in this crate, this guard holds a reference to a
+/// [`ThreadKey`], which is stored in the [`LockContext`]. This ensures that
+/// context cannot be dropped until all guards created with the context are
+/// dropped. The `ThreadKey` can be re-acquired by calling
+/// [`LockContext::unlock`].
+///
+/// [`Mutex`]: `crate::mutex::Mutex`
+/// [`Deref`]: `std::ops::Deref`
+/// [`DerefMut`]: `std::ops::DerefMut`
+/// [`lock`]: `crate::mutex::Mutex::lock`
+/// [`try_lock`]: `crate::Mutex::try_lock`
+pub struct ContextGuard<'a, Guard, Key> {
+ _key: &'a Key,
+ guard: Guard,
+}
+
+#[cfg(test)]
+mod tests {
+ use crate::{
+ collection::OwnedLockCollection, context::iterator::TryLockNextError, Mutex, RwLock,
+ ThreadKey,
+ };
+
+ #[test]
+ fn display_works_for_guard() {
+ let key = ThreadKey::get().unwrap();
+ let collection = OwnedLockCollection::new((Mutex::new("Hello, world!"),));
+ let mut context = collection.context();
+ let tuple = context.tuple(key);
+ let (guard, _) = tuple.lock_0();
+ assert_eq!(guard.to_string(), "Hello, world!".to_string());
+ }
+
+ #[test]
+ fn try_lock_mut_works_single_element_tuple() {
+ let key = ThreadKey::get().unwrap();
+ let collection = OwnedLockCollection::new((RwLock::new(42),));
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+ let mut guard = tuple.try_lock_mut_0().unwrap();
+ **guard = 67;
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+ let guard = tuple.try_read_mut_0();
+ assert!(guard.is_none());
+ });
+ });
+ drop(guard);
+
+ let result = tuple.try_read_mut_0().unwrap();
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+ let guard = tuple.try_lock_mut_0();
+ assert!(guard.is_none());
+ drop(guard);
+ let guard = tuple.try_read_mut_0();
+ assert!(guard.is_some());
+ drop(guard);
+ });
+ });
+ assert_eq!(**result, 67);
+ }
+
+ #[test]
+ fn try_lock_mut_works_double_element_tuple() {
+ let key = ThreadKey::get().unwrap();
+ let collection = OwnedLockCollection::new((RwLock::new(42), RwLock::new(67)));
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+
+ let mut guard = tuple.try_lock_mut_0().unwrap();
+ **guard *= 2;
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+ let guard = tuple.try_read_mut_0();
+ assert!(guard.is_none());
+ });
+ });
+ drop(guard);
+ let mut guard = tuple.try_lock_mut_1().unwrap();
+ **guard *= 2;
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+ let guard = tuple.try_read_mut_1();
+ assert!(guard.is_none());
+ });
+ });
+ drop(guard);
+
+ let result = tuple.try_read_mut_0().unwrap();
+ assert_eq!(**result, 84);
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+ let guard = tuple.try_lock_mut_0();
+ assert!(guard.is_none());
+ drop(guard);
+ let guard = tuple.try_read_mut_0();
+ assert!(guard.is_some());
+ drop(guard);
+ });
+ });
+ drop(result);
+ let result = tuple.try_read_mut_1().unwrap();
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+ let guard = tuple.try_lock_mut_1();
+ assert!(guard.is_none());
+ drop(guard);
+ let guard = tuple.try_read_mut_1();
+ assert!(guard.is_some());
+ drop(guard);
+ });
+ });
+ assert_eq!(**result, 134);
+ }
+
+ #[test]
+ fn try_lock_mut_works_triple_element_tuple() {
+ let key = ThreadKey::get().unwrap();
+ let collection = OwnedLockCollection::new((RwLock::new(1), RwLock::new(2), RwLock::new(3)));
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+
+ let mut guard = tuple.try_lock_mut_0().unwrap();
+ **guard *= 2;
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+ let guard = tuple.try_read_mut_0();
+ assert!(guard.is_none());
+ });
+ });
+ drop(guard);
+ let mut guard = tuple.try_lock_mut_1().unwrap();
+ **guard *= 2;
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+ let guard = tuple.try_read_mut_1();
+ assert!(guard.is_none());
+ });
+ });
+ drop(guard);
+ let mut guard = tuple.try_lock_mut_2().unwrap();
+ **guard *= 2;
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+ let guard = tuple.try_read_mut_2();
+ assert!(guard.is_none());
+ });
+ });
+ drop(guard);
+
+ let result = tuple.try_read_mut_0().unwrap();
+ assert_eq!(**result, 2);
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+ let guard = tuple.try_lock_mut_0();
+ assert!(guard.is_none());
+ drop(guard);
+ let guard = tuple.try_read_mut_0();
+ assert!(guard.is_some());
+ drop(guard);
+ });
+ });
+ drop(result);
+ let result = tuple.try_read_mut_1().unwrap();
+ assert_eq!(**result, 4);
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+ let guard = tuple.try_lock_mut_1();
+ assert!(guard.is_none());
+ drop(guard);
+ let guard = tuple.try_read_mut_1();
+ assert!(guard.is_some());
+ drop(guard);
+ });
+ });
+ drop(result);
+ let result = tuple.try_read_mut_2().unwrap();
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = collection.context();
+ let mut tuple = context.tuple(key);
+ let guard = tuple.try_lock_mut_2();
+ assert!(guard.is_none());
+ drop(guard);
+ let guard = tuple.try_read_mut_2();
+ assert!(guard.is_some());
+ drop(guard);
+ });
+ });
+ assert_eq!(**result, 6);
+ }
+
+ #[test]
+ fn basic_iteration() {
+ let key = ThreadKey::get().unwrap();
+ let data = [Mutex::new(1), Mutex::new(3), Mutex::new(8)];
+ let locks = OwnedLockCollection::new(data);
+ let mut ctx = locks.context();
+ let mut iter = ctx.iter(key);
+
+ let item = iter.lock_next().unwrap();
+ assert_eq!(**item, 1);
+ let item = iter.lock_next().unwrap();
+ assert_eq!(**item, 3);
+ let item = iter.lock_next().unwrap();
+ assert_eq!(**item, 8);
+ assert!(iter.lock_next().is_none());
+ }
+
+ #[test]
+ fn recurse_tuple_of_lists() {
+ let key = ThreadKey::get().unwrap();
+ let data = (
+ [Mutex::new(1), Mutex::new(2), Mutex::new(3)],
+ Mutex::new(true),
+ );
+ let locks = OwnedLockCollection::new(data);
+ let mut ctx = locks.context();
+ let tuple = ctx.tuple(key);
+ let mut iter = tuple.recurse_0_iter();
+
+ let item = iter.lock_next().unwrap();
+ assert_eq!(**item, 1);
+ let item = iter.lock_next().unwrap();
+ assert_eq!(**item, 2);
+ let item = iter.lock_next().unwrap();
+ assert_eq!(**item, 3);
+ assert!(iter.lock_next().is_none());
+
+ let tuple = iter.exit();
+ let (should_assert, _) = tuple.lock_1();
+ assert!(**should_assert);
+ }
+
+ #[test]
+ fn recurse_list_of_lists() {
+ let key = ThreadKey::get().unwrap();
+ let data = [
+ [Mutex::new(1), Mutex::new(2), Mutex::new(3)],
+ [Mutex::new(4), Mutex::new(5), Mutex::new(6)],
+ ];
+ let locks = OwnedLockCollection::new(data);
+ let mut ctx = locks.context();
+ let mut iter = ctx.iter(key);
+
+ let mut list = iter.recurse_next().unwrap();
+ let item = list.lock_next().unwrap();
+ assert_eq!(**item, 1);
+ let item = list.lock_next().unwrap();
+ assert_eq!(**item, 2);
+ let item = list.lock_next().unwrap();
+ assert_eq!(**item, 3);
+ assert!(list.lock_next().is_none());
+ iter = list.exit();
+ let mut list = iter.recurse_next().unwrap();
+ let item = list.lock_next().unwrap();
+ assert_eq!(**item, 4);
+ let item = list.lock_next().unwrap();
+ assert_eq!(**item, 5);
+ let item = list.lock_next().unwrap();
+ assert_eq!(**item, 6);
+ assert!(list.lock_next().is_none());
+ }
+
+ #[test]
+ fn recurse_last_list_of_lists() {
+ let key = ThreadKey::get().unwrap();
+ let data = [
+ [Mutex::new(1), Mutex::new(2), Mutex::new(3)],
+ [Mutex::new(4), Mutex::new(5), Mutex::new(6)],
+ ];
+ let locks = OwnedLockCollection::new(data);
+ let mut ctx = locks.context();
+ let iter = ctx.iter(key);
+
+ let mut list = iter.recurse_last().unwrap();
+ let item = list.lock_next().unwrap();
+ assert_eq!(**item, 4);
+ let item = list.lock_next().unwrap();
+ assert_eq!(**item, 5);
+ let item = list.lock_next().unwrap();
+ assert_eq!(**item, 6);
+ assert!(list.lock_next().is_none());
+ }
+
+ #[test]
+ fn recurse_last_list_of_empty_list() {
+ let key = ThreadKey::get().unwrap();
+ let data: [[Mutex<i32>; 0]; 0] = [];
+ let locks = OwnedLockCollection::new(data);
+ let mut ctx = locks.context();
+ let iter = ctx.iter(key);
+
+ let list = iter.recurse_last();
+ assert!(list.is_none())
+ }
+
+ #[test]
+ fn recurse_list_of_tuples() {
+ let key = ThreadKey::get().unwrap();
+ let data = [
+ (Mutex::new(true), Mutex::new(1)),
+ (Mutex::new(false), Mutex::new(2)),
+ (Mutex::new(true), Mutex::new(3)),
+ ];
+ let locks = OwnedLockCollection::new(data);
+ let mut ctx = locks.context();
+ let mut iter = ctx.iter(key);
+
+ let tuple = iter.recurse_next_tuple().unwrap();
+ let (should_count, mut tuple) = tuple.lock_0();
+ assert!(**should_count);
+ let num = tuple.lock_mut_1();
+ assert_eq!(**num, 1);
+ drop(num);
+ iter = tuple.exit();
+ let tuple = iter.recurse_next_tuple().unwrap();
+ let (should_count, mut tuple) = tuple.lock_0();
+ assert!(!**should_count);
+ let num = tuple.lock_mut_1();
+ assert_eq!(**num, 2);
+ drop(num);
+ iter = tuple.exit();
+ let tuple = iter.recurse_next_tuple().unwrap();
+ let (should_count, mut tuple) = tuple.lock_0();
+ assert!(**should_count);
+ let num = tuple.lock_mut_1();
+ assert_eq!(**num, 3);
+ drop(num);
+ iter = tuple.exit();
+ assert!(iter.recurse_next_tuple().is_none());
+ }
+
+ #[test]
+ fn recurse_last_list_of_tuples() {
+ let key = ThreadKey::get().unwrap();
+ let data = [
+ (Mutex::new(true), Mutex::new(1)),
+ (Mutex::new(false), Mutex::new(2)),
+ (Mutex::new(true), Mutex::new(3)),
+ ];
+ let locks = OwnedLockCollection::new(data);
+ let mut ctx = locks.context();
+ let iter = ctx.iter(key);
+
+ let tuple = iter.recurse_last_tuple().unwrap();
+ let (should_count, mut tuple) = tuple.lock_0();
+ if **should_count {
+ let num = tuple.lock_mut_1();
+ assert_eq!(**num, 3);
+ } else {
+ panic!();
+ }
+ }
+
+ #[test]
+ fn recurse_last_of_empty_list_of_tuples() {
+ let key = ThreadKey::get().unwrap();
+ let data: [(Mutex<bool>, Mutex<i32>); 0] = [];
+ let locks = OwnedLockCollection::new(data);
+ let mut ctx = locks.context();
+ let iter = ctx.iter(key);
+
+ let tuple = iter.recurse_last_tuple();
+ assert!(tuple.is_none());
+ }
+
+ #[test]
+ fn lock_last_of_list() {
+ let key = ThreadKey::get().unwrap();
+ let data = [Mutex::new(1), Mutex::new(3), Mutex::new(8)];
+ let locks = OwnedLockCollection::new(data);
+ let mut ctx = locks.context();
+ let iter = ctx.iter(key);
+ let last = iter.lock_last().unwrap();
+ assert_eq!(**last, 8);
+ }
+
+ #[test]
+ fn try_lock_next_works() {
+ let key = ThreadKey::get().unwrap();
+ let data = [Mutex::new(1), Mutex::new(3), Mutex::new(8)];
+ let locks = OwnedLockCollection::new(data);
+ let mut ctx = locks.context();
+ let mut iter = ctx.iter(key).peekable();
+
+ let item = iter.try_lock_next().unwrap();
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = locks.context();
+ let mut iter = context.iter(key).peekable();
+ let guard = iter.try_lock_next();
+ assert_eq!(guard.unwrap_err(), TryLockNextError::WouldBlock);
+ });
+ });
+ assert_eq!(**item, 1);
+ let item = iter.try_lock_next().unwrap();
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = locks.context();
+ let mut iter = context.iter(key).peekable();
+ iter.skip_next();
+ let guard = iter.try_lock_next();
+ assert_eq!(guard.unwrap_err(), TryLockNextError::WouldBlock);
+ });
+ });
+ assert_eq!(**item, 3);
+ let item = iter.try_lock_next().unwrap();
+ std::thread::scope(|s| {
+ s.spawn(|| {
+ let key = ThreadKey::get().unwrap();
+ let mut context = locks.context();
+ let mut iter = context.iter(key).peekable();
+ iter.skip_mut(2);
+ let guard = iter.try_lock_next();
+ assert_eq!(guard.unwrap_err(), TryLockNextError::WouldBlock);
+ });
+ });
+ assert_eq!(**item, 8);
+ assert_eq!(
+ iter.try_lock_next().unwrap_err(),
+ TryLockNextError::FinishedIteration
+ );
+ }
+}