use serde_json::{Map, Value};
use super::disclosure::{Disclosure, SD_HASH_ALG, digest_of};
use super::error::SdJwtError;
pub(crate) const SD_CLAIM: &str = "_sd";
pub(crate) const SD_ALG_CLAIM: &str = "_sd_alg";
pub fn conceal(
object: &Map<String, Value>,
mut conceal_if: impl FnMut(&str) -> bool,
) -> Result<(Map<String, Value>, Vec<Disclosure>), SdJwtError> {
let mut kept = Map::new();
let mut disclosures = Vec::new();
for (name, value) in object {
if conceal_if(name) {
disclosures.push(Disclosure::new(name.clone(), value.clone())?);
} else {
kept.insert(name.clone(), value.clone());
}
}
if !disclosures.is_empty() {
let mut digests: Vec<Value> = disclosures
.iter()
.map(|d| Value::String(d.digest()))
.collect();
digests.sort_by(|a, b| a.as_str().cmp(&b.as_str()));
kept.insert(SD_CLAIM.to_owned(), Value::Array(digests));
}
Ok((kept, disclosures))
}
#[derive(Debug, Clone)]
pub struct SdJwt {
jwt: String,
disclosures: Vec<Disclosure>,
key_binding_jwt: Option<String>,
}
impl SdJwt {
pub fn new(jwt: String, disclosures: Vec<Disclosure>) -> Self {
Self {
jwt,
disclosures,
key_binding_jwt: None,
}
}
pub fn parse(serialised: &str) -> Result<Self, SdJwtError> {
let segments: Vec<&str> = serialised.split('~').collect();
let [jwt, rest @ .., last] = segments.as_slice() else {
return Err(SdJwtError::NotAnSdJwt);
};
let disclosures = rest
.iter()
.map(|s| Disclosure::parse(s))
.collect::<Result<Vec<_>, _>>()?;
Ok(Self {
jwt: (*jwt).to_owned(),
disclosures,
key_binding_jwt: (!last.is_empty()).then(|| (*last).to_owned()),
})
}
pub fn serialise(&self) -> String {
let mut out = String::from(&self.jwt);
for d in &self.disclosures {
out.push('~');
out.push_str(d.encoded());
}
out.push('~');
if let Some(kb) = &self.key_binding_jwt {
out.push_str(kb);
}
out
}
pub fn jwt(&self) -> &str {
&self.jwt
}
pub fn disclosures(&self) -> &[Disclosure] {
&self.disclosures
}
pub fn has_key_binding(&self) -> bool {
self.key_binding_jwt.is_some()
}
pub fn digests_for_claim(&self, claim_name: &str) -> Vec<String> {
self.disclosures
.iter()
.filter(|d| d.claim_name() == claim_name)
.map(Disclosure::digest)
.collect()
}
pub fn present(&self, reveal_digests: &[&str]) -> Self {
Self {
jwt: self.jwt.clone(),
disclosures: self
.disclosures
.iter()
.filter(|d| reveal_digests.contains(&d.digest().as_str()))
.cloned()
.collect(),
key_binding_jwt: self.key_binding_jwt.clone(),
}
}
pub fn disclosed_payload(&self) -> Result<Map<String, Value>, SdJwtError> {
let payload = decode_jwt_payload(&self.jwt)?;
if let Some(alg) = payload.get(SD_ALG_CLAIM) {
let named = alg.as_str().unwrap_or_default();
if named != SD_HASH_ALG {
return Err(SdJwtError::UnsupportedHashAlg(named.to_owned()));
}
}
let mut by_digest: std::collections::HashMap<String, &Disclosure> =
std::collections::HashMap::new();
for d in &self.disclosures {
let digest = digest_of(d.encoded());
if by_digest.insert(digest.clone(), d).is_some() {
return Err(SdJwtError::DuplicateDigest(digest));
}
}
let mut used = 0usize;
let mut seen_digests = std::collections::HashSet::new();
let mut object = Value::Object(payload);
substitute(&mut object, &by_digest, &mut used, &mut seen_digests)?;
if used != by_digest.len() {
return Err(SdJwtError::UnusedDisclosures(by_digest.len() - used));
}
match object {
Value::Object(map) => Ok(map),
_ => Err(SdJwtError::PayloadNotAnObject),
}
}
}
fn substitute(
value: &mut Value,
by_digest: &std::collections::HashMap<String, &Disclosure>,
used: &mut usize,
seen_digests: &mut std::collections::HashSet<String>,
) -> Result<(), SdJwtError> {
match value {
Value::Object(map) => {
let digests = match map.remove(SD_CLAIM) {
Some(Value::Array(items)) => items,
_ => Vec::new(),
};
map.remove(SD_ALG_CLAIM);
for digest in digests {
let Some(digest) = digest.as_str() else {
continue;
};
if !seen_digests.insert(digest.to_owned()) {
return Err(SdJwtError::DuplicateDigest(digest.to_owned()));
}
if let Some(d) = by_digest.get(digest) {
if map.contains_key(d.claim_name()) {
return Err(SdJwtError::ClaimCollision(d.claim_name().to_owned()));
}
map.insert(d.claim_name().to_owned(), d.claim_value().clone());
*used += 1;
}
}
for (_, v) in map.iter_mut() {
substitute(v, by_digest, used, seen_digests)?;
}
Ok(())
}
Value::Array(items) => {
for item in items {
substitute(item, by_digest, used, seen_digests)?;
}
Ok(())
}
_ => Ok(()),
}
}
fn decode_jwt_payload(jwt: &str) -> Result<Map<String, Value>, SdJwtError> {
use base64::Engine;
let payload_b64 = jwt.split('.').nth(1).ok_or(SdJwtError::MalformedJwt)?;
if jwt.split('.').count() != 3 {
return Err(SdJwtError::MalformedJwt);
}
let bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
.decode(payload_b64)
.map_err(|_| SdJwtError::MalformedJwt)?;
match serde_json::from_slice(&bytes) {
Ok(Value::Object(map)) => Ok(map),
Ok(_) => Err(SdJwtError::PayloadNotAnObject),
Err(_) => Err(SdJwtError::MalformedJwt),
}
}
pub fn build_payload(mut object: Map<String, Value>, concealed_any: bool) -> Value {
if concealed_any {
object.insert(
SD_ALG_CLAIM.to_owned(),
Value::String(SD_HASH_ALG.to_owned()),
);
}
Value::Object(object)
}