Skip to main content

rspamd_client/backend/
async_client.rs

1use 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
55// Temporary structure for making a request
56pub 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            // Check if File header is present - if so, we don't need to send the body
76            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/// Scan an email asynchronously, returning the parsed reply or error.
219/// Example:
220/// ```rust
221/// use rspamd_client::config::Config;
222/// use rspamd_client::scan_async;
223/// use rspamd_client::error::RspamdError;
224/// use bytes::Bytes;
225/// use std::str::FromStr;
226///
227///	#[tokio::main]
228/// async fn main() -> Result<(), RspamdError> {
229/// 	let config = Config::builder()
230/// 		.base_url("http://localhost:11333".to_string())
231/// 		.build();
232/// 	let envelope = Default::default();
233/// 	let email = "...";
234/// 	let response = scan_async(&config, email, envelope).await?;
235/// 	Ok(())
236/// }
237/// ```
238#[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    // Check for Message-Offset header to handle body_block feature
252    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            // Split body into JSON part and rewritten body part
261            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            // Offset is out of bounds, parse entire body as JSON
269            serde_json::from_slice::<RspamdScanReply>(body.as_ref())?
270        }
271    } else {
272        // No Message-Offset header, parse entire body as JSON
273        serde_json::from_slice::<RspamdScanReply>(body.as_ref())?
274    };
275
276    Ok(response)
277}