1use thiserror::Error;
2
3use crate::{Frame, FrameCodecError, Payload, WireMessage};
4
5const PUSH: u8 = 0;
6const REPLY: u8 = 1;
7const BROADCAST: u8 = 2;
8
9pub trait Codec {
11 fn encode(&self, frame: &Frame) -> Result<WireMessage, CodecError>;
13 fn decode(&self, message: WireMessage) -> Result<Frame, CodecError>;
15}
16
17#[derive(Clone, Copy, Debug, Eq, PartialEq)]
19pub struct CodecLimits {
20 pub max_frame_bytes: usize,
22 pub max_binary_payload_bytes: usize,
24}
25
26impl Default for CodecLimits {
27 fn default() -> Self {
28 Self {
29 max_frame_bytes: 16 * 1024 * 1024,
30 max_binary_payload_bytes: 16 * 1024 * 1024,
31 }
32 }
33}
34
35#[derive(Clone, Copy, Debug, Default)]
37pub struct PhoenixV2Codec;
38
39impl PhoenixV2Codec {
40 pub fn limited(limits: CodecLimits) -> LimitedPhoenixV2Codec {
42 LimitedPhoenixV2Codec { limits }
43 }
44}
45
46#[derive(Clone, Copy, Debug)]
48pub struct LimitedPhoenixV2Codec {
49 limits: CodecLimits,
50}
51
52impl LimitedPhoenixV2Codec {
53 pub fn limits(&self) -> CodecLimits {
55 self.limits
56 }
57}
58
59impl Codec for PhoenixV2Codec {
60 fn encode(&self, frame: &Frame) -> Result<WireMessage, CodecError> {
61 match &frame.payload {
62 Payload::Json(_) => Ok(WireMessage::Text(frame.encode_text()?)),
63 Payload::Binary(payload) => Ok(WireMessage::Binary(encode_push(frame, payload)?)),
64 Payload::Reply { .. } => Err(CodecError::InvalidOutboundReplyPayload),
65 }
66 }
67
68 fn decode(&self, message: WireMessage) -> Result<Frame, CodecError> {
69 match message {
70 WireMessage::Text(text) => Ok(Frame::decode_text(&text)?),
71 WireMessage::Binary(bytes) => decode_binary(&bytes),
72 }
73 }
74}
75
76impl Codec for LimitedPhoenixV2Codec {
77 fn encode(&self, frame: &Frame) -> Result<WireMessage, CodecError> {
78 validate_payload_size(&frame.payload, self.limits)?;
79 let message = PhoenixV2Codec.encode(frame)?;
80 validate_frame_size(&message, self.limits)?;
81 Ok(message)
82 }
83
84 fn decode(&self, message: WireMessage) -> Result<Frame, CodecError> {
85 validate_frame_size(&message, self.limits)?;
86 let frame = PhoenixV2Codec.decode(message)?;
87 validate_payload_size(&frame.payload, self.limits)?;
88 Ok(frame)
89 }
90}
91
92fn validate_frame_size(message: &WireMessage, limits: CodecLimits) -> Result<(), CodecError> {
93 let length = match message {
94 WireMessage::Text(text) => text.len(),
95 WireMessage::Binary(bytes) => bytes.len(),
96 };
97 if length > limits.max_frame_bytes {
98 return Err(CodecError::FrameTooLarge {
99 length,
100 maximum: limits.max_frame_bytes,
101 });
102 }
103 Ok(())
104}
105
106fn validate_payload_size(payload: &Payload, limits: CodecLimits) -> Result<(), CodecError> {
107 let length = match payload {
108 Payload::Binary(bytes) => Some(bytes.len()),
109 Payload::Reply { response, .. } => match response.as_ref() {
110 Payload::Binary(bytes) => Some(bytes.len()),
111 _ => None,
112 },
113 Payload::Json(_) => None,
114 };
115 if let Some(length) = length {
116 if length > limits.max_binary_payload_bytes {
117 return Err(CodecError::BinaryPayloadTooLarge {
118 length,
119 maximum: limits.max_binary_payload_bytes,
120 });
121 }
122 }
123 Ok(())
124}
125
126fn encode_push(frame: &Frame, payload: &[u8]) -> Result<Vec<u8>, CodecError> {
127 let join_ref = frame.join_ref.as_deref().unwrap_or_default().as_bytes();
128 let reference = frame.reference.as_deref().unwrap_or_default().as_bytes();
129 let topic = frame.topic.as_bytes();
130 let event = frame.event.as_bytes();
131 let lengths = [
132 field_size(join_ref, "join_ref")?,
133 field_size(reference, "ref")?,
134 field_size(topic, "topic")?,
135 field_size(event, "event")?,
136 ];
137 let mut encoded = Vec::with_capacity(
138 5 + join_ref.len() + reference.len() + topic.len() + event.len() + payload.len(),
139 );
140 encoded.push(PUSH);
141 encoded.extend_from_slice(&lengths);
142 encoded.extend_from_slice(join_ref);
143 encoded.extend_from_slice(reference);
144 encoded.extend_from_slice(topic);
145 encoded.extend_from_slice(event);
146 encoded.extend_from_slice(payload);
147 Ok(encoded)
148}
149
150fn decode_binary(input: &[u8]) -> Result<Frame, CodecError> {
151 let Some(kind) = input.first().copied() else {
152 return Err(CodecError::TruncatedBinaryFrame);
153 };
154 match kind {
155 PUSH => {
156 let sizes = binary_header(input, 4)?;
157 let (fields, payload) = binary_fields(input, 4, sizes)?;
158 Ok(Frame::new(
159 optional_utf8(fields[0], "join_ref")?,
160 None,
161 utf8(fields[1], "topic")?,
162 utf8(fields[2], "event")?,
163 payload.to_vec(),
164 ))
165 }
166 REPLY => {
167 let sizes = binary_header(input, 5)?;
168 let (fields, payload) = binary_fields(input, 5, sizes)?;
169 Ok(Frame::new(
170 optional_utf8(fields[0], "join_ref")?,
171 optional_utf8(fields[1], "ref")?,
172 utf8(fields[2], "topic")?,
173 "phx_reply",
174 Payload::Reply {
175 status: utf8(fields[3], "status")?,
176 response: Box::new(Payload::Binary(payload.to_vec())),
177 },
178 ))
179 }
180 BROADCAST => {
181 let sizes = binary_header(input, 3)?;
182 let (fields, payload) = binary_fields(input, 3, sizes)?;
183 Ok(Frame::new(
184 None,
185 None,
186 utf8(fields[0], "topic")?,
187 utf8(fields[1], "event")?,
188 payload.to_vec(),
189 ))
190 }
191 other => Err(CodecError::UnknownBinaryKind(other)),
192 }
193}
194
195fn binary_header(input: &[u8], header_length: usize) -> Result<&[u8], CodecError> {
196 if input.len() < header_length {
197 return Err(CodecError::TruncatedBinaryFrame);
198 }
199 Ok(&input[1..header_length])
200}
201
202fn binary_fields<'a>(
203 input: &'a [u8],
204 header_length: usize,
205 sizes: &[u8],
206) -> Result<(Vec<&'a [u8]>, &'a [u8]), CodecError> {
207 let metadata_length = sizes.iter().map(|size| usize::from(*size)).sum::<usize>();
208 if input.len() < header_length + metadata_length {
209 return Err(CodecError::TruncatedBinaryFrame);
210 }
211 let mut offset = header_length;
212 let mut fields = Vec::with_capacity(sizes.len());
213 for size in sizes {
214 let end = offset + usize::from(*size);
215 fields.push(&input[offset..end]);
216 offset = end;
217 }
218 Ok((fields, &input[offset..]))
219}
220
221fn optional_utf8(value: &[u8], field: &'static str) -> Result<Option<String>, CodecError> {
222 if value.is_empty() {
223 Ok(None)
224 } else {
225 utf8(value, field).map(Some)
226 }
227}
228
229fn utf8(value: &[u8], field: &'static str) -> Result<String, CodecError> {
230 std::str::from_utf8(value)
231 .map(str::to_owned)
232 .map_err(|_| CodecError::InvalidUtf8(field))
233}
234
235fn field_size(value: &[u8], field: &'static str) -> Result<u8, CodecError> {
236 u8::try_from(value.len()).map_err(|_| CodecError::FieldTooLong {
237 field,
238 length: value.len(),
239 })
240}
241
242#[derive(Debug, Error)]
244pub enum CodecError {
245 #[error(transparent)]
247 Text(#[from] FrameCodecError),
248 #[error("Phoenix binary frame is truncated")]
250 TruncatedBinaryFrame,
251 #[error("unknown Phoenix binary frame kind: {0}")]
253 UnknownBinaryKind(u8),
254 #[error("Phoenix binary frame contains invalid UTF-8 in {0}")]
256 InvalidUtf8(&'static str),
257 #[error("Phoenix binary {field} field exceeds 255 bytes: {length}")]
259 FieldTooLong {
260 field: &'static str,
262 length: usize,
264 },
265 #[error("reply envelope cannot be sent as an application payload")]
267 InvalidOutboundReplyPayload,
268 #[error("Phoenix frame is {length} bytes, exceeding the {maximum}-byte limit")]
270 FrameTooLarge {
271 length: usize,
273 maximum: usize,
275 },
276 #[error("Phoenix binary payload is {length} bytes, exceeding the {maximum}-byte limit")]
278 BinaryPayloadTooLarge {
279 length: usize,
281 maximum: usize,
283 },
284}
285
286#[cfg(test)]
287mod tests {
288 use proptest::prelude::*;
289 use serde_json::json;
290
291 use super::*;
292
293 #[test]
294 fn encodes_client_binary_pushes() {
295 let frame = Frame::new(
296 Some("1".into()),
297 Some("2".into()),
298 "room:lobby",
299 "binary",
300 vec![7, 8, 9],
301 );
302 let WireMessage::Binary(encoded) = PhoenixV2Codec.encode(&frame).unwrap() else {
303 panic!("expected a binary WebSocket message");
304 };
305 assert_eq!(&encoded[..5], &[0, 1, 1, 10, 6]);
306 assert_eq!(&encoded[5..], b"12room:lobbybinary\x07\x08\x09");
307 }
308
309 #[test]
310 fn decodes_binary_push_reply_and_broadcast_frames() {
311 let push = [vec![0, 1, 4, 4], b"1roomping".to_vec(), vec![1, 2]].concat();
312 let frame = PhoenixV2Codec.decode(WireMessage::Binary(push)).unwrap();
313 assert_eq!(frame.topic, "room");
314 assert_eq!(frame.event, "ping");
315 assert_eq!(frame.payload, Payload::Binary(vec![1, 2]));
316
317 let reply = [vec![1, 1, 1, 4, 2], b"12roomok".to_vec(), vec![3, 4]].concat();
318 let frame = PhoenixV2Codec.decode(WireMessage::Binary(reply)).unwrap();
319 assert_eq!(frame.event, "phx_reply");
320 assert_eq!(
321 frame.payload,
322 Payload::Reply {
323 status: "ok".into(),
324 response: Box::new(Payload::Binary(vec![3, 4])),
325 }
326 );
327
328 let broadcast = [vec![2, 4, 5], b"roomevent".to_vec(), vec![5, 6]].concat();
329 let frame = PhoenixV2Codec
330 .decode(WireMessage::Binary(broadcast))
331 .unwrap();
332 assert_eq!(frame.join_ref, None);
333 assert_eq!(frame.payload, Payload::Binary(vec![5, 6]));
334
335 let text = Frame::new(None, None, "room", "event", json!({}));
336 assert!(matches!(
337 PhoenixV2Codec.encode(&text).unwrap(),
338 WireMessage::Text(_)
339 ));
340 }
341
342 #[test]
343 fn enforces_frame_and_binary_payload_limits() {
344 let codec = PhoenixV2Codec::limited(CodecLimits {
345 max_frame_bytes: 64,
346 max_binary_payload_bytes: 2,
347 });
348 let frame = Frame::new(None, None, "room", "binary", vec![1, 2, 3]);
349 assert!(matches!(
350 codec.encode(&frame),
351 Err(CodecError::BinaryPayloadTooLarge { .. })
352 ));
353
354 let codec = PhoenixV2Codec::limited(CodecLimits {
355 max_frame_bytes: 4,
356 max_binary_payload_bytes: 4,
357 });
358 assert!(matches!(
359 codec.decode(WireMessage::Text("[null,null,\"room\",\"e\",{}]".into())),
360 Err(CodecError::FrameTooLarge { .. })
361 ));
362 }
363
364 proptest! {
365 #[test]
366 fn arbitrary_binary_frames_never_panic(input in proptest::collection::vec(any::<u8>(), 0..4096)) {
367 let codec = PhoenixV2Codec::limited(CodecLimits {
368 max_frame_bytes: 2048,
369 max_binary_payload_bytes: 1024,
370 });
371 let _ = codec.decode(WireMessage::Binary(input));
372 }
373
374 #[test]
375 fn arbitrary_text_frames_never_panic(input in ".{0,4096}") {
376 let codec = PhoenixV2Codec::limited(CodecLimits {
377 max_frame_bytes: 2048,
378 max_binary_payload_bytes: 1024,
379 });
380 let _ = codec.decode(WireMessage::Text(input));
381 }
382 }
383}