Skip to main content

xmtp_db/encrypted_store/
association_state.rs

1use diesel::prelude::*;
2
3use super::schema::association_state::{self, dsl};
4use crate::ConnectionExt;
5use crate::DbConnection;
6use crate::{Fetch, StorageError, StoreOrIgnore, impl_fetch, impl_store_or_ignore};
7use prost::Message;
8use xmtp_proto::xmtp::identity::associations::AssociationState as AssociationStateProto;
9
10/// StoredIdentityUpdate holds a serialized IdentityUpdate record
11#[derive(Insertable, Identifiable, Queryable, Debug, Clone, PartialEq, Eq)]
12#[diesel(table_name = association_state)]
13#[diesel(primary_key(inbox_id, sequence_id))]
14pub struct StoredAssociationState {
15    pub inbox_id: String,
16    pub sequence_id: i64,
17    pub state: Vec<u8>,
18}
19impl_fetch!(StoredAssociationState, association_state, (String, i64));
20impl_store_or_ignore!(StoredAssociationState, association_state);
21
22pub trait QueryAssociationStateCache {
23    fn write_to_cache(
24        &self,
25        inbox_id: String,
26        sequence_id: i64,
27        state: AssociationStateProto,
28    ) -> Result<(), StorageError>;
29
30    fn read_from_cache<A: AsRef<str>>(
31        &self,
32        inbox_id: A,
33        sequence_id: i64,
34    ) -> Result<Option<AssociationStateProto>, StorageError>;
35
36    fn batch_read_from_cache(
37        &self,
38        identifiers: Vec<(String, i64)>,
39    ) -> Result<Vec<AssociationStateProto>, StorageError>;
40}
41
42impl<R> QueryAssociationStateCache for &R
43where
44    R: QueryAssociationStateCache,
45{
46    fn write_to_cache(
47        &self,
48        inbox_id: String,
49        sequence_id: i64,
50        state: AssociationStateProto,
51    ) -> Result<(), StorageError> {
52        (**self).write_to_cache(inbox_id, sequence_id, state)
53    }
54
55    fn read_from_cache<A: AsRef<str>>(
56        &self,
57        inbox_id: A,
58        sequence_id: i64,
59    ) -> Result<Option<AssociationStateProto>, StorageError> {
60        (**self).read_from_cache(inbox_id, sequence_id)
61    }
62
63    fn batch_read_from_cache(
64        &self,
65        identifiers: Vec<(String, i64)>,
66    ) -> Result<Vec<AssociationStateProto>, StorageError> {
67        (**self).batch_read_from_cache(identifiers)
68    }
69}
70
71impl<C: ConnectionExt> QueryAssociationStateCache for DbConnection<C> {
72    fn write_to_cache(
73        &self,
74        inbox_id: String,
75        sequence_id: i64,
76        state: AssociationStateProto,
77    ) -> Result<(), StorageError> {
78        let result = StoredAssociationState {
79            inbox_id: inbox_id.clone(),
80            sequence_id,
81            state: state.encode_to_vec(),
82        }
83        .store_or_ignore(self);
84
85        if result.is_ok() {
86            tracing::debug!(
87                "Wrote association state to cache: {} {}",
88                inbox_id,
89                sequence_id
90            );
91        }
92
93        result
94    }
95
96    fn read_from_cache<A: AsRef<str>>(
97        &self,
98        inbox_id: A,
99        sequence_id: i64,
100    ) -> Result<Option<AssociationStateProto>, StorageError> {
101        let inbox_id = inbox_id.as_ref();
102        let stored_state: Option<StoredAssociationState> =
103            self.fetch(&(inbox_id.to_string(), sequence_id))?;
104
105        let result = stored_state
106            .map(|stored_state| stored_state.state)
107            .inspect(|_| {
108                tracing::debug!(
109                    "Loaded association state from cache: {} {}",
110                    inbox_id,
111                    sequence_id
112                )
113            });
114        Ok(result
115            .map(|r| AssociationStateProto::decode(r.as_slice()))
116            .transpose()?)
117    }
118
119    fn batch_read_from_cache(
120        &self,
121        identifiers: Vec<(String, i64)>,
122    ) -> Result<Vec<AssociationStateProto>, StorageError> {
123        if identifiers.is_empty() {
124            return Ok(vec![]);
125        }
126
127        let mut query = dsl::association_state
128            .select((dsl::inbox_id, dsl::sequence_id, dsl::state))
129            .into_boxed();
130
131        for (inbox_id, sequence_id) in &identifiers {
132            let predicate = dsl::inbox_id
133                .eq(inbox_id.clone())
134                .and(dsl::sequence_id.eq(*sequence_id));
135            query = query.or_filter(predicate);
136        }
137
138        let association_states =
139            self.raw_query(|query_conn| query.load::<StoredAssociationState>(query_conn))?;
140
141        association_states
142            .into_iter()
143            .map(|stored_association_state| {
144                Ok(AssociationStateProto::decode(
145                    stored_association_state.state.as_slice(),
146                )?)
147            })
148            .collect::<Result<Vec<_>, _>>()
149    }
150}
151
152#[cfg(test)]
153pub(crate) mod tests {
154    use super::*;
155    use crate::test_utils::with_connection;
156    use serde::{Deserialize, Serialize};
157    use xmtp_proto::xmtp::identity::associations::AssociationState as AssociationStateProto;
158
159    #[derive(Serialize, Deserialize)]
160    pub struct MockState {
161        inbox_id: String,
162    }
163    impl From<StoredAssociationState> for MockState {
164        fn from(v: StoredAssociationState) -> MockState {
165            crate::db_deserialize(&v.state).unwrap()
166        }
167    }
168    impl From<AssociationStateProto> for MockState {
169        fn from(v: AssociationStateProto) -> Self {
170            MockState {
171                inbox_id: v.inbox_id,
172            }
173        }
174    }
175
176    #[xmtp_common::test]
177    fn test_batch_read() {
178        with_connection(|conn| {
179            let mock = AssociationStateProto {
180                inbox_id: "test_id1".into(),
181                members: vec![],
182                ..Default::default()
183            };
184            conn.write_to_cache(mock.inbox_id.clone(), 1, mock.clone())
185                .unwrap();
186            let mock_2 = AssociationStateProto {
187                inbox_id: "test_id2".into(),
188                members: vec![],
189                ..Default::default()
190            };
191
192            conn.write_to_cache(mock_2.inbox_id.clone(), 2, mock_2.clone())
193                .unwrap();
194
195            let first_association_state: Vec<MockState> = conn
196                .batch_read_from_cache(vec![(mock.inbox_id.to_string(), 1)])
197                .unwrap()
198                .into_iter()
199                .map(Into::into)
200                .collect();
201            assert_eq!(first_association_state.len(), 1);
202            assert_eq!(&first_association_state[0].inbox_id, &mock.inbox_id);
203
204            let both_association_states: Vec<MockState> = conn
205                .batch_read_from_cache(vec![
206                    (mock.inbox_id.clone(), 1),
207                    (mock_2.inbox_id.clone(), 2),
208                ])
209                .unwrap()
210                .into_iter()
211                .map(Into::into)
212                .collect();
213
214            assert_eq!(both_association_states.len(), 2);
215
216            let no_results = conn
217                .batch_read_from_cache(vec![(mock.inbox_id.clone(), 2)])
218                .unwrap()
219                .into_iter()
220                .map(Into::into)
221                .collect::<Vec<MockState>>();
222            assert_eq!(no_results.len(), 0);
223        })
224    }
225}