Skip to main content

atomic_websocket/helpers/
common.rs

1use 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/// Log debug macro: prioritizes rinf > debug > default
14#[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/// Log error macro: prioritizes rinf > debug > default
41#[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/// Flattens a `spawn_blocking` join result (a possible `JoinError` wrapping an
68/// inner `String` error) into the `Box<dyn Error>` used by the database helpers.
69///
70/// All native-db (redb) work is synchronous and performs disk I/O — most notably
71/// `commit()` issues an `fsync`. Running that directly on a Tokio worker thread
72/// stalls the async runtime on slow disks (low-spec machines), which drifts the
73/// ping loop timing and causes false disconnections. We therefore run every DB
74/// transaction inside `spawn_blocking` so the fsync never blocks a worker thread.
75#[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/// Create a message from raw bytes without category prefix.
216/// Use this when bebop feature is disabled.
217#[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        // category 100 = 100 % 256 = 100, 100 / 256 = 0
274        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        // Category 10000 = 10000 % 256 = 16, 10000 / 256 = 39
282        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        // Create a valid data with category 100
310        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        // category 10000 = 16 + 39*256
321        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        // Should have category bytes at the start
350        assert!(data.len() > 2);
351        // Category::Ping = 10000 => 16, 39
352        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        // Category::Pong = 10001 => 17, 39
362        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        // Category::Disconnect = 10003 => 19, 39
372        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        // Category::Expired = 10002 => 18, 39
383        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    // ========================================================================
412    // get_setting_by_key, set_setting 데이터베이스 함수 테스트
413    // ========================================================================
414
415    #[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            // 저장된 값 확인
449            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            // 첫 번째 저장
461            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            // 같은 키로 업데이트
468            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            // 업데이트된 값 확인
476            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            // 여러 키 저장
506            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            // 각 키가 독립적으로 조회됨
537            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            // 큰 데이터 저장
579            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}