use std::{fmt::Write as _, time::Duration};
use async_trait::async_trait;
use serde::Deserialize;
use zeroize::Zeroizing;
use crate::{
error::{AuthError, Result},
oidc_provider::validate_oauth_endpoint_url,
provider::{OAuthProvider, TokenResponse, UserInfo},
};
const DISCORD_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
const MAX_DISCORD_RESPONSE_BYTES: usize = 1024 * 1024;
const DEFAULT_BASE_URL: &str = "https://discord.com";
const AVATAR_CDN: &str = "https://cdn.discordapp.com";
const DISCORD_SCOPES: &str = "identify email";
#[derive(Debug, Clone, Deserialize)]
pub struct DiscordUser {
pub id: String,
pub username: String,
#[serde(default)]
pub global_name: Option<String>,
#[serde(default)]
pub email: Option<String>,
#[serde(default)]
pub verified: Option<bool>,
#[serde(default)]
pub avatar: Option<String>,
}
#[derive(Debug, Deserialize)]
struct DiscordTokenResponse {
access_token: String,
#[serde(default)]
token_type: Option<String>,
#[serde(default)]
expires_in: Option<u64>,
#[serde(default)]
refresh_token: Option<String>,
}
pub struct DiscordOAuth {
client_id: String,
client_secret: Zeroizing<String>,
redirect_uri: String,
base_url: String,
client: reqwest::Client,
}
impl std::fmt::Debug for DiscordOAuth {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("DiscordOAuth")
.field("client_id", &self.client_id)
.field("redirect_uri", &self.redirect_uri)
.field("base_url", &self.base_url)
.finish_non_exhaustive() }
}
#[must_use]
pub fn select_linkable_email(user: &DiscordUser) -> (Option<String>, bool) {
let email = user.email.clone().filter(|e| !e.trim().is_empty());
let verified = email.is_some() && user.verified.unwrap_or(false);
(email, verified)
}
impl DiscordOAuth {
pub fn new(client_id: String, client_secret: String, redirect_uri: String) -> Result<Self> {
Self::with_base_url(client_id, client_secret, redirect_uri, DEFAULT_BASE_URL.to_string())
}
pub fn with_base_url(
client_id: String,
client_secret: String,
redirect_uri: String,
base_url: String,
) -> Result<Self> {
validate_oauth_endpoint_url(&base_url)?;
let client =
reqwest::Client::builder()
.timeout(DISCORD_REQUEST_TIMEOUT)
.build()
.map_err(|e| AuthError::ConfigError {
message: format!("Failed to create HTTP client: {e}"),
})?;
Ok(Self {
client_id,
client_secret: Zeroizing::new(client_secret),
redirect_uri,
base_url: base_url.trim_end_matches('/').to_string(),
client,
})
}
async fn post_token(&self, params: &[(&str, &str)]) -> Result<TokenResponse> {
let resp = self
.client
.post(format!("{}/api/oauth2/token", self.base_url))
.form(params)
.send()
.await
.map_err(|e| AuthError::OAuthError {
message: format!("Discord token endpoint request failed: {e}"),
})?;
let status = resp.status();
let bytes = resp.bytes().await.map_err(|e| AuthError::OAuthError {
message: format!("Failed to read Discord token response: {e}"),
})?;
if !status.is_success() {
return Err(AuthError::OAuthError {
message: format!("Discord token endpoint returned HTTP {status}"),
});
}
if bytes.len() > MAX_DISCORD_RESPONSE_BYTES {
return Err(AuthError::OAuthError {
message: format!("Discord token response too large ({} bytes)", bytes.len()),
});
}
let response: DiscordTokenResponse =
serde_json::from_slice(&bytes).map_err(|e| AuthError::OAuthError {
message: format!("Failed to parse Discord token response: {e}"),
})?;
Ok(TokenResponse {
access_token: response.access_token,
refresh_token: response.refresh_token,
expires_in: response.expires_in.unwrap_or(0),
token_type: response.token_type.unwrap_or_else(|| "Bearer".to_string()),
id_token: None,
})
}
pub async fn get_user(&self, access_token: &str) -> Result<DiscordUser> {
let resp = self
.client
.get(format!("{}/api/users/@me", self.base_url))
.header("Authorization", format!("Bearer {access_token}"))
.send()
.await
.map_err(|e| AuthError::OAuthError {
message: format!("Failed to fetch Discord user: {e}"),
})?;
let status = resp.status();
let bytes = resp.bytes().await.map_err(|e| AuthError::OAuthError {
message: format!("Failed to read Discord user response: {e}"),
})?;
if !status.is_success() {
return Err(AuthError::OAuthError {
message: format!("Discord users/@me returned HTTP {status}"),
});
}
if bytes.len() > MAX_DISCORD_RESPONSE_BYTES {
return Err(AuthError::OAuthError {
message: format!("Discord user response too large ({} bytes)", bytes.len()),
});
}
serde_json::from_slice(&bytes).map_err(|e| AuthError::OAuthError {
message: format!("Failed to parse Discord user: {e}"),
})
}
}
#[async_trait]
impl OAuthProvider for DiscordOAuth {
fn name(&self) -> &'static str {
"discord"
}
fn authorization_url(&self, state: &str) -> String {
let mut url = format!("{}/oauth2/authorize", self.base_url);
write!(
url,
"?client_id={}&redirect_uri={}&state={}&response_type=code&scope={}",
urlencoding::encode(&self.client_id),
urlencoding::encode(&self.redirect_uri),
urlencoding::encode(state),
urlencoding::encode(DISCORD_SCOPES),
)
.expect("write to String is infallible");
url
}
async fn exchange_code(&self, code: &str) -> Result<TokenResponse> {
self.post_token(&[
("client_id", self.client_id.as_str()),
("client_secret", self.client_secret.as_str()),
("grant_type", "authorization_code"),
("code", code),
("redirect_uri", self.redirect_uri.as_str()),
])
.await
}
async fn user_info(&self, access_token: &str) -> Result<UserInfo> {
let user = self.get_user(access_token).await?;
let (email, email_verified) = select_linkable_email(&user);
let mut raw_claims = serde_json::Map::new();
raw_claims.insert("discord_id".to_string(), serde_json::json!(user.id));
raw_claims.insert("discord_username".to_string(), serde_json::json!(user.username));
if let Some(ref email) = email {
raw_claims.insert("email".to_string(), serde_json::json!(email));
}
raw_claims.insert("email_verified".to_string(), serde_json::json!(email_verified));
let picture = user
.avatar
.as_ref()
.map(|hash| format!("{AVATAR_CDN}/avatars/{}/{hash}.png", user.id));
Ok(UserInfo {
id: user.id.clone(),
email,
email_verified,
name: user.global_name.clone().or_else(|| Some(user.username.clone())),
picture,
raw_claims: serde_json::Value::Object(raw_claims),
})
}
async fn refresh_token(&self, refresh_token: &str) -> Result<TokenResponse> {
self.post_token(&[
("client_id", self.client_id.as_str()),
("client_secret", self.client_secret.as_str()),
("grant_type", "refresh_token"),
("refresh_token", refresh_token),
])
.await
}
}