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 GspClient {
67 pub fn new() -> Self {
69 Self {
70 seen_requests: BoundedSeen::new(GSP_SEEN_CAP),
71 muted: HashSet::new(),
72 members: HashSet::new(),
73 current_epoch: None,
74 }
75 }
76
77 pub fn send<S: Sealer>(
81 &mut self,
82 node: &mut GroupNode,
83 seal: &mut S,
84 target: MemberId,
85 signal: SignalType,
86 role_claim: u32,
87 request_id: u32,
88 codec: PayloadCodec,
89 ) -> Result<OutboundFrame, GspError> {
90 self.send_with_args(
91 node,
92 seal,
93 target,
94 signal,
95 role_claim,
96 request_id,
97 &[],
98 codec,
99 )
100 }
101
102 pub fn send_with_args<S: Sealer>(
108 &mut self,
109 node: &mut GroupNode,
110 seal: &mut S,
111 target: MemberId,
112 signal: SignalType,
113 role_claim: u32,
114 request_id: u32,
115 args: &[u8],
116 codec: PayloadCodec,
117 ) -> Result<OutboundFrame, GspError> {
118 self.sync_epoch(node.current_epoch);
119 let mut sig = GspSignal::bare(signal as u32, request_id, node.member_id);
120 sig.role_claim = role_claim;
121 sig.args = serde_bytes::ByteBuf::from(args.to_vec());
122 sig.args_length = args.len() as u32;
123 let stream_id = node.member_stream_id(3);
124 Ok(node.send_payload(
125 seal,
126 target,
127 StreamType::Signal,
128 stream_id,
129 GbpFlags::ordered_reliable_ack(),
130 &sig.to_bytes(codec),
131 codec,
132 )?)
133 }
134
135 pub fn accept(
142 &mut self,
143 plaintext: &[u8],
144 current_epoch: u64,
145 codec: PayloadCodec,
146 ) -> Result<GspAccept, GspError> {
147 self.sync_epoch(current_epoch);
148 let s = GspSignal::from_bytes(plaintext, codec)?;
149 let signal = SignalType::try_from(s.signal_type).map_err(GspError::UnknownSignal)?;
150 validate_args(signal, &s.args).map_err(GspError::BadSchema)?;
152 if !self.seen_requests.insert(s.request_id) {
153 return Err(GspError::DuplicateRequest(s.request_id));
154 }
155 match signal {
156 SignalType::Join => {
157 self.members.insert(s.sender_id);
158 }
159 SignalType::Leave => {
160 self.members.remove(&s.sender_id);
161 self.muted.remove(&s.sender_id);
162 }
163 SignalType::Mute => {
164 self.muted.insert(s.sender_id);
165 }
166 SignalType::Unmute => {
167 self.muted.remove(&s.sender_id);
168 }
169 _ => {}
170 }
171 Ok(GspAccept {
172 signal,
173 sender_id: s.sender_id,
174 role_claim: s.role_claim,
175 request_id: s.request_id,
176 })
177 }
178
179 pub fn sync_epoch(&mut self, epoch: u64) {
183 if Some(epoch) != self.current_epoch {
184 self.seen_requests.clear();
185 self.current_epoch = Some(epoch);
186 }
187 }
188
189 pub fn reset(&mut self) {
191 self.seen_requests.clear();
192 self.current_epoch = None;
193 }
194}
195
196#[cfg(test)]
197mod tests {
198 use super::*;
199 use crate::GspSignal;
200
201 fn encode_bare(signal: SignalType, request_id: u32, sender_id: u32) -> Vec<u8> {
202 GspSignal::bare(signal as u32, request_id, sender_id).to_cbor()
203 }
204
205 #[test]
206 fn join_adds_sender_to_members() {
207 let mut c = GspClient::new();
208 let payload = encode_bare(SignalType::Join, 1, 42);
209 let accept = c.accept(&payload, 0, PayloadCodec::Cbor).unwrap();
210 assert_eq!(accept.signal, SignalType::Join);
211 assert!(c.members.contains(&42));
212 }
213
214 #[test]
215 fn leave_removes_sender_from_members() {
216 let mut c = GspClient::new();
217 c.accept(&encode_bare(SignalType::Join, 1, 7), 0, PayloadCodec::Cbor)
218 .unwrap();
219 c.accept(&encode_bare(SignalType::Leave, 2, 7), 0, PayloadCodec::Cbor)
220 .unwrap();
221 assert!(!c.members.contains(&7));
222 }
223
224 #[test]
225 fn leave_also_removes_from_muted() {
226 let mut c = GspClient::new();
227 c.accept(&encode_bare(SignalType::Join, 1, 5), 0, PayloadCodec::Cbor)
228 .unwrap();
229 c.muted.insert(5); c.accept(&encode_bare(SignalType::Leave, 2, 5), 0, PayloadCodec::Cbor)
231 .unwrap();
232 assert!(!c.muted.contains(&5));
233 }
234
235 #[test]
236 fn duplicate_request_id_is_rejected() {
237 let mut c = GspClient::new();
238 c.accept(&encode_bare(SignalType::Join, 99, 1), 0, PayloadCodec::Cbor)
239 .unwrap();
240 let result = c.accept(
241 &encode_bare(SignalType::Leave, 99, 1),
242 0,
243 PayloadCodec::Cbor,
244 );
245 assert!(matches!(result, Err(GspError::DuplicateRequest(99))));
246 }
247
248 #[test]
249 fn epoch_advance_clears_request_seen_set() {
250 let mut c = GspClient::new();
251 let payload = encode_bare(SignalType::Join, 1, 10);
252 c.accept(&payload, 0, PayloadCodec::Cbor).unwrap();
253 let result = c.accept(
255 &encode_bare(SignalType::Leave, 1, 10),
256 1,
257 PayloadCodec::Cbor,
258 );
259 assert!(result.is_ok());
260 }
261
262 #[test]
263 fn reset_clears_state() {
264 let mut c = GspClient::new();
265 c.accept(&encode_bare(SignalType::Join, 1, 3), 0, PayloadCodec::Cbor)
266 .unwrap();
267 c.reset();
268 c.accept(&encode_bare(SignalType::Join, 1, 4), 0, PayloadCodec::Cbor)
270 .unwrap();
271 }
274
275 #[test]
276 fn unknown_signal_type_rejected() {
277 let mut c = GspClient::new();
278 let bad = GspSignal::bare(999, 1, 1).to_cbor();
279 assert!(matches!(
280 c.accept(&bad, 0, PayloadCodec::Cbor),
281 Err(GspError::UnknownSignal(999))
282 ));
283 }
284
285 #[test]
286 fn invalid_cbor_returns_decode_error() {
287 let mut c = GspClient::new();
288 assert!(matches!(
289 c.accept(b"\xFF\xFF", 0, PayloadCodec::Cbor),
290 Err(GspError::Decode(_))
291 ));
292 }
293
294 #[test]
295 fn multiple_members_join_independently() {
296 let mut c = GspClient::new();
297 c.accept(&encode_bare(SignalType::Join, 1, 10), 0, PayloadCodec::Cbor)
298 .unwrap();
299 c.accept(&encode_bare(SignalType::Join, 2, 20), 0, PayloadCodec::Cbor)
300 .unwrap();
301 c.accept(&encode_bare(SignalType::Join, 3, 30), 0, PayloadCodec::Cbor)
302 .unwrap();
303 assert_eq!(c.members.len(), 3);
304 assert!(c.members.contains(&10));
305 assert!(c.members.contains(&20));
306 assert!(c.members.contains(&30));
307 }
308}