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(
377 &self,
378 url: &str,
379 body: serde_json::Value,
380 ) -> Result<reqwest::Response> {
381 let oauth_header = self.build_oauth_header("PATCH", url);
382
383 let mut headers = HeaderMap::new();
384 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
385 headers.insert("Accept", HeaderValue::from_static("application/json"));
386 headers.insert("Content-Type", HeaderValue::from_static("application/json"));
387
388 Ok(self
389 .client
390 .patch(url)
391 .headers(headers)
392 .json(&body)
393 .send()
394 .await?)
395 }
396}
397
398#[cfg(test)]
399mod tests {
400 use super::*;
401
402 #[derive(serde::Deserialize, Debug, PartialEq)]
403 struct Item {
404 #[serde(rename = "Name")]
405 name: String,
406 }
407
408 #[tokio::test]
409 async fn test_get_all_pages_follows_next_page() {
410 let mut server = mockito::Server::new_async().await;
411 let first = server
412 .mock("GET", "/api/v2/thing!items")
413 .match_query(mockito::Matcher::Missing)
414 .with_body(
415 r#"{"Response":{"Item":[{"Name":"a"},{"Name":"b"}],
416 "Pages":{"Total":3,"Start":1,"Count":2,
417 "NextPage":"/api/v2/thing!items?start=3&count=2"}}}"#,
418 )
419 .create_async()
420 .await;
421 let second = server
422 .mock("GET", "/api/v2/thing!items")
423 .match_query(mockito::Matcher::AllOf(vec![
424 mockito::Matcher::UrlEncoded("start".into(), "3".into()),
425 mockito::Matcher::UrlEncoded("count".into(), "2".into()),
426 ]))
427 .match_header(
428 "authorization",
429 mockito::Matcher::Regex("oauth_signature=".to_string()),
430 )
431 .with_body(
432 r#"{"Response":{"Item":[{"Name":"c"}],
433 "Pages":{"Total":3,"Start":3,"Count":1}}}"#,
434 )
435 .create_async()
436 .await;
437
438 let client = create_test_client();
439 let items: Vec<Item> = client
440 .get_all_pages(&format!("{}/api/v2/thing!items", server.url()), "Item")
441 .await
442 .unwrap();
443
444 first.assert_async().await;
445 second.assert_async().await;
446 let names: Vec<&str> = items.iter().map(|i| i.name.as_str()).collect();
447 assert_eq!(names, vec!["a", "b", "c"]);
448 }
449
450 #[tokio::test]
451 async fn rate_limited_requests_are_retried() {
452 let mut server = mockito::Server::new_async().await;
453 let limited = server
454 .mock("GET", "/api/v2/thing!items")
455 .with_status(429)
456 .with_header("retry-after", "0")
457 .expect(1)
458 .create_async()
459 .await;
460 let ok = server
461 .mock("GET", "/api/v2/thing!items")
462 .with_body(r#"{"Response":{"Item":[{"Name":"a"}]}}"#)
463 .expect(1)
464 .create_async()
465 .await;
466
467 let client = create_test_client();
468 let items: Vec<Item> = client
469 .get_all_pages(&format!("{}/api/v2/thing!items", server.url()), "Item")
470 .await
471 .unwrap();
472 limited.assert_async().await;
473 ok.assert_async().await;
474 assert_eq!(items.len(), 1);
475 }
476
477 #[tokio::test]
478 async fn failed_posts_are_not_retried_on_server_errors() {
479 let mut server = mockito::Server::new_async().await;
480 let mock = server
481 .mock("POST", "/api/v2/node/x!children")
482 .with_status(500)
483 .expect(1)
484 .create_async()
485 .await;
486 let client = create_test_client();
487 let response = client
488 .post_with_auth(
489 &format!("{}/api/v2/node/x!children", server.url()),
490 serde_json::json!({}),
491 )
492 .await
493 .unwrap();
494 assert_eq!(response.status().as_u16(), 500);
495 mock.assert_async().await;
496 }
497
498 #[tokio::test]
499 async fn test_get_all_pages_empty_list() {
500 let mut server = mockito::Server::new_async().await;
501 let _mock = server
502 .mock("GET", "/api/v2/thing!items")
503 .with_body(r#"{"Response":{"Pages":{"Total":0,"Start":1,"Count":0}}}"#)
504 .create_async()
505 .await;
506
507 let client = create_test_client();
508 let items: Vec<Item> = client
509 .get_all_pages(&format!("{}/api/v2/thing!items", server.url()), "Item")
510 .await
511 .unwrap();
512 assert!(items.is_empty());
513 }
514
515 #[tokio::test]
516 async fn test_get_all_pages_error_status() {
517 let mut server = mockito::Server::new_async().await;
518 let _mock = server
519 .mock("GET", "/api/v2/thing!items")
520 .with_status(404)
521 .with_body(r#"{"Code":404,"Message":"Not Found"}"#)
522 .create_async()
523 .await;
524
525 let client = create_test_client();
526 let result: Result<Vec<Item>> = client
527 .get_all_pages(&format!("{}/api/v2/thing!items", server.url()), "Item")
528 .await;
529 assert!(result.unwrap_err().to_string().contains("404"));
530 }
531
532 fn create_test_client() -> SmugMugClient {
533 SmugMugClient::new(
534 "test_api_key".to_string(),
535 "test_api_secret".to_string(),
536 "test_access_token".to_string(),
537 "test_access_token_secret".to_string(),
538 )
539 }
540
541 #[test]
542 fn test_smugmug_client_new() {
543 let client = create_test_client();
544 assert_eq!(client.api_key, "test_api_key");
545 assert_eq!(client.api_secret, "test_api_secret");
546 assert_eq!(client.access_token, "test_access_token");
547 assert_eq!(client.access_token_secret, "test_access_token_secret");
548 }
549
550 #[test]
551 fn test_build_oauth_header_get() {
552 let client = create_test_client();
553 let url = "https://api.smugmug.com/api/v2!authuser";
554 let header = client.build_oauth_header("GET", url);
555
556 assert!(header.starts_with("OAuth "));
558 assert!(header.contains("oauth_consumer_key="));
559 assert!(header.contains("oauth_token="));
560 assert!(header.contains("oauth_signature_method="));
561 assert!(header.contains("oauth_timestamp="));
562 assert!(header.contains("oauth_nonce="));
563 assert!(header.contains("oauth_signature="));
564 }
565
566 #[test]
567 fn test_build_oauth_header_post() {
568 let client = create_test_client();
569 let url = "https://api.smugmug.com/api/v2/node/abc123!children";
570 let header = client.build_oauth_header("POST", url);
571
572 assert!(header.starts_with("OAuth "));
574 assert!(header.contains("oauth_consumer_key=\"test_api_key\""));
575 assert!(header.contains("oauth_token=\"test_access_token\""));
576 assert!(header.contains("oauth_signature_method=\"HMAC-SHA1\""));
577 }
578
579 #[test]
580 fn test_build_oauth_header_delete() {
581 let client = create_test_client();
582 let url = "https://api.smugmug.com/api/v2/image/IMG123";
583 let header = client.build_oauth_header("DELETE", url);
584
585 assert!(header.starts_with("OAuth "));
587 assert!(header.contains("oauth_consumer_key=\"test_api_key\""));
588 assert!(header.contains("oauth_token=\"test_access_token\""));
589 assert!(header.contains("oauth_signature_method=\"HMAC-SHA1\""));
590 }
591
592 #[test]
593 fn test_build_oauth_header_patch() {
594 let client = create_test_client();
595 let url = "https://api.smugmug.com/api/v2/image/IMG123";
596 let header = client.build_oauth_header("PATCH", url);
597
598 assert!(header.starts_with("OAuth "));
600 assert!(header.contains("oauth_consumer_key=\"test_api_key\""));
601 assert!(header.contains("oauth_token=\"test_access_token\""));
602 assert!(header.contains("oauth_signature_method=\"HMAC-SHA1\""));
603 }
604
605 #[test]
606 fn test_build_oauth_header_unknown_method() {
607 let client = create_test_client();
608 let url = "https://api.smugmug.com/api/v2!authuser";
609 let header = client.build_oauth_header("PUT", url);
611
612 assert!(header.starts_with("OAuth "));
613 assert!(header.contains("oauth_consumer_key=\"test_api_key\""));
614 }
615
616 #[tokio::test]
617 async fn test_get_auth_user_success() {
618 let mut server = mockito::Server::new_async().await;
619 let mock = server
620 .mock("GET", "/api/v2!authuser")
621 .match_header(
622 "authorization",
623 mockito::Matcher::Regex("OAuth.*".to_string()),
624 )
625 .match_header("accept", "application/json")
626 .with_status(200)
627 .with_header("content-type", "application/json")
628 .with_body(r#"{"Response":{"User":{"Uri":"/api/v2/user/testuser"}}}"#)
629 .create_async()
630 .await;
631
632 let client = SmugMugClient::new(
633 "test_key".to_string(),
634 "test_secret".to_string(),
635 "test_token".to_string(),
636 "test_token_secret".to_string(),
637 );
638
639 drop(mock);
644 }
645
646 #[tokio::test]
647 async fn test_get_auth_user_unauthorized() {
648 let mut server = mockito::Server::new_async().await;
649 let mock = server
650 .mock("GET", "/api/v2!authuser")
651 .match_header(
652 "authorization",
653 mockito::Matcher::Regex("OAuth.*".to_string()),
654 )
655 .with_status(401)
656 .with_body("Unauthorized")
657 .create_async()
658 .await;
659
660 drop(mock);
663 }
664
665 #[test]
666 fn test_node_tree_structure() {
667 let tree = NodeTree {
668 name: "Root".to_string(),
669 node_type: "Folder".to_string(),
670 children: vec![
671 NodeTree {
672 name: "Child1".to_string(),
673 node_type: "Album".to_string(),
674 children: vec![],
675 },
676 NodeTree {
677 name: "Child2".to_string(),
678 node_type: "Folder".to_string(),
679 children: vec![],
680 },
681 ],
682 };
683
684 assert_eq!(tree.name, "Root");
685 assert_eq!(tree.node_type, "Folder");
686 assert_eq!(tree.children.len(), 2);
687 assert_eq!(tree.children[0].name, "Child1");
688 assert_eq!(tree.children[1].name, "Child2");
689 }
690}