From 5c141c6c1b6e76287c5f70e2bd4688f23e12637a Mon Sep 17 00:00:00 2001 From: Adrian Kumpf <8999358+adriankumpf@users.noreply.github.com> Date: Fri, 17 Sep 2021 15:28:57 +0200 Subject: [PATCH] Handle errors --- Cargo.lock | 1 + Cargo.toml | 1 + src/auth.rs | 67 ++++++++++++-------------- src/main.rs | 119 +++++++++++++++++++++++------------------------ views/index.html | 2 - 5 files changed, 89 insertions(+), 101 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 346b0cd..43c8023 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1824,6 +1824,7 @@ dependencies = [ name = "tesla_auth" version = "0.1.0" dependencies = [ + "anyhow", "log", "oauth2", "reqwest", diff --git a/Cargo.toml b/Cargo.toml index 7092475..ec01fa7 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -24,6 +24,7 @@ osx_frameworks = [] # See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html [dependencies] +anyhow = "1.0.44" log = "0.4.14" oauth2 = "4.1.0" reqwest = { version = "0.11.4" , default-features = false, features = ["json", "rustls-tls"]} diff --git a/src/auth.rs b/src/auth.rs index 1c3eb93..b55177b 100644 --- a/src/auth.rs +++ b/src/auth.rs @@ -1,5 +1,7 @@ use std::collections::HashMap; +use anyhow::anyhow; + use oauth2::basic::BasicClient; use oauth2::reqwest::http_client; use oauth2::url::Url; @@ -37,14 +39,15 @@ pub struct Tokens { } pub struct Client { + auth_url: Url, oauth_client: BasicClient, - pkce_verifier: Option, - csrf_token: Option, + pkce_verifier: PkceCodeVerifier, + csrf_token: CsrfToken, } impl Client { - pub fn new() -> Self { - let client = BasicClient::new( + pub fn new() -> Client { + let oauth_client = BasicClient::new( ClientId::new(CLIENT_ID.to_string()), None, AuthUrl::new(AUTH_URL.to_string()).unwrap(), @@ -53,18 +56,9 @@ impl Client { .set_auth_type(AuthType::RequestBody) .set_redirect_uri(RedirectUrl::new(REDIRECT_URL.to_string()).unwrap()); - Client { - oauth_client: client, - pkce_verifier: None, - csrf_token: None, - } - } - - pub fn authorization_url(&mut self) -> Url { let (pkce_challenge, pkce_verifier) = PkceCodeChallenge::new_random_sha256(); - let (auth_url, csrf_token) = self - .oauth_client + let (auth_url, csrf_token) = oauth_client .authorize_url(CsrfToken::new_random) .add_scope(Scope::new("openid".to_string())) .add_scope(Scope::new("email".to_string())) @@ -72,39 +66,40 @@ impl Client { .set_pkce_challenge(pkce_challenge) .url(); - self.pkce_verifier = Some(pkce_verifier); - self.csrf_token = Some(csrf_token); - - auth_url + Client { + oauth_client, + auth_url, + pkce_verifier, + csrf_token, + } } - pub fn verify_csrf_state(&self, state: String) { - let csrf_token = self.csrf_token.as_ref().unwrap(); - assert_eq!(&state, csrf_token.secret()); + pub fn authorize_url(&self) -> Url { + self.auth_url.clone() } - pub fn retrieve_tokens(&mut self, code: &str) -> Tokens { - let code = AuthorizationCode::new(code.to_string()); - let pkce_verifier = self.pkce_verifier.take().unwrap(); + pub fn retrieve_tokens(self, code: &str, state: &str) -> anyhow::Result { + if state != self.csrf_token.secret() { + return Err(anyhow!("CSRF state does not match!")); + } let tokens = self .oauth_client - .exchange_code(code) - .set_pkce_verifier(pkce_verifier) - .request(http_client) - .unwrap(); + .exchange_code(AuthorizationCode::new(code.to_string())) + .set_pkce_verifier(self.pkce_verifier) + .request(http_client)?; - let access_token = exchange_sso_access_token(tokens.access_token()); + let access_token = exchange_sso_access_token(tokens.access_token())?; let refresh_token = tokens.refresh_token().unwrap().secret().to_string(); - Tokens { + Ok(Tokens { access: access_token, refresh: refresh_token, - } + }) } } -fn exchange_sso_access_token(access_token: &AccessToken) -> String { +fn exchange_sso_access_token(access_token: &AccessToken) -> anyhow::Result { let mut body = HashMap::new(); body.insert("grant_type", "urn:ietf:params:oauth:grant-type:jwt-bearer"); body.insert("client_id", SSO_CLIENT_ID); @@ -114,10 +109,8 @@ fn exchange_sso_access_token(access_token: &AccessToken) -> String { .post(SSO_TOKEN_URL) .header(AUTHORIZATION, format!("Bearer {}", access_token.secret())) .json(&body) - .send() - .unwrap() - .json() - .unwrap(); + .send()? + .json()?; - tokens.access_token + Ok(tokens.access_token) } diff --git a/src/main.rs b/src/main.rs index a8d137e..b152bfa 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,19 +1,21 @@ mod auth; use std::collections::HashMap; -use std::sync::mpsc::{channel, Receiver}; +use std::sync::mpsc::channel; use std::thread; -use log::{debug, info, LevelFilter}; +use anyhow::anyhow; + +use log::{debug, error, info, LevelFilter}; use simple_logger::SimpleLogger; use oauth2::url::Url; use wry::application::accelerator::{Accelerator, SysMods}; use wry::application::event::{Event, WindowEvent}; -use wry::application::event_loop::{ControlFlow, EventLoop, EventLoopProxy}; +use wry::application::event_loop::{ControlFlow, EventLoop}; use wry::application::keyboard::KeyCode; -use wry::application::menu::{CustomMenuItem, MenuBar, MenuItem, MenuItemAttributes, MenuType}; +use wry::application::menu::{MenuBar, MenuId, MenuItem, MenuItemAttributes, MenuType}; use wry::application::window::{Window, WindowBuilder}; use wry::http::{Request, Response, ResponseBuilder}; use wry::webview::{RpcRequest, WebViewBuilder}; @@ -25,9 +27,8 @@ const INITIALIZATION_SCRIPT: &str = r#" if (url.startsWith("https://auth.tesla.com/void/callback")) { location.replace("wry://index.html?access=loading...&refresh=loading..."); + rpc.call('url', url); } - - rpc.call('url', url); }); "#; @@ -36,36 +37,50 @@ enum CustomEvent { Tokens(auth::Tokens), } -fn main() -> wry::Result<()> { +fn main() -> anyhow::Result<()> { SimpleLogger::new() .with_level(LevelFilter::Off) .with_module_level("reqwest", LevelFilter::Debug) .with_module_level("tesla_auth", LevelFilter::Debug) - .init() - .unwrap(); + .init()?; let event_loop = EventLoop::::with_user_event(); let event_proxy = event_loop.create_proxy(); + let client = auth::Client::new(); + let auth_url = client.authorize_url(); + let (tx, rx) = channel(); let handler = move |_window: &Window, req: RpcRequest| { - if req.method == "url" { - let url = parse_url(req.params.unwrap()); - tx.send(url).unwrap(); + if let ("url", Some(params)) = (req.method.as_str(), req.params) { + if let Ok(url) = parse_url(params) { + tx.send(url).unwrap(); + } } None }; - let mut client = auth::Client::new(); - let auth_url = client.authorization_url(); - thread::spawn(move || { - handle_url_changes(rx, client, event_proxy); + while let Ok(url) = rx.recv() { + if auth::is_redirect_url(&url) { + let query: HashMap<_, _> = url.query_pairs().collect(); + + let state = query.get("state").expect("No state parameter found"); + let code = query.get("code").expect("No code parameter found"); + + match client.retrieve_tokens(code, state) { + Ok(tokens) => event_proxy.send_event(CustomEvent::Tokens(tokens)).unwrap(), + Err(e) => error!("{}", e), + }; + + break; + } + } }); - let (menu, quit_item) = build_menu(); + let (menu, quit_id) = build_menu(); let window = WindowBuilder::new() .with_title("Tesla Auth") @@ -89,6 +104,7 @@ fn main() -> wry::Result<()> { event: WindowEvent::CloseRequested, .. } => *control_flow = ControlFlow::Exit, + Event::UserEvent(CustomEvent::Tokens(tokens)) => { info!("Received tokens: {:#?}", tokens); @@ -99,22 +115,26 @@ fn main() -> wry::Result<()> { webview.evaluate_script(&url).unwrap(); } + Event::MenuEvent { menu_id, origin: MenuType::MenuBar, .. } => { - if menu_id == quit_item.clone().id() { - *control_flow = ControlFlow::Exit; - } - println!("Clicked on {:?}", menu_id); + debug!("Clicked on {:?}", menu_id); + + match menu_id { + id if id == quit_id => *control_flow = ControlFlow::Exit, + _ => (), + }; } + _ => (), } }); } -fn build_menu() -> (MenuBar, CustomMenuItem) { +fn build_menu() -> (MenuBar, MenuId) { let mut menu_bar_menu = MenuBar::new(); let mut menu = MenuBar::new(); @@ -131,35 +151,7 @@ fn build_menu() -> (MenuBar, CustomMenuItem) { menu_bar_menu.add_submenu("First menu", true, menu); - (menu_bar_menu, quit_item) -} - -fn handle_url_changes( - rx: Receiver, - mut client: auth::Client, - event_proxy: EventLoopProxy, -) { - let mut tokens_retrieved = false; - - while let Ok(url) = rx.recv() { - if !auth::is_redirect_url(&url) || tokens_retrieved { - debug!("URL changed: {}", &url); - continue; - } - - let query: HashMap<_, _> = url.query_pairs().collect(); - - let state = query.get("state").expect("No state parameter found"); - let code = query.get("code").expect("No code parameter found"); - - client.verify_csrf_state(state.to_string()); - - let tokens = client.retrieve_tokens(code); - - tokens_retrieved = true; - - event_proxy.send_event(CustomEvent::Tokens(tokens)).unwrap(); - } + (menu_bar_menu, quit_item.id()) } fn protocol_handler(request: &Request) -> wry::Result { @@ -169,23 +161,26 @@ fn protocol_handler(request: &Request) -> wry::Result { Some("index.html") => { let query = url.query_pairs().collect::>(); - let (access, refresh) = (query.get("access").unwrap(), query.get("refresh").unwrap()); + let content = match (query.get("access"), query.get("refresh")) { + (Some(access), Some(refresh)) => include_str!("../views/index.html") + .replace("{access_token}", access) + .replace("{refresh_token}", refresh) + .as_bytes() + .to_vec(), - let content = include_str!("../views/index.html") - .replace("{access_token}", access) - .replace("{refresh_token}", refresh); + (_, _) => vec![], + }; - ResponseBuilder::new() - .mimetype("text/html") - .body(content.as_bytes().to_vec()) + ResponseBuilder::new().mimetype("text/html").body(content) } domain => unimplemented!("Cannot open {:?}", domain), } } -fn parse_url(params: Value) -> Url { - let args = serde_json::from_value::>(params).unwrap(); - let url = args.first().unwrap(); - Url::parse(url).expect("Invalid URL") +fn parse_url(params: Value) -> anyhow::Result { + match &serde_json::from_value::>(params)?[..] { + [url] => Ok(Url::parse(url)?), + _ => Err(anyhow!("Invalid url param!")), + } } diff --git a/views/index.html b/views/index.html index 12047d3..e7452ce 100644 --- a/views/index.html +++ b/views/index.html @@ -32,8 +32,6 @@
-

Tokens generated!

-