pub mod cache;
pub mod did;
pub mod jwt;
pub mod ssrf;
use std::num::NonZeroUsize;
use std::sync::Arc;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use proto_blue_crypto::{K256Keypair, Verifier as _, k256_compress_pubkey, parse_multikey};
use cache::{CachedResolve, DidDocCache, JtiCache};
use did::{DidResolver, HttpDidResolver, ResolveError};
const MODERATOR_KEY_FRAGMENT: &str = "#atproto";
const ACCEPTED_JWT_ALG: &str = "ES256K";
#[derive(Debug, Clone)]
pub struct AuthConfig {
pub service_did: String,
pub plc_directory_url: String,
pub resolver_timeout: Duration,
pub positive_cache_ttl: Duration,
pub negative_cache_ttl: Duration,
pub doc_cache_size: NonZeroUsize,
pub jti_cache_size: NonZeroUsize,
pub clock_skew: Duration,
pub max_iat_future: Duration,
}
impl Default for AuthConfig {
fn default() -> Self {
Self {
service_did: String::new(),
plc_directory_url: "https://plc.directory".to_string(),
resolver_timeout: Duration::from_secs(5),
positive_cache_ttl: Duration::from_secs(60),
negative_cache_ttl: Duration::from_secs(5),
doc_cache_size: NonZeroUsize::new(1024).expect("non-zero"),
jti_cache_size: NonZeroUsize::new(100_000).expect("non-zero"),
clock_skew: Duration::from_secs(30),
max_iat_future: Duration::from_secs(30),
}
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct VerifiedCaller {
pub iss: String,
pub key_id: String,
}
#[derive(Debug, thiserror::Error)]
pub enum AuthError {
#[error("JWT alg not allowed: {0}")]
AlgRejected(String),
#[error("JWT structure invalid: {0}")]
Structural(#[from] jwt::JwtParseError),
#[error("resolver error: {0}")]
Resolve(#[from] ResolveError),
#[error("DID doc has no `{MODERATOR_KEY_FRAGMENT}` verification method")]
NoVerificationMethod,
#[error("verification method has wrong key type (expected Multikey ES256K)")]
WrongKeyType,
#[error("signature verification failed")]
SignatureInvalid,
#[error("claim mismatch: {0}")]
ClaimMismatch(&'static str),
#[error("replay detected (iss={iss}, jti={jti})")]
Replay {
iss: String,
jti: String,
},
#[error("crypto error: {0}")]
Crypto(#[from] proto_blue_crypto::CryptoError),
#[error("system clock before unix epoch")]
Clock,
}
pub struct AuthContext {
config: AuthConfig,
resolver: Arc<dyn DidResolver>,
doc_cache: DidDocCache,
jti_cache: JtiCache,
}
impl AuthContext {
pub fn new(config: AuthConfig) -> Self {
let resolver: Arc<dyn DidResolver> = Arc::new(HttpDidResolver::new(
config.plc_directory_url.clone(),
config.resolver_timeout,
));
Self::with_resolver(config, resolver)
}
pub fn with_resolver(config: AuthConfig, resolver: Arc<dyn DidResolver>) -> Self {
let clock: Arc<dyn cache::Clock> = Arc::new(cache::SystemClock);
let doc_cache = DidDocCache::new(
config.doc_cache_size,
config.positive_cache_ttl,
config.negative_cache_ttl,
clock.clone(),
);
let jti_cache = JtiCache::new(config.jti_cache_size, clock);
Self {
config,
resolver,
doc_cache,
jti_cache,
}
}
pub async fn verify_service_auth(
&self,
token: &str,
expected_lxm: &str,
) -> Result<VerifiedCaller, AuthError> {
let parsed = jwt::parse(token)?;
if parsed.header.alg != ACCEPTED_JWT_ALG {
return Err(AuthError::AlgRejected(parsed.header.alg));
}
let doc = self.resolve_did(&parsed.payload.iss).await?;
let vm = doc
.find_verification_method(MODERATOR_KEY_FRAGMENT)
.ok_or(AuthError::NoVerificationMethod)?;
let parsed_key = parse_multikey(&vm.public_key_multibase)?;
if parsed_key.jwt_alg != ACCEPTED_JWT_ALG {
return Err(AuthError::WrongKeyType);
}
let compressed = k256_compress_pubkey(&parsed_key.key_bytes)?;
let verifier = K256Keypair::verifier_from_compressed(&compressed)?;
let sig_ok = verifier.verify(&parsed.signing_input, &parsed.signature)?;
if !sig_ok {
return Err(AuthError::SignatureInvalid);
}
if parsed.payload.aud != self.config.service_did {
return Err(AuthError::ClaimMismatch("aud"));
}
let now = unix_seconds()?;
let skew = self.config.clock_skew.as_secs() as i64;
if parsed.payload.exp < now.saturating_sub(skew) {
return Err(AuthError::ClaimMismatch("exp"));
}
let max_iat_future = self.config.max_iat_future.as_secs() as i64;
if parsed.payload.iat > now.saturating_add(max_iat_future) {
return Err(AuthError::ClaimMismatch("iat"));
}
if parsed.payload.lxm != expected_lxm {
return Err(AuthError::ClaimMismatch("lxm"));
}
let ttl_secs = (parsed.payload.exp - now).max(0) as u64;
let expires_at = Instant::now() + Duration::from_secs(ttl_secs);
self.jti_cache
.check_and_record(&parsed.payload.iss, &parsed.payload.jti, expires_at)
.map_err(|_| AuthError::Replay {
iss: parsed.payload.iss.clone(),
jti: parsed.payload.jti.clone(),
})?;
Ok(VerifiedCaller {
iss: parsed.payload.iss,
key_id: vm.id.clone(),
})
}
async fn resolve_did(&self, did: &str) -> Result<did::DidDocument, AuthError> {
if let Some(cached) = self.doc_cache.get(did) {
return match cached {
CachedResolve::Ok(doc) => Ok(doc),
CachedResolve::Err => Err(AuthError::Resolve(ResolveError::Network(
"negatively cached".into(),
))),
};
}
match self.resolver.resolve(did).await {
Ok(doc) => {
self.doc_cache.insert_ok(did.to_owned(), doc.clone());
Ok(doc)
}
Err(e) => {
self.doc_cache.insert_err(did.to_owned());
Err(AuthError::Resolve(e))
}
}
}
}
fn unix_seconds() -> Result<i64, AuthError> {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs() as i64)
.map_err(|_| AuthError::Clock)
}