use std::{
fmt::Write as _,
sync::RwLock,
time::{Duration, SystemTime, UNIX_EPOCH},
};
use async_trait::async_trait;
use base64::Engine as _;
use serde::{Deserialize, Serialize};
use zeroize::Zeroizing;
use crate::{
error::{AuthError, Result},
oidc_provider::validate_oauth_endpoint_url,
provider::{OAuthProvider, TokenResponse, UserInfo},
};
const APPLE_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
const MAX_APPLE_RESPONSE_BYTES: usize = 1024 * 1024;
const DEFAULT_BASE_URL: &str = "https://appleid.apple.com";
const CLIENT_SECRET_AUDIENCE: &str = "https://appleid.apple.com";
const APPLE_SCOPES: &str = "name email";
const CLIENT_SECRET_TTL_SECS: u64 = 3600;
const CLIENT_SECRET_RENEW_SKEW_SECS: u64 = 300;
const PRIVATE_RELAY_DOMAIN: &str = "@privaterelay.appleid.com";
#[derive(Debug, Serialize)]
struct ClientSecretClaims<'a> {
iss: &'a str,
iat: u64,
exp: u64,
aud: &'a str,
sub: &'a str,
}
#[derive(Debug, Clone)]
struct CachedClientSecret {
assertion: String,
expires_at: u64,
}
#[derive(Debug, Deserialize)]
#[serde(untagged)]
enum Audience {
One(String),
Many(Vec<String>),
}
impl Audience {
fn contains(&self, value: &str) -> bool {
match self {
Self::One(a) => a == value,
Self::Many(all) => all.iter().any(|a| a == value),
}
}
}
#[derive(Debug, Deserialize)]
#[serde(untagged)]
enum AppleBool {
Bool(bool),
Str(String),
}
impl AppleBool {
fn as_bool(&self) -> bool {
match self {
Self::Bool(b) => *b,
Self::Str(s) => s.eq_ignore_ascii_case("true"),
}
}
}
#[derive(Debug, Deserialize)]
struct AppleIdTokenClaims {
iss: String,
aud: Audience,
sub: String,
exp: u64,
#[serde(default)]
email: Option<String>,
#[serde(default)]
email_verified: Option<AppleBool>,
#[serde(default)]
is_private_email: Option<AppleBool>,
}
#[derive(Debug, Deserialize)]
struct AppleTokenResponse {
access_token: String,
#[serde(default)]
token_type: Option<String>,
#[serde(default)]
expires_in: Option<u64>,
#[serde(default)]
refresh_token: Option<String>,
#[serde(default)]
id_token: Option<String>,
}
#[derive(Debug, Clone, Default, Deserialize)]
pub struct AppleFirstAuthUser {
#[serde(default)]
pub name: Option<AppleFirstAuthName>,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AppleFirstAuthName {
#[serde(default)]
pub first_name: Option<String>,
#[serde(default)]
pub last_name: Option<String>,
}
impl AppleFirstAuthUser {
#[must_use]
pub fn parse(raw: &str) -> Option<Self> {
serde_json::from_str(raw).ok()
}
#[must_use]
pub fn display_name(&self) -> Option<String> {
let name = self.name.as_ref()?;
let joined = [name.first_name.as_deref(), name.last_name.as_deref()]
.into_iter()
.flatten()
.map(str::trim)
.filter(|part| !part.is_empty())
.collect::<Vec<_>>()
.join(" ");
(!joined.is_empty()).then_some(joined)
}
}
#[must_use]
pub fn is_private_relay_email(email: &str) -> bool {
email.trim().to_ascii_lowercase().ends_with(PRIVATE_RELAY_DOMAIN)
}
pub struct AppleOAuth {
client_id: String,
team_id: String,
key_id: String,
private_key: Zeroizing<String>,
redirect_uri: String,
base_url: String,
client: reqwest::Client,
cached_secret: RwLock<Option<CachedClientSecret>>,
}
impl std::fmt::Debug for AppleOAuth {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AppleOAuth")
.field("client_id", &self.client_id)
.field("team_id", &self.team_id)
.field("key_id", &self.key_id)
.field("redirect_uri", &self.redirect_uri)
.field("base_url", &self.base_url)
.finish_non_exhaustive() }
}
fn now_secs() -> Result<u64> {
SystemTime::now().duration_since(UNIX_EPOCH).map(|d| d.as_secs()).map_err(|e| {
AuthError::ConfigError {
message: format!("system clock is before the Unix epoch: {e}"),
}
})
}
impl AppleOAuth {
pub fn new(
client_id: String,
team_id: String,
key_id: String,
private_key_pem: String,
redirect_uri: String,
) -> Result<Self> {
Self::with_base_url(
client_id,
team_id,
key_id,
private_key_pem,
redirect_uri,
DEFAULT_BASE_URL.to_string(),
)
}
pub fn with_base_url(
client_id: String,
team_id: String,
key_id: String,
private_key_pem: String,
redirect_uri: String,
base_url: String,
) -> Result<Self> {
validate_oauth_endpoint_url(&base_url)?;
jsonwebtoken::EncodingKey::from_ec_pem(private_key_pem.as_bytes()).map_err(|e| {
AuthError::ConfigError {
message: format!(
"[auth.social.apple] private key is not a usable ES256 key — Apple issues a \
PKCS#8 PEM `.p8` file: {e}"
),
}
})?;
let client =
reqwest::Client::builder().timeout(APPLE_REQUEST_TIMEOUT).build().map_err(|e| {
AuthError::ConfigError {
message: format!("Failed to create HTTP client: {e}"),
}
})?;
Ok(Self {
client_id,
team_id,
key_id,
private_key: Zeroizing::new(private_key_pem),
redirect_uri,
base_url: base_url.trim_end_matches('/').to_string(),
client,
cached_secret: RwLock::new(None),
})
}
fn mint_client_secret(&self, now: u64) -> Result<CachedClientSecret> {
let expires_at = now + CLIENT_SECRET_TTL_SECS;
let mut header = jsonwebtoken::Header::new(jsonwebtoken::Algorithm::ES256);
header.kid = Some(self.key_id.clone());
let claims = ClientSecretClaims {
iss: &self.team_id,
iat: now,
exp: expires_at,
aud: CLIENT_SECRET_AUDIENCE,
sub: &self.client_id,
};
let key =
jsonwebtoken::EncodingKey::from_ec_pem(self.private_key.as_bytes()).map_err(|e| {
AuthError::ConfigError {
message: format!(
"[auth.social.apple] private key is not a usable ES256 key: {e}"
),
}
})?;
let assertion =
jsonwebtoken::encode(&header, &claims, &key).map_err(|e| AuthError::ConfigError {
message: format!(
"[auth.social.apple] client-secret assertion could not be signed: {e}"
),
})?;
Ok(CachedClientSecret {
assertion,
expires_at,
})
}
pub fn client_secret(&self) -> Result<String> {
let now = now_secs()?;
if let Ok(guard) = self.cached_secret.read() {
if let Some(cached) = guard.as_ref() {
if cached.expires_at > now + CLIENT_SECRET_RENEW_SKEW_SECS {
return Ok(cached.assertion.clone());
}
}
}
let minted = self.mint_client_secret(now)?;
if let Ok(mut guard) = self.cached_secret.write() {
*guard = Some(minted.clone());
}
Ok(minted.assertion)
}
async fn post_token(&self, params: &[(&str, &str)]) -> Result<TokenResponse> {
let resp = self
.client
.post(format!("{}/auth/token", self.base_url))
.form(params)
.send()
.await
.map_err(|e| AuthError::OAuthError {
message: format!("Apple token endpoint request failed: {e}"),
})?;
let status = resp.status();
let bytes = resp.bytes().await.map_err(|e| AuthError::OAuthError {
message: format!("Failed to read Apple token response: {e}"),
})?;
if !status.is_success() {
return Err(AuthError::OAuthError {
message: format!("Apple token endpoint returned HTTP {status}"),
});
}
if bytes.len() > MAX_APPLE_RESPONSE_BYTES {
return Err(AuthError::OAuthError {
message: format!("Apple token response too large ({} bytes)", bytes.len()),
});
}
let response: AppleTokenResponse =
serde_json::from_slice(&bytes).map_err(|e| AuthError::OAuthError {
message: format!("Failed to parse Apple 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: response.id_token,
})
}
fn decode_id_token(&self, id_token: &str, now: u64) -> Result<AppleIdTokenClaims> {
let payload = id_token.split('.').nth(1).ok_or_else(|| AuthError::OAuthError {
message: "Apple id_token is not a JWT".to_string(),
})?;
let bytes =
base64::engine::general_purpose::URL_SAFE_NO_PAD.decode(payload).map_err(|e| {
AuthError::OAuthError {
message: format!("Apple id_token payload is not base64url: {e}"),
}
})?;
let claims: AppleIdTokenClaims =
serde_json::from_slice(&bytes).map_err(|e| AuthError::OAuthError {
message: format!("Apple id_token claims do not parse: {e}"),
})?;
if claims.iss.trim_end_matches('/') != self.base_url {
return Err(AuthError::OAuthError {
message: format!(
"Apple id_token issuer {} does not match the configured Apple ID service",
claims.iss
),
});
}
if !claims.aud.contains(&self.client_id) {
return Err(AuthError::OAuthError {
message: "Apple id_token audience does not name this client".to_string(),
});
}
if claims.exp <= now {
return Err(AuthError::OAuthError {
message: "Apple id_token has expired".to_string(),
});
}
Ok(claims)
}
pub fn user_info_from_id_token(&self, id_token: &str) -> Result<UserInfo> {
let now = now_secs()?;
let claims = self.decode_id_token(id_token, now)?;
let email_verified = claims.email_verified.as_ref().is_some_and(AppleBool::as_bool);
let is_private_email = claims.is_private_email.as_ref().is_some_and(AppleBool::as_bool);
let email = claims.email.filter(|e| !e.trim().is_empty());
let mut raw_claims = serde_json::Map::new();
raw_claims.insert("apple_sub".to_string(), serde_json::json!(claims.sub));
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));
raw_claims.insert("is_private_email".to_string(), serde_json::json!(is_private_email));
Ok(UserInfo {
id: claims.sub,
email,
email_verified,
name: None,
picture: None,
raw_claims: serde_json::Value::Object(raw_claims),
})
}
}
#[async_trait]
impl OAuthProvider for AppleOAuth {
fn name(&self) -> &'static str {
"apple"
}
fn authorization_url(&self, state: &str) -> String {
let mut url = format!("{}/auth/authorize", self.base_url);
write!(
url,
"?client_id={}&redirect_uri={}&state={}&response_type=code&scope={}\
&response_mode=form_post",
urlencoding::encode(&self.client_id),
urlencoding::encode(&self.redirect_uri),
urlencoding::encode(state),
urlencoding::encode(APPLE_SCOPES),
)
.expect("write to String is infallible");
url
}
async fn exchange_code(&self, code: &str) -> Result<TokenResponse> {
let secret = self.client_secret()?;
self.post_token(&[
("client_id", self.client_id.as_str()),
("client_secret", secret.as_str()),
("code", code),
("grant_type", "authorization_code"),
("redirect_uri", self.redirect_uri.as_str()),
])
.await
}
async fn user_info(&self, _access_token: &str) -> Result<UserInfo> {
Err(AuthError::OAuthError {
message: "Apple exposes no userinfo endpoint: the identity is in the token \
endpoint's id_token — use OAuthProvider::user_info_from_tokens"
.to_string(),
})
}
async fn user_info_from_tokens(&self, tokens: &TokenResponse) -> Result<UserInfo> {
let id_token = tokens.id_token.as_deref().ok_or_else(|| AuthError::OAuthError {
message: "Apple token response carried no id_token — there is no other source of \
identity"
.to_string(),
})?;
self.user_info_from_id_token(id_token)
}
async fn refresh_token(&self, refresh_token: &str) -> Result<TokenResponse> {
let secret = self.client_secret()?;
self.post_token(&[
("client_id", self.client_id.as_str()),
("client_secret", secret.as_str()),
("grant_type", "refresh_token"),
("refresh_token", refresh_token),
])
.await
}
}