1use gbp::CodecError;
4use gbp_core::PayloadCodec;
5use serde::{Deserialize, Serialize};
6use serde_bytes::ByteBuf;
7
8#[repr(u8)]
10#[derive(Copy, Clone, Debug, PartialEq, Eq)]
11pub enum GtpContentType {
12 Plain = 0,
14 Markdown = 1,
16 Binary = 2,
18 AttachmentRef = 3,
20}
21
22#[derive(Clone, Debug, Serialize, Deserialize)]
24pub struct GtpMessage {
25 #[serde(rename = "mid")]
27 pub message_id: u64,
28 #[serde(rename = "sid")]
30 pub sender_id: u32,
31 #[serde(rename = "ts")]
33 pub timestamp_ms: u64,
34 #[serde(rename = "rid")]
36 pub request_id: u32,
37 #[serde(rename = "fl")]
39 pub flags: u8,
40 #[serde(rename = "ct")]
42 pub content_type: u8,
43 #[serde(rename = "len")]
45 pub content_length: u32,
46 #[serde(rename = "body")]
48 pub content: ByteBuf,
49}
50
51impl GtpMessage {
52 pub fn plain(sender_id: u32, message_id: u64, text: &str) -> Self {
54 let body = text.as_bytes().to_vec();
55 Self {
56 message_id,
57 sender_id,
58 timestamp_ms: 0,
59 request_id: 0,
60 flags: 0x01,
61 content_type: GtpContentType::Plain as u8,
62 content_length: body.len() as u32,
63 content: ByteBuf::from(body),
64 }
65 }
66
67 pub fn to_cbor(&self) -> Vec<u8> {
69 let mut buf = Vec::new();
70 ciborium::into_writer(self, &mut buf).expect("cbor encode");
71 buf
72 }
73
74 pub fn from_cbor(data: &[u8]) -> Result<Self, CodecError> {
76 let m: Self = ciborium::from_reader(data).map_err(|e| CodecError::Decode(e.to_string()))?;
77 if m.content_length as usize != m.content.len() {
78 return Err(CodecError::PayloadSizeMismatch);
79 }
80 Ok(m)
81 }
82
83 pub fn text(&self) -> Option<&str> {
85 std::str::from_utf8(&self.content).ok()
86 }
87
88 pub fn to_bytes(&self, codec: PayloadCodec) -> Vec<u8> {
90 match codec {
91 PayloadCodec::Cbor => self.to_cbor(),
92 PayloadCodec::Protobuf => {
93 use prost::Message as _;
94 gbp_proto::gtp::GtpMessage::from(self).encode_to_vec()
95 }
96 PayloadCodec::FlatBuffers => {
97 let mut b = gbp_flat::planus::Builder::new();
98 b.finish(gbp_flat::gtp::GtpMessage::from(self), None)
99 .to_vec()
100 }
101 }
102 }
103
104 pub fn from_bytes(data: &[u8], codec: PayloadCodec) -> Result<Self, CodecError> {
106 match codec {
107 PayloadCodec::Cbor => Self::from_cbor(data),
108 PayloadCodec::Protobuf => {
109 use prost::Message as _;
110 let p = gbp_proto::gtp::GtpMessage::decode(data)
111 .map_err(|e| CodecError::Decode(e.to_string()))?;
112 Self::try_from(p).map_err(|_| CodecError::PayloadSizeMismatch)
113 }
114 PayloadCodec::FlatBuffers => {
115 use gbp_flat::planus::ReadAsRoot as _;
116 let r = gbp_flat::gtp::GtpMessageRef::read_as_root(data)
117 .map_err(|e| CodecError::Decode(e.to_string()))?;
118 Self::try_from(r).map_err(|_| CodecError::PayloadSizeMismatch)
119 }
120 }
121 }
122}
123
124impl From<&GtpMessage> for gbp_proto::gtp::GtpMessage {
127 fn from(m: &GtpMessage) -> Self {
128 Self {
129 message_id: m.message_id,
130 sender_id: m.sender_id,
131 timestamp_ms: m.timestamp_ms,
132 request_id: m.request_id,
133 flags: m.flags as u32,
134 content_type: m.content_type as u32,
135 content_length: m.content_length,
136 content: m.content.to_vec(),
137 }
138 }
139}
140
141impl TryFrom<gbp_proto::gtp::GtpMessage> for GtpMessage {
142 type Error = ();
143 fn try_from(p: gbp_proto::gtp::GtpMessage) -> Result<Self, ()> {
144 if p.content_length as usize != p.content.len() {
145 return Err(());
146 }
147 Ok(Self {
148 message_id: p.message_id,
149 sender_id: p.sender_id,
150 timestamp_ms: p.timestamp_ms,
151 request_id: p.request_id,
152 flags: p.flags as u8,
153 content_type: p.content_type as u8,
154 content_length: p.content_length,
155 content: ByteBuf::from(p.content),
156 })
157 }
158}
159
160impl From<&GtpMessage> for gbp_flat::gtp::GtpMessage {
163 fn from(m: &GtpMessage) -> Self {
164 Self {
165 message_id: m.message_id,
166 sender_id: m.sender_id,
167 timestamp_ms: m.timestamp_ms,
168 request_id: m.request_id,
169 flags: m.flags as u32,
170 content_type: m.content_type as u32,
171 content_length: m.content_length,
172 content: Some(m.content.to_vec()),
173 }
174 }
175}
176
177impl<'a> TryFrom<gbp_flat::gtp::GtpMessageRef<'a>> for GtpMessage {
178 type Error = ();
179 fn try_from(r: gbp_flat::gtp::GtpMessageRef<'a>) -> Result<Self, ()> {
180 let content = r.content().map_err(|_| ())?.unwrap_or(&[]).to_vec();
181 let content_length = r.content_length().map_err(|_| ())?;
182 if content_length as usize != content.len() {
183 return Err(());
184 }
185 Ok(Self {
186 message_id: r.message_id().map_err(|_| ())?,
187 sender_id: r.sender_id().map_err(|_| ())?,
188 timestamp_ms: r.timestamp_ms().map_err(|_| ())?,
189 request_id: r.request_id().map_err(|_| ())?,
190 flags: r.flags().map_err(|_| ())? as u8,
191 content_type: r.content_type().map_err(|_| ())? as u8,
192 content_length,
193 content: ByteBuf::from(content),
194 })
195 }
196}
197
198#[cfg(test)]
199mod tests {
200 use super::*;
201
202 fn sample() -> GtpMessage {
203 GtpMessage::plain(42, 0xDEAD_BEEF, "codec roundtrip")
204 }
205
206 #[test]
207 fn cbor_roundtrip() {
208 let orig = sample();
209 let bytes = orig.to_bytes(PayloadCodec::Cbor);
210 let decoded = GtpMessage::from_bytes(&bytes, PayloadCodec::Cbor).unwrap();
211 assert_eq!(decoded.message_id, orig.message_id);
212 assert_eq!(decoded.sender_id, orig.sender_id);
213 assert_eq!(decoded.text().unwrap(), "codec roundtrip");
214 }
215
216 #[test]
217 fn protobuf_roundtrip() {
218 let orig = sample();
219 let bytes = orig.to_bytes(PayloadCodec::Protobuf);
220 let decoded = GtpMessage::from_bytes(&bytes, PayloadCodec::Protobuf).unwrap();
221 assert_eq!(decoded.message_id, orig.message_id);
222 assert_eq!(decoded.sender_id, orig.sender_id);
223 assert_eq!(decoded.text().unwrap(), "codec roundtrip");
224 }
225
226 #[test]
227 fn flatbuffers_roundtrip() {
228 let orig = sample();
229 let bytes = orig.to_bytes(PayloadCodec::FlatBuffers);
230 let decoded = GtpMessage::from_bytes(&bytes, PayloadCodec::FlatBuffers).unwrap();
231 assert_eq!(decoded.message_id, orig.message_id);
232 assert_eq!(decoded.sender_id, orig.sender_id);
233 assert_eq!(decoded.text().unwrap(), "codec roundtrip");
234 }
235
236 #[test]
237 fn codec_bytes_differ() {
238 let msg = sample();
239 let cbor = msg.to_bytes(PayloadCodec::Cbor);
240 let proto = msg.to_bytes(PayloadCodec::Protobuf);
241 let flat = msg.to_bytes(PayloadCodec::FlatBuffers);
242 assert_ne!(cbor, proto);
243 assert_ne!(cbor, flat);
244 assert_ne!(proto, flat);
245 }
246}