1use std::cell::Cell;
2
3use crate::codec::{Decoder, Encoder};
4use crate::util::{BytePage, BytePages, ByteString, Bytes, BytesMut};
5
6use super::error::ProtocolError;
7use super::frame::Parser;
8use super::proto::{CloseReason, OpCode};
9
10#[derive(Debug, Clone, PartialEq, Eq)]
12pub enum Message {
13 Text(ByteString),
15 Binary(Bytes),
17 Continuation(Item),
19 Ping(Bytes),
21 Pong(Bytes),
23 Close(Option<CloseReason>),
25}
26
27#[derive(Debug, Clone, PartialEq, Eq)]
29pub enum Frame {
30 Text(Bytes),
34 Binary(Bytes),
36 Continuation(Item),
38 Ping(Bytes),
40 Pong(Bytes),
42 Close(Option<CloseReason>),
44}
45
46#[derive(Debug, Clone, PartialEq, Eq)]
48pub enum Item {
49 FirstText(Bytes),
51 FirstBinary(Bytes),
53 Continue(Bytes),
55 Last(Bytes),
57}
58
59#[derive(Debug, Clone)]
60pub struct Codec {
62 flags: Cell<Flags>,
63 max_size: usize,
64}
65
66bitflags::bitflags! {
67 #[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
68 struct Flags: u8 {
69 const SERVER = 0b0000_0001;
70 const R_CONTINUATION = 0b0000_0010;
71 const W_CONTINUATION = 0b0000_0100;
72 const CLOSED = 0b0000_1000;
73 const R_CLOSED = 0b0001_0000;
74 }
75}
76
77impl Codec {
78 #[must_use]
80 pub fn new() -> Codec {
81 Codec {
82 max_size: 65_536,
83 flags: Cell::new(Flags::SERVER),
84 }
85 }
86
87 #[must_use]
91 pub fn max_size(mut self, size: usize) -> Self {
92 self.max_size = size;
93 self
94 }
95
96 #[must_use]
101 pub fn set_client_mode(self) -> Self {
102 self.remove_flags(Flags::SERVER);
103 self
104 }
105
106 pub fn is_closed(&self) -> bool {
108 self.flags().contains(Flags::CLOSED)
109 }
110
111 fn flags(&self) -> Flags {
112 self.flags.get()
113 }
114
115 fn insert_flags(&self, f: Flags) {
116 self.flags.set(self.flags() | f);
117 }
118
119 fn remove_flags(&self, f: Flags) {
120 self.flags.set(self.flags() - f);
121 }
122
123 pub fn encode_page(&self, page: BytePage, dst: &mut BytePages) -> Result<(), ProtocolError> {
131 if self.is_closed() {
132 return Err(ProtocolError::Closed);
133 }
134 if self.flags().contains(Flags::W_CONTINUATION) {
135 return Err(ProtocolError::ContinuationStarted);
136 }
137 Parser::write_message(
138 dst,
139 page,
140 OpCode::Binary,
141 true,
142 !self.flags().contains(Flags::SERVER),
143 )
144 .expect("binary frames are always valid");
145 Ok(())
146 }
147}
148
149impl Default for Codec {
150 fn default() -> Self {
151 Self::new()
152 }
153}
154
155impl Encoder for Codec {
156 type Item = Message;
157 type Error = ProtocolError;
158
159 fn encode(&self, item: Message, dst: &mut BytePages) -> Result<(), Self::Error> {
160 if self.is_closed() {
161 return Err(ProtocolError::Closed);
162 }
163
164 match item {
165 Message::Text(txt) => {
166 if self.flags().contains(Flags::W_CONTINUATION) {
167 return Err(ProtocolError::ContinuationStarted);
168 }
169 Parser::write_message(
170 dst,
171 txt,
172 OpCode::Text,
173 true,
174 !self.flags().contains(Flags::SERVER),
175 )?;
176 }
177 Message::Binary(bin) => {
178 if self.flags().contains(Flags::W_CONTINUATION) {
179 return Err(ProtocolError::ContinuationStarted);
180 }
181 Parser::write_message(
182 dst,
183 bin,
184 OpCode::Binary,
185 true,
186 !self.flags().contains(Flags::SERVER),
187 )?;
188 }
189 Message::Ping(txt) => Parser::write_message(
190 dst,
191 txt,
192 OpCode::Ping,
193 true,
194 !self.flags().contains(Flags::SERVER),
195 )?,
196 Message::Pong(txt) => Parser::write_message(
197 dst,
198 txt,
199 OpCode::Pong,
200 true,
201 !self.flags().contains(Flags::SERVER),
202 )?,
203 Message::Close(reason) => {
204 Parser::write_close(dst, reason, !self.flags().contains(Flags::SERVER))?;
205 self.insert_flags(Flags::CLOSED);
206 }
207 Message::Continuation(cont) => match cont {
208 Item::FirstText(data) => {
209 if self.flags().contains(Flags::W_CONTINUATION) {
210 return Err(ProtocolError::ContinuationStarted);
211 }
212 self.insert_flags(Flags::W_CONTINUATION);
213 Parser::write_message(
214 dst,
215 data,
216 OpCode::Text,
217 false,
218 !self.flags().contains(Flags::SERVER),
219 )?;
220 }
221 Item::FirstBinary(data) => {
222 if self.flags().contains(Flags::W_CONTINUATION) {
223 return Err(ProtocolError::ContinuationStarted);
224 }
225 self.insert_flags(Flags::W_CONTINUATION);
226 Parser::write_message(
227 dst,
228 data,
229 OpCode::Binary,
230 false,
231 !self.flags().contains(Flags::SERVER),
232 )?;
233 }
234 Item::Continue(data) => {
235 if self.flags().contains(Flags::W_CONTINUATION) {
236 Parser::write_message(
237 dst,
238 data,
239 OpCode::Continue,
240 false,
241 !self.flags().contains(Flags::SERVER),
242 )?;
243 } else {
244 return Err(ProtocolError::ContinuationNotStarted);
245 }
246 }
247 Item::Last(data) => {
248 if self.flags().contains(Flags::W_CONTINUATION) {
249 self.remove_flags(Flags::W_CONTINUATION);
250 Parser::write_message(
251 dst,
252 data,
253 OpCode::Continue,
254 true,
255 !self.flags().contains(Flags::SERVER),
256 )?;
257 } else {
258 return Err(ProtocolError::ContinuationNotStarted);
259 }
260 }
261 },
262 }
263 Ok(())
264 }
265}
266
267impl Decoder for Codec {
268 type Item = Frame;
269 type Error = ProtocolError;
270
271 fn decode(&self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
272 if self.flags().contains(Flags::R_CLOSED) {
274 src.clear();
275 return Ok(None);
276 }
277
278 match Parser::parse(src, self.flags().contains(Flags::SERVER), self.max_size) {
279 Ok(Some((finished, opcode, payload))) => {
280 if finished {
282 match opcode {
283 OpCode::Continue => {
284 if self.flags().contains(Flags::R_CONTINUATION) {
285 self.remove_flags(Flags::R_CONTINUATION);
286 let payload = payload.unwrap_or_default();
287 Ok(Some(Frame::Continuation(Item::Last(payload))))
288 } else {
289 Err(ProtocolError::ContinuationNotStarted)
290 }
291 }
292 OpCode::Close => {
293 let reason = if let Some(pl) = payload {
294 Parser::parse_close_payload(&pl)?
295 } else {
296 None
297 };
298 self.insert_flags(Flags::R_CLOSED);
299 Ok(Some(Frame::Close(reason)))
300 }
301 OpCode::Ping => Ok(Some(Frame::Ping(payload.unwrap_or_default()))),
302 OpCode::Pong => Ok(Some(Frame::Pong(payload.unwrap_or_default()))),
303 OpCode::Binary => {
304 if self.flags().contains(Flags::R_CONTINUATION) {
305 Err(ProtocolError::ContinuationStarted)
306 } else {
307 Ok(Some(Frame::Binary(payload.unwrap_or_else(Bytes::new))))
308 }
309 }
310 OpCode::Text => {
311 if self.flags().contains(Flags::R_CONTINUATION) {
312 Err(ProtocolError::ContinuationStarted)
313 } else {
314 Ok(Some(Frame::Text(payload.unwrap_or_else(Bytes::new))))
315 }
316 }
317 }
318 } else {
319 match opcode {
320 OpCode::Continue => {
321 if self.flags().contains(Flags::R_CONTINUATION) {
322 Ok(Some(Frame::Continuation(Item::Continue(
323 payload.unwrap_or_else(Bytes::new),
324 ))))
325 } else {
326 Err(ProtocolError::ContinuationNotStarted)
327 }
328 }
329 OpCode::Binary => {
330 if self.flags().contains(Flags::R_CONTINUATION) {
331 Err(ProtocolError::ContinuationStarted)
332 } else {
333 self.insert_flags(Flags::R_CONTINUATION);
334 Ok(Some(Frame::Continuation(Item::FirstBinary(
335 payload.unwrap_or_else(Bytes::new),
336 ))))
337 }
338 }
339 OpCode::Text => {
340 if self.flags().contains(Flags::R_CONTINUATION) {
341 Err(ProtocolError::ContinuationStarted)
342 } else {
343 self.insert_flags(Flags::R_CONTINUATION);
344 Ok(Some(Frame::Continuation(Item::FirstText(
345 payload.unwrap_or_else(Bytes::new),
346 ))))
347 }
348 }
349 OpCode::Ping | OpCode::Pong | OpCode::Close => {
351 Err(ProtocolError::FragmentedControlFrame(opcode))
352 }
353 }
354 }
355 }
356 Ok(None) => Ok(None),
357 Err(e) => Err(e),
358 }
359 }
360}
361
362#[cfg(test)]
363mod tests {
364 use super::*;
365 use crate::ws::CloseCode;
366
367 #[test]
368 fn text_payload_is_not_validated() {
369 let codec = Codec::new().set_client_mode();
370 let mut frame = BytesMut::from(&[0x81, 0x01, 0xff][..]);
371 assert!(matches!(
372 codec.decode(&mut frame),
373 Ok(Some(Frame::Text(data))) if data == Bytes::from_static(&[0xff])
374 ));
375 }
376
377 #[test]
378 fn input_after_close_is_discarded() {
379 let codec = Codec::new().set_client_mode();
380 let mut src = BytesMut::from(&[0x88, 0x00, 0x81, 0x01, b'a'][..]);
382 assert!(matches!(
383 codec.decode(&mut src),
384 Ok(Some(Frame::Close(None)))
385 ));
386 assert!(matches!(codec.decode(&mut src), Ok(None)));
387 assert!(src.is_empty());
388
389 src.extend_from_slice(&[0x81, 0x01, b'b']);
390 assert!(matches!(codec.decode(&mut src), Ok(None)));
391 assert!(src.is_empty());
392 }
393
394 #[test]
395 fn accepts_extension_close_code() {
396 let codec = Codec::new().set_client_mode();
398 let mut close = BytesMut::from(&[0x88, 0x02, 0x03, 0xf2][..]);
399 assert!(matches!(
400 codec.decode(&mut close),
401 Ok(Some(Frame::Close(Some(CloseReason {
402 code: CloseCode::Extension,
403 ..
404 }))))
405 ));
406
407 let codec = Codec::new();
409 let mut dst = BytePages::default();
410 assert!(matches!(
411 codec.encode(Message::Close(Some(CloseCode::Extension.into())), &mut dst),
412 Err(ProtocolError::InvalidCloseCode(1010))
413 ));
414 }
415
416 #[test]
417 fn validates_outgoing_continuations() {
418 let codec = Codec::new();
419 let mut dst = BytePages::default();
420 codec
421 .encode(
422 Message::Continuation(Item::FirstBinary(Bytes::new())),
423 &mut dst,
424 )
425 .unwrap();
426 assert!(matches!(
427 codec.encode(Message::Text("text".into()), &mut dst),
428 Err(ProtocolError::ContinuationStarted)
429 ));
430 assert!(matches!(
431 codec.encode_page(BytePage::from(Bytes::new()), &mut dst),
432 Err(ProtocolError::ContinuationStarted)
433 ));
434
435 codec
436 .encode(Message::Continuation(Item::Last(Bytes::new())), &mut dst)
437 .unwrap();
438 codec
439 .encode_page(BytePage::from(Bytes::new()), &mut dst)
440 .unwrap();
441 }
442
443 #[test]
444 fn rejects_messages_after_close() {
445 let codec = Codec::new();
446 let mut dst = BytePages::default();
447 codec.encode(Message::Close(None), &mut dst).unwrap();
448
449 assert!(matches!(
450 codec.encode(Message::Text("text".into()), &mut dst),
451 Err(ProtocolError::Closed)
452 ));
453 assert!(matches!(
454 codec.encode_page(BytePage::from(Bytes::new()), &mut dst),
455 Err(ProtocolError::Closed)
456 ));
457 }
458
459 #[test]
460 fn encode_errors() {
461 let codec = Codec::new();
462 let mut dst = BytePages::default();
463 let big = Bytes::from(vec![0; 126]);
464 assert!(matches!(
465 codec.encode(Message::Ping(big.clone()), &mut dst),
466 Err(ProtocolError::InvalidLength(126))
467 ));
468 assert!(matches!(
469 codec.encode(Message::Pong(big), &mut dst),
470 Err(ProtocolError::InvalidLength(126))
471 ));
472
473 codec
474 .encode(Message::Continuation(Item::FirstText("a".into())), &mut dst)
475 .unwrap();
476 assert!(matches!(
477 codec.encode(Message::Binary("b".into()), &mut dst),
478 Err(ProtocolError::ContinuationStarted)
479 ));
480 assert!(matches!(
481 codec.encode(
482 Message::Continuation(Item::FirstBinary("b".into())),
483 &mut dst
484 ),
485 Err(ProtocolError::ContinuationStarted)
486 ));
487 codec
488 .encode(Message::Continuation(Item::Last("c".into())), &mut dst)
489 .unwrap();
490 assert!(matches!(
491 codec.encode(Message::Continuation(Item::Continue("d".into())), &mut dst),
492 Err(ProtocolError::ContinuationNotStarted)
493 ));
494 }
495
496 fn decode(codec: &Codec, frame: &[u8]) -> Result<Option<Frame>, ProtocolError> {
497 codec.decode(&mut BytesMut::from(frame))
498 }
499
500 #[test]
501 fn decode_continuation_errors() {
502 let codec = Codec::new().set_client_mode();
503 assert!(matches!(
505 decode(&codec, &[0x80, 0x01, b'a']),
506 Err(ProtocolError::ContinuationNotStarted)
507 ));
508 assert!(matches!(
509 decode(&codec, &[0x00, 0x01, b'a']),
510 Err(ProtocolError::ContinuationNotStarted)
511 ));
512
513 assert!(matches!(
515 decode(&codec, &[0x02, 0x01, b'a']),
516 Ok(Some(Frame::Continuation(Item::FirstBinary(_))))
517 ));
518 for frame in [
519 &[0x82, 0x01, b'a'],
520 &[0x81, 0x01, b'a'],
521 &[0x02, 0x01, b'a'],
522 &[0x01, 0x01, b'a'],
523 ] {
524 assert!(matches!(
525 decode(&codec, frame),
526 Err(ProtocolError::ContinuationStarted)
527 ));
528 }
529 assert!(matches!(
530 decode(&codec, &[0x00, 0x01, b'b']),
531 Ok(Some(Frame::Continuation(Item::Continue(_))))
532 ));
533 assert!(matches!(
534 decode(&codec, &[0x80, 0x00]),
535 Ok(Some(Frame::Continuation(Item::Last(data)))) if data.is_empty()
536 ));
537 }
538}