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);
}