1use 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 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 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
92pub 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) .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}