diff --git a/examples/web_server.rs b/examples/web_server.rs index a7fa91f..940e45e 100644 --- a/examples/web_server.rs +++ b/examples/web_server.rs @@ -9,11 +9,11 @@ //! //! The user must have API access enabled to be able to make API calls to Salesforce. //! - use std::process::Command; + use oauth2::http::header::AUTHORIZATION; use oauth2::http::{HeaderMap, HeaderValue}; -use rustsf_auth::{SalesforceCredentials}; +use rustsf_auth::OAuthWebService; pub const CONNECT_TIMEOUT: u64 = 15; pub const REQUEST_TIMEOUT: u64 = 30; @@ -42,7 +42,7 @@ fn open_browser(url: &str) { } #[tokio::main] -async fn main(){ +async fn main() { // The client id, urls and scopes. let login_url = "https://example.my.salesforce.com/"; let client_id = "PlatformCLI"; // The connected app the SFDX CLI is using @@ -57,7 +57,18 @@ async fn main(){ client_secret, redirect_url, scopes, - ); + ) + // Optionally: add a custom callback response the user will see after authentication + .with_callback_response( + r#" + + Authenticated + +

Authentication complete

+

You can close this tab.

+ +"#, + ); let auth_url = web_service.authorization_url().await.unwrap(); // Ask user to authenticate themselves @@ -65,10 +76,6 @@ async fn main(){ open_browser(&auth_url); let session = web_service.connect().await.unwrap(); - // Constructing the authentication session and connecting to Salesforce - // This will open the default browser and wait for the user to complete the authentication process. - let session = config.connect().await.unwrap(); - // Build the headers to include the access token let mut headers = HeaderMap::new(); let auth_value = format!("Bearer {}", session.access_token().await.unwrap()); diff --git a/src/credentials/access_token.rs b/src/credentials/access_token.rs index 894d722..0c3c5d9 100644 --- a/src/credentials/access_token.rs +++ b/src/credentials/access_token.rs @@ -44,8 +44,6 @@ impl SalesforceCredentials { access_token: Some(access_token.into()), refresh_token, instance_url: Some(instance_url.into()), - redirect_uri: None, - scopes: vec![], } } diff --git a/src/credentials/client_credentials.rs b/src/credentials/client_credentials.rs index f1e58ab..6bfe676 100644 --- a/src/credentials/client_credentials.rs +++ b/src/credentials/client_credentials.rs @@ -35,8 +35,6 @@ impl SalesforceCredentials { access_token: None, refresh_token: None, instance_url: None, - redirect_uri: None, - scopes: vec![], } } diff --git a/src/credentials/jwt_bearer.rs b/src/credentials/jwt_bearer.rs index c9ebb50..e48a585 100644 --- a/src/credentials/jwt_bearer.rs +++ b/src/credentials/jwt_bearer.rs @@ -66,8 +66,6 @@ impl SalesforceCredentials { access_token: None, refresh_token: None, instance_url: None, - redirect_uri: None, - scopes: vec![], } } diff --git a/src/credentials/mod.rs b/src/credentials/mod.rs index 144c82d..d7682c9 100644 --- a/src/credentials/mod.rs +++ b/src/credentials/mod.rs @@ -10,7 +10,7 @@ mod access_token; mod client_credentials; mod jwt_bearer; pub(crate) mod sfdx_auth_url; -mod web_server; +pub mod web_server; /// Supported Salesforce OAuth authentication flows. /// @@ -36,9 +36,6 @@ pub enum SalesforceAuthFlow { /// Authenticate from an SFDX auth URL. SfdxUrl, - /// Authenticate using the OAuth 2.0 Web Server Flow. - WebServer, - /// Use an already available Salesforce access token. AccessToken, } @@ -95,12 +92,6 @@ pub struct SalesforceCredentials { /// Salesforce instance URL, for example `https://example.my.salesforce.com`. pub instance_url: Option, - - /// OAuth redirect URI used by the Web Server Flow. - pub redirect_uri: Option, - - /// OAuth scopes requested by the Web Server Flow. - pub scopes: Vec, } impl SalesforceCredentials { @@ -180,7 +171,6 @@ impl SalesforceCredentials { SalesforceAuthFlow::ClientCredentials => self.connect_client_credentials().await, SalesforceAuthFlow::JwtBearer => self.connect_jwt().await, SalesforceAuthFlow::SfdxUrl => self.connect_sfdx_url().await, - SalesforceAuthFlow::WebServer => self.connect_web_server().await, } } diff --git a/src/credentials/sfdx_auth_url.rs b/src/credentials/sfdx_auth_url.rs index 0e14926..1032fa2 100644 --- a/src/credentials/sfdx_auth_url.rs +++ b/src/credentials/sfdx_auth_url.rs @@ -91,8 +91,6 @@ impl SalesforceCredentials { access_token: None, refresh_token: Some(refresh_token), instance_url: None, - redirect_uri: None, - scopes: vec![], }) } diff --git a/src/credentials/web_server.rs b/src/credentials/web_server.rs index 332e2b1..4184e75 100644 --- a/src/credentials/web_server.rs +++ b/src/credentials/web_server.rs @@ -1,51 +1,55 @@ use std::collections::HashMap; -use std::process::Command; +use std::sync::RwLock; use std::time::{SystemTime, UNIX_EPOCH}; use log::trace; +use oauth2::TokenUrl; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::TcpListener; use url::Url; -use crate::credentials::{http_client, required, SalesforceAuthFlow, SalesforceCredentials}; +use crate::credentials::{ + http_client, parse_salesforce_identity_ids, required, SalesforceAuthFlow, SalesforceCredentials, +}; use crate::salesforce_token_response::SalesforceTokenResponse; -use crate::SalesforceAuthError; +use crate::{SalesforceAuthError, SalesforceAuthSession, SalesforceAuthToken}; const DEFAULT_WEB_SERVER_SCOPES: &[&str] = &["api", "refresh_token", "offline_access"]; -impl SalesforceCredentials { - /// Creates a configuration for the OAuth 2.0 Web Server Flow. +const DEFAULT_CALLBACK_RESPONSE: &str = concat!( +"", +"", +"Salesforce Login Complete", +"", +"

Salesforce login complete

", +"

You can close this browser window and return to your application.

", +"", +"" +); + +/// OAuth 2.0 Web Server Flow helper. +/// +/// This type prepares a Salesforce authorization URL, lets the caller decide how +/// to open it, then listens for the OAuth callback and exchanges the received +/// authorization code for a Salesforce session. +#[derive(Debug, Clone)] +pub struct OAuthWebService { + login_url: String, + client_id: String, + client_secret: Option, + redirect_uri: String, + scopes: Vec, + state: String, + callback_response: Option, +} + +impl OAuthWebService { + /// Creates a new OAuth Web Server Flow helper. /// - /// This flow opens the Salesforce authorization URL in the user's browser, - /// starts a temporary local callback server, receives the authorization code, - /// and exchanges it for an access token and refresh token. - /// - /// The connected app must have a callback URL matching `redirect_uri`, for example: - /// - /// `http://localhost:1717/OauthRedirect` - /// - /// # Examples - /// - /// ```rust,no_run - /// use rustsf_auth::{SalesforceAuthFlow, SalesforceCredentials}; - /// - /// # async fn example() -> Result<(), rustsf_auth::SalesforceAuthError> { - /// let config = SalesforceCredentials::web_server( - /// "https://login.salesforce.com", - /// "client-id", - /// Some("client-secret".to_string()), - /// "http://localhost:1717/OauthRedirect", - /// None, - /// ); - /// - /// assert_eq!(config.flow, SalesforceAuthFlow::WebServer); - /// - /// let session = config.connect().await?; - /// println!("{}", session.access_token().await?); - /// # Ok(()) - /// # } - /// ``` - pub fn web_server( + /// The caller should call [`OAuthWebService::authorization_url`] first, open + /// the returned URL in a browser, then call [`OAuthWebService::connect`] to + /// wait for the callback and receive a [`SalesforceAuthSession`]. + pub fn new( login_url: impl Into, client_id: impl Into, client_secret: Option, @@ -53,59 +57,57 @@ impl SalesforceCredentials { scopes: Option>, ) -> Self { Self { - flow: SalesforceAuthFlow::WebServer, - login_url: Some(login_url.into()), - client_id: Some(client_id.into()), + login_url: login_url.into(), + client_id: client_id.into(), client_secret, - username: None, - private_key_pem: None, - access_token: None, - refresh_token: None, - instance_url: None, - redirect_uri: Some(redirect_uri.into()), + redirect_uri: redirect_uri.into(), scopes: scopes.unwrap_or_else(|| { DEFAULT_WEB_SERVER_SCOPES .iter() .map(|scope| (*scope).to_string()) .collect() }), + state: create_state(), + callback_response: None, } } - /// Builds the Salesforce authorization URL for the Web Server Flow. + /// Builds the Salesforce authorization URL. /// - /// This is useful if callers want to present or open the URL themselves. - pub fn web_server_authorization_url(&self, state: &str) -> Result { - let login_url = required(self.login_url.as_deref(), "login_url")?; - let client_id = required(self.client_id.as_deref(), "client_id")?; - let redirect_uri = required(self.redirect_uri.as_deref(), "redirect_uri")?; + /// This method does not open a browser. The caller is responsible for opening + /// the returned URL or presenting it to the user. + pub async fn authorization_url(&self) -> Result { + let normalized = self.login_url.trim_end_matches('/'); + let authorize_url = format!("{normalized}/services/oauth2/authorize"); - let normalized = login_url.trim_end_matches('/'); - let mut url = Url::parse(&format!("{normalized}/services/oauth2/authorize")) + let mut url = Url::parse(&authorize_url) .map_err(|source| SalesforceAuthError::InvalidUrl { - url: format!("{normalized}/services/oauth2/authorize"), + url: authorize_url, source, })?; url.query_pairs_mut() .append_pair("response_type", "code") - .append_pair("client_id", client_id) - .append_pair("redirect_uri", redirect_uri) + .append_pair("client_id", &self.client_id) + .append_pair("redirect_uri", &self.redirect_uri) .append_pair("scope", &self.scopes.join(" ")) - .append_pair("state", state) + .append_pair("state", &self.state) .append_pair("prompt", "login"); Ok(url.to_string()) } - pub(crate) async fn connect_web_server(&self) -> Result { - let state = create_state(); - let auth_url = self.web_server_authorization_url(&state)?; - let redirect_uri = required(self.redirect_uri.as_deref(), "redirect_uri")?; - let callback_url = Url::parse(redirect_uri).map_err(|source| SalesforceAuthError::InvalidUrl { - url: redirect_uri.to_string(), - source, - })?; + /// Starts listening for the OAuth callback and exchanges the authorization + /// code for a Salesforce authentication session. + /// + /// Call [`OAuthWebService::authorization_url`] first and open that URL in a + /// browser before awaiting this method. + pub async fn connect(&self) -> Result { + let callback_url = Url::parse(&self.redirect_uri) + .map_err(|source| SalesforceAuthError::InvalidUrl { + url: self.redirect_uri.clone(), + source, + })?; let host = callback_url.host_str().unwrap_or("127.0.0.1"); let port = callback_url @@ -115,18 +117,16 @@ impl SalesforceCredentials { let bind_host = if host == "localhost" { "127.0.0.1" } else { host }; let listener = TcpListener::bind((bind_host, port)).await?; - trace!("Web Server Flow authorization URL: {}", auth_url); - open_browser(&auth_url); - println!("Open this URL in your browser if it did not open automatically:\n{auth_url}"); + let callback = receive_oauth_callback( + listener, + self.callback_response.as_deref(), + ).await?; - - - let callback = receive_oauth_callback(listener).await?; let callback_state = callback .get("state") .ok_or(SalesforceAuthError::InvalidOAuthCallback)?; - if callback_state != &state { + if callback_state != &self.state { return Err(SalesforceAuthError::OAuthStateMismatch); } @@ -134,23 +134,32 @@ impl SalesforceCredentials { .get("code") .ok_or(SalesforceAuthError::InvalidOAuthCallback)?; - self.exchange_authorization_code(code).await + let token_response = self.exchange_authorization_code(code).await?; + + self.session_from_token_response(token_response) + } + + /// Sets the HTML response returned to the browser after Salesforce redirects + /// back to the local OAuth callback listener. + /// + /// If this method is not called, a default "Salesforce login complete" page is + /// returned. + pub fn with_callback_response(mut self, response: impl Into) -> Self { + self.callback_response = Some(response.into()); + self } async fn exchange_authorization_code( &self, code: &str, ) -> Result { - let client_id = required(self.client_id.as_deref(), "client_id")?; - let redirect_uri = required(self.redirect_uri.as_deref(), "redirect_uri")?; - let url = self.token_url()?.url().clone(); let mut data = vec![ ("grant_type", "authorization_code"), ("code", code), - ("client_id", client_id), - ("redirect_uri", redirect_uri), + ("client_id", self.client_id.as_str()), + ("redirect_uri", self.redirect_uri.as_str()), ]; if let Some(client_secret) = self.client_secret.as_deref() { @@ -177,10 +186,56 @@ impl SalesforceCredentials { Ok(serde_json::from_str::(&body) .map_err(|e| SalesforceAuthError::TokenExchange(e.to_string()))?) } + + fn session_from_token_response( + &self, + token_response: SalesforceTokenResponse, + ) -> Result { + let (org_id, user_id) = parse_salesforce_identity_ids(token_response.id.as_deref()); + + let instance_url = token_response + .instance_url + .clone() + .unwrap_or_else(|| self.login_url.trim_end_matches('/').to_string()); + + let credentials = SalesforceCredentials { + flow: SalesforceAuthFlow::AccessToken, + login_url: Some(self.login_url.clone()), + client_id: Some(self.client_id.clone()), + client_secret: self.client_secret.clone(), + username: None, + private_key_pem: None, + access_token: Some(token_response.access_token.clone()), + refresh_token: token_response.refresh_token.clone(), + instance_url: Some(instance_url.clone()), + }; + + Ok(SalesforceAuthSession { + token: RwLock::new(SalesforceAuthToken { + access_token: token_response.access_token, + token_type: token_response.token_type, + issued_at: token_response.issued_at, + signature: token_response.signature, + }), + credentials, + instance_url, + org_id, + user_id, + }) + } + + fn token_url(&self) -> Result { + let normalized = self.login_url.trim_end_matches('/'); + let url = format!("{normalized}/services/oauth2/token"); + + Ok(TokenUrl::new(url.clone()) + .map_err(|source| SalesforceAuthError::InvalidUrl { url, source })?) + } } async fn receive_oauth_callback( listener: TcpListener, + callback_response: Option<&str>, ) -> Result, SalesforceAuthError> { let (mut stream, _) = listener.accept().await?; @@ -206,19 +261,18 @@ async fn receive_oauth_callback( .map(|(key, value)| (key.to_string(), value.to_string())) .collect::>(); - let response = concat!( - "HTTP/1.1 200 OK\r\n", - "Content-Type: text/html; charset=utf-8\r\n", - "Connection: close\r\n", - "\r\n", - "", - "", - "Salesforce Login Complete", - "", - "

Salesforce login complete

", - "

You can close this browser window and return to your application.

", - "", - "" + let body = callback_response.unwrap_or(DEFAULT_CALLBACK_RESPONSE); + let response = format!( + concat!( + "HTTP/1.1 200 OK\r\n", + "Content-Type: text/html; charset=utf-8\r\n", + "Content-Length: {}\r\n", + "Connection: close\r\n", + "\r\n", + "{}" + ), + body.len(), + body ); stream.write_all(response.as_bytes()).await?; @@ -234,27 +288,4 @@ fn create_state() -> String { .unwrap_or_default(); format!("rustsf-auth-{nanos}") -} - -fn open_browser(url: &str) { - #[cfg(target_os = "windows")] - { - let _ = Command::new("cmd") - .args(["/C", "start", "", url]) - .spawn(); - } - - #[cfg(target_os = "macos")] - { - let _ = Command::new("open") - .arg(url) - .spawn(); - } - - #[cfg(all(unix, not(target_os = "macos")))] - { - let _ = Command::new("xdg-open") - .arg(url) - .spawn(); - } } \ No newline at end of file diff --git a/src/lib.rs b/src/lib.rs index 77386df..6cf27ea 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -10,6 +10,7 @@ use self::salesforce_auth_token::SalesforceAuthToken; pub use self::credentials::{SalesforceAuthFlow, SalesforceCredentials}; pub use self::credentials::sfdx_auth_url::SfdxAuthJson; +pub use self::credentials::web_server::OAuthWebService; /// The default Salesforce production login URL. /// diff --git a/src/main.rs b/src/main.rs index cdac932..5039608 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,7 +1,9 @@ +use std::process::Command; use anyhow::{Context, Result}; use log::LevelFilter; use reqwest::header::{AUTHORIZATION, HeaderMap, HeaderValue}; use rustsf_auth::credentials::SalesforceCredentials; +use rustsf_auth::OAuthWebService; pub const CONNECT_TIMEOUT: u64 = 15; pub const REQUEST_TIMEOUT: u64 = 30; @@ -15,6 +17,29 @@ pub fn get_http_client() -> Result { .context("Failed to build HTTP client")?) } +fn open_browser(url: &str) { + #[cfg(target_os = "windows")] + { + let _ = Command::new("cmd") + .args(["/C", "start", "", url]) + .spawn(); + } + + #[cfg(target_os = "macos")] + { + let _ = Command::new("open") + .arg(url) + .spawn(); + } + + #[cfg(all(unix, not(target_os = "macos")))] + { + let _ = Command::new("xdg-open") + .arg(url) + .spawn(); + } +} + #[tokio::main] async fn main() { println!("Hello, world!"); @@ -53,18 +78,33 @@ async fn main() { ); */ // WEB Flow - let config = SalesforceCredentials::web_server( + let web_service = OAuthWebService::new( "https://computing-platform-9537--qa.sandbox.my.salesforce.com", "PlatformCLI", None, "http://localhost:1717/OauthRedirect", None, - ); + ) + .with_callback_response( + r#" + + Authenticated + +

Authentication complete

+

You can close this tab.

+ +"#, + ); + let auth_url = web_service.authorization_url().await.unwrap(); + // Ask user to authenticate themselves + println!("Open this URL in your browser if it did not open automatically:\n{auth_url}"); + open_browser(&auth_url); + let session = web_service.connect().await.unwrap(); - println!("Config: {:?}", config); +// println!("Config: {:?}", config); - let session = config.connect().await.unwrap(); +// let session = config.connect().await.unwrap(); println!("Instance URL: {}", session.instance_url); println!("Access token: {}", session.access_token().await.unwrap());