Skip to main content

xmtp_id/associations/
state.rs

1//! [`AssociationState`] describes a single point in time for an Inbox where it contains a set of
2//! associated [`MemberIdentifier`]'s, which may be one of [`MemberKind::Address`]
3//! or[`MemberKind::Installation`]. A diff between two states can be calculated to determine
4//! a change of membership between two periods of time. [XIP-46](https://github.com/xmtp/XIPs/pull/53)
5
6use 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    /// Converts the [`AssociationState`] to a diff that represents all members
229    /// of the inbox at the current state.
230    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}