Skip to main content

pidge_client/graph/
batch.rs

1//! Graph `$batch`: up to 20 sub-requests per round-trip.
2//!
3//! Used for bulk mutations (mark-read, move, delete, categorize) where pidge
4//! previously issued one HTTP call per message. Sub-request throttling
5//! (per-item 429 with a `retryAfter` hint) is honored by re-batching the
6//! throttled subset.
7
8use serde::{Deserialize, Serialize};
9use serde_json::Value;
10
11use crate::error::ClientError;
12
13pub const MAX_BATCH: usize = 20;
14const MAX_ROUNDS: usize = 4;
15
16#[derive(Debug, Clone, Serialize)]
17pub struct BatchRequest {
18    pub id: String,
19    pub method: &'static str,
20    pub url: String,
21    #[serde(skip_serializing_if = "Option::is_none")]
22    pub body: Option<Value>,
23    #[serde(skip_serializing_if = "Option::is_none")]
24    pub headers: Option<serde_json::Map<String, Value>>,
25}
26
27impl BatchRequest {
28    /// A JSON-bodied request (adds the Content-Type header Graph requires).
29    pub fn json(
30        id: impl Into<String>,
31        method: &'static str,
32        url: impl Into<String>,
33        body: Value,
34    ) -> Self {
35        let mut headers = serde_json::Map::new();
36        headers.insert(
37            "Content-Type".into(),
38            Value::String("application/json".into()),
39        );
40        Self {
41            id: id.into(),
42            method,
43            url: url.into(),
44            body: Some(body),
45            headers: Some(headers),
46        }
47    }
48
49    /// A body-less request (POST /move-style actions take json; DELETE takes none).
50    pub fn bare(id: impl Into<String>, method: &'static str, url: impl Into<String>) -> Self {
51        Self {
52            id: id.into(),
53            method,
54            url: url.into(),
55            body: None,
56            headers: None,
57        }
58    }
59}
60
61#[derive(Debug, Clone, Deserialize)]
62pub struct BatchResponse {
63    pub id: String,
64    pub status: u16,
65    #[serde(default)]
66    pub body: Value,
67}
68
69impl BatchResponse {
70    pub fn is_success(&self) -> bool {
71        (200..300).contains(&self.status)
72    }
73
74    fn retry_after(&self) -> Option<u64> {
75        self.body
76            .pointer("/error/retryAfterSeconds")
77            .and_then(Value::as_str)
78            .and_then(|s| s.parse().ok())
79            .or_else(|| {
80                self.body
81                    .pointer("/error/retryAfterSeconds")
82                    .and_then(Value::as_u64)
83            })
84    }
85}
86
87#[derive(Debug, Deserialize)]
88struct BatchEnvelope {
89    responses: Vec<BatchResponse>,
90}
91
92/// Execute all `requests` (chunked ≤20), re-batching per-item 429s up to
93/// [`MAX_ROUNDS`] times. Returns one response per request id (order not
94/// guaranteed; correlate by id). Items still throttled after the final
95/// round are returned with their last 429 response.
96pub async fn batch_all(
97    http: &reqwest::Client,
98    base_url: &str,
99    access_token: &str,
100    requests: Vec<BatchRequest>,
101) -> Result<Vec<BatchResponse>, ClientError> {
102    let mut pending = requests;
103    let mut done: Vec<BatchResponse> = Vec::new();
104
105    for round in 0..MAX_ROUNDS {
106        if pending.is_empty() {
107            break;
108        }
109        let mut next_pending: Vec<BatchRequest> = Vec::new();
110        let mut max_retry_after = 0u64;
111
112        for chunk in pending.chunks(MAX_BATCH) {
113            let payload = serde_json::json!({ "requests": chunk });
114            let req = http
115                .post(format!("{base_url}/$batch"))
116                .bearer_auth(access_token)
117                .json(&payload);
118            let resp = super::send_with_retry(req).await?;
119            let status = resp.status();
120            if !status.is_success() {
121                let text = resp.text().await.unwrap_or_default();
122                return Err(ClientError::Graph {
123                    status: status.as_u16(),
124                    message: text,
125                });
126            }
127            let envelope: BatchEnvelope = resp.json().await?;
128            for item in envelope.responses {
129                if item.status == 429 && round + 1 < MAX_ROUNDS {
130                    max_retry_after = max_retry_after.max(item.retry_after().unwrap_or(1));
131                    if let Some(original) = chunk.iter().find(|r| r.id == item.id) {
132                        next_pending.push(original.clone());
133                        continue;
134                    }
135                }
136                done.push(item);
137            }
138        }
139
140        pending = next_pending;
141        if !pending.is_empty() {
142            tokio::time::sleep(std::time::Duration::from_secs(max_retry_after.min(30))).await;
143        }
144    }
145    Ok(done)
146}
147
148#[cfg(test)]
149mod tests {
150    use super::*;
151    use wiremock::matchers::{method, path};
152    use wiremock::{Mock, MockServer, Request, Respond, ResponseTemplate};
153
154    #[tokio::test]
155    async fn chunks_over_20_into_multiple_posts() {
156        let server = MockServer::start().await;
157        struct Echo;
158        impl Respond for Echo {
159            fn respond(&self, req: &Request) -> ResponseTemplate {
160                let body: Value = serde_json::from_slice(&req.body).unwrap();
161                let responses: Vec<Value> = body["requests"]
162                    .as_array()
163                    .unwrap()
164                    .iter()
165                    .map(|r| serde_json::json!({"id": r["id"], "status": 204, "body": null}))
166                    .collect();
167                ResponseTemplate::new(200)
168                    .set_body_json(serde_json::json!({"responses": responses}))
169            }
170        }
171        Mock::given(method("POST"))
172            .and(path("/$batch"))
173            .respond_with(Echo)
174            .expect(2) // 25 requests → 20 + 5
175            .mount(&server)
176            .await;
177
178        let requests: Vec<BatchRequest> = (0..25)
179            .map(|i| BatchRequest::bare(i.to_string(), "DELETE", format!("/me/messages/{i}")))
180            .collect();
181        let http = reqwest::Client::new();
182        let out = batch_all(&http, &server.uri(), "tok", requests)
183            .await
184            .unwrap();
185        assert_eq!(out.len(), 25);
186        assert!(out.iter().all(BatchResponse::is_success));
187    }
188
189    #[tokio::test]
190    async fn per_item_429_is_rebatched() {
191        let server = MockServer::start().await;
192        struct ThrottleOnce {
193            hits: std::sync::atomic::AtomicU32,
194        }
195        impl Respond for ThrottleOnce {
196            fn respond(&self, req: &Request) -> ResponseTemplate {
197                let first = self.hits.fetch_add(1, std::sync::atomic::Ordering::SeqCst) == 0;
198                let body: Value = serde_json::from_slice(&req.body).unwrap();
199                let responses: Vec<Value> = body["requests"]
200                    .as_array()
201                    .unwrap()
202                    .iter()
203                    .map(|r| {
204                        if first && r["id"] == "1" {
205                            serde_json::json!({"id": r["id"], "status": 429,
206                                "body": {"error": {"code": "TooManyRequests", "retryAfterSeconds": 0}}})
207                        } else {
208                            serde_json::json!({"id": r["id"], "status": 200, "body": {}})
209                        }
210                    })
211                    .collect();
212                ResponseTemplate::new(200)
213                    .set_body_json(serde_json::json!({"responses": responses}))
214            }
215        }
216        Mock::given(method("POST"))
217            .and(path("/$batch"))
218            .respond_with(ThrottleOnce {
219                hits: std::sync::atomic::AtomicU32::new(0),
220            })
221            .expect(2)
222            .mount(&server)
223            .await;
224
225        let requests = vec![
226            BatchRequest::json(
227                "0",
228                "PATCH",
229                "/me/messages/a",
230                serde_json::json!({"isRead": true}),
231            ),
232            BatchRequest::json(
233                "1",
234                "PATCH",
235                "/me/messages/b",
236                serde_json::json!({"isRead": true}),
237            ),
238        ];
239        let http = reqwest::Client::new();
240        let out = batch_all(&http, &server.uri(), "tok", requests)
241            .await
242            .unwrap();
243        assert_eq!(out.len(), 2);
244        assert!(
245            out.iter().all(BatchResponse::is_success),
246            "throttled item retried"
247        );
248    }
249
250    #[tokio::test]
251    async fn per_item_404_surfaces_without_failing_batch() {
252        let server = MockServer::start().await;
253        Mock::given(method("POST"))
254            .and(path("/$batch"))
255            .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
256                "responses": [
257                    {"id": "0", "status": 204, "body": null},
258                    {"id": "1", "status": 404, "body": {"error": {"code": "ErrorItemNotFound"}}}
259                ]
260            })))
261            .mount(&server)
262            .await;
263
264        let requests = vec![
265            BatchRequest::bare("0", "DELETE", "/me/messages/a"),
266            BatchRequest::bare("1", "DELETE", "/me/messages/b"),
267        ];
268        let http = reqwest::Client::new();
269        let out = batch_all(&http, &server.uri(), "tok", requests)
270            .await
271            .unwrap();
272        assert_eq!(out.iter().filter(|r| r.is_success()).count(), 1);
273        assert_eq!(out.iter().find(|r| r.id == "1").unwrap().status, 404);
274    }
275}