1use anyhow::{Context, Result};
7use oauth1_request as oauth;
8
9const OAUTH_BASE: &str = "https://api.smugmug.com/services/oauth/1.0a";
10
11#[derive(Debug, Clone, PartialEq)]
13pub struct TokenPair {
14 pub token: String,
15 pub secret: String,
16}
17
18pub struct OAuthEndpoints {
21 pub request_token_url: String,
22 pub access_token_url: String,
23 pub authorize_url: String,
24}
25
26impl Default for OAuthEndpoints {
27 fn default() -> Self {
28 OAuthEndpoints {
29 request_token_url: format!("{}/getRequestToken", OAUTH_BASE),
30 access_token_url: format!("{}/getAccessToken", OAUTH_BASE),
31 authorize_url: format!("{}/authorize", OAUTH_BASE),
32 }
33 }
34}
35
36pub async fn get_request_token(
39 endpoints: &OAuthEndpoints,
40 api_key: &str,
41 api_secret: &str,
42) -> Result<TokenPair> {
43 let mut builder: oauth::Builder<'_, _, &str> = oauth::Builder::new(
44 oauth::Credentials::new(api_key, api_secret),
45 oauth::HmacSha1::new(),
46 );
47 builder.callback("oob");
48 let header = builder.get(&endpoints.request_token_url, &());
49
50 fetch_token_pair(&endpoints.request_token_url, header)
51 .await
52 .context("Failed to get a request token (check your API key and secret)")
53}
54
55pub fn authorize_url(endpoints: &OAuthEndpoints, request_token: &TokenPair) -> String {
58 let mut url = url::Url::parse(&endpoints.authorize_url).expect("valid authorize URL");
59 url.query_pairs_mut()
60 .append_pair("oauth_token", &request_token.token)
61 .append_pair("Access", "Full")
62 .append_pair("Permissions", "Modify");
63 url.to_string()
64}
65
66pub async fn get_access_token(
69 endpoints: &OAuthEndpoints,
70 api_key: &str,
71 api_secret: &str,
72 request_token: &TokenPair,
73 verifier: &str,
74) -> Result<TokenPair> {
75 let token = oauth::Token::from_parts(
76 api_key,
77 api_secret,
78 request_token.token.as_str(),
79 request_token.secret.as_str(),
80 );
81 let mut builder = oauth::Builder::with_token(token, oauth::HmacSha1::new());
82 builder.verifier(verifier);
83 let header = builder.get(&endpoints.access_token_url, &());
84
85 fetch_token_pair(&endpoints.access_token_url, header)
86 .await
87 .context("Failed to exchange the verifier code for an access token")
88}
89
90async fn fetch_token_pair(url: &str, authorization: String) -> Result<TokenPair> {
91 let response = reqwest::Client::new()
92 .get(url)
93 .header(reqwest::header::AUTHORIZATION, authorization)
94 .send()
95 .await?;
96 let status = response.status();
97 let body = response.text().await?;
98 if !status.is_success() {
99 anyhow::bail!("HTTP {}: {}", status.as_u16(), body.trim());
100 }
101 parse_token_response(&body)
102}
103
104fn parse_token_response(body: &str) -> Result<TokenPair> {
106 let mut token = None;
107 let mut secret = None;
108 for (key, value) in url::form_urlencoded::parse(body.trim().as_bytes()) {
109 match key.as_ref() {
110 "oauth_token" => token = Some(value.into_owned()),
111 "oauth_token_secret" => secret = Some(value.into_owned()),
112 _ => {}
113 }
114 }
115 match (token, secret) {
116 (Some(token), Some(secret)) => Ok(TokenPair { token, secret }),
117 _ => anyhow::bail!("Unexpected token response: {}", body.trim()),
118 }
119}
120
121#[cfg(test)]
122mod tests {
123 use super::*;
124
125 fn endpoints(server: &mockito::Server) -> OAuthEndpoints {
126 OAuthEndpoints {
127 request_token_url: format!("{}/getRequestToken", server.url()),
128 access_token_url: format!("{}/getAccessToken", server.url()),
129 authorize_url: format!("{}/authorize", server.url()),
130 }
131 }
132
133 #[test]
134 fn test_parse_token_response() {
135 let pair = parse_token_response(
136 "oauth_token=abc&oauth_token_secret=d%2Fef&oauth_callback_confirmed=true",
137 )
138 .unwrap();
139 assert_eq!(pair.token, "abc");
140 assert_eq!(pair.secret, "d/ef");
141 }
142
143 #[test]
144 fn test_parse_token_response_missing_secret() {
145 assert!(parse_token_response("oauth_problem=signature_invalid").is_err());
146 }
147
148 #[test]
149 fn test_authorize_url() {
150 let url = authorize_url(
151 &OAuthEndpoints::default(),
152 &TokenPair {
153 token: "req token".to_string(),
154 secret: "s".to_string(),
155 },
156 );
157 assert_eq!(
158 url,
159 "https://api.smugmug.com/services/oauth/1.0a/authorize?oauth_token=req+token&Access=Full&Permissions=Modify"
160 );
161 }
162
163 #[tokio::test]
164 async fn test_request_token_sends_oob_callback() {
165 let mut server = mockito::Server::new_async().await;
166 let mock = server
167 .mock("GET", "/getRequestToken")
168 .match_header(
169 "authorization",
170 mockito::Matcher::AllOf(vec![
171 mockito::Matcher::Regex("oauth_callback=\"oob\"".to_string()),
172 mockito::Matcher::Regex("oauth_consumer_key=\"key\"".to_string()),
173 ]),
174 )
175 .with_body("oauth_token=req&oauth_token_secret=reqsecret&oauth_callback_confirmed=true")
176 .create_async()
177 .await;
178
179 let pair = get_request_token(&endpoints(&server), "key", "secret")
180 .await
181 .unwrap();
182 mock.assert_async().await;
183 assert_eq!(
184 pair,
185 TokenPair {
186 token: "req".to_string(),
187 secret: "reqsecret".to_string()
188 }
189 );
190 }
191
192 #[tokio::test]
193 async fn test_access_token_sends_verifier_and_request_token() {
194 let mut server = mockito::Server::new_async().await;
195 let mock = server
196 .mock("GET", "/getAccessToken")
197 .match_header(
198 "authorization",
199 mockito::Matcher::AllOf(vec![
200 mockito::Matcher::Regex("oauth_verifier=\"123456\"".to_string()),
201 mockito::Matcher::Regex("oauth_token=\"req\"".to_string()),
202 ]),
203 )
204 .with_body("oauth_token=access&oauth_token_secret=accesssecret")
205 .create_async()
206 .await;
207
208 let request_token = TokenPair {
209 token: "req".to_string(),
210 secret: "reqsecret".to_string(),
211 };
212 let pair = get_access_token(
213 &endpoints(&server),
214 "key",
215 "secret",
216 &request_token,
217 "123456",
218 )
219 .await
220 .unwrap();
221 mock.assert_async().await;
222 assert_eq!(pair.token, "access");
223 assert_eq!(pair.secret, "accesssecret");
224 }
225
226 #[tokio::test]
227 async fn test_request_token_error_status() {
228 let mut server = mockito::Server::new_async().await;
229 let _mock = server
230 .mock("GET", "/getRequestToken")
231 .with_status(401)
232 .with_body("oauth_problem=consumer_key_unknown")
233 .create_async()
234 .await;
235
236 let err = get_request_token(&endpoints(&server), "bad", "bad")
237 .await
238 .unwrap_err();
239 assert!(format!("{:#}", err).contains("consumer_key_unknown"));
240 }
241}