rspamd_client/backend/
async_client.rs1use crate::backend::traits::*;
2use crate::config::{Config, EnvelopeData};
3use crate::error::RspamdError;
4use crate::protocol::commands::{RspamdCommand, RspamdEndpoint};
5use crate::protocol::encryption::{httpcrypt_decrypt, httpcrypt_encrypt, make_key_header};
6use crate::protocol::RspamdScanReply;
7use bytes::{Bytes, BytesMut};
8use reqwest::header::{HeaderName, HeaderValue};
9use reqwest::Client;
10use std::str::FromStr;
11use std::time::Duration;
12use url::Url;
13use zstd::zstd_safe::WriteBuf;
14
15pub struct AsyncClient<'a> {
16 config: &'a Config,
17 inner: Client,
18}
19
20#[cfg(feature = "async")]
21pub fn async_client(options: &Config) -> Result<AsyncClient<'_>, RspamdError> {
22 let client = Client::builder().timeout(Duration::from_secs_f64(options.timeout));
23
24 let client = if let Some(ref proxy) = options.proxy_config {
25 let proxy = reqwest::Proxy::all(proxy.proxy_url.clone())
26 .map_err(|e| RspamdError::HttpError(e.to_string()))?;
27 client.proxy(proxy)
28 } else {
29 client
30 };
31 let client = if let Some(ref tls) = options.tls_settings {
32 if let Some(ca_path) = tls.ca_path.as_ref() {
33 client.add_root_certificate(
34 reqwest::Certificate::from_pem(
35 &std::fs::read(std::fs::canonicalize(ca_path.as_str()).unwrap())
36 .map_err(|e| RspamdError::ConfigError(e.to_string()))?,
37 )
38 .map_err(|e| RspamdError::HttpError(e.to_string()))?,
39 )
40 } else {
41 client
42 }
43 } else {
44 client
45 };
46
47 Ok(AsyncClient {
48 inner: client
49 .build()
50 .map_err(|e| RspamdError::HttpError(e.to_string()))?,
51 config: options,
52 })
53}
54
55pub struct ReqwestRequest<'a, B> {
57 endpoint: RspamdEndpoint<'a>,
58 client: AsyncClient<'a>,
59 body: B,
60 envelope_data: Option<EnvelopeData>,
61}
62
63#[maybe_async::maybe_async]
64impl<'a, B: AsRef<[u8]> + Send> Request for ReqwestRequest<'a, B> {
65 type Body = Bytes;
66 type HeaderMap = reqwest::header::HeaderMap;
67
68 async fn response(mut self) -> Result<(Self::HeaderMap, Self::Body), RspamdError> {
69 let mut retry_cnt = self.client.config.retries;
70 let mut maybe_sk = Default::default();
71 let extra_hdrs: Vec<(String, String)> =
72 self.envelope_data.take().unwrap().into_iter().collect();
73
74 let response = loop {
75 let has_file_header = extra_hdrs.iter().any(|(k, _)| k == "File");
77 let need_body = self.endpoint.need_body && !has_file_header;
78 let method = if need_body {
79 reqwest::Method::POST
80 } else {
81 reqwest::Method::GET
82 };
83
84 let mut url = Url::from_str(self.client.config.base_url.as_str())
85 .map_err(|e| RspamdError::HttpError(e.to_string()))?;
86 url.set_path(self.endpoint.url);
87 let mut req = self.client.inner.request(method, url.clone());
88
89 if let Some(ref password) = self.client.config.password {
90 req = req.header("Password", password);
91 }
92
93 if self.client.config.zstd && need_body {
94 req = req.header("Content-Encoding", "zstd");
95 req = req.header("Compression", "zstd");
96 }
97
98 for (k, v) in extra_hdrs.iter() {
99 req = req.header(k, v);
100 }
101
102 if let Some(ref encryption_key) = self.client.config.encryption_key {
103 let inner_req = req
104 .build()
105 .map_err(|e| RspamdError::HttpError(e.to_string()))?;
106 let body = if need_body {
107 if self.client.config.zstd {
108 zstd::encode_all(self.body.as_ref(), 0)?
109 } else {
110 self.body.as_ref().to_vec()
111 }
112 } else {
113 Vec::new()
114 };
115 let encrypted = httpcrypt_encrypt(
116 url.path(),
117 body.as_slice(),
118 inner_req.headers(),
119 encryption_key.as_bytes(),
120 )?;
121 req = self.client.inner.request(reqwest::Method::POST, url);
122 let key_header =
123 make_key_header(encryption_key.as_str(), encrypted.peer_key.as_str())?;
124 req = req.header("Key", key_header);
125 req = req.body(encrypted.body);
126 maybe_sk = Some(encrypted.shared_key);
127 } else if need_body {
128 req = if self.client.config.zstd {
129 req.body(reqwest::Body::from(zstd::encode_all(
130 self.body.as_ref(),
131 0,
132 )?))
133 } else {
134 req.body(Bytes::copy_from_slice(self.body.as_ref()))
135 };
136 }
137
138 let req = req.timeout(Duration::from_secs_f64(self.client.config.timeout));
139 let req = req
140 .build()
141 .map_err(|e| RspamdError::HttpError(e.to_string()))?;
142
143 match self.client.inner.execute(req).await {
144 Ok(v) => break Ok(v),
145 Err(e) => {
146 if (retry_cnt - 1) == 0 {
147 break Err(e);
148 }
149 retry_cnt -= 1;
150 let delay = Duration::from_secs_f64(self.client.config.timeout);
151 tokio::time::sleep(delay).await;
152 continue;
153 }
154 };
155 }
156 .map_err(|e| RspamdError::HttpError(e.to_string()))?;
157
158 if !response.status().is_success() {
159 return Err(RspamdError::HttpError(format!(
160 "Status: {}",
161 response.status()
162 )));
163 }
164
165 if let Some(sk) = maybe_sk {
166 let mut body = BytesMut::from(
167 response
168 .bytes()
169 .await
170 .map_err(|e| RspamdError::HttpError(e.to_string()))?,
171 );
172 let decrypted_offset = httpcrypt_decrypt(body.as_mut(), sk)?;
173 let mut hdrs = [httparse::EMPTY_HEADER; 64];
174 let mut parsed = httparse::Response::new(&mut hdrs);
175
176 let body_offset = parsed
177 .parse(&body.as_slice()[decrypted_offset..])
178 .map_err(|s| RspamdError::HttpError(s.to_string()))?;
179 let mut output_hdrs = reqwest::header::HeaderMap::with_capacity(parsed.headers.len());
180 for hdr in parsed.headers.iter_mut() {
181 output_hdrs.insert(
182 HeaderName::from_str(hdr.name)?,
183 HeaderValue::from_str(std::str::from_utf8(hdr.value)?)?,
184 );
185 }
186 let body = if output_hdrs
187 .get("Compression")
188 .is_some_and(|hv| hv == "zstd")
189 {
190 zstd::decode_all(&body.as_slice()[body_offset.unwrap() + decrypted_offset..])?
191 } else {
192 body.as_slice()[body_offset.unwrap() + decrypted_offset..].to_vec()
193 };
194 Ok((output_hdrs, body.into()))
195 } else {
196 Ok((response.headers().clone(), response.bytes().await?))
197 }
198 }
199}
200
201#[maybe_async::maybe_async]
202impl<'a, B: AsRef<[u8]> + Send> ReqwestRequest<'a, B> {
203 pub async fn new(
204 client: AsyncClient<'a>,
205 body: B,
206 command: RspamdCommand,
207 envelope_data: EnvelopeData,
208 ) -> Result<ReqwestRequest<'a, B>, RspamdError> {
209 Ok(Self {
210 endpoint: RspamdEndpoint::from_command(command),
211 client,
212 body,
213 envelope_data: Some(envelope_data),
214 })
215 }
216}
217
218#[maybe_async::maybe_async]
239pub async fn scan_async<B: AsRef<[u8]> + Send>(
240 options: &Config,
241 body: B,
242 envelope_data: EnvelopeData,
243) -> Result<RspamdScanReply, RspamdError> {
244 let client = async_client(options)?;
245 let request = ReqwestRequest::new(client, body, RspamdCommand::Scan, envelope_data).await?;
246 let (headers, body) = request
247 .response()
248 .await
249 .map_err(|e| RspamdError::HttpError(e.to_string()))?;
250
251 let response = if let Some(offset_header) = headers.get("Message-Offset") {
253 let offset = offset_header
254 .to_str()
255 .map_err(|e| RspamdError::HttpError(format!("Invalid Message-Offset header: {}", e)))?
256 .parse::<usize>()
257 .map_err(|e| RspamdError::HttpError(format!("Invalid Message-Offset value: {}", e)))?;
258
259 if offset < body.len() {
260 let json_part = &body[..offset];
262 let body_part = &body[offset..];
263
264 let mut response = serde_json::from_slice::<RspamdScanReply>(json_part)?;
265 response.rewritten_body = Some(body_part.to_vec());
266 response
267 } else {
268 serde_json::from_slice::<RspamdScanReply>(body.as_ref())?
270 }
271 } else {
272 serde_json::from_slice::<RspamdScanReply>(body.as_ref())?
274 };
275
276 Ok(response)
277}