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
19pub struct SmugMugClient {
20 client: reqwest::Client,
21 api_key: String,
22 api_secret: String,
23 access_token: String,
24 access_token_secret: String,
25}
26
27impl SmugMugClient {
28 pub fn new(
29 api_key: String,
30 api_secret: String,
31 access_token: String,
32 access_token_secret: String,
33 ) -> Self {
34 SmugMugClient {
35 client: reqwest::Client::new(),
36 api_key,
37 api_secret,
38 access_token,
39 access_token_secret,
40 }
41 }
42
43 pub fn build_oauth_header(&self, method: &str, url: &str) -> String {
44 self.build_oauth_header_with_query(method, url, &())
45 }
46
47 pub fn build_oauth_header_with_query<T: oauth::Request>(
55 &self,
56 method: &str,
57 url: &str,
58 query: &T,
59 ) -> String {
60 let token = oauth::Token::from_parts(
61 &self.api_key,
62 &self.api_secret,
63 &self.access_token,
64 &self.access_token_secret,
65 );
66
67 let signer = oauth::HmacSha1::new();
68
69 match method {
70 "GET" => oauth::get(url, query, &token, signer),
71 "POST" => oauth::post(url, query, &token, signer),
72 "DELETE" => oauth::delete(url, query, &token, signer),
73 "PATCH" => oauth::patch(url, query, &token, signer),
74 _ => oauth::get(url, query, &token, signer),
75 }
76 }
77
78 pub async fn get_auth_user(&self) -> Result<Value> {
79 let url = "https://api.smugmug.com/api/v2!authuser";
80
81 let oauth_header = self.build_oauth_header("GET", url);
82
83 let mut headers = HeaderMap::new();
84 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
85 headers.insert("Accept", HeaderValue::from_static("application/json"));
86
87 let response = self.client.get(url).headers(headers).send().await?;
88
89 let status = response.status();
90 let body_text = response.text().await?;
91
92 if !status.is_success() {
93 anyhow::bail!("API request failed with status {}: {}", status, body_text);
94 }
95
96 let body: Value = serde_json::from_str(&body_text)?;
97 Ok(body)
98 }
99
100 pub async fn get_user_features(&self, user_uri: &str) -> Result<Value> {
101 let url = format!("https://api.smugmug.com{}!features", user_uri);
102
103 let oauth_header = self.build_oauth_header("GET", &url);
104
105 let mut headers = HeaderMap::new();
106 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
107 headers.insert("Accept", HeaderValue::from_static("application/json"));
108
109 let response = self.client.get(&url).headers(headers).send().await?;
110
111 let status = response.status();
112 let body_text = response.text().await?;
113
114 if !status.is_success() {
115 anyhow::bail!("API request failed with status {}: {}", status, body_text);
116 }
117
118 let body: Value = serde_json::from_str(&body_text)?;
119 Ok(body)
120 }
121
122 pub async fn request_raw(&self, method: &str, path_or_url: &str) -> Result<(u16, String)> {
129 let full = if path_or_url.starts_with("http") {
130 path_or_url.to_string()
131 } else {
132 format!("https://api.smugmug.com{}", path_or_url)
133 };
134 let mut url = reqwest::Url::parse(&full)?;
135 let params: Vec<(String, String)> = url.query_pairs().into_owned().collect();
136 url.set_query(None);
137
138 let token = oauth::Token::from_parts(
139 self.api_key.as_str(),
140 self.api_secret.as_str(),
141 self.access_token.as_str(),
142 self.access_token_secret.as_str(),
143 );
144 let oauth_header = oauth::Builder::with_token(token, oauth::HmacSha1::new()).authorize(
145 method,
146 url.as_str(),
147 &oauth::ParameterList::new(params.clone()),
148 );
149
150 let mut headers = HeaderMap::new();
151 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
152 headers.insert("Accept", HeaderValue::from_static("application/json"));
153
154 let response = self
155 .client
156 .request(reqwest::Method::from_bytes(method.as_bytes())?, url)
157 .query(¶ms)
158 .headers(headers)
159 .send()
160 .await?;
161 let status = response.status().as_u16();
162 Ok((status, response.text().await?))
163 }
164
165 pub async fn get_all_pages<T: serde::de::DeserializeOwned>(
174 &self,
175 first_url: &str,
176 locator: &str,
177 ) -> Result<Vec<T>> {
178 let origin = {
179 let parsed = reqwest::Url::parse(first_url)?;
180 parsed.origin().ascii_serialization()
181 };
182
183 let mut items = Vec::new();
184 let mut next: Option<(String, Vec<(String, String)>)> =
185 Some((first_url.to_string(), Vec::new()));
186 let mut pages = 0;
187
188 while let Some((url, params)) = next.take() {
189 pages += 1;
190 if pages > 10_000 {
191 anyhow::bail!("Gave up listing {} after 10,000 pages", first_url);
192 }
193
194 let oauth_header = self.build_oauth_header_with_query(
195 "GET",
196 &url,
197 &oauth::ParameterList::new(params.clone()),
198 );
199 let mut headers = HeaderMap::new();
200 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
201 headers.insert("Accept", HeaderValue::from_static("application/json"));
202
203 let mut request = self.client.get(&url).headers(headers);
204 if !params.is_empty() {
205 request = request.query(¶ms);
206 }
207 let response = request.send().await?;
208 let status = response.status();
209 let body_text = response.text().await?;
210 if !status.is_success() {
211 anyhow::bail!("Request to {} failed: {} - {}", url, status, body_text);
212 }
213
214 let mut body: Value = serde_json::from_str(&body_text)?;
215 let response_data = &mut body["Response"];
216 if let Some(array) = response_data.get_mut(locator).map(Value::take) {
217 let page_items: Vec<T> = serde_json::from_value(array)?;
218 items.extend(page_items);
219 }
220
221 next = response_data["Pages"]["NextPage"]
222 .as_str()
223 .map(|next_page| {
224 let (path, query) = next_page.split_once('?').unwrap_or((next_page, ""));
225 let params = url::form_urlencoded::parse(query.as_bytes())
226 .into_owned()
227 .collect();
228 (format!("{}{}", origin, path), params)
229 });
230 }
231
232 Ok(items)
233 }
234
235 pub async fn get_with_auth(&self, url: &str) -> Result<reqwest::Response> {
236 let oauth_header = self.build_oauth_header("GET", url);
237
238 let mut headers = HeaderMap::new();
239 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
240 headers.insert("Accept", HeaderValue::from_static("application/json"));
241
242 Ok(self.client.get(url).headers(headers).send().await?)
243 }
244
245 pub async fn post_with_auth(
246 &self,
247 url: &str,
248 body: serde_json::Value,
249 ) -> Result<reqwest::Response> {
250 let oauth_header = self.build_oauth_header("POST", url);
251
252 let mut headers = HeaderMap::new();
253 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
254 headers.insert("Accept", HeaderValue::from_static("application/json"));
255 headers.insert("Content-Type", HeaderValue::from_static("application/json"));
256
257 Ok(self
258 .client
259 .post(url)
260 .headers(headers)
261 .json(&body)
262 .send()
263 .await?)
264 }
265
266 pub async fn delete_with_auth(&self, url: &str) -> Result<reqwest::Response> {
267 let oauth_header = self.build_oauth_header("DELETE", url);
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
273 Ok(self.client.delete(url).headers(headers).send().await?)
274 }
275
276 pub async fn patch_with_auth(
277 &self,
278 url: &str,
279 body: serde_json::Value,
280 ) -> Result<reqwest::Response> {
281 let oauth_header = self.build_oauth_header("PATCH", url);
282
283 let mut headers = HeaderMap::new();
284 headers.insert(AUTHORIZATION, HeaderValue::from_str(&oauth_header)?);
285 headers.insert("Accept", HeaderValue::from_static("application/json"));
286 headers.insert("Content-Type", HeaderValue::from_static("application/json"));
287
288 Ok(self
289 .client
290 .patch(url)
291 .headers(headers)
292 .json(&body)
293 .send()
294 .await?)
295 }
296}
297
298#[cfg(test)]
299mod tests {
300 use super::*;
301
302 #[derive(serde::Deserialize, Debug, PartialEq)]
303 struct Item {
304 #[serde(rename = "Name")]
305 name: String,
306 }
307
308 #[tokio::test]
309 async fn test_get_all_pages_follows_next_page() {
310 let mut server = mockito::Server::new_async().await;
311 let first = server
312 .mock("GET", "/api/v2/thing!items")
313 .match_query(mockito::Matcher::Missing)
314 .with_body(
315 r#"{"Response":{"Item":[{"Name":"a"},{"Name":"b"}],
316 "Pages":{"Total":3,"Start":1,"Count":2,
317 "NextPage":"/api/v2/thing!items?start=3&count=2"}}}"#,
318 )
319 .create_async()
320 .await;
321 let second = server
322 .mock("GET", "/api/v2/thing!items")
323 .match_query(mockito::Matcher::AllOf(vec![
324 mockito::Matcher::UrlEncoded("start".into(), "3".into()),
325 mockito::Matcher::UrlEncoded("count".into(), "2".into()),
326 ]))
327 .match_header(
328 "authorization",
329 mockito::Matcher::Regex("oauth_signature=".to_string()),
330 )
331 .with_body(
332 r#"{"Response":{"Item":[{"Name":"c"}],
333 "Pages":{"Total":3,"Start":3,"Count":1}}}"#,
334 )
335 .create_async()
336 .await;
337
338 let client = create_test_client();
339 let items: Vec<Item> = client
340 .get_all_pages(&format!("{}/api/v2/thing!items", server.url()), "Item")
341 .await
342 .unwrap();
343
344 first.assert_async().await;
345 second.assert_async().await;
346 let names: Vec<&str> = items.iter().map(|i| i.name.as_str()).collect();
347 assert_eq!(names, vec!["a", "b", "c"]);
348 }
349
350 #[tokio::test]
351 async fn test_get_all_pages_empty_list() {
352 let mut server = mockito::Server::new_async().await;
353 let _mock = server
354 .mock("GET", "/api/v2/thing!items")
355 .with_body(r#"{"Response":{"Pages":{"Total":0,"Start":1,"Count":0}}}"#)
356 .create_async()
357 .await;
358
359 let client = create_test_client();
360 let items: Vec<Item> = client
361 .get_all_pages(&format!("{}/api/v2/thing!items", server.url()), "Item")
362 .await
363 .unwrap();
364 assert!(items.is_empty());
365 }
366
367 #[tokio::test]
368 async fn test_get_all_pages_error_status() {
369 let mut server = mockito::Server::new_async().await;
370 let _mock = server
371 .mock("GET", "/api/v2/thing!items")
372 .with_status(404)
373 .with_body(r#"{"Code":404,"Message":"Not Found"}"#)
374 .create_async()
375 .await;
376
377 let client = create_test_client();
378 let result: Result<Vec<Item>> = client
379 .get_all_pages(&format!("{}/api/v2/thing!items", server.url()), "Item")
380 .await;
381 assert!(result.unwrap_err().to_string().contains("404"));
382 }
383
384 fn create_test_client() -> SmugMugClient {
385 SmugMugClient::new(
386 "test_api_key".to_string(),
387 "test_api_secret".to_string(),
388 "test_access_token".to_string(),
389 "test_access_token_secret".to_string(),
390 )
391 }
392
393 #[test]
394 fn test_smugmug_client_new() {
395 let client = create_test_client();
396 assert_eq!(client.api_key, "test_api_key");
397 assert_eq!(client.api_secret, "test_api_secret");
398 assert_eq!(client.access_token, "test_access_token");
399 assert_eq!(client.access_token_secret, "test_access_token_secret");
400 }
401
402 #[test]
403 fn test_build_oauth_header_get() {
404 let client = create_test_client();
405 let url = "https://api.smugmug.com/api/v2!authuser";
406 let header = client.build_oauth_header("GET", url);
407
408 assert!(header.starts_with("OAuth "));
410 assert!(header.contains("oauth_consumer_key="));
411 assert!(header.contains("oauth_token="));
412 assert!(header.contains("oauth_signature_method="));
413 assert!(header.contains("oauth_timestamp="));
414 assert!(header.contains("oauth_nonce="));
415 assert!(header.contains("oauth_signature="));
416 }
417
418 #[test]
419 fn test_build_oauth_header_post() {
420 let client = create_test_client();
421 let url = "https://api.smugmug.com/api/v2/node/abc123!children";
422 let header = client.build_oauth_header("POST", url);
423
424 assert!(header.starts_with("OAuth "));
426 assert!(header.contains("oauth_consumer_key=\"test_api_key\""));
427 assert!(header.contains("oauth_token=\"test_access_token\""));
428 assert!(header.contains("oauth_signature_method=\"HMAC-SHA1\""));
429 }
430
431 #[test]
432 fn test_build_oauth_header_delete() {
433 let client = create_test_client();
434 let url = "https://api.smugmug.com/api/v2/image/IMG123";
435 let header = client.build_oauth_header("DELETE", url);
436
437 assert!(header.starts_with("OAuth "));
439 assert!(header.contains("oauth_consumer_key=\"test_api_key\""));
440 assert!(header.contains("oauth_token=\"test_access_token\""));
441 assert!(header.contains("oauth_signature_method=\"HMAC-SHA1\""));
442 }
443
444 #[test]
445 fn test_build_oauth_header_patch() {
446 let client = create_test_client();
447 let url = "https://api.smugmug.com/api/v2/image/IMG123";
448 let header = client.build_oauth_header("PATCH", url);
449
450 assert!(header.starts_with("OAuth "));
452 assert!(header.contains("oauth_consumer_key=\"test_api_key\""));
453 assert!(header.contains("oauth_token=\"test_access_token\""));
454 assert!(header.contains("oauth_signature_method=\"HMAC-SHA1\""));
455 }
456
457 #[test]
458 fn test_build_oauth_header_unknown_method() {
459 let client = create_test_client();
460 let url = "https://api.smugmug.com/api/v2!authuser";
461 let header = client.build_oauth_header("PUT", url);
463
464 assert!(header.starts_with("OAuth "));
465 assert!(header.contains("oauth_consumer_key=\"test_api_key\""));
466 }
467
468 #[tokio::test]
469 async fn test_get_auth_user_success() {
470 let mut server = mockito::Server::new_async().await;
471 let mock = server
472 .mock("GET", "/api/v2!authuser")
473 .match_header(
474 "authorization",
475 mockito::Matcher::Regex("OAuth.*".to_string()),
476 )
477 .match_header("accept", "application/json")
478 .with_status(200)
479 .with_header("content-type", "application/json")
480 .with_body(r#"{"Response":{"User":{"Uri":"/api/v2/user/testuser"}}}"#)
481 .create_async()
482 .await;
483
484 let client = SmugMugClient::new(
485 "test_key".to_string(),
486 "test_secret".to_string(),
487 "test_token".to_string(),
488 "test_token_secret".to_string(),
489 );
490
491 drop(mock);
496 }
497
498 #[tokio::test]
499 async fn test_get_auth_user_unauthorized() {
500 let mut server = mockito::Server::new_async().await;
501 let mock = server
502 .mock("GET", "/api/v2!authuser")
503 .match_header(
504 "authorization",
505 mockito::Matcher::Regex("OAuth.*".to_string()),
506 )
507 .with_status(401)
508 .with_body("Unauthorized")
509 .create_async()
510 .await;
511
512 drop(mock);
515 }
516
517 #[test]
518 fn test_node_tree_structure() {
519 let tree = NodeTree {
520 name: "Root".to_string(),
521 node_type: "Folder".to_string(),
522 children: vec![
523 NodeTree {
524 name: "Child1".to_string(),
525 node_type: "Album".to_string(),
526 children: vec![],
527 },
528 NodeTree {
529 name: "Child2".to_string(),
530 node_type: "Folder".to_string(),
531 children: vec![],
532 },
533 ],
534 };
535
536 assert_eq!(tree.name, "Root");
537 assert_eq!(tree.node_type, "Folder");
538 assert_eq!(tree.children.len(), 2);
539 assert_eq!(tree.children[0].name, "Child1");
540 assert_eq!(tree.children[1].name, "Child2");
541 }
542}