use std::sync::Arc;
use openidconnect::core::{
CoreAuthenticationFlow, CoreClient, CoreIdTokenClaims, CoreJwsSigningAlgorithm,
CoreProviderMetadata, CoreResponseType, CoreSubjectIdentifierType,
};
use openidconnect::reqwest;
use openidconnect::url::form_urlencoded;
use openidconnect::{
AuthType, AuthUrl, AuthorizationCode, ClientId, ClientSecret, CsrfToken,
EmptyAdditionalProviderMetadata, EndpointMaybeSet, EndpointNotSet, EndpointSet, IssuerUrl,
JsonWebKeySetUrl, Nonce, PkceCodeChallenge, RedirectUrl, ResponseTypes, Scope, TokenResponse,
TokenUrl,
};
use serde::Deserialize;
use super::client_secret::{AppleClientSecret, ClientSecretError};
use crate::providers::oidc_generic::{
OidcBegin, OidcError, OidcFlowState, OidcFlowStore, VerifiedIdToken, DEFAULT_FLOW_TTL_SECONDS,
};
pub const APPLE_ISSUER: &str = "https://appleid.apple.com";
pub const APPLE_AUTHORIZATION_ENDPOINT: &str = "https://appleid.apple.com/auth/authorize";
pub const APPLE_TOKEN_ENDPOINT: &str = "https://appleid.apple.com/auth/token";
pub const APPLE_JWKS_URI: &str = "https://appleid.apple.com/auth/keys";
pub const APPLE_DEFAULT_SCOPES: &[&str] = &["name", "email"];
pub fn apple_provider_metadata() -> CoreProviderMetadata {
CoreProviderMetadata::new(
IssuerUrl::new(APPLE_ISSUER.to_owned()).expect("Apple issuer URL parses"),
AuthUrl::new(APPLE_AUTHORIZATION_ENDPOINT.to_owned())
.expect("Apple authorization endpoint URL parses"),
JsonWebKeySetUrl::new(APPLE_JWKS_URI.to_owned()).expect("Apple JWKS URL parses"),
vec![ResponseTypes::new(vec![CoreResponseType::Code])],
vec![CoreSubjectIdentifierType::Pairwise],
vec![CoreJwsSigningAlgorithm::RsaSsaPkcs1V15Sha256],
EmptyAdditionalProviderMetadata {},
)
.set_token_endpoint(Some(
TokenUrl::new(APPLE_TOKEN_ENDPOINT.to_owned()).expect("Apple token endpoint URL parses"),
))
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum AppleRedirectError {
#[error("oidc: {0}")]
Oidc(#[from] OidcError),
#[error("apple form missing required field `{0}`")]
MissingFormField(&'static str),
#[error("apple provider returned error: {0}")]
Provider(String),
#[error("apple client_secret: {0}")]
ClientSecret(#[from] ClientSecretError),
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct AppleCallbackForm {
pub code: AuthorizationCode,
pub state: CsrfToken,
pub user_json: Option<String>,
pub error: Option<String>,
}
impl AppleCallbackForm {
pub fn parse(body: &str) -> Result<Self, AppleRedirectError> {
let mut code: Option<String> = None;
let mut state: Option<String> = None;
let mut user_json: Option<String> = None;
let mut error: Option<String> = None;
for (k, v) in form_urlencoded::parse(body.as_bytes()) {
match k.as_ref() {
"code" => code = Some(v.into_owned()),
"state" => state = Some(v.into_owned()),
"user" => user_json = Some(v.into_owned()),
"error" => error = Some(v.into_owned()),
_ => {}
}
}
if let Some(err) = error {
return Err(AppleRedirectError::Provider(err));
}
let code = code
.map(AuthorizationCode::new)
.ok_or(AppleRedirectError::MissingFormField("code"))?;
let state = state
.map(CsrfToken::new)
.ok_or(AppleRedirectError::MissingFormField("state"))?;
Ok(Self {
code,
state,
user_json,
error: None,
})
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct FirstLoginName(Option<String>);
#[derive(Debug, Deserialize)]
struct AppleUserField {
name: Option<AppleUserName>,
}
#[derive(Debug, Deserialize)]
struct AppleUserName {
#[serde(rename = "firstName")]
first_name: Option<String>,
#[serde(rename = "lastName")]
last_name: Option<String>,
}
impl FirstLoginName {
pub fn empty() -> Self {
Self(None)
}
pub fn as_str(&self) -> Option<&str> {
self.0.as_deref()
}
pub fn is_some(&self) -> bool {
self.0.is_some()
}
pub fn into_inner(self) -> Option<String> {
self.0
}
pub fn from_apple_user_field(s: Option<&str>) -> Self {
let Some(s) = s.filter(|x| !x.is_empty()) else {
return Self::empty();
};
let Ok(parsed) = serde_json::from_str::<AppleUserField>(s) else {
return Self::empty();
};
let Some(name) = parsed.name else {
return Self::empty();
};
let combined = match (
name.first_name.as_deref().map(str::trim).filter(|x| !x.is_empty()),
name.last_name.as_deref().map(str::trim).filter(|x| !x.is_empty()),
) {
(Some(f), Some(l)) => Some(format!("{f} {l}")),
(Some(f), None) => Some(f.to_owned()),
(None, Some(l)) => Some(l.to_owned()),
(None, None) => None,
};
Self(combined)
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct AppleVerified {
pub id_token: VerifiedIdToken,
pub first_login: FirstLoginName,
}
type AppleCoreClient = CoreClient<
EndpointSet,
EndpointNotSet,
EndpointNotSet,
EndpointNotSet,
EndpointMaybeSet,
EndpointMaybeSet,
>;
pub struct AppleRedirectProvider<S> {
metadata: CoreProviderMetadata,
client_id: ClientId,
redirect_uri: RedirectUrl,
secret_gen: Arc<AppleClientSecret>,
scopes: Vec<Scope>,
flow_ttl_seconds: i64,
flows: S,
}
impl<S> AppleRedirectProvider<S> {
pub fn from_provider_metadata(
metadata: CoreProviderMetadata,
client_id: ClientId,
redirect_uri: RedirectUrl,
secret_gen: Arc<AppleClientSecret>,
flows: S,
) -> Self {
Self {
metadata,
client_id,
redirect_uri,
secret_gen,
scopes: APPLE_DEFAULT_SCOPES
.iter()
.map(|s| Scope::new((*s).to_owned()))
.collect(),
flow_ttl_seconds: DEFAULT_FLOW_TTL_SECONDS,
flows,
}
}
pub async fn discover(
client_id: ClientId,
redirect_uri: RedirectUrl,
secret_gen: Arc<AppleClientSecret>,
flows: S,
http: &reqwest::Client,
) -> Result<Self, AppleRedirectError> {
let issuer = IssuerUrl::new(APPLE_ISSUER.to_owned())
.map_err(|e| OidcError::Discovery(format!("invalid Apple issuer URL: {e}")))?;
let metadata = CoreProviderMetadata::discover_async(issuer, http)
.await
.map_err(|e| OidcError::Discovery(format!("{e}")))?;
Ok(Self::from_provider_metadata(
metadata,
client_id,
redirect_uri,
secret_gen,
flows,
))
}
pub fn with_scopes<I, V>(mut self, scopes: I) -> Self
where
I: IntoIterator<Item = V>,
V: Into<String>,
{
self.scopes = scopes.into_iter().map(|s| Scope::new(s.into())).collect();
self
}
pub fn with_flow_ttl_seconds(mut self, ttl: i64) -> Self {
self.flow_ttl_seconds = ttl;
self
}
pub fn scopes(&self) -> &[Scope] {
&self.scopes
}
pub fn flow_ttl_seconds(&self) -> i64 {
self.flow_ttl_seconds
}
pub fn flows(&self) -> &S {
&self.flows
}
pub fn secret_gen(&self) -> &AppleClientSecret {
&self.secret_gen
}
pub fn invalidate_client_secret(&self) {
self.secret_gen.invalidate();
}
fn build_client(
&self,
client_secret: Option<ClientSecret>,
) -> AppleCoreClient {
CoreClient::from_provider_metadata(
self.metadata.clone(),
self.client_id.clone(),
client_secret,
)
.set_redirect_uri(self.redirect_uri.clone())
.set_auth_type(AuthType::RequestBody)
}
}
impl<S: OidcFlowStore> AppleRedirectProvider<S> {
pub async fn begin(&self, now: i64) -> Result<OidcBegin, AppleRedirectError> {
let client = self.build_client(None);
let (pkce_challenge, pkce_verifier) = PkceCodeChallenge::new_random_sha256();
let mut req = client.authorize_url(
CoreAuthenticationFlow::AuthorizationCode,
CsrfToken::new_random,
Nonce::new_random,
);
for scope in &self.scopes {
req = req.add_scope(scope.clone());
}
req = req.add_extra_param("response_mode", "form_post");
let (authorize_url, csrf_state, nonce) = req.set_pkce_challenge(pkce_challenge).url();
let flow_state = OidcFlowState::from_parts(
csrf_state.clone(),
nonce,
pkce_verifier,
now.saturating_add(self.flow_ttl_seconds),
);
self.flows
.put(csrf_state.secret(), flow_state)
.await
.map_err(OidcError::Store)?;
Ok(OidcBegin {
authorize_url,
csrf_state,
})
}
pub async fn finish_form_post(
&self,
form: AppleCallbackForm,
http: &reqwest::Client,
now: i64,
) -> Result<AppleVerified, AppleRedirectError> {
let flow_state = self
.flows
.take(form.state.secret())
.await
.map_err(OidcError::Store)?
.ok_or(OidcError::UnknownFlow)?;
let (csrf_token, nonce, pkce_verifier, expires_at) = flow_state.into_parts();
if csrf_token.secret() != form.state.secret() {
return Err(OidcError::StateMismatch.into());
}
if expires_at <= now {
return Err(OidcError::FlowExpired.into());
}
let jwt = self.secret_gen.current(now)?;
let client = self.build_client(Some(ClientSecret::new(jwt)));
let token_response = client
.exchange_code(form.code)
.map_err(|e| OidcError::Config(format!("{e}")))?
.set_pkce_verifier(pkce_verifier)
.request_async(http)
.await
.map_err(|e| OidcError::Http(format!("{e}")))?;
let id_token = token_response
.id_token()
.ok_or(OidcError::MissingIdToken)?;
let verifier = client.id_token_verifier();
let claims: &CoreIdTokenClaims = id_token
.claims(&verifier, &nonce)
.map_err(|e| OidcError::IdToken(format!("{e}")))?;
let first_login = FirstLoginName::from_apple_user_field(form.user_json.as_deref());
Ok(AppleVerified {
id_token: VerifiedIdToken {
issuer: claims.issuer().as_str().to_owned(),
subject: claims.subject().as_str().to_owned(),
email: claims.email().map(|e| e.as_str().to_owned()),
email_verified: claims.email_verified(),
name: claims
.name()
.and_then(|n| n.get(None).or_else(|| n.iter().next().map(|(_, v)| v)))
.map(|v| v.as_str().to_owned()),
},
first_login,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::{Duration, Utc};
use openidconnect::core::{
CoreIdToken, CoreIdTokenClaims, CoreIdTokenFields, CoreJsonWebKeySet,
CoreRsaPrivateSigningKey, CoreTokenResponse, CoreTokenType,
};
use openidconnect::{
AccessToken, Audience, EmptyAdditionalClaims, EmptyExtraTokenFields, EndUserEmail,
EndUserName, JsonWebKeyId, LocalizedClaim, PrivateSigningKey, StandardClaims,
SubjectIdentifier,
};
use wiremock::matchers::{body_string_contains, method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use crate::providers::apple::client_secret::{
AppleClientSecret, APPLE_AUDIENCE, DEFAULT_TOKEN_TTL_SECONDS,
};
use crate::providers::oidc_generic::{MemoryOidcFlowStore, OidcFlowStore};
const TEST_RSA_PEM: &str = concat!(
"-----BEGIN RSA PRIVATE KEY-----\n",
"MIIEowIBAAKCAQEAsRMj0YYjy7du6v1gWyKSTJx3YjBzZTG0XotRP0IaObw0k+68\n",
"30dXadjL5jVhSWNdcg9OyMyTGWfdNqfdrS6ppBqlQNgjZJdloIqL9zOLBZrDm7G4\n",
"+qN4KeZ4/5TyEilq2zOHHGFEzXpOq/UxqVnm3J4fhjqCNaS2nKd7HVVXGBQQ+4+F\n",
"dVT+MyJXemw5maz2F/h324TQi6XoUPEwUddxBwLQFSOlzWnHYMc4/lcyZJ8MpTXC\n",
"MPe/YJFNtb9CaikKUdf8x4mzwH7usSf8s2d6R4dQITzKrjrEJ0u3w3eGkBBapoMV\n",
"FBGPjP3Haz5FsVtHc5VEN3FZVIDF6HrbJH1C4QIDAQABAoIBAHSS3izM+3nc7Bel\n",
"8S5uRxRKmcm5je6b11u6qiVUFkHWJmMRc6QmqmSThkCq+b4/vUAe1cYZ7+l02Exo\n",
"HOcrZiEULaDP6hUKGqyjKVv3wdlRtt8kFFxlC/HBufzAiNDuFVvzw0oquwnvMCXC\n",
"yQvtlK+/JY/PqvM32cSt+b4o9apySsHqAtdsoHHohK82jsQqIfCi1v8XYV/xRBJB\n",
"cQMCaA0Ls3tFpmJv3JdikyyQxio4kZ5tswghC63znCp1iL+qDq1wjjKzjick9MDb\n",
"Qzb95X09QQP201l1FPWN7Kbhj4ybg6PJGz/VHQcvILcBCoYIc0UY/OMSBt9VN9yD\n",
"wr1WlbECgYEA37difsTMcLmUEN57sicFe1q4lxH6eqnUBjmoKBflx4oMIIyRnfjF\n",
"Jwsu9yIiBkJfBCP85nl2tZdcV0wfZLf6amxB/KMtdfW6r8eoTDzE472OYxSIg1F5\n",
"dI4qn2nBI0Dou0g58xj+Kv0iLaym0pxtyJkSg/rxZGwKb9a+x5WAs50CgYEAyqC0\n",
"NcZs2BRIiT5kEOF6+MeUvarbKh1mangKHKcTdXRrvoJ+Z5izm7FifBixo/79MYpt\n",
"0VofW0IzYKtAI9KZDq2JcozEbZ+lt/ZPH5QEXO4T39QbDoAG8BbOmEP7l+6m+7QO\n",
"PiQ0WSNjDnwk3W7Zihgg31DH7hyxsxQCapKLcxUCgYAwERXPiPcoDSd8DGFlYK7z\n",
"1wUsKEe6DT0p7T9tBd1v5wA+ChXLbETn46Y+oQ3QbHg/yn+vAU/5KkFD3G4uVL0w\n",
"Gnx/DIxa+OYYmHxXjQL8r6ClNycxl9LRsS4FPFKsAWk/u///dFI/6E1spNjfDY8k\n",
"94ab5tHwsqn3Z5tsBHo3nQKBgFUmxbSXh2Qi2fy6+GhTqU7k6G/wXhvLsR9rBKzX\n",
"1YiVfTXZNu+oL0ptd/q4keZeIN7x0oaY/fZm0pp8PP8Q4HtXmBxIZb+/yG+Pld6q\n",
"YE8BSd7VDu3ABapdm0JHx3Iou4mpOBcLNeiDw3vx1bgsfkTXMPFHzE0XR+H+tak9\n",
"nlalAoGBALAmAF7WBGdOt43Rj8hPaKOM/ahj+6z3CNwVreToNsVBHoyNmiO8q7MC\n",
"+tRo4jgdrzk1pzs66OIHfbx5P1mXKPtgPZhvI5omAY8WqXEgeNqSL1Ksp6LZ2ql/\n",
"ouZns5xwKc9+aRL+GWoAGNzwzcjE8cP52sBy/r0rYXTs/sZo5kgV\n",
"-----END RSA PRIVATE KEY-----\n",
);
const TEST_KID: &str = "apple-test-key";
const CLIENT_ID: &str = "com.example.signin";
const REDIRECT_URI: &str = "https://app.example/auth/callback/apple";
fn signing_key() -> CoreRsaPrivateSigningKey {
CoreRsaPrivateSigningKey::from_pem(
TEST_RSA_PEM,
Some(JsonWebKeyId::new(TEST_KID.into())),
)
.expect("test RSA PEM parses")
}
fn dummy_http() -> reqwest::Client {
reqwest::ClientBuilder::new()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("reqwest builds")
}
fn p8_pem() -> String {
use p256::pkcs8::{EncodePrivateKey, LineEnding};
let mut bytes = [0u8; 32];
bytes[31] = 0x42;
let sk = p256::SecretKey::from_slice(&bytes).expect("valid P-256 scalar");
sk.to_pkcs8_pem(LineEnding::LF)
.expect("P-256 PKCS#8 PEM")
.to_string()
}
fn apple_secret() -> Arc<AppleClientSecret> {
Arc::new(
AppleClientSecret::from_p8_pem(
"TEAM123ABC",
"KEYID45678",
CLIENT_ID,
p8_pem().as_bytes(),
)
.expect("p8 pem parses"),
)
}
#[test]
fn form_parse_happy_path_no_user() {
let body = "code=ABC123&state=STATE-XYZ";
let form = AppleCallbackForm::parse(body).unwrap();
assert_eq!(form.code.secret(), "ABC123");
assert_eq!(form.state.secret(), "STATE-XYZ");
assert!(form.user_json.is_none());
assert!(form.error.is_none());
}
#[test]
fn form_parse_extracts_user_field_verbatim() {
let body =
"code=C&state=S&user=%7B%22name%22%3A%7B%22firstName%22%3A%22Ada%22%2C%22lastName%22%3A%22Lovelace%22%7D%7D";
let form = AppleCallbackForm::parse(body).unwrap();
assert_eq!(form.code.secret(), "C");
assert_eq!(form.state.secret(), "S");
assert_eq!(
form.user_json.as_deref(),
Some(r#"{"name":{"firstName":"Ada","lastName":"Lovelace"}}"#)
);
}
#[test]
fn form_parse_provider_error_short_circuits() {
let body = "error=access_denied";
let err = AppleCallbackForm::parse(body).unwrap_err();
match err {
AppleRedirectError::Provider(s) => assert_eq!(s, "access_denied"),
other => panic!("expected Provider error, got {other:?}"),
}
}
#[test]
fn form_parse_missing_code_errors() {
let body = "state=ONLY";
let err = AppleCallbackForm::parse(body).unwrap_err();
assert!(matches!(err, AppleRedirectError::MissingFormField("code")));
}
#[test]
fn form_parse_missing_state_errors() {
let body = "code=ONLY";
let err = AppleCallbackForm::parse(body).unwrap_err();
assert!(matches!(err, AppleRedirectError::MissingFormField("state")));
}
#[test]
fn form_parse_ignores_unknown_fields() {
let body = "code=C&state=S&extra=junk&id_token=ignored";
let form = AppleCallbackForm::parse(body).unwrap();
assert_eq!(form.code.secret(), "C");
assert_eq!(form.state.secret(), "S");
}
#[test]
fn first_login_name_combines_first_and_last() {
let name = FirstLoginName::from_apple_user_field(Some(
r#"{"name":{"firstName":"Ada","lastName":"Lovelace"},"email":"a@l.com"}"#,
));
assert_eq!(name.as_str(), Some("Ada Lovelace"));
assert!(name.is_some());
}
#[test]
fn first_login_name_first_only() {
let name = FirstLoginName::from_apple_user_field(Some(
r#"{"name":{"firstName":"Cher"}}"#,
));
assert_eq!(name.as_str(), Some("Cher"));
}
#[test]
fn first_login_name_last_only() {
let name = FirstLoginName::from_apple_user_field(Some(
r#"{"name":{"lastName":"Lovelace"}}"#,
));
assert_eq!(name.as_str(), Some("Lovelace"));
}
#[test]
fn first_login_name_trims_whitespace_components() {
let name = FirstLoginName::from_apple_user_field(Some(
r#"{"name":{"firstName":" Ada ","lastName":" "}}"#,
));
assert_eq!(name.as_str(), Some("Ada"));
}
#[test]
fn first_login_name_empty_when_absent() {
assert!(FirstLoginName::from_apple_user_field(None).as_str().is_none());
assert!(FirstLoginName::from_apple_user_field(Some("")).as_str().is_none());
assert!(
FirstLoginName::from_apple_user_field(Some(r#"{"email":"x@y.com"}"#))
.as_str()
.is_none()
);
assert!(
FirstLoginName::from_apple_user_field(Some(r#"{"name":{}}"#))
.as_str()
.is_none()
);
}
#[test]
fn first_login_name_malformed_json_is_empty() {
assert!(
FirstLoginName::from_apple_user_field(Some("not json"))
.as_str()
.is_none()
);
assert!(
FirstLoginName::from_apple_user_field(Some(r#"{"name":"plain string"}"#))
.as_str()
.is_none()
);
}
#[test]
fn apple_provider_metadata_constants_match_published_urls() {
let meta = apple_provider_metadata();
assert_eq!(meta.issuer().as_str(), APPLE_ISSUER);
assert_eq!(
meta.authorization_endpoint().as_str(),
APPLE_AUTHORIZATION_ENDPOINT
);
assert_eq!(
meta.token_endpoint().expect("token endpoint").as_str(),
APPLE_TOKEN_ENDPOINT
);
assert_eq!(meta.jwks_uri().as_str(), APPLE_JWKS_URI);
}
fn provider_pinned_to_apple() -> AppleRedirectProvider<MemoryOidcFlowStore> {
AppleRedirectProvider::from_provider_metadata(
apple_provider_metadata(),
ClientId::new(CLIENT_ID.into()),
RedirectUrl::new(REDIRECT_URI.into()).unwrap(),
apple_secret(),
MemoryOidcFlowStore::new(),
)
}
#[tokio::test]
async fn begin_url_targets_apple_authorize_with_form_post_mode() {
let p = provider_pinned_to_apple();
let begin = p.begin(1_000).await.unwrap();
let u = &begin.authorize_url;
assert_eq!(u.scheme(), "https");
assert_eq!(u.host_str(), Some("appleid.apple.com"));
assert_eq!(u.path(), "/auth/authorize");
let qs: Vec<(String, String)> = u
.query_pairs()
.map(|(k, v)| (k.into_owned(), v.into_owned()))
.collect();
let lookup = |k: &str| qs.iter().find(|(kk, _)| kk == k).map(|(_, v)| v.clone());
assert_eq!(lookup("response_type").as_deref(), Some("code"));
assert_eq!(lookup("response_mode").as_deref(), Some("form_post"));
assert_eq!(lookup("client_id").as_deref(), Some(CLIENT_ID));
assert_eq!(lookup("redirect_uri").as_deref(), Some(REDIRECT_URI));
let scope = lookup("scope").expect("scope param");
let parts: Vec<&str> = scope.split(' ').collect();
assert!(parts.contains(&"name"));
assert!(parts.contains(&"email"));
assert_eq!(lookup("code_challenge_method").as_deref(), Some("S256"));
assert!(!lookup("code_challenge").unwrap_or_default().is_empty());
assert!(!lookup("state").unwrap_or_default().is_empty());
assert!(!lookup("nonce").unwrap_or_default().is_empty());
}
#[tokio::test]
async fn begin_stashes_flow_keyed_by_csrf_state() {
let p = provider_pinned_to_apple();
let begin = p.begin(1_000).await.unwrap();
assert_eq!(p.flows().len(), 1);
let stashed = p
.flows()
.take(begin.csrf_state.secret())
.await
.unwrap()
.expect("stashed");
assert_eq!(stashed.csrf_token().secret(), begin.csrf_state.secret());
assert_eq!(stashed.expires_at(), 1_000 + DEFAULT_FLOW_TTL_SECONDS);
}
#[tokio::test]
async fn finish_with_unknown_state_returns_unknown_flow() {
let p = provider_pinned_to_apple();
let form = AppleCallbackForm {
code: AuthorizationCode::new("C".into()),
state: CsrfToken::new("never-stashed".into()),
user_json: None,
error: None,
};
let err = p
.finish_form_post(form, &dummy_http(), 1_000)
.await
.unwrap_err();
match err {
AppleRedirectError::Oidc(OidcError::UnknownFlow) => {}
other => panic!("expected UnknownFlow, got {other:?}"),
}
}
#[tokio::test]
async fn finish_with_expired_flow_returns_flow_expired() {
let p = provider_pinned_to_apple().with_flow_ttl_seconds(60);
let begin = p.begin(1_000).await.unwrap();
let form = AppleCallbackForm {
code: AuthorizationCode::new("C".into()),
state: begin.csrf_state,
user_json: None,
error: None,
};
let err = p
.finish_form_post(form, &dummy_http(), 1_000 + 60)
.await
.unwrap_err();
match err {
AppleRedirectError::Oidc(OidcError::FlowExpired) => {}
other => panic!("expected FlowExpired, got {other:?}"),
}
}
async fn mount_apple_discovery_and_jwks(server: &MockServer, base: &str) {
let metadata = CoreProviderMetadata::new(
IssuerUrl::new(base.to_owned()).unwrap(),
AuthUrl::new(format!("{base}/auth/authorize")).unwrap(),
JsonWebKeySetUrl::new(format!("{base}/auth/keys")).unwrap(),
vec![ResponseTypes::new(vec![CoreResponseType::Code])],
vec![CoreSubjectIdentifierType::Pairwise],
vec![CoreJwsSigningAlgorithm::RsaSsaPkcs1V15Sha256],
EmptyAdditionalProviderMetadata {},
)
.set_token_endpoint(Some(TokenUrl::new(format!("{base}/auth/token")).unwrap()));
Mock::given(method("GET"))
.and(path("/.well-known/openid-configuration"))
.respond_with(ResponseTemplate::new(200).set_body_json(&metadata))
.mount(server)
.await;
let jwks = CoreJsonWebKeySet::new(vec![signing_key().as_verification_key()]);
Mock::given(method("GET"))
.and(path("/auth/keys"))
.respond_with(ResponseTemplate::new(200).set_body_json(&jwks))
.mount(server)
.await;
}
fn build_id_token(issuer: &str, nonce: &Nonce, name: Option<&str>) -> CoreIdToken {
let now = Utc::now();
let mut std_claims =
StandardClaims::new(SubjectIdentifier::new("apple-sub-abcdef".to_owned()))
.set_email(Some(EndUserEmail::new(
"abc@privaterelay.appleid.com".to_owned(),
)))
.set_email_verified(Some(true));
if let Some(n) = name {
let mut lc: LocalizedClaim<EndUserName> = LocalizedClaim::default();
lc.insert(None, EndUserName::new(n.to_owned()));
std_claims = std_claims.set_name(Some(lc));
}
let claims = CoreIdTokenClaims::new(
IssuerUrl::new(issuer.to_owned()).unwrap(),
vec![Audience::new(CLIENT_ID.to_owned())],
now + Duration::seconds(600),
now,
std_claims,
EmptyAdditionalClaims {},
)
.set_nonce(Some(nonce.clone()));
CoreIdToken::new(
claims,
&signing_key(),
CoreJwsSigningAlgorithm::RsaSsaPkcs1V15Sha256,
None,
None,
)
.expect("ID token signs")
}
async fn mount_token_endpoint(server: &MockServer, id_token: CoreIdToken) {
let resp = CoreTokenResponse::new(
AccessToken::new("test-access-token".to_owned()),
CoreTokenType::Bearer,
CoreIdTokenFields::new(Some(id_token), EmptyExtraTokenFields {}),
);
Mock::given(method("POST"))
.and(path("/auth/token"))
.respond_with(ResponseTemplate::new(200).set_body_json(&resp))
.mount(server)
.await;
}
async fn peek_stashed_nonce(
provider: &AppleRedirectProvider<MemoryOidcFlowStore>,
csrf_state_secret: &str,
) -> Nonce {
let st = provider
.flows()
.take(csrf_state_secret)
.await
.expect("store ok")
.expect("flow stashed");
let nonce = st.nonce().clone();
let put_back = OidcFlowState::from_parts(
st.csrf_token().clone(),
st.nonce().clone(),
openidconnect::PkceCodeVerifier::new(st.pkce_verifier().secret().clone()),
st.expires_at(),
);
provider
.flows()
.put(csrf_state_secret, put_back)
.await
.expect("re-put");
nonce
}
async fn build_provider_via_discovery(
server: &MockServer,
http: &reqwest::Client,
) -> AppleRedirectProvider<MemoryOidcFlowStore> {
let issuer = IssuerUrl::new(server.uri()).unwrap();
let metadata = CoreProviderMetadata::discover_async(issuer, http)
.await
.expect("wiremock discovery succeeds");
AppleRedirectProvider::from_provider_metadata(
metadata,
ClientId::new(CLIENT_ID.into()),
RedirectUrl::new(REDIRECT_URI.into()).unwrap(),
apple_secret(),
MemoryOidcFlowStore::new(),
)
}
#[tokio::test]
async fn form_post_round_trip_captures_first_login_name() {
let http = dummy_http();
let server = MockServer::start().await;
let base = server.uri();
mount_apple_discovery_and_jwks(&server, &base).await;
let provider = build_provider_via_discovery(&server, &http).await;
let now_seconds = Utc::now().timestamp();
let begin = provider.begin(now_seconds).await.unwrap();
let nonce = peek_stashed_nonce(&provider, begin.csrf_state.secret()).await;
let id_token = build_id_token(&base, &nonce, None); mount_token_endpoint(&server, id_token).await;
let user_json = r#"{"name":{"firstName":"Ada","lastName":"Lovelace"},"email":"ada@example.com"}"#;
let body = form_urlencoded::Serializer::new(String::new())
.append_pair("code", "the-code")
.append_pair("state", begin.csrf_state.secret())
.append_pair("user", user_json)
.finish();
let form = AppleCallbackForm::parse(&body).expect("parse form-post body");
let verified = provider
.finish_form_post(form, &http, now_seconds)
.await
.expect("form_post round-trip");
assert_eq!(verified.id_token.issuer, base);
assert_eq!(verified.id_token.subject, "apple-sub-abcdef");
assert_eq!(
verified.id_token.email.as_deref(),
Some("abc@privaterelay.appleid.com")
);
assert_eq!(verified.id_token.email_verified, Some(true));
assert_eq!(verified.first_login.as_str(), Some("Ada Lovelace"));
}
#[tokio::test]
async fn form_post_round_trip_returning_user_has_no_first_login_name() {
let http = dummy_http();
let server = MockServer::start().await;
let base = server.uri();
mount_apple_discovery_and_jwks(&server, &base).await;
let provider = build_provider_via_discovery(&server, &http).await;
let now_seconds = Utc::now().timestamp();
let begin = provider.begin(now_seconds).await.unwrap();
let nonce = peek_stashed_nonce(&provider, begin.csrf_state.secret()).await;
let id_token = build_id_token(&base, &nonce, None);
mount_token_endpoint(&server, id_token).await;
let body = format!("code=C&state={}", begin.csrf_state.secret());
let form = AppleCallbackForm::parse(&body).unwrap();
let verified = provider
.finish_form_post(form, &http, now_seconds)
.await
.unwrap();
assert!(verified.first_login.as_str().is_none());
}
#[tokio::test]
async fn token_endpoint_receives_es256_client_secret_jwt() {
use jsonwebtoken::{Algorithm, EncodingKey, Header};
use serde::Serialize;
let http = dummy_http();
let server = MockServer::start().await;
let base = server.uri();
mount_apple_discovery_and_jwks(&server, &base).await;
let provider = build_provider_via_discovery(&server, &http).await;
let now_seconds = Utc::now().timestamp();
let begin = provider.begin(now_seconds).await.unwrap();
let nonce = peek_stashed_nonce(&provider, begin.csrf_state.secret()).await;
let id_token = build_id_token(&base, &nonce, None);
#[derive(Serialize)]
struct ExpectedClaims<'a> {
iss: &'a str,
iat: i64,
exp: i64,
aud: &'a str,
sub: &'a str,
}
let mut header = Header::new(Algorithm::ES256);
header.kid = Some("KEYID45678".to_owned());
let probe_pem = p8_pem();
let probe_key = EncodingKey::from_ec_pem(probe_pem.as_bytes()).unwrap();
let probe = jsonwebtoken::encode(
&header,
&ExpectedClaims {
iss: "TEAM123ABC",
iat: now_seconds,
exp: now_seconds + DEFAULT_TOKEN_TTL_SECONDS,
aud: APPLE_AUDIENCE,
sub: CLIENT_ID,
},
&probe_key,
)
.unwrap();
let prefix: String = probe.split('.').take(2).collect::<Vec<_>>().join(".");
let needle = format!("client_secret={prefix}.");
let token_resp = CoreTokenResponse::new(
AccessToken::new("at".to_owned()),
CoreTokenType::Bearer,
CoreIdTokenFields::new(Some(id_token), EmptyExtraTokenFields {}),
);
Mock::given(method("POST"))
.and(path("/auth/token"))
.and(body_string_contains(&needle))
.respond_with(ResponseTemplate::new(200).set_body_json(&token_resp))
.expect(1)
.mount(&server)
.await;
let body = format!("code=C&state={}", begin.csrf_state.secret());
let form = AppleCallbackForm::parse(&body).unwrap();
provider
.finish_form_post(form, &http, now_seconds)
.await
.expect("token endpoint sees the JWT");
}
}