use super::BuildEntityErrorKind;
use crate::jwt::Token;
use smol_str::{SmolStr, ToSmolStr};
use std::fmt::Display;
pub(super) enum EntityIdSrc<'a> {
Token { token: &'a Token, claim: &'a str },
String(&'a str),
}
pub(super) fn get_first_valid_entity_id(
id_srcs: &[EntityIdSrc],
) -> Result<SmolStr, BuildEntityErrorKind> {
let mut errors = Vec::new();
for src in id_srcs {
match src {
EntityIdSrc::Token {
token,
claim: claim_name,
} => {
let Some(claim) = token.get_claim_val(claim_name) else {
errors.push(GetEntityIdError {
token: token.name.clone(),
claim: (*claim_name).to_string(),
reason: GetEntityIdErrorReason::MissingClaim,
});
continue;
};
let claim = claim.to_string();
let id = claim.trim_matches('"');
if id.is_empty() {
errors.push(GetEntityIdError {
token: token.name.clone(),
claim: (*claim_name).to_string(),
reason: GetEntityIdErrorReason::EmptyString,
});
continue;
}
return Ok(id.to_smolstr());
},
EntityIdSrc::String(eid) => return Ok(eid.to_smolstr()),
}
}
Err(BuildEntityErrorKind::MissingEntityId(errors.into()))
}
#[derive(Debug, thiserror::Error, PartialEq)]
#[error("failed to use {claim} from {token} since {reason}")]
pub struct GetEntityIdError {
token: String,
claim: String,
reason: GetEntityIdErrorReason,
}
#[derive(Debug, thiserror::Error, PartialEq)]
pub(super) enum GetEntityIdErrorReason {
#[error("the claim cannot be an empty string")]
EmptyString,
#[error("the claim was not present in the token")]
MissingClaim,
}
#[derive(Debug)]
pub struct GetEntityIdErrors(Vec<GetEntityIdError>);
impl From<Vec<GetEntityIdError>> for GetEntityIdErrors {
fn from(errors: Vec<GetEntityIdError>) -> Self {
Self(errors)
}
}
impl Display for GetEntityIdErrors {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{:?}", self.0.iter().map(ToString::to_string))
}
}
#[cfg(test)]
impl GetEntityIdErrors {
pub fn contains(&self, err: &GetEntityIdError) -> bool {
self.0.contains(err)
}
pub fn len(&self) -> usize {
self.0.len()
}
}
#[cfg(test)]
mod test {
use super::*;
use serde_json::json;
use std::collections::HashMap;
#[test]
fn can_get_first_valid_eid() {
let expected_aud = "some_aud";
let expected_client_id = "some_client_id";
let token = Token::new(
"test_token",
HashMap::from([("aud".into(), json!(expected_aud))]).into(),
None,
);
let id = get_first_valid_entity_id(&[EntityIdSrc::Token {
token: &token,
claim: "aud",
}])
.expect("should get entity id from token's aud");
assert_eq!(id, expected_aud);
let id = get_first_valid_entity_id(&[
EntityIdSrc::Token {
token: &token,
claim: "client_id",
},
EntityIdSrc::Token {
token: &token,
claim: "aud",
},
])
.expect("should get entity id from token's aud");
assert_eq!(id, expected_aud);
let token = Token::new(
"test_token",
HashMap::from([
("aud".into(), json!(expected_aud)),
("client_id".into(), json!(expected_client_id)),
])
.into(),
None,
);
let id = get_first_valid_entity_id(&[
EntityIdSrc::Token {
token: &token,
claim: "client_id",
},
EntityIdSrc::Token {
token: &token,
claim: "aud",
},
])
.expect("should get entity id from token's client_id");
assert_eq!(id, expected_client_id);
let token = Token::new(
"test_token",
HashMap::from([("empty".into(), json!(""))]).into(),
None,
);
let err = get_first_valid_entity_id(&[
EntityIdSrc::Token {
token: &token,
claim: "empty",
},
EntityIdSrc::Token {
token: &token,
claim: "missing",
},
])
.expect_err("should error while getting id");
let expected_errs = vec![
GetEntityIdError {
token: "test_token".into(),
claim: "empty".into(),
reason: GetEntityIdErrorReason::EmptyString,
},
GetEntityIdError {
token: "test_token".into(),
claim: "missing".into(),
reason: GetEntityIdErrorReason::MissingClaim,
},
];
assert!(
matches!(
err,
BuildEntityErrorKind::MissingEntityId(GetEntityIdErrors(ref errs))
if *errs == expected_errs
),
"expected: {:?}\nbut got: {:?}",
BuildEntityErrorKind::MissingEntityId(GetEntityIdErrors(expected_errs)),
err,
);
}
}