Execute a callback when a mutable reference is dropped

Viewed 77

I want to implement a "counter map" that keeps track of how many A(s) are in the map as shown in below code snippet. However, the get_mut method can cause the counter to go out-of-sync by allowing the caller to replace a B with a A using the returned mutable reference.

If I do want to keep this method (get_mut), is it possible to let the mutable reference updates the counter when it is dropped perhaps by using some smart pointer that customize drop()?

I have sketched a solution for this that seems to work for now. But what is the idiomatic way to do this in Rust?

Playgound link for the code example below.

use std::{
    collections::HashMap,
    ops::{Deref, DerefMut},
};

#[derive(Debug, PartialEq)]
enum State {
    A,
    B,
}

struct StateMutRef<'a> {
    was_a: bool,
    state_mut_ref: &'a mut State,
    counter: &'a mut usize,
}

impl<'a> Drop for StateMutRef<'a> {
    fn drop(&mut self) {
        let was_a = self.was_a;
        let is_a = *self.state_mut_ref == State::A;
        match (was_a, is_a) {
            (true, false) => *self.counter -= 1,
            (false, true) => *self.counter += 1,
            (_, _) => (),
        };
    }
}

impl<'a> Deref for StateMutRef<'a> {
    type Target = State;
    fn deref(&self) -> &Self::Target {
        self.state_mut_ref
    }
}

impl<'a> DerefMut for StateMutRef<'a> {
    fn deref_mut(&mut self) -> &mut Self::Target {
        self.state_mut_ref
    }
}

impl<'a> StateMutRef<'a> {
    fn new(state_mut_ref: &'a mut State, counter: &'a mut usize) -> Self {
        Self {
            was_a: *state_mut_ref == State::A,
            state_mut_ref,
            counter,
        }
    }
}

struct CounterMap {
    inner: HashMap<u32, State>,
    // Nnumber of A(s) we have in the map
    num_of_a: usize,
}

impl CounterMap {
    fn new() -> Self {
        Self {
            inner: HashMap::new(),
            num_of_a: 0,
        }
    }

    fn insert(&mut self, id: u32, val: State) {
        let is_a = val == State::A;
        if let Some(old_val) = self.inner.insert(id, val) {
            let was_a = old_val == State::A;
            match (was_a, is_a) {
                (true, false) => self.num_of_a -= 1,
                (false, true) => self.num_of_a += 1,
                (_, _) => (),
            };
        } else {
            if is_a {
                self.num_of_a += 1;
            }
        }
    }

    fn remove(&mut self, id: &u32) {
        if let Some(v) = self.inner.remove(id) {
            if v == State::A {
                self.num_of_a -= 1;
            }
        }
    }

    fn get_mut(&mut self, id: &u32) -> Option<StateMutRef> {
        let counter_mut = &mut self.num_of_a;
        self.inner
            .get_mut(id)
            .map(move |s| StateMutRef::new(s, counter_mut))
    }

    fn num_of_a(&self) -> usize {
        self.num_of_a
    }
}

fn main() {
    let mut map = CounterMap::new();
    assert_eq!(map.num_of_a(), 0);

    map.insert(0, State::B);
    assert_eq!(map.num_of_a(), 0);

    map.insert(0, State::A); // replace B with A
    assert_eq!(map.num_of_a(), 1);

    map.insert(0, State::A); // replace A with A
    assert_eq!(map.num_of_a(), 1);

    map.insert(0, State::B); // replace A with B
    assert_eq!(map.num_of_a(), 0);

    map.insert(1, State::A); // add another A with different id
    assert_eq!(map.num_of_a(), 1);

    map.remove(&0);
    assert_eq!(map.num_of_a(), 1);
    map.remove(&1);
    assert_eq!(map.num_of_a(), 0);

    map.insert(0, State::B);
    assert_eq!(map.num_of_a(), 0);

    if let Some(mut mut_ref) = map.get_mut(&0) {
        *mut_ref = State::A;
    }
    assert_eq!(map.num_of_a(), 1);

    if let Some(mut mut_ref) = map.get_mut(&0) {
        *mut_ref = State::B;
    }
    assert_eq!(map.num_of_a(), 0);
}
0 Answers
Related