Skip to main content

toolcraft_request/
client.rs

1use futures_util::StreamExt;
2use reqwest::{Client, multipart};
3use url::Url;
4
5use crate::{
6    error::{Error, Result},
7    header_map::HeaderMap,
8    response::{ByteStream, Response},
9};
10
11/// An HTTP request builder and executor with base URL and default headers.
12#[derive(Debug)]
13pub struct Request {
14    client: Client,
15    base_url: Option<Url>,
16    default_headers: HeaderMap,
17}
18
19impl Request {
20    /// Create a new Request client.
21    pub fn new() -> Result<Self> {
22        let client = Client::builder()
23            .build()
24            .map_err(|e| Error::ErrorMessage(e.to_string().into()))?;
25        Ok(Request {
26            client,
27            base_url: None,
28            default_headers: HeaderMap::new(),
29        })
30    }
31
32    pub fn with_timeout(timeout_sec: u64) -> Result<Self> {
33        let client = Client::builder()
34            .timeout(std::time::Duration::from_secs(timeout_sec))
35            .build()
36            .map_err(|e| Error::ErrorMessage(e.to_string().into()))?;
37        Ok(Request {
38            client,
39            base_url: None,
40            default_headers: HeaderMap::new(),
41        })
42    }
43
44    /// Set the base URL for all requests.
45    pub fn set_base_url(&mut self, base_url: &str) -> Result<()> {
46        let mut url_str = base_url.to_string();
47        if !url_str.ends_with('/') {
48            url_str.push('/');
49        }
50        let url = Url::parse(&url_str)?;
51        self.base_url = Some(url);
52        Ok(())
53    }
54
55    /// Set default headers to be applied on all requests.
56    pub fn set_default_headers(&mut self, headers: HeaderMap) {
57        self.default_headers = headers;
58    }
59
60    /// Send a GET request.
61    pub async fn get(
62        &self,
63        endpoint: &str,
64        query: Option<Vec<(String, String)>>,
65        headers: Option<HeaderMap>,
66    ) -> Result<Response> {
67        let url = self.build_url(endpoint, query)?;
68        let mut request = self.client.get(url.as_str());
69
70        let mut combined_headers = self.default_headers.clone();
71        if let Some(custom_headers) = headers {
72            combined_headers.merge(custom_headers);
73        }
74        request = request.headers(combined_headers.inner().clone());
75
76        let response = request.send().await?;
77        Ok(response.into())
78    }
79
80    /// Send a POST request with JSON body.
81    pub async fn post(
82        &self,
83        endpoint: &str,
84        body: &serde_json::Value,
85        headers: Option<HeaderMap>,
86    ) -> Result<Response> {
87        let url = self.build_url(endpoint, None)?;
88        let mut request = self.client.post(url).json(body);
89
90        let mut combined_headers = self.default_headers.clone();
91        if let Some(custom_headers) = headers {
92            combined_headers.merge(custom_headers);
93        }
94        request = request.headers(combined_headers.inner().clone());
95
96        let response = request.send().await?;
97        Ok(response.into())
98    }
99
100    /// Send a PUT request with JSON body.
101    pub async fn put(
102        &self,
103        endpoint: &str,
104        body: &serde_json::Value,
105        headers: Option<HeaderMap>,
106    ) -> Result<Response> {
107        let url = self.build_url(endpoint, None)?;
108        let mut request = self.client.put(url).json(body);
109
110        let mut combined_headers = self.default_headers.clone();
111        if let Some(custom_headers) = headers {
112            combined_headers.merge(custom_headers);
113        }
114        request = request.headers(combined_headers.inner().clone());
115
116        let response = request.send().await?;
117        Ok(response.into())
118    }
119
120    /// Send a PUT request with raw bytes body.
121    pub async fn put_bytes(
122        &self,
123        endpoint: &str,
124        body: impl Into<bytes::Bytes>,
125        headers: Option<HeaderMap>,
126    ) -> Result<Response> {
127        let url = self.build_url(endpoint, None)?;
128        let mut request = self.client.put(url).body(body.into());
129
130        let mut combined_headers = self.default_headers.clone();
131        if let Some(custom_headers) = headers {
132            combined_headers.merge(custom_headers);
133        }
134        request = request.headers(combined_headers.inner().clone());
135
136        let response = request.send().await?;
137        Ok(response.into())
138    }
139
140    /// Send a DELETE request.
141    pub async fn delete(&self, endpoint: &str, headers: Option<HeaderMap>) -> Result<Response> {
142        let url = self.build_url(endpoint, None)?;
143        let mut request = self.client.delete(url);
144
145        let mut combined_headers = self.default_headers.clone();
146        if let Some(custom_headers) = headers {
147            combined_headers.merge(custom_headers);
148        }
149        request = request.headers(combined_headers.inner().clone());
150
151        let response = request.send().await?;
152        Ok(response.into())
153    }
154
155    /// Send a HEAD request.
156    pub async fn head(&self, endpoint: &str, headers: Option<HeaderMap>) -> Result<Response> {
157        let url = self.build_url(endpoint, None)?;
158        let mut request = self.client.head(url);
159
160        let mut combined_headers = self.default_headers.clone();
161        if let Some(custom_headers) = headers {
162            combined_headers.merge(custom_headers);
163        }
164        request = request.headers(combined_headers.inner().clone());
165
166        let response = request.send().await?;
167        Ok(response.into())
168    }
169
170    /// Send a POST request with multipart/form-data.
171    ///
172    /// # Arguments
173    /// * `endpoint` - The URL endpoint
174    /// * `form_fields` - Vector of form fields (text or file)
175    /// * `headers` - Optional custom headers
176    ///
177    /// # Important
178    /// The `Content-Type` header will be automatically removed from default and custom headers
179    /// to allow reqwest to set the correct `multipart/form-data` with boundary.
180    ///
181    /// # Example
182    /// ```no_run
183    /// use toolcraft_request::{FormField, Request};
184    ///
185    /// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
186    /// let client = Request::new()?;
187    /// let fields = vec![
188    ///     FormField::text("name", "John"),
189    ///     FormField::file("avatar", "/path/to/image.jpg").await?,
190    /// ];
191    /// let response = client.post_form("/upload", fields, None).await?;
192    /// # Ok(())
193    /// # }
194    /// ```
195    pub async fn post_form(
196        &self,
197        endpoint: &str,
198        form_fields: Vec<FormField>,
199        headers: Option<HeaderMap>,
200    ) -> Result<Response> {
201        let url = self.build_url(endpoint, None)?;
202
203        let mut form = multipart::Form::new();
204        for field in form_fields {
205            match field {
206                FormField::Text { name, value } => {
207                    form = form.text(name, value);
208                }
209                FormField::File {
210                    name,
211                    filename,
212                    content,
213                } => {
214                    let part = multipart::Part::bytes(content).file_name(filename);
215                    form = form.part(name, part);
216                }
217            }
218        }
219
220        let mut combined_headers = self.default_headers.clone();
221        if let Some(custom_headers) = headers {
222            combined_headers.merge(custom_headers);
223        }
224
225        // Remove Content-Type to let reqwest set the correct multipart/form-data with boundary
226        combined_headers.remove("Content-Type");
227        combined_headers.remove("content-type");
228
229        let mut request = self.client.post(url).multipart(form);
230        request = request.headers(combined_headers.inner().clone());
231
232        let response = request.send().await?;
233        Ok(response.into())
234    }
235
236    /// Send a streaming POST request and return the response stream.
237    pub async fn post_stream(
238        &self,
239        endpoint: &str,
240        body: &serde_json::Value,
241        headers: Option<HeaderMap>,
242    ) -> Result<ByteStream> {
243        let url = self.build_url(endpoint, None)?;
244        let mut request = self.client.post(url).json(body);
245
246        let mut combined_headers = self.default_headers.clone();
247        if let Some(custom_headers) = headers {
248            combined_headers.merge(custom_headers);
249        }
250        request = request.headers(combined_headers.inner().clone());
251
252        let response = request.send().await?;
253        if !response.status().is_success() {
254            return Err(Error::ErrorMessage(
255                format!("Unexpected status: {}", response.status()).into(),
256            ));
257        }
258
259        let stream = response
260            .bytes_stream()
261            .map(|chunk_result| chunk_result.map_err(Error::from));
262        Ok(Box::pin(stream))
263    }
264
265    /// Build a full URL by combining base URL, endpoint, and optional query parameters.
266    fn build_url(&self, endpoint: &str, query: Option<Vec<(String, String)>>) -> Result<Url> {
267        let mut url = if let Some(base_url) = &self.base_url {
268            base_url.join(endpoint)?
269        } else {
270            Url::parse(endpoint)?
271        };
272
273        if let Some(query_params) = query {
274            let query_pairs: Vec<(String, String)> = query_params.into_iter().collect();
275            url.query_pairs_mut().extend_pairs(query_pairs);
276        }
277
278        Ok(url)
279    }
280}
281
282/// Parse a full URL with optional query parameters.
283pub fn parse_url(url: &str, query: Option<Vec<(String, String)>>) -> Result<Url> {
284    let mut url = Url::parse(url)?;
285    if let Some(query_params) = query {
286        let query_pairs: Vec<(String, String)> = query_params.into_iter().collect();
287        url.query_pairs_mut().extend_pairs(query_pairs);
288    }
289    Ok(url)
290}
291
292/// Represents a field in a multipart/form-data request.
293#[derive(Debug, Clone)]
294pub enum FormField {
295    /// A text field.
296    Text { name: String, value: String },
297    /// A file field.
298    File {
299        name: String,
300        filename: String,
301        content: Vec<u8>,
302    },
303}
304
305impl FormField {
306    /// Create a text field.
307    ///
308    /// # Example
309    /// ```
310    /// use toolcraft_request::FormField;
311    /// let field = FormField::text("username", "john_doe");
312    /// ```
313    pub fn text(name: impl Into<String>, value: impl Into<String>) -> Self {
314        FormField::Text {
315            name: name.into(),
316            value: value.into(),
317        }
318    }
319
320    /// Create a file field from bytes.
321    ///
322    /// # Example
323    /// ```
324    /// use toolcraft_request::FormField;
325    /// let data = b"file content".to_vec();
326    /// let field = FormField::file_from_bytes("avatar", "photo.jpg", data);
327    /// ```
328    pub fn file_from_bytes(
329        name: impl Into<String>,
330        filename: impl Into<String>,
331        content: Vec<u8>,
332    ) -> Self {
333        FormField::File {
334            name: name.into(),
335            filename: filename.into(),
336            content,
337        }
338    }
339
340    /// Create a file field by reading from a file path.
341    ///
342    /// # Example
343    /// ```no_run
344    /// use toolcraft_request::FormField;
345    ///
346    /// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
347    /// let field = FormField::file("avatar", "/path/to/image.jpg").await?;
348    /// # Ok(())
349    /// # }
350    /// ```
351    pub async fn file(name: impl Into<String>, path: impl AsRef<std::path::Path>) -> Result<Self> {
352        let path = path.as_ref();
353        let filename = path
354            .file_name()
355            .and_then(|n| n.to_str())
356            .ok_or_else(|| Error::ErrorMessage("Invalid file path".into()))?
357            .to_string();
358
359        let content = tokio::fs::read(path)
360            .await
361            .map_err(|e| Error::ErrorMessage(format!("Failed to read file: {}", e).into()))?;
362
363        Ok(FormField::File {
364            name: name.into(),
365            filename,
366            content,
367        })
368    }
369}