Skip to main content

ntex/ws/
cfg.rs

1use std::{fmt, net};
2
3use base64::{Engine, engine::general_purpose::STANDARD as base64};
4#[cfg(feature = "cookie")]
5use coo_kie::{Cookie, CookieJar};
6
7use crate::http::error::HttpError;
8use crate::http::header::{
9    self, AUTHORIZATION, HeaderMap, HeaderName, HeaderValue, InvalidHeaderValue,
10};
11use crate::service::cfg::{CfgContext, Configuration};
12use crate::time::Millis;
13
14/// Configuration for a WebSocket client connection.
15///
16/// Store this value in [`SharedCfg`](crate::SharedCfg) and pass the resulting
17/// configuration to [`WsClient::new`](super::WsClient::new).
18#[derive(Debug)]
19pub struct WsClientConfig {
20    pub(super) addr: Option<net::SocketAddr>,
21    pub(super) max_size: usize,
22    pub(super) timeout: Millis,
23    pub(super) close_timeout: Millis,
24    pub(super) headers: HeaderMap,
25    pub(super) server_mode: bool,
26    #[cfg(feature = "cookie")]
27    pub(super) cookies: Option<CookieJar>,
28
29    config: CfgContext,
30}
31
32impl Default for WsClientConfig {
33    fn default() -> Self {
34        Self::new()
35    }
36}
37
38impl Configuration for WsClientConfig {
39    const NAME: &str = "WebSocket client configuration";
40
41    fn ctx(&self) -> &CfgContext {
42        &self.config
43    }
44
45    fn set_ctx(&mut self, ctx: CfgContext) {
46        self.config = ctx;
47    }
48}
49
50impl WsClientConfig {
51    #[must_use]
52    /// Creates a WebSocket client configuration with default values.
53    pub fn new() -> WsClientConfig {
54        let mut headers = HeaderMap::new();
55        headers.insert(header::UPGRADE, HeaderValue::from_static("websocket"));
56        headers.insert(
57            header::SEC_WEBSOCKET_VERSION,
58            HeaderValue::from_static("13"),
59        );
60
61        WsClientConfig {
62            headers,
63            addr: None,
64            max_size: 65_536,
65            server_mode: false,
66            timeout: Millis(5_000),
67            close_timeout: Millis(5_000),
68            #[cfg(feature = "cookie")]
69            cookies: None,
70            config: CfgContext::default(),
71        }
72    }
73
74    #[must_use]
75    /// Sets the server socket address.
76    ///
77    /// This address is used instead of resolving the URI host name.
78    pub fn set_address(mut self, addr: net::SocketAddr) -> Self {
79        self.addr = Some(addr);
80        self
81    }
82
83    /// Sets the WebSocket subprotocols offered to the server.
84    ///
85    /// This replaces the current `Sec-WebSocket-Protocol` header. An empty
86    /// iterator removes the header.
87    ///
88    /// # Errors
89    ///
90    /// Returns [`HttpError`] if a protocol is not a valid HTTP token or the
91    /// resulting list is not a valid HTTP header value.
92    pub fn set_protocols<U, V>(mut self, protos: U) -> Result<Self, HttpError>
93    where
94        U: IntoIterator<Item = V>,
95        V: AsRef<str>,
96    {
97        let mut values = Vec::new();
98        for proto in protos {
99            let proto = proto.as_ref();
100            if !is_token(proto) {
101                return Err(InvalidHeaderValue::default().into());
102            }
103            values.push(proto.to_owned());
104        }
105        let protos = values.join(",");
106
107        if protos.is_empty() {
108            self.headers.remove(header::SEC_WEBSOCKET_PROTOCOL);
109        } else {
110            self.headers.insert(
111                header::SEC_WEBSOCKET_PROTOCOL,
112                HeaderValue::try_from(protos.as_str())?,
113            );
114        }
115        Ok(self)
116    }
117
118    #[must_use]
119    #[cfg(feature = "cookie")]
120    /// Adds a cookie to the opening handshake.
121    pub fn set_cookie<C>(mut self, cookie: C) -> Self
122    where
123        C: Into<Cookie<'static>>,
124    {
125        if let Some(cookies) = &mut self.cookies {
126            cookies.add(cookie.into());
127        } else {
128            let mut jar = CookieJar::new();
129            jar.add(cookie.into());
130            self.cookies = Some(jar);
131        }
132        self
133    }
134
135    /// Sets the `Origin` header for the opening handshake.
136    pub fn set_origin<V, E>(mut self, origin: V) -> Result<Self, HttpError>
137    where
138        HeaderValue: TryFrom<V, Error = E>,
139        HttpError: From<E>,
140    {
141        self.headers
142            .insert(header::ORIGIN, HeaderValue::try_from(origin)?);
143        Ok(self)
144    }
145
146    #[must_use]
147    /// Sets the maximum accepted frame payload size.
148    ///
149    /// The default is 64 KiB.
150    pub fn set_max_frame_size(mut self, size: usize) -> Self {
151        self.max_size = size;
152        self
153    }
154
155    #[must_use]
156    /// Configures the connection to use server-side masking rules.
157    ///
158    /// By default, the client masks outgoing frames and expects unmasked
159    /// incoming frames. Server mode reverses those rules.
160    pub fn set_server_mode(mut self) -> Self {
161        self.server_mode = true;
162        self
163    }
164
165    /// Sets a header for the opening handshake.
166    ///
167    /// This replaces any existing value with the same name.
168    pub fn set_header<K, V>(mut self, key: K, value: V) -> Result<Self, HttpError>
169    where
170        HeaderName: TryFrom<K>,
171        HeaderValue: TryFrom<V>,
172        <HeaderName as TryFrom<K>>::Error: Into<HttpError>,
173        <HeaderValue as TryFrom<V>>::Error: Into<HttpError>,
174    {
175        let key = HeaderName::try_from(key).map_err(Into::into)?;
176        let value = HeaderValue::try_from(value).map_err(Into::into)?;
177        self.headers.insert(key, value);
178        Ok(self)
179    }
180
181    /// Sets a handshake header if it is not already present.
182    pub fn set_header_if_none<K, V>(mut self, key: K, value: V) -> Result<Self, HttpError>
183    where
184        HeaderName: TryFrom<K>,
185        HeaderValue: TryFrom<V>,
186        <HeaderName as TryFrom<K>>::Error: Into<HttpError>,
187        <HeaderValue as TryFrom<V>>::Error: Into<HttpError>,
188    {
189        let key = HeaderName::try_from(key).map_err(Into::into)?;
190        if !self.headers.contains_key(&key) {
191            self.headers
192                .insert(key, HeaderValue::try_from(value).map_err(Into::into)?);
193        }
194        Ok(self)
195    }
196
197    /// Sets the HTTP Basic authentication header.
198    pub fn set_basic_auth(
199        self,
200        username: impl fmt::Display,
201        password: Option<&str>,
202    ) -> Result<Self, HttpError> {
203        let auth = match password {
204            Some(password) => format!("{username}:{password}"),
205            None => format!("{username}:"),
206        };
207        self.set_header(AUTHORIZATION, format!("Basic {}", base64.encode(auth)))
208    }
209
210    /// Sets the HTTP bearer authentication header.
211    pub fn set_bearer_auth(self, token: impl fmt::Display) -> Result<Self, HttpError> {
212        self.set_header(AUTHORIZATION, format!("Bearer {token}"))
213    }
214
215    #[must_use]
216    /// Sets the opening-handshake timeout.
217    ///
218    /// The timeout covers sending the upgrade request and receiving the
219    /// response after a connection has been established. The default is
220    /// 5 seconds. A zero duration disables the timeout.
221    pub fn set_handshake_timeout(mut self, timeout: impl Into<Millis>) -> Self {
222        self.timeout = timeout.into();
223        self
224    }
225
226    #[must_use]
227    /// Sets the closing-handshake timeout.
228    ///
229    /// After sending a close frame, the client waits this long for the peer's
230    /// close response before shutting down the connection. The default is
231    /// 5 seconds. A zero duration disables the timeout.
232    pub fn set_close_timeout(mut self, timeout: impl Into<Millis>) -> Self {
233        self.close_timeout = timeout.into();
234        self
235    }
236}
237
238pub(crate) fn is_token(value: &str) -> bool {
239    !value.is_empty()
240        && value.bytes().all(|byte| {
241            byte.is_ascii_alphanumeric()
242                || matches!(
243                    byte,
244                    b'!' | b'#'
245                        | b'$'
246                        | b'%'
247                        | b'&'
248                        | b'\''
249                        | b'*'
250                        | b'+'
251                        | b'-'
252                        | b'.'
253                        | b'^'
254                        | b'_'
255                        | b'`'
256                        | b'|'
257                        | b'~'
258                )
259        })
260}