atomic_websocket/helpers/
common.rs1use std::error::Error;
2
3#[cfg(feature = "bebop")]
4use bebop::Record;
5use tokio_tungstenite::tungstenite::{Bytes, Message};
6
7#[cfg(feature = "bebop")]
8use crate::schema::{Category, Data, Disconnect, Expired, Ping};
9use crate::Settings;
10
11use super::types::DB;
12
13#[cfg(feature = "rinf")]
15#[macro_export]
16macro_rules! log_debug {
17 ($($rest:tt)*) => {
18 rinf::debug_print!($($rest)*)
19 };
20}
21
22#[cfg(all(not(feature = "rinf"), feature = "debug"))]
23#[macro_export]
24macro_rules! log_debug {
25 ($($rest:tt)*) => {
26 log::debug!($($rest)*)
27 };
28}
29
30#[cfg(all(not(feature = "rinf"), not(feature = "debug")))]
31#[macro_export]
32macro_rules! log_debug {
33 ($($rest:tt)*) => {
34 if cfg!(debug_assertions) {
35 println!($($rest)*)
36 }
37 };
38}
39
40#[cfg(feature = "rinf")]
42#[macro_export]
43macro_rules! log_error {
44 ($($rest:tt)*) => {
45 rinf::debug_print!($($rest)*)
46 };
47}
48
49#[cfg(all(not(feature = "rinf"), feature = "debug"))]
50#[macro_export]
51macro_rules! log_error {
52 ($($rest:tt)*) => {
53 log::error!($($rest)*)
54 };
55}
56
57#[cfg(all(not(feature = "rinf"), not(feature = "debug")))]
58#[macro_export]
59macro_rules! log_error {
60 ($($rest:tt)*) => {
61 if cfg!(debug_assertions) {
62 println!($($rest)*)
63 }
64 };
65}
66
67#[cfg(feature = "native-db")]
76fn flatten_join<T>(
77 res: Result<Result<T, String>, tokio::task::JoinError>,
78) -> Result<T, Box<dyn Error>> {
79 match res {
80 Ok(Ok(value)) => Ok(value),
81 Ok(Err(e)) => Err(<Box<dyn Error>>::from(e)),
82 Err(e) => Err(<Box<dyn Error>>::from(e.to_string())),
83 }
84}
85
86#[cfg(feature = "native-db")]
87pub async fn get_setting_by_key(db: DB, key: String) -> Result<Option<Settings>, Box<dyn Error>> {
88 let res = tokio::task::spawn_blocking(move || -> Result<Option<Settings>, String> {
89 let db = db.blocking_lock();
90 let reader = db.r_transaction().map_err(|e| e.to_string())?;
91 reader
92 .get()
93 .primary::<Settings>(key)
94 .map_err(|e| e.to_string())
95 })
96 .await;
97 flatten_join(res)
98}
99
100#[cfg(not(feature = "native-db"))]
101pub async fn get_setting_by_key(db: DB, key: String) -> Result<Option<Settings>, Box<dyn Error>> {
102 let db = db.lock().await;
103 Ok(db.get(&key).map(|v| Settings {
104 key,
105 value: v.clone(),
106 }))
107}
108
109#[cfg(feature = "native-db")]
110pub async fn set_setting(db: DB, settings: Settings) -> Result<bool, Box<dyn Error>> {
111 let res = tokio::task::spawn_blocking(move || -> Result<bool, String> {
112 let db = db.blocking_lock();
113 let reader = db.r_transaction().map_err(|e| e.to_string())?;
114 let setting = reader
115 .get()
116 .primary::<Settings>(settings.key.clone())
117 .map_err(|e| e.to_string())?;
118 drop(reader);
119
120 let writer = db.rw_transaction().map_err(|e| e.to_string())?;
121 match setting {
122 Some(setting) => writer
123 .update::<Settings>(setting, settings)
124 .map_err(|e| e.to_string())?,
125 None => writer
126 .insert::<Settings>(settings)
127 .map_err(|e| e.to_string())?,
128 }
129 writer.commit().map_err(|e| e.to_string())?;
130
131 Ok(true)
132 })
133 .await;
134 flatten_join(res)
135}
136
137#[cfg(not(feature = "native-db"))]
138pub async fn set_setting(db: DB, settings: Settings) -> Result<bool, Box<dyn Error>> {
139 let mut db = db.lock().await;
140 db.insert(settings.key, settings.value);
141 Ok(true)
142}
143
144#[cfg(feature = "native-db")]
145#[allow(dead_code)]
146pub async fn remove_setting(db: DB, key: String) -> Result<bool, Box<dyn Error>> {
147 let res = tokio::task::spawn_blocking(move || -> Result<bool, String> {
148 let db = db.blocking_lock();
149 let reader = db.r_transaction().map_err(|e| e.to_string())?;
150 let setting = reader
151 .get()
152 .primary::<Settings>(key)
153 .map_err(|e| e.to_string())?;
154 drop(reader);
155
156 if let Some(setting) = setting {
157 let writer = db.rw_transaction().map_err(|e| e.to_string())?;
158 writer
159 .remove::<Settings>(setting)
160 .map_err(|e| e.to_string())?;
161 writer.commit().map_err(|e| e.to_string())?;
162 }
163 Ok(true)
164 })
165 .await;
166 flatten_join(res)
167}
168
169#[cfg(not(feature = "native-db"))]
170#[allow(dead_code)]
171pub async fn remove_setting(db: DB, key: String) -> Result<bool, Box<dyn Error>> {
172 let mut db = db.lock().await;
173 db.remove(&key);
174 Ok(true)
175}
176
177#[cfg(feature = "bebop")]
178pub fn make_ping_message(peer: &str) -> Message {
179 let mut datas = Vec::with_capacity(64);
180 Ping {
181 peer,
182 activations: 0,
183 }
184 .serialize(&mut datas)
185 .expect("Ping serialization should never fail");
186 make_response_message(Category::Ping, datas)
187}
188
189#[cfg(not(feature = "bebop"))]
190pub fn make_ping_message(_peer: &str) -> Message {
191 Message::Binary(Bytes::new())
192}
193
194#[cfg(feature = "bebop")]
195pub fn get_data_schema(data: &[u8]) -> Result<Data<'_>, Box<dyn Error>> {
196 if data.len() < 2 {
197 return Err("Data length is too short".into());
198 }
199 Ok(Data {
200 category: data[0] as u16 + data[1] as u16 * 256,
201 datas: bebop::SliceWrapper::from_raw(&data[2..]),
202 })
203}
204
205pub fn make_atomic_message(category: u16, mut datas: Vec<u8>) -> Message {
206 let mut byte = {
207 let quotient = category / 256;
208 let remainder = category % 256;
209 vec![remainder as u8, quotient as u8]
210 };
211 byte.append(&mut datas);
212 Message::Binary(Bytes::from(byte))
213}
214
215#[allow(dead_code)]
218pub fn make_raw_message(data: &[u8]) -> Message {
219 Message::Binary(Bytes::copy_from_slice(data))
220}
221
222#[cfg(feature = "bebop")]
223pub fn make_response_message(category: Category, datas: Vec<u8>) -> Message {
224 make_atomic_message(category as u16, datas)
225}
226
227#[cfg(feature = "bebop")]
228pub fn make_disconnect_message(peer: &str) -> Message {
229 let mut datas = Vec::with_capacity(64);
230 Disconnect { peer }
231 .serialize(&mut datas)
232 .expect("Disconnect serialization should never fail");
233 make_response_message(Category::Disconnect, datas)
234}
235
236#[cfg(not(feature = "bebop"))]
237pub fn make_disconnect_message(_peer: &str) -> Message {
238 Message::Binary(Bytes::new())
239}
240
241#[cfg(feature = "bebop")]
242pub fn make_pong_message() -> Message {
243 make_response_message(Category::Pong, Vec::new())
244}
245
246#[cfg(not(feature = "bebop"))]
247pub fn make_pong_message() -> Message {
248 Message::Binary(Bytes::new())
249}
250
251#[cfg(feature = "bebop")]
252pub fn make_expired_output_message() -> Message {
253 let mut datas = Vec::with_capacity(16);
254 Expired { is_expired: true }
255 .serialize(&mut datas)
256 .expect("Expired serialization should never fail");
257 make_response_message(Category::Expired, datas)
258}
259
260#[cfg(not(feature = "bebop"))]
261pub fn make_expired_output_message() -> Message {
262 Message::Binary(Bytes::new())
263}
264
265#[cfg(test)]
266mod tests {
267 use super::*;
268
269 #[test]
270 fn test_make_atomic_message_basic() {
271 let msg = make_atomic_message(100, vec![1, 2, 3]);
272 let data = msg.into_data();
273 assert_eq!(data[0], 100);
275 assert_eq!(data[1], 0);
276 assert_eq!(&data[2..], &[1, 2, 3]);
277 }
278
279 #[test]
280 fn test_make_atomic_message_large_category() {
281 let msg = make_atomic_message(10000, vec![]);
283 let data = msg.into_data();
284 assert_eq!(data[0], 16);
285 assert_eq!(data[1], 39);
286 assert_eq!(data.len(), 2);
287 }
288
289 #[test]
290 fn test_make_atomic_message_empty_data() {
291 let msg = make_atomic_message(0, vec![]);
292 let data = msg.into_data();
293 assert_eq!(data.len(), 2);
294 assert_eq!(data[0], 0);
295 assert_eq!(data[1], 0);
296 }
297
298 #[test]
299 fn test_make_raw_message() {
300 let raw_data = vec![10, 20, 30, 40];
301 let msg = make_raw_message(&raw_data);
302 let data = msg.into_data();
303 assert_eq!(data, raw_data);
304 }
305
306 #[cfg(feature = "bebop")]
307 #[test]
308 fn test_get_data_schema_valid() {
309 let data = vec![100, 0, 1, 2, 3];
311 let result = get_data_schema(&data);
312 assert!(result.is_ok());
313 let schema = result.unwrap();
314 assert_eq!(schema.category, 100);
315 }
316
317 #[cfg(feature = "bebop")]
318 #[test]
319 fn test_get_data_schema_large_category() {
320 let data = vec![16, 39, 1, 2, 3];
322 let result = get_data_schema(&data);
323 assert!(result.is_ok());
324 let schema = result.unwrap();
325 assert_eq!(schema.category, 10000);
326 }
327
328 #[cfg(feature = "bebop")]
329 #[test]
330 fn test_get_data_schema_too_short() {
331 let data = vec![1];
332 let result = get_data_schema(&data);
333 assert!(result.is_err());
334 }
335
336 #[cfg(feature = "bebop")]
337 #[test]
338 fn test_get_data_schema_empty() {
339 let data: Vec<u8> = vec![];
340 let result = get_data_schema(&data);
341 assert!(result.is_err());
342 }
343
344 #[cfg(feature = "bebop")]
345 #[test]
346 fn test_make_ping_message() {
347 let msg = make_ping_message("test-peer");
348 let data = msg.into_data();
349 assert!(data.len() > 2);
351 assert_eq!(data[0], 16);
353 assert_eq!(data[1], 39);
354 }
355
356 #[cfg(feature = "bebop")]
357 #[test]
358 fn test_make_pong_message() {
359 let msg = make_pong_message();
360 let data = msg.into_data();
361 assert_eq!(data[0], 17);
363 assert_eq!(data[1], 39);
364 }
365
366 #[cfg(feature = "bebop")]
367 #[test]
368 fn test_make_disconnect_message() {
369 let msg = make_disconnect_message("peer-123");
370 let data = msg.into_data();
371 assert_eq!(data[0], 19);
373 assert_eq!(data[1], 39);
374 assert!(data.len() > 2);
375 }
376
377 #[cfg(feature = "bebop")]
378 #[test]
379 fn test_make_expired_output_message() {
380 let msg = make_expired_output_message();
381 let data = msg.into_data();
382 assert_eq!(data[0], 18);
384 assert_eq!(data[1], 39);
385 }
386
387 #[cfg(not(feature = "bebop"))]
388 #[test]
389 fn test_make_ping_message_no_bebop() {
390 let msg = make_ping_message("test-peer");
391 let data = msg.into_data();
392 assert!(data.is_empty());
393 }
394
395 #[cfg(not(feature = "bebop"))]
396 #[test]
397 fn test_make_pong_message_no_bebop() {
398 let msg = make_pong_message();
399 let data = msg.into_data();
400 assert!(data.is_empty());
401 }
402
403 #[cfg(not(feature = "bebop"))]
404 #[test]
405 fn test_make_disconnect_message_no_bebop() {
406 let msg = make_disconnect_message("peer-123");
407 let data = msg.into_data();
408 assert!(data.is_empty());
409 }
410
411 #[cfg(not(feature = "native-db"))]
416 mod db_tests {
417 use super::*;
418 use crate::helpers::types::InMemoryStorage;
419 use std::sync::Arc;
420 use tokio::sync::Mutex;
421
422 fn create_test_db() -> DB {
423 Arc::new(Mutex::new(InMemoryStorage::new()))
424 }
425
426 #[tokio::test]
427 async fn test_get_setting_by_key_not_found() {
428 let db = create_test_db();
429
430 let result = get_setting_by_key(db, "nonexistent_key".to_string()).await;
431 assert!(result.is_ok());
432 assert!(result.unwrap().is_none());
433 }
434
435 #[tokio::test]
436 async fn test_set_setting_insert_new() {
437 let db = create_test_db();
438
439 let settings = Settings {
440 key: "test_key".to_string(),
441 value: vec![1, 2, 3, 4, 5],
442 };
443
444 let result = set_setting(db.clone(), settings).await;
445 assert!(result.is_ok());
446 assert!(result.unwrap());
447
448 let retrieved = get_setting_by_key(db, "test_key".to_string()).await;
450 assert!(retrieved.is_ok());
451 let settings = retrieved.unwrap();
452 assert!(settings.is_some());
453 assert_eq!(settings.unwrap().value, vec![1, 2, 3, 4, 5]);
454 }
455
456 #[tokio::test]
457 async fn test_set_setting_update_existing() {
458 let db = create_test_db();
459
460 let settings1 = Settings {
462 key: "update_key".to_string(),
463 value: vec![1, 2, 3],
464 };
465 set_setting(db.clone(), settings1).await.unwrap();
466
467 let settings2 = Settings {
469 key: "update_key".to_string(),
470 value: vec![10, 20, 30, 40],
471 };
472 let result = set_setting(db.clone(), settings2).await;
473 assert!(result.is_ok());
474
475 let retrieved = get_setting_by_key(db, "update_key".to_string()).await;
477 assert!(retrieved.is_ok());
478 let settings = retrieved.unwrap().unwrap();
479 assert_eq!(settings.value, vec![10, 20, 30, 40]);
480 }
481
482 #[tokio::test]
483 async fn test_get_setting_by_key_returns_correct_key() {
484 let db = create_test_db();
485
486 let settings = Settings {
487 key: "my_key".to_string(),
488 value: vec![100],
489 };
490 set_setting(db.clone(), settings).await.unwrap();
491
492 let retrieved = get_setting_by_key(db, "my_key".to_string())
493 .await
494 .unwrap()
495 .unwrap();
496
497 assert_eq!(retrieved.key, "my_key");
498 assert_eq!(retrieved.value, vec![100]);
499 }
500
501 #[tokio::test]
502 async fn test_multiple_keys_independent() {
503 let db = create_test_db();
504
505 set_setting(
507 db.clone(),
508 Settings {
509 key: "key1".to_string(),
510 value: vec![1],
511 },
512 )
513 .await
514 .unwrap();
515
516 set_setting(
517 db.clone(),
518 Settings {
519 key: "key2".to_string(),
520 value: vec![2],
521 },
522 )
523 .await
524 .unwrap();
525
526 set_setting(
527 db.clone(),
528 Settings {
529 key: "key3".to_string(),
530 value: vec![3],
531 },
532 )
533 .await
534 .unwrap();
535
536 let v1 = get_setting_by_key(db.clone(), "key1".to_string())
538 .await
539 .unwrap()
540 .unwrap();
541 let v2 = get_setting_by_key(db.clone(), "key2".to_string())
542 .await
543 .unwrap()
544 .unwrap();
545 let v3 = get_setting_by_key(db.clone(), "key3".to_string())
546 .await
547 .unwrap()
548 .unwrap();
549
550 assert_eq!(v1.value, vec![1]);
551 assert_eq!(v2.value, vec![2]);
552 assert_eq!(v3.value, vec![3]);
553 }
554
555 #[tokio::test]
556 async fn test_set_setting_empty_value() {
557 let db = create_test_db();
558
559 let settings = Settings {
560 key: "empty_key".to_string(),
561 value: vec![],
562 };
563
564 let result = set_setting(db.clone(), settings).await;
565 assert!(result.is_ok());
566
567 let retrieved = get_setting_by_key(db, "empty_key".to_string())
568 .await
569 .unwrap()
570 .unwrap();
571 assert!(retrieved.value.is_empty());
572 }
573
574 #[tokio::test]
575 async fn test_set_setting_large_value() {
576 let db = create_test_db();
577
578 let large_value: Vec<u8> = (0..10000).map(|i| (i % 256) as u8).collect();
580
581 let settings = Settings {
582 key: "large_key".to_string(),
583 value: large_value.clone(),
584 };
585
586 let result = set_setting(db.clone(), settings).await;
587 assert!(result.is_ok());
588
589 let retrieved = get_setting_by_key(db, "large_key".to_string())
590 .await
591 .unwrap()
592 .unwrap();
593 assert_eq!(retrieved.value.len(), 10000);
594 assert_eq!(retrieved.value, large_value);
595 }
596 }
597}