Skip to main content

wsio_core/packet/codecs/
mod.rs

1use std::{
2    fmt::{
3        Debug as FmtDebug,
4        Formatter,
5        Result as FmtResult,
6    },
7    sync::Arc,
8};
9
10use anyhow::Result;
11use bytes::Bytes;
12use erased_serde::deserialize;
13use serde::{
14    Serialize,
15    de::DeserializeOwned,
16};
17
18#[cfg(feature = "packet-codec-cbor")]
19mod cbor;
20pub mod custom;
21mod msgpack;
22#[cfg(feature = "packet-codec-postcard")]
23mod postcard;
24
25#[cfg(feature = "packet-codec-cbor")]
26use self::cbor::WsIoPacketCborCodec;
27#[cfg(feature = "packet-codec-postcard")]
28use self::postcard::WsIoPacketPostcardCodec;
29use self::{
30    custom::WsIoPacketCustomCodec,
31    msgpack::WsIoPacketMsgpackCodec,
32};
33use super::WsIoPacket;
34
35// Enums
36#[derive(Clone)]
37pub enum WsIoPacketCodec {
38    #[cfg(feature = "packet-codec-cbor")]
39    Cbor,
40    Custom(Arc<dyn WsIoPacketCustomCodec>),
41    Msgpack,
42
43    #[cfg(feature = "packet-codec-postcard")]
44    Postcard,
45}
46
47impl FmtDebug for WsIoPacketCodec {
48    fn fmt(&self, f: &mut Formatter<'_>) -> FmtResult {
49        match self {
50            #[cfg(feature = "packet-codec-cbor")]
51            Self::Cbor => f.write_str("WsIoPacketCodec::Cbor"),
52            Self::Msgpack => f.write_str("WsIoPacketCodec::Msgpack"),
53            Self::Custom(_) => f.write_str("WsIoPacketCodec::Custom(<codec>)"),
54            #[cfg(feature = "packet-codec-postcard")]
55            Self::Postcard => f.write_str("WsIoPacketCodec::Postcard"),
56        }
57    }
58}
59
60impl WsIoPacketCodec {
61    #[inline]
62    pub fn decode(&self, bytes: &[u8]) -> Result<WsIoPacket> {
63        match self {
64            #[cfg(feature = "packet-codec-cbor")]
65            Self::Cbor => WsIoPacketCborCodec::decode(bytes),
66            Self::Custom(codec) => codec.decode(bytes),
67            Self::Msgpack => WsIoPacketMsgpackCodec::decode(bytes),
68
69            #[cfg(feature = "packet-codec-postcard")]
70            Self::Postcard => WsIoPacketPostcardCodec::decode(bytes),
71        }
72    }
73
74    #[inline]
75    pub fn decode_data<D: DeserializeOwned>(&self, bytes: &[u8]) -> Result<D> {
76        match self {
77            #[cfg(feature = "packet-codec-cbor")]
78            Self::Cbor => WsIoPacketCborCodec::decode_data(bytes),
79            Self::Custom(codec) => {
80                let mut deserializer = codec.decode_data(bytes)?;
81                Ok(deserialize(&mut *deserializer)?)
82            },
83            Self::Msgpack => WsIoPacketMsgpackCodec::decode_data(bytes),
84
85            #[cfg(feature = "packet-codec-postcard")]
86            Self::Postcard => WsIoPacketPostcardCodec::decode_data(bytes),
87        }
88    }
89
90    #[inline]
91    pub fn encode(&self, packet: &WsIoPacket) -> Result<Bytes> {
92        match self {
93            #[cfg(feature = "packet-codec-cbor")]
94            Self::Cbor => WsIoPacketCborCodec::encode(packet),
95            Self::Custom(codec) => codec.encode(packet),
96            Self::Msgpack => WsIoPacketMsgpackCodec::encode(packet),
97
98            #[cfg(feature = "packet-codec-postcard")]
99            Self::Postcard => WsIoPacketPostcardCodec::encode(packet),
100        }
101    }
102
103    #[inline]
104    pub fn encode_data<D: Serialize>(&self, data: &D) -> Result<Bytes> {
105        match self {
106            #[cfg(feature = "packet-codec-cbor")]
107            Self::Cbor => WsIoPacketCborCodec::encode_data(data),
108            Self::Custom(codec) => codec.encode_data(data),
109            Self::Msgpack => WsIoPacketMsgpackCodec::encode_data(data),
110
111            #[cfg(feature = "packet-codec-postcard")]
112            Self::Postcard => WsIoPacketPostcardCodec::encode_data(data),
113        }
114    }
115}
116
117#[cfg(test)]
118mod tests {
119    use bytes::Bytes;
120    use serde::{
121        Deserialize,
122        Serialize,
123    };
124
125    use super::*;
126    use crate::packet::{
127        WsIoPacket,
128        WsIoPacketType,
129    };
130
131    #[derive(Debug, Deserialize, PartialEq, Serialize)]
132    struct TestPayload {
133        id: u32,
134        message: String,
135    }
136
137    macro_rules! test_codec {
138        ($codec:expr, $name:ident) => {
139            #[test]
140            fn $name() {
141                let codec = $codec;
142
143                // 1. Test encoding/decoding raw data
144                let original_data = TestPayload {
145                    id: 42,
146                    message: "hello world".to_string(),
147                };
148
149                let encoded_data = codec.encode_data(&original_data).expect("Failed to encode data");
150                let decoded_data: TestPayload = codec.decode_data(&encoded_data).expect("Failed to decode data");
151                assert_eq!(
152                    original_data, decoded_data,
153                    "Data decoding did not match original"
154                );
155
156                // 2. Test encoding/decoding an Event packet with data
157                let packet = WsIoPacket::new_event("chat", Some(encoded_data.clone()));
158                let encoded_packet = codec.encode(&packet).expect("Failed to encode packet");
159                let decoded_packet = codec.decode(&encoded_packet).expect("Failed to decode packet");
160
161                assert!(
162                    matches!(decoded_packet.r#type, WsIoPacketType::Event),
163                    "Packet type mismatch"
164                );
165
166                assert_eq!(decoded_packet.key.as_deref(), Some("chat"), "Packet key mismatch");
167                assert_eq!(
168                    decoded_packet.data.as_deref(),
169                    Some(&encoded_data[..]),
170                    "Packet data mismatch"
171                );
172
173                // 3. Test encoding/decoding a Disconnect packet (no data, no key)
174                let packet = WsIoPacket::new_disconnect();
175                let encoded_packet = codec.encode(&packet).expect("Failed to encode disconnect packet");
176                let decoded_packet = codec
177                    .decode(&encoded_packet)
178                    .expect("Failed to decode disconnect packet");
179
180                assert!(
181                    matches!(decoded_packet.r#type, WsIoPacketType::Disconnect),
182                    "Packet type mismatch"
183                );
184
185                assert_eq!(decoded_packet.key, None, "Packet key should be None");
186                assert_eq!(decoded_packet.data, None, "Packet data should be None");
187            }
188        };
189    }
190
191    #[cfg(feature = "packet-codec-cbor")]
192    test_codec!(WsIoPacketCodec::Cbor, test_cbor_codec);
193    test_codec!(WsIoPacketCodec::Msgpack, test_msgpack_codec);
194
195    #[test]
196    fn test_msgpack_packet_data_uses_binary_payload() {
197        let payload = Bytes::from_static(&[0, 1, 255]);
198        let packet = WsIoPacket::new_event("chat", Some(payload.clone()));
199        let encoded_packet = WsIoPacketCodec::Msgpack
200            .encode(&packet)
201            .expect("Failed to encode packet");
202
203        assert!(encoded_packet.windows(2).any(|window| window == [0xc4, 3]));
204
205        let decoded_packet = WsIoPacketCodec::Msgpack
206            .decode(&encoded_packet)
207            .expect("Failed to decode packet");
208        assert_eq!(decoded_packet.data.as_deref(), Some(&payload[..]));
209    }
210
211    #[cfg(feature = "packet-codec-postcard")]
212    test_codec!(WsIoPacketCodec::Postcard, test_postcard_codec);
213}