1use anyhow::Result;
2use oauth1_request as oauth;
3use reqwest::header::{AUTHORIZATION, HeaderMap, HeaderValue};
4use serde_json::Value;
5
6pub mod albums;
7pub mod comments;
8pub mod images;
9pub mod oauth_flow;
10pub mod upload;
11
12#[derive(Debug)]
13pub struct NodeTree {
14 pub name: String,
15 pub node_type: String,
16 pub children: Vec<NodeTree>,
17}
18
19const CONNECT_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
21
22const READ_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(300);
26
27const MAX_ATTEMPTS: u32 = 5;
29
30fn backoff(attempt: u32) -> std::time::Duration {
33 let jitter = std::time::SystemTime::now()
34 .duration_since(std::time::UNIX_EPOCH)
35 .map(|d| d.subsec_millis() % 500)
36 .unwrap_or(0);
37 std::time::Duration::from_millis(
38 1000 * 2u64.pow(attempt.saturating_sub(1).min(6)) + jitter as u64,
39 )
40}
41
42fn retry_after(response: &reqwest::Response) -> Option<std::time::Duration> {
44 let secs: u64 = response
45 .headers()
46 .get(reqwest::header::RETRY_AFTER)?
47 .to_str()
48 .ok()?
49 .trim()
50 .parse()
51 .ok()?;
52 Some(std::time::Duration::from_secs(secs.min(600)))
53}
54
55pub struct SmugMugClient {
56 client: reqwest::Client,
57 api_key: String,
58 api_secret: String,
59 access_token: String,
60 access_token_secret: String,
61 auth_user: tokio::sync::OnceCell<AuthUser>,
63}
64
65#[derive(Debug, Clone)]
67pub struct AuthUser {
68 pub nickname: String,
69 pub root_node_uri: String,
71}
72
73impl SmugMugClient {
74 pub fn new(
75 api_key: String,
76 api_secret: String,
77 access_token: String,
78 access_token_secret: String,
79 ) -> Self {
80 SmugMugClient {
81 client: reqwest::Client::builder()
82 .connect_timeout(CONNECT_TIMEOUT)
83 .read_timeout(READ_TIMEOUT)
84 .build()
85 .expect("Failed to set up the HTTP client (TLS initialization failed)"),
88 api_key,
89 api_secret,
90 access_token,
91 access_token_secret,
92 auth_user: tokio::sync::OnceCell::new(),
93 }
94 }
95
96 pub(crate) fn http(&self) -> &reqwest::Client {
98 &self.client
99 }
100
101 pub fn build_oauth_header(&self, method: &str, url: &str) -> String {
102 self.build_oauth_header_with_query(method, url, &())
103 }
104
105 pub fn build_oauth_header_with_query<T: oauth::Request>(
113 &self,
114 method: &str,
115 url: &str,
116 query: &T,
117 ) -> String {
118 let token = oauth::Token::from_parts(
119 &self.api_key,
120 &self.api_secret,
121 &self.access_token,
122 &self.access_token_secret,
123 );
124
125 let signer = oauth::HmacSha1::new();
126
127 match method {
128 "GET" => oauth::get(url, query, &token, signer),
129 "POST" => oauth::post(url, query, &token, signer),
130 "DELETE" => oauth::delete(url, query, &token, signer),
131 "PATCH" => oauth::patch(url, query, &token, signer),
132 _ => oauth::get(url, query, &token, signer),
133 }
134 }
135
136 pub async fn get_auth_user(&self) -> Result<Value> {
137 let url = "https://api.smugmug.com/api/v2!authuser";
138
139 let oauth_header = self.build_oauth_header("GET", url);
140
141 let mut headers = HeaderMap::new();
142 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
143 headers.insert("Accept", HeaderValue::from_static("application/json"));
144
145 let response = self.client.get(url).headers(headers).send().await?;
146
147 let status = response.status();
148 let body_text = response.text().await?;
149
150 if !status.is_success() {
151 anyhow::bail!("API request failed with status {}: {}", status, body_text);
152 }
153
154 let body: Value = serde_json::from_str(&body_text)?;
155 Ok(body)
156 }
157
158 pub async fn get_user_features(&self, user_uri: &str) -> Result<Value> {
159 let url = format!("https://api.smugmug.com{}!features", user_uri);
160
161 let oauth_header = self.build_oauth_header("GET", &url);
162
163 let mut headers = HeaderMap::new();
164 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
165 headers.insert("Accept", HeaderValue::from_static("application/json"));
166
167 let response = self.client.get(&url).headers(headers).send().await?;
168
169 let status = response.status();
170 let body_text = response.text().await?;
171
172 if !status.is_success() {
173 anyhow::bail!("API request failed with status {}: {}", status, body_text);
174 }
175
176 let body: Value = serde_json::from_str(&body_text)?;
177 Ok(body)
178 }
179
180 pub async fn request_raw(
188 &self,
189 method: &str,
190 path_or_url: &str,
191 body: Option<serde_json::Value>,
192 ) -> Result<(u16, String)> {
193 let full = if path_or_url.starts_with("http") {
194 path_or_url.to_string()
195 } else {
196 format!("https://api.smugmug.com{}", path_or_url)
197 };
198 let mut url = reqwest::Url::parse(&full)?;
199 let params: Vec<(String, String)> = url.query_pairs().into_owned().collect();
200 url.set_query(None);
201
202 let token = oauth::Token::from_parts(
203 self.api_key.as_str(),
204 self.api_secret.as_str(),
205 self.access_token.as_str(),
206 self.access_token_secret.as_str(),
207 );
208 let oauth_header = oauth::Builder::with_token(token, oauth::HmacSha1::new()).authorize(
209 method,
210 url.as_str(),
211 &oauth::ParameterList::new(params.clone()),
212 );
213
214 let mut headers = HeaderMap::new();
215 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
216 headers.insert("Accept", HeaderValue::from_static("application/json"));
217
218 let mut request = self
219 .client
220 .request(reqwest::Method::from_bytes(method.as_bytes())?, url)
221 .query(¶ms)
222 .headers(headers);
223 if let Some(body) = body {
224 request = request.json(&body);
225 }
226 let response = request.send().await?;
227 let status = response.status().as_u16();
228 Ok((status, response.text().await?))
229 }
230
231 pub async fn get_all_pages<T: serde::de::DeserializeOwned>(
240 &self,
241 first_url: &str,
242 locator: &str,
243 ) -> Result<Vec<T>> {
244 let mut parsed = reqwest::Url::parse(first_url)?;
245 let origin = parsed.origin().ascii_serialization();
246 let first_params: Vec<(String, String)> = parsed.query_pairs().into_owned().collect();
248 parsed.set_query(None);
249
250 let mut items = Vec::new();
251 let mut next: Option<(String, Vec<(String, String)>)> =
252 Some((parsed.to_string(), first_params));
253 let mut pages = 0;
254
255 while let Some((url, params)) = next.take() {
256 pages += 1;
257 if pages > 10_000 {
258 anyhow::bail!("Gave up listing {} after 10,000 pages", first_url);
259 }
260
261 let response = self
262 .send_retrying(true, || {
263 let oauth_header = self.build_oauth_header_with_query(
265 "GET",
266 &url,
267 &oauth::ParameterList::new(params.clone()),
268 );
269 let mut headers = HeaderMap::new();
270 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
271 headers.insert("Accept", HeaderValue::from_static("application/json"));
272 let mut request = self.client.get(&url).headers(headers);
273 if !params.is_empty() {
274 request = request.query(¶ms);
275 }
276 Ok(request)
277 })
278 .await?;
279 let status = response.status();
280 let body_text = response.text().await?;
281 if !status.is_success() {
282 anyhow::bail!("Request to {} failed: {} - {}", url, status, body_text);
283 }
284
285 let mut body: Value = serde_json::from_str(&body_text)?;
286 let response_data = &mut body["Response"];
287 if let Some(array) = response_data.get_mut(locator).map(Value::take) {
288 let page_items: Vec<T> = serde_json::from_value(array)?;
289 items.extend(page_items);
290 }
291
292 next = response_data["Pages"]["NextPage"]
293 .as_str()
294 .map(|next_page| {
295 let (path, query) = next_page.split_once('?').unwrap_or((next_page, ""));
296 let params = url::form_urlencoded::parse(query.as_bytes())
297 .into_owned()
298 .collect();
299 (format!("{}{}", origin, path), params)
300 });
301 }
302
303 Ok(items)
304 }
305
306 pub async fn get_with_auth(&self, url: &str) -> Result<reqwest::Response> {
307 self.send_retrying(true, || {
308 let oauth_header = self.build_oauth_header("GET", url);
309 let mut headers = HeaderMap::new();
310 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
311 headers.insert("Accept", HeaderValue::from_static("application/json"));
312 Ok(self.client.get(url).headers(headers))
313 })
314 .await
315 }
316
317 pub async fn post_with_auth(
318 &self,
319 url: &str,
320 body: serde_json::Value,
321 ) -> Result<reqwest::Response> {
322 self.send_retrying(false, || {
323 let oauth_header = self.build_oauth_header("POST", url);
324 let mut headers = HeaderMap::new();
325 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
326 headers.insert("Accept", HeaderValue::from_static("application/json"));
327 headers.insert("Content-Type", HeaderValue::from_static("application/json"));
328 Ok(self.client.post(url).headers(headers).json(&body))
329 })
330 .await
331 }
332
333 async fn send_retrying(
340 &self,
341 idempotent: bool,
342 make: impl Fn() -> Result<reqwest::RequestBuilder>,
343 ) -> Result<reqwest::Response> {
344 let mut attempt = 1;
345 loop {
346 let last = attempt >= MAX_ATTEMPTS;
347 match make()?.send().await {
348 Ok(response) => {
349 let status = response.status().as_u16();
350 let retry = status == 429 || (idempotent && status >= 500);
351 if !retry || last {
352 return Ok(response);
353 }
354 let wait = retry_after(&response).unwrap_or_else(|| backoff(attempt));
355 tokio::time::sleep(wait).await;
356 }
357 Err(e) if !last && (e.is_connect() || (idempotent && e.is_timeout())) => {
358 tokio::time::sleep(backoff(attempt)).await;
359 }
360 Err(e) => return Err(e.into()),
361 }
362 attempt += 1;
363 }
364 }
365
366 pub async fn delete_with_auth(&self, url: &str) -> Result<reqwest::Response> {
367 let oauth_header = self.build_oauth_header("DELETE", url);
368
369 let mut headers = HeaderMap::new();
370 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
371 headers.insert("Accept", HeaderValue::from_static("application/json"));
372
373 Ok(self.client.delete(url).headers(headers).send().await?)
374 }
375
376 pub async fn patch_with_auth(
378 &self,
379 url: &str,
380 body: serde_json::Value,
381 ) -> Result<reqwest::Response> {
382 self.send_retrying(true, || {
383 let oauth_header = self.build_oauth_header("PATCH", url);
384 let mut headers = HeaderMap::new();
385 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
386 headers.insert("Accept", HeaderValue::from_static("application/json"));
387 headers.insert("Content-Type", HeaderValue::from_static("application/json"));
388 Ok(self.client.patch(url).headers(headers).json(&body))
389 })
390 .await
391 }
392}
393
394#[cfg(test)]
395mod tests {
396 use super::*;
397
398 #[derive(serde::Deserialize, Debug, PartialEq)]
399 struct Item {
400 #[serde(rename = "Name")]
401 name: String,
402 }
403
404 #[tokio::test]
405 async fn test_get_all_pages_follows_next_page() {
406 let mut server = mockito::Server::new_async().await;
407 let first = server
408 .mock("GET", "/api/v2/thing!items")
409 .match_query(mockito::Matcher::Missing)
410 .with_body(
411 r#"{"Response":{"Item":[{"Name":"a"},{"Name":"b"}],
412 "Pages":{"Total":3,"Start":1,"Count":2,
413 "NextPage":"/api/v2/thing!items?start=3&count=2"}}}"#,
414 )
415 .create_async()
416 .await;
417 let second = server
418 .mock("GET", "/api/v2/thing!items")
419 .match_query(mockito::Matcher::AllOf(vec![
420 mockito::Matcher::UrlEncoded("start".into(), "3".into()),
421 mockito::Matcher::UrlEncoded("count".into(), "2".into()),
422 ]))
423 .match_header(
424 "authorization",
425 mockito::Matcher::Regex("oauth_signature=".to_string()),
426 )
427 .with_body(
428 r#"{"Response":{"Item":[{"Name":"c"}],
429 "Pages":{"Total":3,"Start":3,"Count":1}}}"#,
430 )
431 .create_async()
432 .await;
433
434 let client = create_test_client();
435 let items: Vec<Item> = client
436 .get_all_pages(&format!("{}/api/v2/thing!items", server.url()), "Item")
437 .await
438 .unwrap();
439
440 first.assert_async().await;
441 second.assert_async().await;
442 let names: Vec<&str> = items.iter().map(|i| i.name.as_str()).collect();
443 assert_eq!(names, vec!["a", "b", "c"]);
444 }
445
446 #[tokio::test]
447 async fn rate_limited_requests_are_retried() {
448 let mut server = mockito::Server::new_async().await;
449 let limited = server
450 .mock("GET", "/api/v2/thing!items")
451 .with_status(429)
452 .with_header("retry-after", "0")
453 .expect(1)
454 .create_async()
455 .await;
456 let ok = server
457 .mock("GET", "/api/v2/thing!items")
458 .with_body(r#"{"Response":{"Item":[{"Name":"a"}]}}"#)
459 .expect(1)
460 .create_async()
461 .await;
462
463 let client = create_test_client();
464 let items: Vec<Item> = client
465 .get_all_pages(&format!("{}/api/v2/thing!items", server.url()), "Item")
466 .await
467 .unwrap();
468 limited.assert_async().await;
469 ok.assert_async().await;
470 assert_eq!(items.len(), 1);
471 }
472
473 #[tokio::test]
474 async fn failed_posts_are_not_retried_on_server_errors() {
475 let mut server = mockito::Server::new_async().await;
476 let mock = server
477 .mock("POST", "/api/v2/node/x!children")
478 .with_status(500)
479 .expect(1)
480 .create_async()
481 .await;
482 let client = create_test_client();
483 let response = client
484 .post_with_auth(
485 &format!("{}/api/v2/node/x!children", server.url()),
486 serde_json::json!({}),
487 )
488 .await
489 .unwrap();
490 assert_eq!(response.status().as_u16(), 500);
491 mock.assert_async().await;
492 }
493
494 #[tokio::test]
495 async fn test_get_all_pages_empty_list() {
496 let mut server = mockito::Server::new_async().await;
497 let _mock = server
498 .mock("GET", "/api/v2/thing!items")
499 .with_body(r#"{"Response":{"Pages":{"Total":0,"Start":1,"Count":0}}}"#)
500 .create_async()
501 .await;
502
503 let client = create_test_client();
504 let items: Vec<Item> = client
505 .get_all_pages(&format!("{}/api/v2/thing!items", server.url()), "Item")
506 .await
507 .unwrap();
508 assert!(items.is_empty());
509 }
510
511 #[tokio::test]
512 async fn test_get_all_pages_error_status() {
513 let mut server = mockito::Server::new_async().await;
514 let _mock = server
515 .mock("GET", "/api/v2/thing!items")
516 .with_status(404)
517 .with_body(r#"{"Code":404,"Message":"Not Found"}"#)
518 .create_async()
519 .await;
520
521 let client = create_test_client();
522 let result: Result<Vec<Item>> = client
523 .get_all_pages(&format!("{}/api/v2/thing!items", server.url()), "Item")
524 .await;
525 assert!(result.unwrap_err().to_string().contains("404"));
526 }
527
528 fn create_test_client() -> SmugMugClient {
529 SmugMugClient::new(
530 "test_api_key".to_string(),
531 "test_api_secret".to_string(),
532 "test_access_token".to_string(),
533 "test_access_token_secret".to_string(),
534 )
535 }
536
537 #[test]
538 fn test_smugmug_client_new() {
539 let client = create_test_client();
540 assert_eq!(client.api_key, "test_api_key");
541 assert_eq!(client.api_secret, "test_api_secret");
542 assert_eq!(client.access_token, "test_access_token");
543 assert_eq!(client.access_token_secret, "test_access_token_secret");
544 }
545
546 #[test]
547 fn test_build_oauth_header_get() {
548 let client = create_test_client();
549 let url = "https://api.smugmug.com/api/v2!authuser";
550 let header = client.build_oauth_header("GET", url);
551
552 assert!(header.starts_with("OAuth "));
554 assert!(header.contains("oauth_consumer_key="));
555 assert!(header.contains("oauth_token="));
556 assert!(header.contains("oauth_signature_method="));
557 assert!(header.contains("oauth_timestamp="));
558 assert!(header.contains("oauth_nonce="));
559 assert!(header.contains("oauth_signature="));
560 }
561
562 #[test]
563 fn test_build_oauth_header_post() {
564 let client = create_test_client();
565 let url = "https://api.smugmug.com/api/v2/node/abc123!children";
566 let header = client.build_oauth_header("POST", url);
567
568 assert!(header.starts_with("OAuth "));
570 assert!(header.contains("oauth_consumer_key=\"test_api_key\""));
571 assert!(header.contains("oauth_token=\"test_access_token\""));
572 assert!(header.contains("oauth_signature_method=\"HMAC-SHA1\""));
573 }
574
575 #[test]
576 fn test_build_oauth_header_delete() {
577 let client = create_test_client();
578 let url = "https://api.smugmug.com/api/v2/image/IMG123";
579 let header = client.build_oauth_header("DELETE", url);
580
581 assert!(header.starts_with("OAuth "));
583 assert!(header.contains("oauth_consumer_key=\"test_api_key\""));
584 assert!(header.contains("oauth_token=\"test_access_token\""));
585 assert!(header.contains("oauth_signature_method=\"HMAC-SHA1\""));
586 }
587
588 #[test]
589 fn test_build_oauth_header_patch() {
590 let client = create_test_client();
591 let url = "https://api.smugmug.com/api/v2/image/IMG123";
592 let header = client.build_oauth_header("PATCH", url);
593
594 assert!(header.starts_with("OAuth "));
596 assert!(header.contains("oauth_consumer_key=\"test_api_key\""));
597 assert!(header.contains("oauth_token=\"test_access_token\""));
598 assert!(header.contains("oauth_signature_method=\"HMAC-SHA1\""));
599 }
600
601 #[test]
602 fn test_build_oauth_header_unknown_method() {
603 let client = create_test_client();
604 let url = "https://api.smugmug.com/api/v2!authuser";
605 let header = client.build_oauth_header("PUT", url);
607
608 assert!(header.starts_with("OAuth "));
609 assert!(header.contains("oauth_consumer_key=\"test_api_key\""));
610 }
611
612 #[tokio::test]
613 async fn test_get_auth_user_success() {
614 let mut server = mockito::Server::new_async().await;
615 let mock = server
616 .mock("GET", "/api/v2!authuser")
617 .match_header(
618 "authorization",
619 mockito::Matcher::Regex("OAuth.*".to_string()),
620 )
621 .match_header("accept", "application/json")
622 .with_status(200)
623 .with_header("content-type", "application/json")
624 .with_body(r#"{"Response":{"User":{"Uri":"/api/v2/user/testuser"}}}"#)
625 .create_async()
626 .await;
627
628 let client = SmugMugClient::new(
629 "test_key".to_string(),
630 "test_secret".to_string(),
631 "test_token".to_string(),
632 "test_token_secret".to_string(),
633 );
634
635 drop(mock);
640 }
641
642 #[tokio::test]
643 async fn test_get_auth_user_unauthorized() {
644 let mut server = mockito::Server::new_async().await;
645 let mock = server
646 .mock("GET", "/api/v2!authuser")
647 .match_header(
648 "authorization",
649 mockito::Matcher::Regex("OAuth.*".to_string()),
650 )
651 .with_status(401)
652 .with_body("Unauthorized")
653 .create_async()
654 .await;
655
656 drop(mock);
659 }
660
661 #[test]
662 fn test_node_tree_structure() {
663 let tree = NodeTree {
664 name: "Root".to_string(),
665 node_type: "Folder".to_string(),
666 children: vec![
667 NodeTree {
668 name: "Child1".to_string(),
669 node_type: "Album".to_string(),
670 children: vec![],
671 },
672 NodeTree {
673 name: "Child2".to_string(),
674 node_type: "Folder".to_string(),
675 children: vec![],
676 },
677 ],
678 };
679
680 assert_eq!(tree.name, "Root");
681 assert_eq!(tree.node_type, "Folder");
682 assert_eq!(tree.children.len(), 2);
683 assert_eq!(tree.children[0].name, "Child1");
684 assert_eq!(tree.children[1].name, "Child2");
685 }
686}