Skip to main content

smugmug_cli/api/
oauth_flow.rs

1//! OAuth 1.0a three-legged sign-in for SmugMug, used to obtain an access
2//! token pair from just the API key and secret. Uses the out-of-band ("oob")
3//! callback: the user approves access in a browser and SmugMug shows a
4//! 6-digit verifier code to type back into the CLI.
5
6use anyhow::{Context, Result};
7use oauth1_request as oauth;
8
9const OAUTH_BASE: &str = "https://api.smugmug.com/services/oauth/1.0a";
10
11/// A token/secret pair returned by one of the OAuth token endpoints.
12#[derive(Debug, Clone, PartialEq)]
13pub struct TokenPair {
14    pub token: String,
15    pub secret: String,
16}
17
18/// Endpoints for the sign-in flow; overridable so tests can point at a mock
19/// server.
20pub 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
36/// Step 1: get a temporary request token, signed with only the API key and
37/// secret.
38pub 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
55/// Step 2: the URL where the user approves access. Full access with modify
56/// permission is needed for uploads and deletes.
57pub 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
66/// Step 3: exchange the request token plus the verifier code the user was
67/// shown for a long-lived access token.
68pub 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
104/// Parse an `oauth_token=...&oauth_token_secret=...` form-encoded body.
105fn 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}