1use crate::{GspSignal, args::validate_args};
4use gbp::CodecError;
5use gbp_core::{BoundedSeen, GbpFlags, MemberId, PayloadCodec, SignalType, StreamType};
6use gbp_node::{GroupNode, NodeError, OutboundFrame, Sealer};
7use std::collections::HashSet;
8
9#[derive(Debug, thiserror::Error)]
11pub enum GspError {
12 #[error("decode: {0}")]
14 Decode(#[from] CodecError),
15 #[error("unknown signal_type: {0}")]
17 UnknownSignal(u32),
18 #[error("duplicate request_id: {0}")]
20 DuplicateRequest(u32),
21 #[error("bad args schema: {0}")]
23 BadSchema(&'static str),
24 #[error("node: {0}")]
26 Node(#[from] NodeError),
27}
28
29#[derive(Debug, Clone)]
32pub struct GspAccept {
33 pub signal: SignalType,
35 pub sender_id: MemberId,
37 pub role_claim: u32,
39 pub request_id: u32,
41}
42
43const GSP_SEEN_CAP: usize = 10_000;
45
46pub struct GspClient {
58 seen_requests: BoundedSeen<u32>,
59 pub muted: HashSet<MemberId>,
61 pub members: HashSet<MemberId>,
63 current_epoch: Option<u64>,
64}
65
66impl Default for GspClient {
67 fn default() -> Self {
68 Self::new()
69 }
70}
71
72impl GspClient {
73 pub fn new() -> Self {
75 Self {
76 seen_requests: BoundedSeen::new(GSP_SEEN_CAP),
77 muted: HashSet::new(),
78 members: HashSet::new(),
79 current_epoch: None,
80 }
81 }
82
83 #[allow(clippy::too_many_arguments)]
89 pub fn send<S: Sealer>(
90 &mut self,
91 node: &mut GroupNode,
92 seal: &mut S,
93 target: MemberId,
94 signal: SignalType,
95 role_claim: u32,
96 request_id: u32,
97 codec: PayloadCodec,
98 ) -> Result<OutboundFrame, GspError> {
99 self.send_with_args(
100 node,
101 seal,
102 target,
103 signal,
104 role_claim,
105 request_id,
106 &[],
107 codec,
108 )
109 }
110
111 #[allow(clippy::too_many_arguments)]
119 pub fn send_with_args<S: Sealer>(
120 &mut self,
121 node: &mut GroupNode,
122 seal: &mut S,
123 target: MemberId,
124 signal: SignalType,
125 role_claim: u32,
126 request_id: u32,
127 args: &[u8],
128 codec: PayloadCodec,
129 ) -> Result<OutboundFrame, GspError> {
130 self.sync_epoch(node.current_epoch);
131 let mut sig = GspSignal::bare(signal as u32, request_id, node.member_id);
132 sig.role_claim = role_claim;
133 sig.args = serde_bytes::ByteBuf::from(args.to_vec());
134 sig.args_length = args.len() as u32;
135 let stream_id = node.member_stream_id(3);
136 Ok(node.send_payload(
137 seal,
138 target,
139 StreamType::Signal,
140 stream_id,
141 GbpFlags::ordered_reliable_ack(),
142 &sig.to_bytes(codec),
143 codec,
144 )?)
145 }
146
147 pub fn accept(
154 &mut self,
155 plaintext: &[u8],
156 current_epoch: u64,
157 codec: PayloadCodec,
158 ) -> Result<GspAccept, GspError> {
159 self.sync_epoch(current_epoch);
160 let s = GspSignal::from_bytes(plaintext, codec)?;
161 let signal = SignalType::try_from(s.signal_type).map_err(GspError::UnknownSignal)?;
162 validate_args(signal, &s.args).map_err(GspError::BadSchema)?;
164 if !self.seen_requests.insert(s.request_id) {
165 return Err(GspError::DuplicateRequest(s.request_id));
166 }
167 match signal {
168 SignalType::Join => {
169 self.members.insert(s.sender_id);
170 }
171 SignalType::Leave => {
172 self.members.remove(&s.sender_id);
173 self.muted.remove(&s.sender_id);
174 }
175 SignalType::Mute => {
176 self.muted.insert(s.sender_id);
177 }
178 SignalType::Unmute => {
179 self.muted.remove(&s.sender_id);
180 }
181 _ => {}
182 }
183 Ok(GspAccept {
184 signal,
185 sender_id: s.sender_id,
186 role_claim: s.role_claim,
187 request_id: s.request_id,
188 })
189 }
190
191 pub fn sync_epoch(&mut self, epoch: u64) {
195 if Some(epoch) != self.current_epoch {
196 self.seen_requests.clear();
197 self.current_epoch = Some(epoch);
198 }
199 }
200
201 pub fn reset(&mut self) {
203 self.seen_requests.clear();
204 self.current_epoch = None;
205 }
206}
207
208#[cfg(test)]
209mod tests {
210 use super::*;
211 use crate::GspSignal;
212
213 fn encode_bare(signal: SignalType, request_id: u32, sender_id: u32) -> Vec<u8> {
214 GspSignal::bare(signal as u32, request_id, sender_id).to_cbor()
215 }
216
217 #[test]
218 fn join_adds_sender_to_members() {
219 let mut c = GspClient::new();
220 let payload = encode_bare(SignalType::Join, 1, 42);
221 let accept = c.accept(&payload, 0, PayloadCodec::Cbor).unwrap();
222 assert_eq!(accept.signal, SignalType::Join);
223 assert!(c.members.contains(&42));
224 }
225
226 #[test]
227 fn leave_removes_sender_from_members() {
228 let mut c = GspClient::new();
229 c.accept(&encode_bare(SignalType::Join, 1, 7), 0, PayloadCodec::Cbor)
230 .unwrap();
231 c.accept(&encode_bare(SignalType::Leave, 2, 7), 0, PayloadCodec::Cbor)
232 .unwrap();
233 assert!(!c.members.contains(&7));
234 }
235
236 #[test]
237 fn leave_also_removes_from_muted() {
238 let mut c = GspClient::new();
239 c.accept(&encode_bare(SignalType::Join, 1, 5), 0, PayloadCodec::Cbor)
240 .unwrap();
241 c.muted.insert(5); c.accept(&encode_bare(SignalType::Leave, 2, 5), 0, PayloadCodec::Cbor)
243 .unwrap();
244 assert!(!c.muted.contains(&5));
245 }
246
247 #[test]
248 fn duplicate_request_id_is_rejected() {
249 let mut c = GspClient::new();
250 c.accept(&encode_bare(SignalType::Join, 99, 1), 0, PayloadCodec::Cbor)
251 .unwrap();
252 let result = c.accept(
253 &encode_bare(SignalType::Leave, 99, 1),
254 0,
255 PayloadCodec::Cbor,
256 );
257 assert!(matches!(result, Err(GspError::DuplicateRequest(99))));
258 }
259
260 #[test]
261 fn epoch_advance_clears_request_seen_set() {
262 let mut c = GspClient::new();
263 let payload = encode_bare(SignalType::Join, 1, 10);
264 c.accept(&payload, 0, PayloadCodec::Cbor).unwrap();
265 let result = c.accept(
267 &encode_bare(SignalType::Leave, 1, 10),
268 1,
269 PayloadCodec::Cbor,
270 );
271 assert!(result.is_ok());
272 }
273
274 #[test]
275 fn reset_clears_state() {
276 let mut c = GspClient::new();
277 c.accept(&encode_bare(SignalType::Join, 1, 3), 0, PayloadCodec::Cbor)
278 .unwrap();
279 c.reset();
280 c.accept(&encode_bare(SignalType::Join, 1, 4), 0, PayloadCodec::Cbor)
282 .unwrap();
283 }
286
287 #[test]
288 fn unknown_signal_type_rejected() {
289 let mut c = GspClient::new();
290 let bad = GspSignal::bare(999, 1, 1).to_cbor();
291 assert!(matches!(
292 c.accept(&bad, 0, PayloadCodec::Cbor),
293 Err(GspError::UnknownSignal(999))
294 ));
295 }
296
297 #[test]
298 fn invalid_cbor_returns_decode_error() {
299 let mut c = GspClient::new();
300 assert!(matches!(
301 c.accept(b"\xFF\xFF", 0, PayloadCodec::Cbor),
302 Err(GspError::Decode(_))
303 ));
304 }
305
306 #[test]
307 fn multiple_members_join_independently() {
308 let mut c = GspClient::new();
309 c.accept(&encode_bare(SignalType::Join, 1, 10), 0, PayloadCodec::Cbor)
310 .unwrap();
311 c.accept(&encode_bare(SignalType::Join, 2, 20), 0, PayloadCodec::Cbor)
312 .unwrap();
313 c.accept(&encode_bare(SignalType::Join, 3, 30), 0, PayloadCodec::Cbor)
314 .unwrap();
315 assert_eq!(c.members.len(), 3);
316 assert!(c.members.contains(&10));
317 assert!(c.members.contains(&20));
318 assert!(c.members.contains(&30));
319 }
320}