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#[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}