Handle errors

This commit is contained in:
Adrian Kumpf
2021-09-17 15:28:57 +02:00
parent 457746b6ca
commit 5c141c6c1b
5 changed files with 89 additions and 101 deletions
Generated
+1
View File
@@ -1824,6 +1824,7 @@ dependencies = [
name = "tesla_auth"
version = "0.1.0"
dependencies = [
"anyhow",
"log",
"oauth2",
"reqwest",
+1
View File
@@ -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"]}
+30 -37
View File
@@ -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<PkceCodeVerifier>,
csrf_token: Option<CsrfToken>,
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<Tokens> {
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<String> {
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)
}
+57 -62
View File
@@ -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::<CustomEvent>::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<Url>,
mut client: auth::Client,
event_proxy: EventLoopProxy<CustomEvent>,
) {
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<Response> {
@@ -169,23 +161,26 @@ fn protocol_handler(request: &Request) -> wry::Result<Response> {
Some("index.html") => {
let query = url.query_pairs().collect::<HashMap<_, _>>();
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::<Vec<String>>(params).unwrap();
let url = args.first().unwrap();
Url::parse(url).expect("Invalid URL")
fn parse_url(params: Value) -> anyhow::Result<Url> {
match &serde_json::from_value::<Vec<String>>(params)?[..] {
[url] => Ok(Url::parse(url)?),
_ => Err(anyhow!("Invalid url param!")),
}
}
-2
View File
@@ -32,8 +32,6 @@
</head>
<body>
<main>
<h1>Tokens generated!</h1>
<form>
<div>
<label for="access_token">Access Token:</label><br />