1use rama_core::error::BoxErrorExt as _;
7use std::{fmt, str::FromStr};
8
9use rama_core::error::{BoxError, ErrorContext as _, ErrorExt};
10use rama_core::extensions::Extension as ExtensionTrait;
11use rama_core::telemetry::tracing;
12use rama_utils::str::arcstr::ArcStr;
13
14derive_non_empty_flat_csv_header! {
15 #[header(name = SEC_WEBSOCKET_EXTENSIONS, sep = Comma)]
16 #[derive(Clone, Debug, PartialEq, Eq)]
20 pub struct SecWebSocketExtensions(pub NonEmptySmallVec<3, Extension>);
21}
22
23impl SecWebSocketExtensions {
24 #[inline]
25 #[must_use]
28 pub fn per_message_deflate() -> Self {
29 Self::per_message_deflate_with_config(Default::default())
30 }
31
32 #[inline]
33 #[must_use]
36 pub fn per_message_deflate_with_config(config: PerMessageDeflateConfig) -> Self {
37 Self::new(Extension::PerMessageDeflate(config))
38 }
39}
40
41impl SecWebSocketExtensions {
42 rama_utils::macros::generate_set_and_with! {
43 pub fn extra_extension(mut self, ext: impl Into<Extension>) -> Self {
45 self.0.push(ext.into());
46 self
47 }
48 }
49
50 rama_utils::macros::generate_set_and_with! {
51 pub fn extra_extensions(mut self, ext_it: impl IntoIterator<Item = impl Into<Extension>>) -> Self {
53 self.0.extend(ext_it.into_iter().map(Into::into));
54 self
55 }
56 }
57}
58
59#[derive(Debug, Clone, PartialEq, Eq)]
60pub enum Extension {
64 PerMessageDeflate(PerMessageDeflateConfig),
73
74 Empty,
76
77 Unknown(ArcStr),
81}
82
83impl ExtensionTrait for Extension {}
84
85impl Extension {
86 #[must_use]
87 pub fn into_header(self) -> SecWebSocketExtensions {
90 SecWebSocketExtensions::new(self)
91 }
92}
93
94impl From<Extension> for SecWebSocketExtensions {
95 fn from(value: Extension) -> Self {
96 Self::new(value)
97 }
98}
99
100impl From<PerMessageDeflateConfig> for Extension {
101 fn from(value: PerMessageDeflateConfig) -> Self {
102 Self::PerMessageDeflate(value)
103 }
104}
105
106impl From<PerMessageDeflateIdentifier> for Extension {
107 fn from(value: PerMessageDeflateIdentifier) -> Self {
108 Self::PerMessageDeflate(PerMessageDeflateConfig::from(value))
109 }
110}
111
112rama_utils::macros::enums::enum_builder! {
113 #[derive(Default)]
115 @String
116 pub enum PerMessageDeflateIdentifier {
117 #[default]
118 PerMessageDeflate => "permessage-deflate",
120 PerFrameDeflate => "perframe-deflate",
122 XWebKitDeflateFrame => "x-webkit-deflate-frame",
124 }
125}
126
127#[derive(Debug, Clone, Default, PartialEq, Eq)]
128pub struct PerMessageDeflateConfig {
129 pub identifier: PerMessageDeflateIdentifier,
134
135 pub server_no_context_takeover: bool,
148
149 pub client_no_context_takeover: bool,
165
166 pub server_max_window_bits: Option<u8>,
180
181 pub client_max_window_bits: Option<u8>,
202}
203
204impl From<PerMessageDeflateIdentifier> for PerMessageDeflateConfig {
205 fn from(identifier: PerMessageDeflateIdentifier) -> Self {
206 Self {
207 identifier,
208 ..Default::default()
209 }
210 }
211}
212
213impl FromStr for Extension {
214 type Err = BoxError;
215
216 fn from_str(s: &str) -> Result<Self, Self::Err> {
217 let mut parts = s.split(';').map(|s| s.trim());
218 let identifier = parts
219 .next()
220 .context("empty WebSocket Extension is invalid")?;
221 if let Some(identifier) = PerMessageDeflateIdentifier::strict_parse(identifier) {
222 let mut config = PerMessageDeflateConfig {
223 identifier,
224 ..Default::default()
225 };
226 for part in parts {
227 if part.eq_ignore_ascii_case("server_no_context_takeover") {
228 if std::mem::replace(&mut config.server_no_context_takeover, true) {
229 return Err(BoxError::from_static_str(
230 "duplicate extension param: server_no_context_takeover",
231 ));
232 }
233 } else if part.eq_ignore_ascii_case("client_no_context_takeover") {
234 if std::mem::replace(&mut config.client_no_context_takeover, true) {
235 return Err(BoxError::from_static_str(
236 "duplicate extension param: client_no_context_takeover",
237 ));
238 }
239 } else if part.eq_ignore_ascii_case("server_max_window_bits") {
240 if config.server_max_window_bits.replace(0).is_some() {
241 return Err(BoxError::from_static_str(
242 "duplicate extension param: server_max_window_bits",
243 ));
244 }
245 } else if part.eq_ignore_ascii_case("client_max_window_bits") {
246 if config.client_max_window_bits.replace(0).is_some() {
247 return Err(BoxError::from_static_str(
248 "duplicate extension param: client_max_window_bits",
249 ));
250 }
251 } else if let Some((k, v)) = part.split_once('=') {
252 let k = k.trim();
253
254 let v = v.trim();
256 let v = v
257 .strip_prefix('"')
258 .and_then(|v| v.strip_suffix('"'))
259 .unwrap_or(v);
260 let v = v.trim();
261
262 if k.eq_ignore_ascii_case("server_max_window_bits") {
263 match v.trim().parse::<u8>() {
264 Ok(v) => {
265 if !(8..=15).contains(&v) {
266 tracing::debug!(
267 "fail per-message-deflate config value for server max windows bits: {v} not in [8,15] range"
268 );
269 return Err(BoxError::from_static_str(
270 "invalid server max windows bits (OOB)",
271 )
272 .context_field("value", v));
273 }
274 if config.server_max_window_bits.replace(v).is_some() {
275 return Err(BoxError::from_static_str(
276 "duplicate extension param: server_max_window_bits",
277 ));
278 }
279 }
280 Err(err) => {
281 tracing::debug!(
282 "fail per-message-deflate config with invalid value for server max windows bits: {k} = {v}; err = {err}"
283 );
284 return Err(err.context("invalid per-message-deflate config value for server max windows bits"));
285 }
286 }
287 } else if k.eq_ignore_ascii_case("client_max_window_bits") {
288 match v.trim().parse::<u8>() {
289 Ok(v) => {
290 if !(8..=15).contains(&v) {
291 tracing::debug!(
292 "fail per-message-deflate config value for client max windows bits: {v} not in [8,15] range"
293 );
294 return Err(BoxError::from_static_str(
295 "invalid client max windows bits (OOB)",
296 )
297 .context_field("value", v));
298 }
299 if config.client_max_window_bits.replace(v).is_some() {
300 return Err(BoxError::from_static_str(
301 "duplicate extension param: client_max_window_bits",
302 ));
303 }
304 }
305 Err(err) => {
306 tracing::debug!(
307 "fail per-message-deflate config with invalid value for client max windows bits: {k} = {v}; err = {err}"
308 );
309 return Err(err.context("invalid per-message-deflate config value for client max windows bits"));
310 }
311 }
312 } else {
313 tracing::debug!(
314 "fail per-message-deflate config with unknown permessage-deflate config parameter: {k} = {v}"
315 );
316 return Err(
317 BoxError::from_static_str("value not expected for given key")
318 .context_str_field("key", k)
319 .context_str_field("value", v),
320 );
321 }
322 } else {
323 tracing::debug!(
324 "received unknown permessage-deflate config parameter part: {part}"
325 );
326 return Err(BoxError::from_static_str(
327 "key not expected for permessage-deflate config",
328 )
329 .context_str_field("key", part));
330 }
331 }
332 Ok(Self::PerMessageDeflate(config))
333 } else if s.trim().is_empty() {
334 Ok(Self::Empty)
335 } else {
336 tracing::trace!(
337 "received unknown extension with identifier: {identifier} (full: {s}); store as unkown"
338 );
339 Ok(Self::Unknown(s.into()))
340 }
341 }
342}
343
344impl fmt::Display for Extension {
345 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
346 match self {
347 Self::PerMessageDeflate(config) => {
348 write!(f, "{}", config.identifier)?;
349 if config.server_no_context_takeover {
350 write!(f, "; server_no_context_takeover")?
351 }
352 if let Some(log) = config.server_max_window_bits {
353 if log == 0 {
354 write!(f, "; server_max_window_bits")?
355 } else {
356 write!(f, "; server_max_window_bits={log}")?
357 }
358 }
359 if config.client_no_context_takeover {
360 write!(f, "; client_no_context_takeover")?
361 }
362 if let Some(log) = config.client_max_window_bits {
363 if log == 0 {
364 write!(f, "; client_max_window_bits")?
365 } else {
366 write!(f, "; client_max_window_bits={log}")?
367 }
368 }
369 Ok(())
370 }
371 Self::Empty => Ok(()), Self::Unknown(v) => write!(f, "{v}"),
373 }
374 }
375}
376
377#[cfg(test)]
378mod tests {
379 use super::super::{test_decode, test_encode};
380 use super::{
381 Extension, PerMessageDeflateConfig, PerMessageDeflateIdentifier, SecWebSocketExtensions,
382 };
383
384 #[test]
385 fn decode_sec_websocket_extensions() {
386 for (name, input, expected_output) in [
387 (
389 "single extension",
390 vec!["permessage-deflate"],
391 Some(SecWebSocketExtensions::per_message_deflate()),
392 ),
393 (
394 "valueless client_max_window_bits",
395 vec!["permessage-deflate; client_max_window_bits"],
396 Some(SecWebSocketExtensions::per_message_deflate_with_config(
397 PerMessageDeflateConfig {
398 client_max_window_bits: Some(0), ..Default::default()
400 },
401 )),
402 ),
403 (
404 "x-webkit-deflate-frame identifier",
405 vec!["x-webkit-deflate-frame"],
406 Some(SecWebSocketExtensions::per_message_deflate_with_config(
407 PerMessageDeflateConfig {
408 identifier: super::PerMessageDeflateIdentifier::XWebKitDeflateFrame,
409 ..Default::default()
410 },
411 )),
412 ),
413 (
415 "client and server no context takeover",
416 vec!["permessage-deflate; client_no_context_takeover; server_no_context_takeover"],
417 Some(SecWebSocketExtensions::per_message_deflate_with_config(
418 PerMessageDeflateConfig {
419 client_no_context_takeover: true,
420 server_no_context_takeover: true,
421 ..Default::default()
422 },
423 )),
424 ),
425 (
426 "valued client and server max window bits",
427 vec!["permessage-deflate; client_max_window_bits=10; server_max_window_bits=11"],
428 Some(SecWebSocketExtensions::per_message_deflate_with_config(
429 PerMessageDeflateConfig {
430 client_max_window_bits: Some(10),
431 server_max_window_bits: Some(11),
432 ..Default::default()
433 },
434 )),
435 ),
436 (
437 "all parameters mixed",
438 vec![
439 "permessage-deflate; server_no_context_takeover; client_max_window_bits=12; client_no_context_takeover",
440 ],
441 Some(SecWebSocketExtensions::per_message_deflate_with_config(
442 PerMessageDeflateConfig {
443 server_no_context_takeover: true,
444 client_no_context_takeover: true,
445 client_max_window_bits: Some(12),
446 ..Default::default()
447 },
448 )),
449 ),
450 (
452 "multiple headers, duplicates allowed",
453 vec![
454 "permessage-deflate; client_no_context_takeover",
455 "x-webkit-deflate-frame",
456 ],
457 Some(
458 SecWebSocketExtensions::per_message_deflate_with_config(
459 PerMessageDeflateConfig {
460 client_no_context_takeover: true,
461 ..Default::default()
462 },
463 )
464 .with_extra_extension(Extension::PerMessageDeflate(
465 PerMessageDeflateConfig {
466 identifier: PerMessageDeflateIdentifier::XWebKitDeflateFrame,
467 ..Default::default()
468 },
469 )),
470 ),
471 ),
472 (
473 "multiple headers, preserve unknown extensions",
474 vec![
475 "unknown-extension, another-one",
476 "permessage-deflate; server_max_window_bits=14",
477 ],
478 Some(
479 SecWebSocketExtensions::new(Extension::Unknown("unknown-extension".into()))
480 .with_extra_extension(Extension::Unknown("another-one".into()))
481 .with_extra_extension(Extension::PerMessageDeflate(
482 PerMessageDeflateConfig {
483 server_max_window_bits: Some(14),
484 ..Default::default()
485 },
486 )),
487 ),
488 ),
489 (
491 "multiple headers, preserve unknown extensions",
492 vec![
493 "unknown-extension, another-one",
494 "permessage-deflate; server_max_window_bits=\"14\"",
495 ],
496 Some(
497 SecWebSocketExtensions::new(Extension::Unknown("unknown-extension".into()))
498 .with_extra_extension(Extension::Unknown("another-one".into()))
499 .with_extra_extension(Extension::PerMessageDeflate(
500 PerMessageDeflateConfig {
501 server_max_window_bits: Some(14),
502 ..Default::default()
503 },
504 )),
505 ),
506 ),
507 (
509 "leading/trailing whitespace",
510 vec![" permessage-deflate ; client_no_context_takeover "],
511 Some(SecWebSocketExtensions::per_message_deflate_with_config(
512 PerMessageDeflateConfig {
513 client_no_context_takeover: true,
514 ..Default::default()
515 },
516 )),
517 ),
518 (
519 "case-insensitive name and params",
520 vec!["PerMessage-Deflate; Client_No_Context_Takeover; SERVER_MAX_WINDOW_BITS=8"],
521 Some(SecWebSocketExtensions::per_message_deflate_with_config(
522 PerMessageDeflateConfig {
523 client_no_context_takeover: true,
524 server_max_window_bits: Some(8),
525 ..Default::default()
526 },
527 )),
528 ),
529 (
531 "invalid duplicate client_max_window_bits",
532 vec!["permessage-deflate; client_max_window_bits=15; client_max_window_bits=14"],
533 None,
534 ),
535 (
536 "invalid duplicate server_max_window_bits",
537 vec!["permessage-deflate; server_max_window_bits=15; server_max_window_bits=14"],
538 None,
539 ),
540 (
541 "invalid duplicate client_max_window_bits w/o value",
542 vec!["permessage-deflate; client_max_window_bits=15; client_max_window_bits"],
543 None,
544 ),
545 (
546 "invalid duplicate server_max_window_bits w/o value",
547 vec!["permessage-deflate; server_max_window_bits=15; server_max_window_bits"],
548 None,
549 ),
550 (
551 "invalid duplicate server_no_context_takeover",
552 vec![
553 "permessage-deflate; server_no_context_takeover; client_no_context_takeover; server_no_context_takeover",
554 ],
555 None,
556 ),
557 (
558 "invalid duplicate client_no_context_takeover",
559 vec![
560 "permessage-deflate; client_no_context_takeover; server_no_context_takeover; client_no_context_takeover",
561 ],
562 None,
563 ),
564 (
566 "empty header",
567 vec![""],
568 Some(SecWebSocketExtensions::new(Extension::Empty)),
569 ),
570 (
571 "whitespace only header",
572 vec![" "],
573 Some(SecWebSocketExtensions::new(Extension::Empty)),
574 ),
575 (
576 "unknown extension",
577 vec!["super-zip"],
578 Some(SecWebSocketExtensions::new(Extension::Unknown(
579 "super-zip".into(),
580 ))),
581 ),
582 (
584 "windows bits OOB: client: underflow",
585 vec!["permessage-deflate; client_max_window_bits=7"],
586 None,
587 ),
588 (
589 "windows bits OOB: client: overflow",
590 vec!["permessage-deflate; client_max_window_bits=16"],
591 None,
592 ),
593 (
594 "windows bits OOB: server: underflow",
595 vec!["permessage-deflate; server_max_window_bits=7"],
596 None,
597 ),
598 (
599 "windows bits OOB: server: overflow",
600 vec!["permessage-deflate; server_max_window_bits=16"],
601 None,
602 ),
603 (
604 "invalid parameter format",
605 vec!["permessage-deflate; client_max_window_bits_15"],
606 None,
607 ),
608 (
609 "invalid parameter value",
610 vec!["permessage-deflate; client_max_window_bits=abc"],
611 None,
612 ),
613 (
614 "parameter with empty value",
615 vec!["permessage-deflate; client_max_window_bits="],
616 None,
617 ),
618 (
620 "malformed header with comma",
621 vec!["permessage-deflate, client_max_window_bits"],
622 Some(
623 SecWebSocketExtensions::per_message_deflate()
624 .with_extra_extension(Extension::Unknown("client_max_window_bits".into())),
625 ),
626 ),
627 (
628 "multiple conflicting headers",
629 vec![
630 "permessage-deflate; client_max_window_bits=10",
631 "permessage-deflate; client_max_window_bits=11",
632 ],
633 Some(
634 SecWebSocketExtensions::per_message_deflate_with_config(
635 PerMessageDeflateConfig {
636 client_max_window_bits: Some(10),
637 ..Default::default()
638 },
639 )
640 .with_extra_extension(Extension::PerMessageDeflate(
641 PerMessageDeflateConfig {
642 client_max_window_bits: Some(11),
643 ..Default::default()
644 },
645 )),
646 ),
647 ),
648 ] {
649 assert_eq!(
650 test_decode::<SecWebSocketExtensions>(&input),
651 expected_output,
652 "Failed test case: {name}",
653 );
654 }
655 }
656
657 #[test]
658 fn encode_sec_websocket_extensions_extended() {
659 for (name, input, expected_output) in [
660 (
662 "default permessage-deflate",
663 SecWebSocketExtensions::per_message_deflate(),
664 "permessage-deflate",
665 ),
666 (
667 "valueless client_max_window_bits (chromium style)",
668 SecWebSocketExtensions::per_message_deflate_with_config(PerMessageDeflateConfig {
669 client_max_window_bits: Some(0), ..Default::default()
671 }),
672 "permessage-deflate; client_max_window_bits",
673 ),
674 (
675 "x-webkit-deflate-frame identifier (safari style)",
676 SecWebSocketExtensions::per_message_deflate_with_config(PerMessageDeflateConfig {
677 identifier: PerMessageDeflateIdentifier::XWebKitDeflateFrame,
678 ..Default::default()
679 }),
680 "x-webkit-deflate-frame",
681 ),
682 (
684 "client_no_context_takeover enabled",
685 SecWebSocketExtensions::per_message_deflate_with_config(PerMessageDeflateConfig {
686 client_no_context_takeover: true,
687 ..Default::default()
688 }),
689 "permessage-deflate; client_no_context_takeover",
690 ),
691 (
692 "server_no_context_takeover enabled",
693 SecWebSocketExtensions::per_message_deflate_with_config(PerMessageDeflateConfig {
694 server_no_context_takeover: true,
695 ..Default::default()
696 }),
697 "permessage-deflate; server_no_context_takeover",
698 ),
699 (
700 "both no_context_takeover flags enabled",
701 SecWebSocketExtensions::per_message_deflate_with_config(PerMessageDeflateConfig {
702 client_no_context_takeover: true,
703 server_no_context_takeover: true,
704 ..Default::default()
705 }),
706 "permessage-deflate; server_no_context_takeover; client_no_context_takeover",
709 ),
710 (
712 "specific client_max_window_bits",
713 SecWebSocketExtensions::per_message_deflate_with_config(PerMessageDeflateConfig {
714 client_max_window_bits: Some(12),
715 ..Default::default()
716 }),
717 "permessage-deflate; client_max_window_bits=12",
718 ),
719 (
720 "specific server_max_window_bits",
721 SecWebSocketExtensions::per_message_deflate_with_config(PerMessageDeflateConfig {
722 server_max_window_bits: Some(10),
723 ..Default::default()
724 }),
725 "permessage-deflate; server_max_window_bits=10",
726 ),
727 (
729 "all parameters configured",
730 SecWebSocketExtensions::per_message_deflate_with_config(PerMessageDeflateConfig {
731 client_no_context_takeover: true,
732 server_no_context_takeover: true,
733 client_max_window_bits: Some(15),
734 server_max_window_bits: Some(15),
735 ..Default::default()
736 }),
737 "permessage-deflate; server_no_context_takeover; server_max_window_bits=15; client_no_context_takeover; client_max_window_bits=15",
739 ),
740 (
741 "mixed valued and boolean parameters",
742 SecWebSocketExtensions::per_message_deflate_with_config(PerMessageDeflateConfig {
743 client_no_context_takeover: true,
744 server_max_window_bits: Some(11),
745 ..Default::default()
746 }),
747 "permessage-deflate; server_max_window_bits=11; client_no_context_takeover",
748 ),
749 (
750 "webkit identifier with parameters",
751 SecWebSocketExtensions::per_message_deflate_with_config(PerMessageDeflateConfig {
752 identifier: PerMessageDeflateIdentifier::XWebKitDeflateFrame,
753 client_max_window_bits: Some(10),
754 client_no_context_takeover: true,
755 ..Default::default()
756 }),
757 "x-webkit-deflate-frame; client_no_context_takeover; client_max_window_bits=10",
758 ),
759 ] {
760 let headers = test_encode(input);
761 assert_eq!(
762 headers["sec-websocket-extensions"], expected_output,
763 "Failed test case: {name}",
764 );
765 }
766 }
767}