Skip to main content

dbx_tools_databricks_auth/
oauth.rs

1use std::{collections::HashMap, net::SocketAddr, time::Duration};
2
3use oauth2::{
4    basic::BasicClient, AuthUrl, AuthorizationCode, ClientId, CsrfToken, EndpointNotSet,
5    EndpointSet, PkceCodeChallenge, RedirectUrl, RefreshToken, Scope, TokenUrl,
6};
7use tokio::{
8    io::{AsyncReadExt, AsyncWriteExt},
9    net::TcpListener,
10};
11use url::Url;
12
13use crate::{
14    oauth_endpoints, token::OAuthTokenResponse, Error, OAuthTemplate, OAuthTemplateContext,
15    Profile, Result, Token,
16};
17
18const DEFAULT_PORT: u16 = 8020;
19const MAX_PORT: u16 = 8040;
20
21type OAuthClient =
22    BasicClient<EndpointSet, EndpointNotSet, EndpointNotSet, EndpointNotSet, EndpointSet>;
23
24pub struct OAuthFlow {
25    profile: Profile,
26    http: reqwest::Client,
27    template: OAuthTemplate,
28}
29
30impl OAuthFlow {
31    pub fn new(profile: Profile) -> Result<Self> {
32        let http = reqwest::Client::builder()
33            .redirect(reqwest::redirect::Policy::none())
34            .build()?;
35        Ok(Self {
36            profile,
37            http,
38            template: OAuthTemplate::default(),
39        })
40    }
41
42    /// Set the branding used by the browser callback page.
43    pub fn with_template(mut self, template: OAuthTemplate) -> Self {
44        self.template = template;
45        self
46    }
47
48    pub async fn login(&self, timeout: Duration) -> Result<Token> {
49        let endpoints = oauth_endpoints::resolve(&self.profile, &self.http).await?;
50        let (listener, address) = bind_callback().await?;
51        let redirect = format!("http://localhost:{}", address.port());
52        let client = self.client(&endpoints, &redirect)?;
53        let (challenge, verifier) = PkceCodeChallenge::new_random_sha256();
54        let mut request = client
55            .authorize_url(CsrfToken::new_random)
56            .set_pkce_challenge(challenge);
57        for scope in self.profile.effective_scopes() {
58            request = request.add_scope(Scope::new(scope));
59        }
60        let (authorization_url, csrf) = request.url();
61        if open::that(authorization_url.as_str()).is_err() {
62            eprintln!("Open this URL in a browser:\n{authorization_url}");
63        }
64
65        let callback = tokio::time::timeout(
66            timeout,
67            receive_callback(
68                listener,
69                &redirect,
70                &self.template,
71                Some(self.profile.host.as_str()),
72            ),
73        )
74        .await
75        .map_err(|_| Error::OAuth("timed out waiting for browser authorization".into()))??;
76        if callback.state.as_deref() != Some(csrf.secret()) {
77            return Err(Error::OAuth("OAuth state did not match".into()));
78        }
79        if let Some(error) = callback.error {
80            return Err(Error::OAuth(match callback.error_description {
81                Some(description) => format!("{error}: {description}"),
82                None => error,
83            }));
84        }
85        let code = callback
86            .code
87            .ok_or_else(|| Error::OAuth("authorization callback contained no code".into()))?;
88        let response = client
89            .exchange_code(AuthorizationCode::new(code))
90            .set_pkce_verifier(verifier)
91            .request_async(&self.http)
92            .await
93            .map_err(|error| {
94                Error::OAuth(format!("authorization-code exchange failed: {error}"))
95            })?;
96        Token::from_response(&response, time::OffsetDateTime::now_utc(), None)
97    }
98
99    pub async fn refresh(&self, token: &Token) -> Result<Token> {
100        let endpoints = oauth_endpoints::resolve(&self.profile, &self.http).await?;
101        let client = self.client(&endpoints, "http://localhost:8020")?;
102        let refresh = token
103            .refresh_token()
104            .ok_or_else(|| Error::LoginRequired(self.profile.name.clone()))?;
105        let response: OAuthTokenResponse = client
106            .exchange_refresh_token(&RefreshToken::new(refresh.secret().to_owned()))
107            .request_async(&self.http)
108            .await
109            .map_err(|error| Error::OAuth(format!("refresh-token exchange failed: {error}")))?;
110        Token::from_response(&response, time::OffsetDateTime::now_utc(), Some(token))
111    }
112
113    fn client(
114        &self,
115        endpoints: &oauth_endpoints::AuthorizationServer,
116        redirect: &str,
117    ) -> Result<OAuthClient> {
118        Ok(
119            BasicClient::new(ClientId::new(self.profile.client_id.clone()))
120                .set_auth_uri(
121                    AuthUrl::new(endpoints.authorization_endpoint.clone())
122                        .map_err(|error| Error::OAuth(error.to_string()))?,
123                )
124                .set_token_uri(
125                    TokenUrl::new(endpoints.token_endpoint.clone())
126                        .map_err(|error| Error::OAuth(error.to_string()))?,
127                )
128                .set_redirect_uri(
129                    RedirectUrl::new(redirect.to_owned())
130                        .map_err(|error| Error::OAuth(error.to_string()))?,
131                ),
132        )
133    }
134}
135
136#[derive(Debug)]
137struct Callback {
138    code: Option<String>,
139    state: Option<String>,
140    error: Option<String>,
141    error_description: Option<String>,
142}
143
144async fn bind_callback() -> Result<(TcpListener, SocketAddr)> {
145    for port in DEFAULT_PORT..=MAX_PORT {
146        if let Ok(listener) = TcpListener::bind(("127.0.0.1", port)).await {
147            let address = listener.local_addr()?;
148            return Ok((listener, address));
149        }
150    }
151    Err(Error::OAuth(format!(
152        "no callback port available from {DEFAULT_PORT} through {MAX_PORT}"
153    )))
154}
155
156async fn receive_callback(
157    listener: TcpListener,
158    redirect: &str,
159    template: &OAuthTemplate,
160    host: Option<&str>,
161) -> Result<Callback> {
162    let (mut stream, _) = listener.accept().await?;
163    let mut buffer = vec![0; 16 * 1024];
164    let read = stream.read(&mut buffer).await?;
165    let request = std::str::from_utf8(&buffer[..read])
166        .map_err(|error| Error::OAuth(format!("invalid callback request: {error}")))?;
167    let target = request
168        .lines()
169        .next()
170        .and_then(|line| line.split_whitespace().nth(1))
171        .ok_or_else(|| Error::OAuth("invalid callback request line".into()))?;
172    let url = Url::parse(&format!("{redirect}{target}"))?;
173    let values: HashMap<String, String> = url
174        .query_pairs()
175        .map(|(key, value)| (key.into_owned(), value.into_owned()))
176        .collect();
177    let callback = Callback {
178        code: values.get("code").cloned(),
179        state: values.get("state").cloned(),
180        error: values.get("error").cloned(),
181        error_description: values.get("error_description").cloned(),
182    };
183    let successful = callback.error.is_none() && callback.code.is_some();
184    let body = callback_response(template, host, &callback, successful);
185    let status = if successful {
186        "200 OK"
187    } else {
188        "400 Bad Request"
189    };
190    let response = format!(
191        "HTTP/1.1 {status}\r\nContent-Type: text/html; charset=utf-8\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
192        body.len(),
193    );
194    stream.write_all(response.as_bytes()).await?;
195    Ok(callback)
196}
197
198fn callback_response(
199    template: &OAuthTemplate,
200    host: Option<&str>,
201    callback: &Callback,
202    successful: bool,
203) -> String {
204    let default_error = (!successful && callback.error.is_none()).then_some("authorization_failed");
205    let error = callback.error.as_deref().or(default_error);
206    template.render(OAuthTemplateContext {
207        host,
208        error,
209        error_description: callback.error_description.as_deref(),
210    })
211}