dbx_tools_databricks_auth/
oauth.rs1use 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 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}