use async_trait::async_trait;
use error_stack::{Report, ResultExt};
use oauth2::{
basic::{BasicClient, BasicTokenResponse},
reqwest::async_http_client,
AuthorizationCode, CsrfToken, PkceCodeVerifier,
};
use url::Url;
use super::OAuthError;
use crate::config::prefixed_env_var;
mod github;
mod google;
mod twitter;
pub use github::*;
pub use google::*;
pub use twitter::*;
#[derive(Default, Debug)]
pub struct OAuthUserDetails {
pub login_id: String,
pub name: Option<String>,
pub email: Option<String>,
pub avatar_url: Option<Url>,
pub twitter_id: Option<String>,
}
pub struct AuthorizeUrl {
pub url: Url,
pub state: CsrfToken,
pub pkce_verifier: Option<PkceCodeVerifier>,
}
#[async_trait]
pub trait OAuthProvider: Send + Sync + 'static {
fn name(&self) -> &'static str;
fn client(&self) -> &BasicClient;
fn authorize_url(&self) -> AuthorizeUrl;
async fn fetch_access_token(
&self,
authorization_code: String,
pkce_verifier: String,
) -> Result<BasicTokenResponse, Report<OAuthError>>;
async fn fetch_user_details(
&self,
client: reqwest::Client,
access_token: &str,
) -> Result<OAuthUserDetails, reqwest::Error>;
}
pub async fn fetch_access_token_simple(
client: &BasicClient,
authorization_code: String,
) -> Result<BasicTokenResponse, Report<OAuthError>> {
client
.exchange_code(AuthorizationCode::new(authorization_code))
.request_async(async_http_client)
.await
.change_context(OAuthError::ExchangeError)
}
pub async fn fetch_access_token_with_pkce(
client: &BasicClient,
authorization_code: String,
pkce_verifier: String,
) -> Result<BasicTokenResponse, Report<OAuthError>> {
let verifier = PkceCodeVerifier::new(pkce_verifier);
client
.exchange_code(AuthorizationCode::new(authorization_code))
.set_pkce_verifier(verifier)
.request_async(async_http_client)
.await
.change_context(OAuthError::ExchangeError)
}
pub fn build_redirect_url(base: &str, provider_name: &str) -> String {
format!("{base}/{provider_name}/callback")
}
pub fn create_supported_providers(
env_prefix: &str,
redirect_base_url: &str,
) -> Vec<Box<dyn OAuthProvider>> {
let github_provider = match (
prefixed_env_var(env_prefix, "OAUTH_GITHUB_CLIENT_ID"),
prefixed_env_var(env_prefix, "OAUTH_GITHUB_CLIENT_SECRET"),
) {
(Ok(client_id), Ok(client_secret)) => Some(Box::new(GitHubOAuthProvider::new(
client_id,
client_secret,
redirect_base_url,
)) as Box<dyn OAuthProvider>),
_ => None,
};
let google_provider = match (
prefixed_env_var(env_prefix, "OAUTH_GOOGLE_CLIENT_ID"),
prefixed_env_var(env_prefix, "OAUTH_GOOGLE_CLIENT_SECRET"),
) {
(Ok(client_id), Ok(client_secret)) => Some(Box::new(GoogleOAuthProvider::new(
client_id,
client_secret,
redirect_base_url,
)) as Box<dyn OAuthProvider>),
_ => None,
};
let twitter_provider = match (
prefixed_env_var(env_prefix, "OAUTH_TWITTER_CLIENT_ID"),
prefixed_env_var(env_prefix, "OAUTH_TWITTER_CLIENT_SECRET"),
) {
(Ok(client_id), Ok(client_secret)) => Some(Box::new(TwitterOAuthProvider::new(
client_id,
client_secret,
redirect_base_url,
)) as Box<dyn OAuthProvider>),
_ => None,
};
[github_provider, google_provider, twitter_provider]
.into_iter()
.flatten()
.collect()
}