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#[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 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 pub fn set_address(mut self, addr: net::SocketAddr) -> Self {
79 self.addr = Some(addr);
80 self
81 }
82
83 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 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 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 pub fn set_max_frame_size(mut self, size: usize) -> Self {
151 self.max_size = size;
152 self
153 }
154
155 #[must_use]
156 pub fn set_server_mode(mut self) -> Self {
161 self.server_mode = true;
162 self
163 }
164
165 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 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 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 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 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 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}