1use super::{
7 AssociationError, MemberIdentifier, MemberKind, ident,
8 member::{Identifier, Member},
9};
10use crate::InboxIdRef;
11use std::{
12 collections::{HashMap, HashSet},
13 fmt::{Debug, Write},
14};
15
16#[derive(Debug, Clone)]
17pub struct AssociationStateDiff {
18 pub new_members: Vec<MemberIdentifier>,
19 pub removed_members: Vec<MemberIdentifier>,
20}
21
22#[derive(Debug)]
23pub struct Installation {
24 pub id: Vec<u8>,
25 pub client_timestamp_ns: Option<u64>,
26}
27
28impl AssociationStateDiff {
29 pub fn new_installations(&self) -> Vec<Vec<u8>> {
30 self.new_members
31 .iter()
32 .filter_map(|member| match member {
33 MemberIdentifier::Installation(ident::Installation(key)) => Some(key.clone()),
34 _ => None,
35 })
36 .collect()
37 }
38
39 pub fn removed_installations(&self) -> Vec<Vec<u8>> {
40 self.removed_members
41 .iter()
42 .filter_map(|member| match member {
43 MemberIdentifier::Installation(ident::Installation(key)) => Some(key.clone()),
44 _ => None,
45 })
46 .collect()
47 }
48}
49
50#[derive(Clone)]
51pub struct AssociationState {
52 pub(crate) inbox_id: String,
53 pub(crate) members: HashMap<MemberIdentifier, Member>,
54 pub(crate) recovery_identifier: Identifier,
55 pub(crate) seen_signatures: HashSet<Vec<u8>>,
56}
57
58impl std::fmt::Debug for AssociationState {
59 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
60 let mut members = String::new();
61 for member in self.members.keys() {
62 write!(members, "{:?}", member)?;
63 write!(members, ",")?;
64 }
65
66 let mut signatures = String::new();
67 for signature in self.seen_signatures.iter() {
68 write!(
69 signatures,
70 "{}",
71 xmtp_common::fmt::truncate_hex(hex::encode(signature))
72 )?;
73 write!(signatures, ",")?;
74 }
75
76 write!(
77 f,
78 "AssociationState {{ inbox_id: {}, members: {}, recovery: {}, seen_signatures: {} }}",
79 self.inbox_id, members, self.recovery_identifier, signatures
80 )
81 }
82}
83
84impl TryFrom<MemberIdentifier> for Identifier {
85 type Error = AssociationError;
86 fn try_from(ident: MemberIdentifier) -> Result<Self, Self::Error> {
87 let ident = match ident {
88 MemberIdentifier::Ethereum(eth) => Self::Ethereum(eth),
89 MemberIdentifier::Passkey(passkey) => Self::Passkey(passkey),
90 MemberIdentifier::Installation(_) => {
91 return Err(AssociationError::NotIdentifier(
92 "Installation Keys".to_string(),
93 ));
94 }
95 };
96 Ok(ident)
97 }
98}
99
100impl AssociationState {
101 pub fn add(&self, member: Member) -> Self {
102 let mut new_state = self.clone();
103 let _ = new_state.members.insert(member.identifier.clone(), member);
104
105 new_state
106 }
107
108 pub fn remove(&self, identifier: &MemberIdentifier) -> Self {
109 let mut new_state = self.clone();
110 let _ = new_state.members.remove(identifier);
111
112 new_state
113 }
114
115 pub fn set_recovery_identifier(&self, recovery_identifier: Identifier) -> Self {
116 let mut new_state = self.clone();
117 new_state.recovery_identifier = recovery_identifier;
118
119 new_state
120 }
121
122 pub fn get(&self, identifier: &MemberIdentifier) -> Option<&Member> {
123 self.members.get(identifier)
124 }
125
126 pub fn add_seen_signatures(&self, signatures: Vec<Vec<u8>>) -> Self {
127 let mut new_state = self.clone();
128 new_state.seen_signatures.extend(signatures);
129
130 new_state
131 }
132
133 pub fn has_seen(&self, signature: &Vec<u8>) -> bool {
134 self.seen_signatures.contains(signature)
135 }
136
137 pub fn members(&self) -> Vec<Member> {
138 let mut sorted_members: Vec<_> = self.members.values().cloned().collect();
139 sorted_members.sort_by_key(|m| m.client_timestamp_ns.unwrap_or(u64::MAX));
140 sorted_members
141 }
142
143 pub fn inbox_id(&self) -> InboxIdRef<'_> {
144 &self.inbox_id
145 }
146
147 pub fn recovery_identifier(&self) -> &Identifier {
148 &self.recovery_identifier
149 }
150
151 pub fn members_by_parent(&self, parent_id: &MemberIdentifier) -> Vec<Member> {
152 self.members
153 .values()
154 .filter(|e| e.added_by_entity.eq(&Some(parent_id.clone())))
155 .cloned()
156 .collect()
157 }
158
159 pub fn members_by_kind(&self, kind: MemberKind) -> Vec<Member> {
160 self.members
161 .values()
162 .filter(|e| e.kind() == kind)
163 .cloned()
164 .collect()
165 }
166
167 pub fn identifiers(&self) -> Vec<Identifier> {
168 let mut address_members: Vec<_> = self.members.values().cloned().collect();
169
170 address_members.sort_by_key(|m| m.client_timestamp_ns.unwrap_or(u64::MAX));
171
172 address_members
173 .into_iter()
174 .filter_map(|member| match member.identifier {
175 MemberIdentifier::Ethereum(eth) => Some(Identifier::Ethereum(eth)),
176 MemberIdentifier::Passkey(pk) => Some(Identifier::Passkey(pk)),
177 _ => None,
178 })
179 .collect()
180 }
181
182 pub fn installation_ids(&self) -> Vec<Vec<u8>> {
183 self.members_by_kind(MemberKind::Installation)
184 .into_iter()
185 .filter_map(|member| match member.identifier {
186 MemberIdentifier::Installation(ident::Installation(key)) => Some(key),
187 _ => None,
188 })
189 .collect()
190 }
191
192 pub fn installations(&self) -> Vec<Installation> {
193 self.members()
194 .into_iter()
195 .filter_map(|member| match member.identifier {
196 MemberIdentifier::Installation(ident::Installation(id)) => Some(Installation {
197 id,
198 client_timestamp_ns: member.client_timestamp_ns,
199 }),
200 _ => None,
201 })
202 .collect()
203 }
204
205 pub fn diff(&self, new_state: &Self) -> AssociationStateDiff {
206 let new_members: Vec<MemberIdentifier> = new_state
207 .members
208 .keys()
209 .filter(|new_member_identifier| !self.members.contains_key(new_member_identifier))
210 .cloned()
211 .collect();
212
213 let removed_members: Vec<MemberIdentifier> = self
214 .members
215 .keys()
216 .filter(|existing_member_identifier| {
217 !new_state.members.contains_key(existing_member_identifier)
218 })
219 .cloned()
220 .collect();
221
222 AssociationStateDiff {
223 new_members,
224 removed_members,
225 }
226 }
227
228 pub fn as_diff(&self) -> AssociationStateDiff {
231 AssociationStateDiff {
232 new_members: self.members.keys().cloned().collect(),
233 removed_members: vec![],
234 }
235 }
236
237 pub fn new(
238 account_identifier: Identifier,
239 nonce: u64,
240 chain_id: Option<u64>,
241 ) -> Result<Self, AssociationError> {
242 let member_identifier: MemberIdentifier = account_identifier.clone().into();
243
244 let inbox_id = account_identifier.inbox_id(nonce)?;
245 let new_member = Member::new(member_identifier.clone(), None, None, chain_id);
246 Ok(Self {
247 members: HashMap::from_iter([(member_identifier, new_member)]),
248 seen_signatures: HashSet::new(),
249 recovery_identifier: account_identifier,
250 inbox_id,
251 })
252 }
253}
254
255#[cfg(test)]
256pub(crate) mod tests {
257 use super::*;
258
259 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
260 #[cfg_attr(not(target_arch = "wasm32"), test)]
261 fn can_add_remove() {
262 let starting_state = AssociationState::new(Identifier::rand_ethereum(), 0, None).unwrap();
263 let new_entity = Member::default();
264 let with_add = starting_state.add(new_entity.clone());
265 assert!(with_add.get(&new_entity.identifier).is_some());
266 assert!(starting_state.get(&new_entity.identifier).is_none());
267 }
268
269 #[cfg_attr(target_arch = "wasm32", wasm_bindgen_test::wasm_bindgen_test)]
270 #[cfg_attr(not(target_arch = "wasm32"), test)]
271 fn can_diff() {
272 let starting_state = AssociationState::new(Identifier::rand_ethereum(), 0, None).unwrap();
273 let entity_1 = Member::default();
274 let entity_2 = Member::default();
275 let entity_3 = Member::default();
276
277 let state_1 = starting_state.add(entity_1.clone()).add(entity_2.clone());
278 let state_2 = state_1.remove(&entity_1.identifier).add(entity_3.clone());
279
280 let diff = state_1.diff(&state_2);
281
282 assert_eq!(diff.new_members, vec![entity_3.identifier]);
283 assert_eq!(diff.removed_members, vec![entity_1.identifier]);
284 }
285}