use std::borrow::Borrow;
use std::cell::Cell;
use std::cmp::Ordering;
use std::fmt;
use std::hash::{Hash, Hasher};
use std::ops::Deref;
use std::ptr::NonNull;
use crate::ptr::takeable::Takeable;
type IsZero = bool;
#[derive(PartialEq, Debug)]
struct Node<T> {
prev: Option<NonNull<Node<T>>>,
value: Takeable<T>,
count: Cell<usize>,
next: Option<NonNull<Node<T>>>,
}
impl<T> Node<T> {
fn new(value: T) -> Self {
Node {
prev: None,
value: Takeable::new(value),
count: Cell::new(1),
next: None,
}
}
fn into_not_null(self) -> NonNull<Self> {
unsafe { NonNull::new_unchecked(Box::into_raw(Box::new(self))) }
}
fn get_count(&self) -> usize {
self.count.get()
}
fn inc_count(&self) {
let mut count = self.count.get();
count += 1;
self.count.set(count);
}
fn dec_count(&self) -> IsZero {
let mut count = self.count.get();
count -= 1;
self.count.set(count);
count == 0
}
}
unsafe fn decrement_and_possibly_deallocate<T>(node: NonNull<Node<T>>) {
if node.as_ref().dec_count() {
if let Some(prev) = (*node.as_ptr()).prev.as_mut() {
prev.as_mut().next = (*node.as_ptr()).next.take();
}
if let Some(next) = (*node.as_ptr()).next.as_mut() {
next.as_mut().prev = (*node.as_ptr()).prev.take();
}
std::ptr::drop_in_place(node.as_ptr());
}
}
pub struct Lrc<T> {
head: Option<NonNull<Node<T>>>,
}
#[allow(clippy::len_without_is_empty)] impl<T> Lrc<T> {
pub fn new(value: T) -> Self {
let node = Node::new(value);
Lrc {
head: Some(node.into_not_null()),
}
}
pub fn set(&mut self, value: T) {
if self.is_exclusive() {
*self.get_mut_head_node().value.as_mut() = value;
} else {
self.push_head(Node::new(value));
}
}
pub fn get_mut(&mut self) -> Option<&mut T> {
if self.is_exclusive() {
let node = self.get_mut_head_node();
Some(node.value.as_mut())
} else {
None
}
}
pub fn try_unwrap(self) -> Result<T, Self> {
if self.is_exclusive() {
let head: NonNull<Node<T>> = self.head.unwrap();
unsafe {
let value = (*head.as_ptr()).value.take();
if let Some(prev) = (*head.as_ptr()).prev.as_mut() {
prev.as_mut().next = (*head.as_ptr()).next.take();
}
if let Some(next) = (*head.as_ptr()).next.as_mut() {
next.as_mut().prev = (*head.as_ptr()).prev.take();
}
std::ptr::drop_in_place(head.as_ptr());
Ok(value)
}
} else {
Err(self)
}
}
pub fn has_prev(&self) -> bool {
self.get_ref_head_node().prev.is_some()
}
pub fn has_next(&self) -> bool {
self.get_ref_head_node().next.is_some()
}
pub fn update(&mut self) -> bool {
let did_update = self.has_prev();
while let Some(prev) = self.next_back() {
*self = prev;
}
did_update
}
pub fn advance_next(&mut self) -> bool {
unsafe {
let head_node: &mut NonNull<Node<T>> = self.head.as_mut().unwrap();
let next: Option<NonNull<Node<T>>> = (*head_node.as_ptr()).next;
if let Some(next) = next {
decrement_and_possibly_deallocate(*head_node);
next.as_ref().inc_count();
self.head = Some(next);
true
} else {
false
}
}
}
pub fn advance_back(&mut self) -> bool {
unsafe {
let head_node: &mut NonNull<Node<T>> = self.head.as_mut().unwrap();
let prev: Option<NonNull<Node<T>>> = (*head_node.as_ptr()).prev;
if let Some(prev) = prev {
decrement_and_possibly_deallocate(*head_node);
prev.as_ref().inc_count();
self.head = Some(prev);
true
} else {
false
}
}
}
pub fn ptr_eq(lhs: &Self, rhs: &Self) -> bool {
lhs.head.unwrap().eq(&rhs.head.unwrap())
}
fn push_head(&mut self, mut node: Node<T>) {
self.update();
node.next = self.head;
let node = Some(node.into_not_null());
let head = self.head.unwrap();
unsafe {
(*head.as_ptr()).prev = node;
decrement_and_possibly_deallocate(head)
}
self.head = node;
}
pub fn get_count(&self) -> usize {
self.get_ref_head_node().get_count()
}
pub fn is_exclusive(&self) -> bool {
self.get_count() == 1
}
pub fn len(&self) -> usize {
1 + self.next_len() + self.prev_len()
}
pub fn next_len(&self) -> usize {
let mut count = 0;
unsafe {
let mut node = self.get_ref_head_node();
while let Some(next_node) = node.next.as_ref() {
count += 1;
node = next_node.as_ref()
}
}
count
}
pub fn prev_len(&self) -> usize {
let mut count = 0;
unsafe {
let mut node = self.get_ref_head_node();
while let Some(prev_node) = node.prev.as_ref() {
count += 1;
node = prev_node.as_ref()
}
}
count
}
fn get_mut_head_node(&mut self) -> &mut Node<T> {
unsafe { self.head.as_mut().unwrap().as_mut() }
}
fn get_ref_head_node(&self) -> &Node<T> {
unsafe { self.head.as_ref().unwrap().as_ref() }
}
}
impl<T: Clone> Lrc<T> {
pub fn make_mut(&mut self) -> &mut T {
if !self.is_exclusive() {
let cloned_value: T = self.clone_inner();
self.push_head(Node::new(cloned_value))
}
self.get_mut_head_node().value.as_mut()
}
pub fn clone_unwrap(self) -> T {
if self.is_exclusive() {
let head: NonNull<Node<T>> = self.head.unwrap();
unsafe {
let value = (*head.as_ptr()).value.take();
if let Some(prev) = (*head.as_ptr()).prev.as_mut() {
prev.as_mut().next = (*head.as_ptr()).next.take();
}
if let Some(next) = (*head.as_ptr()).next.as_mut() {
next.as_mut().prev = (*head.as_ptr()).prev.take();
}
std::ptr::drop_in_place(head.as_ptr());
value
}
} else {
self.clone_inner()
}
}
pub fn clone_inner(&self) -> T {
self.get_ref_head_node().value.as_ref().clone()
}
}
impl<T: PartialEq> Lrc<T> {
pub fn neq_set(&mut self, value: T) -> bool {
if self.get_ref_head_node().value.as_ref() != &value {
self.set(value);
true
} else {
false
}
}
}
impl<T> Drop for Lrc<T> {
fn drop(&mut self) {
let head = self.head.expect("Head should always be present.");
unsafe {
decrement_and_possibly_deallocate(head);
}
}
}
impl<T> Clone for Lrc<T> {
fn clone(&self) -> Self {
if let Some(head) = self.head {
unsafe {
head.as_ref().inc_count();
}
}
Lrc { head: self.head }
}
}
impl<T: PartialEq> PartialEq for Lrc<T> {
fn eq(&self, other: &Self) -> bool {
unsafe {
match (self.head, other.head) {
(Some(lhs), Some(rhs)) => lhs.as_ref().value.eq(&rhs.as_ref().value),
_ => false,
}
}
}
}
impl<T: Eq> Eq for Lrc<T> {}
impl<T: PartialOrd> PartialOrd for Lrc<T> {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
self.get_ref_head_node()
.value
.partial_cmp(&other.get_ref_head_node().value)
}
}
impl<T: Ord> Ord for Lrc<T> {
fn cmp(&self, other: &Self) -> Ordering {
self.get_ref_head_node()
.value
.cmp(&other.get_ref_head_node().value)
}
}
impl<T: Hash> Hash for Lrc<T> {
fn hash<H: Hasher>(&self, state: &mut H) {
self.get_ref_head_node().value.hash(state)
}
}
impl<T> AsRef<T> for Lrc<T> {
fn as_ref(&self) -> &T {
&self.get_ref_head_node().value.as_ref()
}
}
impl<T> Deref for Lrc<T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.get_ref_head_node().value.as_ref()
}
}
impl<T> Borrow<T> for Lrc<T> {
fn borrow(&self) -> &T {
&self.get_ref_head_node().value.as_ref()
}
}
impl<T: fmt::Debug> fmt::Debug for Lrc<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_tuple("Lrc").field(&self.head).finish()
}
}
impl<T> Iterator for Lrc<T> {
type Item = Lrc<T>;
fn next(&mut self) -> Option<Self::Item> {
self.get_ref_head_node().next.map(|ptr| {
unsafe {
ptr.as_ref().inc_count();
}
Lrc { head: Some(ptr) }
})
}
}
impl<T> DoubleEndedIterator for Lrc<T> {
fn next_back(&mut self) -> Option<Self::Item> {
self.get_ref_head_node().prev.map(|ptr| {
unsafe {
ptr.as_ref().inc_count();
}
Lrc { head: Some(ptr) }
})
}
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn lrc_new() {
let lrc = Lrc::new(25);
assert_eq!(
lrc,
Lrc {
head: Some(Node::new(25).into_not_null())
}
);
assert_eq!(lrc.as_ref(), &25)
}
#[allow(clippy::redundant_clone)]
#[test]
fn clone_lrc() {
let lrc = Lrc::new(25);
let copy = lrc.clone();
assert_eq!(copy.as_ref(), &25)
}
#[test]
fn set_lrc() {
let mut lrc = Lrc::new(25);
lrc.set(30);
assert_eq!(lrc.as_ref(), &30)
}
#[test]
fn len_not_changed_by_setting_exclusive_lrc() {
let mut lrc = Lrc::new(25);
lrc.set(24);
assert_eq!(lrc.len(), 1);
}
#[test]
fn make_mut_will_clone_if_shared() {
let mut lrc = Lrc::new(0);
let _shared = lrc.clone();
lrc.make_mut();
assert_eq!(lrc.len(), 2);
}
#[test]
fn exclusive_set_equivalent_to_exclusive_make_mut() {
let mut lrc = Lrc::new(0);
lrc.set(1);
assert_eq!(lrc.as_ref(), &1);
assert_eq!(lrc.len(), 1);
assert_eq!(lrc.get_count(), 1);
let mut lrc = Lrc::new(0);
*lrc.make_mut() = 1;
assert_eq!(lrc.as_ref(), &1);
assert_eq!(lrc.len(), 1);
assert_eq!(lrc.get_count(), 1);
}
#[test]
fn droping_middle_connects_prev_and_next() {
let mut lrc = Lrc::new(0);
assert_eq!(
lrc.get_ref_head_node().count,
Cell::new(1),
"exclusive ownership"
);
let _og_clone = lrc.clone();
assert_eq!(
lrc.get_ref_head_node().count,
Cell::new(2),
"shared ownership"
);
lrc.set(1);
assert_eq!(lrc.get_ref_head_node().prev, None);
assert_eq!(lrc.get_ref_head_node().value.as_ref(), &1);
assert_eq!(lrc.get_ref_head_node().count, Cell::new(1));
assert!(
lrc.get_ref_head_node().next.is_some(),
"Should have pointer to previous head"
);
unsafe {
let lrcs_next = lrc
.get_ref_head_node()
.next
.as_ref()
.expect("Should have next node")
.as_ref();
let lrcs_nexts_prev = lrcs_next
.prev
.as_ref()
.expect("next.prev should be some")
.as_ref();
assert_eq!(lrcs_next.value.as_ref(), &0);
assert_eq!(
lrcs_next.count,
Cell::new(1),
"Should still be owned by the Og Clone"
);
assert!(lrcs_next.prev.is_some(), "Should point to head");
assert_eq!(
lrcs_nexts_prev,
lrc.get_ref_head_node(),
"the head's next ptr's prev ptr should point back to the head"
);
}
let cloned_lrc = lrc.clone();
assert_eq!(lrc.len(), 2);
assert_eq!(cloned_lrc.get_ref_head_node().prev, None);
assert_eq!(cloned_lrc.get_ref_head_node().value.as_ref(), &1);
assert_eq!(cloned_lrc.get_ref_head_node().count, Cell::new(2));
assert!(
cloned_lrc.get_ref_head_node().next.is_some(),
"Should have pointer to previous head"
);
lrc.set(2);
assert_eq!(lrc.get_ref_head_node().prev, None);
assert_eq!(
lrc.get_ref_head_node().value.as_ref(),
&2,
"value should now be updated to 2"
);
assert_eq!(
lrc.get_ref_head_node().count,
Cell::new(1),
"there should only be one owner of this node"
);
assert!(
lrc.get_ref_head_node().next.is_some(),
"Should have pointer to previous head"
);
unsafe {
let cloned_lrcs_heads_prev_value = cloned_lrc
.get_ref_head_node()
.prev
.as_ref()
.expect("Should point to head")
.as_ref();
assert_eq!(cloned_lrcs_heads_prev_value, lrc.get_ref_head_node());
}
assert_eq!(lrc.len(), 3);
std::mem::drop(cloned_lrc);
assert_eq!(lrc.len(), 2);
unsafe {
let lrcs_next = lrc
.get_ref_head_node()
.next
.as_ref()
.expect("Should have next node")
.as_ref();
assert_eq!(lrcs_next.value.as_ref(), &0);
}
}
#[test]
fn single_node_older_yeilds_none() {
let mut lrc = Lrc::new(25);
let older = lrc.next();
assert_eq!(older, None)
}
#[test]
fn single_node_newer_yeilds_none() {
let mut lrc = Lrc::new(25);
let newer = lrc.next_back();
assert_eq!(newer, None)
}
#[test]
fn older_traverses_to_previous_lrc() {
let mut lrc = Lrc::new(25);
let _clone = lrc.clone();
lrc.set(26);
let older = lrc.next();
assert_eq!(older, Some(Lrc::new(25)))
}
#[test]
fn newer_traverses_back_to_original_head_lrc() {
let mut lrc = Lrc::new(25);
let _clone = lrc.clone();
lrc.set(26);
let older = lrc.next();
assert_eq!(older, Some(Lrc::new(25)));
let newer = older.unwrap().next_back();
assert_eq!(newer, Some(lrc));
}
#[test]
fn attempt_to_dangle_ref() {
let lrc = Lrc::new(vec![25]);
let mut cloned_lrc = lrc.clone();
let first_item_ref = &lrc.as_ref()[0];
cloned_lrc.set(vec![22, 23]);
assert_eq!(first_item_ref, &25)
}
#[test]
fn ptr_eq_positive() {
let lrc = Lrc::new(24);
let cloned_lrc = lrc.clone();
assert!(Lrc::ptr_eq(&lrc, &cloned_lrc));
}
#[test]
fn ptr_eq_negative() {
let lrc = Lrc::new(24);
let other_lrc = Lrc::new(24);
assert!(!Lrc::ptr_eq(&lrc, &other_lrc));
}
#[test]
fn update_sets_lrc_to_have_newest_value() {
let mut lrc = Lrc::new(0);
let mut cloned_lrc = lrc.clone();
cloned_lrc.set(1);
assert_eq!(cloned_lrc.as_ref(), &1);
assert_eq!(lrc.as_ref(), &0);
let did_update = lrc.update();
assert!(did_update);
assert_eq!(lrc.as_ref(), &1);
}
#[test]
fn advance_back() {
let mut lrc = Lrc::new(0);
let mut clone = lrc.clone();
lrc.set(1);
let did_advance = clone.advance_back();
assert!(did_advance);
assert_eq!(clone.as_ref(), &1);
assert_eq!(clone.len(), 1);
assert_eq!(clone.get_count(), 2);
let did_advance = clone.advance_back();
assert!(!did_advance, "No newer values to advance to.");
let did_advance = clone.advance_next();
assert!(
!did_advance,
"can't restore old value, as it has be dropped."
);
}
#[test]
fn advance_next() {
let mut lrc = Lrc::new(0);
let mut clone = lrc.clone();
lrc.set(1);
let did_advance = lrc.advance_next();
assert!(did_advance);
assert_eq!(lrc.as_ref(), &0);
assert_eq!(lrc.len(), 1);
assert_eq!(lrc.get_count(), 2);
let did_advance = clone.advance_next();
assert!(!did_advance, "No older values to advance to.");
let did_advance = clone.advance_back();
assert!(
!did_advance,
"Can't restore old value, as it has be dropped."
);
}
#[test]
fn size_of_node_overhead() {
let lrc = Lrc::new(());
let node_size_overhead_bytes = std::mem::size_of_val(lrc.get_ref_head_node());
let usize_size = std::mem::size_of::<usize>();
assert_eq!(node_size_overhead_bytes, usize_size * 4);
}
#[test]
fn size_of_node_for_not_null() {
let lrc = Lrc::new(Box::new(0));
let node_size_overhead_bytes = std::mem::size_of_val(lrc.get_ref_head_node());
let usize_size = std::mem::size_of::<usize>();
assert_eq!(node_size_overhead_bytes, usize_size * 4);
}
#[test]
fn size_of_node_for_usize() {
let lrc = Lrc::new(0usize);
let node_size_overhead_bytes = std::mem::size_of_val(lrc.get_ref_head_node());
let usize_size = std::mem::size_of::<usize>();
assert_eq!(node_size_overhead_bytes, usize_size * 5);
}
}