use base64::Engine;
use dashmap::DashMap;
use hickory_resolver::TokioResolver;
use hickory_resolver::config::ResolverOpts;
use jsonwebtoken::DecodingKey;
use moka::sync::Cache;
use p256::elliptic_curve::sec1::ToEncodedPoint;
use serde::Deserialize;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
#[derive(Clone, Debug)]
pub enum DidError {
Authoritative(String),
Transient(String),
}
impl DidError {
pub fn message(&self) -> &str {
match self {
DidError::Authoritative(s) | DidError::Transient(s) => s,
}
}
pub fn is_transient(&self) -> bool {
matches!(self, DidError::Transient(_))
}
}
impl std::fmt::Display for DidError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
self.message().fmt(f)
}
}
const DEFAULT_CACHE_TTL: Duration = Duration::from_secs(300);
const NEGATIVE_CACHE_TTL: Duration = Duration::from_secs(60);
const DEFAULT_PLC_DIRECTORY: &str = "https://plc.directory";
const DID_FETCH_TIMEOUT: Duration = Duration::from_secs(10);
const MAX_DID_DOCUMENT_SIZE: usize = 256 * 1024;
const MAX_CACHE_ENTRIES: u64 = 10_000;
const MAX_DOMAIN_CLIENTS: u64 = 100;
const DOMAIN_FETCH_CONCURRENCY_LIMIT: u64 = 10;
const GLOBAL_DIDWEB_FETCH_CONCURRENCY_LIMIT: u64 = 50;
#[derive(Clone)]
pub enum ResolvedKey {
P256(DecodingKey),
K256(k256::PublicKey),
}
#[derive(Clone)]
pub struct DidResolver {
client: reqwest::Client,
cache: moka::future::Cache<String, ResolvedKey>,
negative_cache: Cache<String, String>,
domain_clients: Cache<String, reqwest::Client>,
domain_fetch_counts: Arc<DashMap<String, Arc<AtomicU64>>>,
global_didweb_fetch_count: Arc<AtomicU64>,
async_resolver: TokioResolver,
plc_directory: String,
}
fn split_did_web_domain_path(raw: &str) -> Result<(&str, Option<&str>), String> {
if raw.starts_with('[') {
let bracket_end = raw
.find(']')
.ok_or_else(|| "did:web contains opening '[' with no closing ']'".to_string())?;
let after_bracket = &raw[bracket_end + 1..];
let rest = if let Some(stripped) = after_bracket
.strip_prefix("%3A")
.or_else(|| after_bracket.strip_prefix("%3a"))
{
let port_end = stripped
.find(':')
.map(|i| bracket_end + 1 + 3 + i) .unwrap_or(raw.len());
let domain = &raw[..port_end];
let path = if port_end < raw.len() {
Some(&raw[port_end + 1..]) } else {
None
};
return Ok((domain, path));
} else if after_bracket.starts_with(':') {
let domain = &raw[..bracket_end + 1];
let path = Some(&raw[bracket_end + 2..]); return Ok((domain, path));
} else {
&raw[bracket_end + 1..]
};
if rest.is_empty() {
Ok((raw, None))
} else {
Err(format!(
"unexpected characters after IPv6 bracket in did:web: '{raw}'"
))
}
} else {
match raw.find(':') {
Some(i) => Ok((&raw[..i], Some(&raw[i + 1..]))),
None => Ok((raw, None)),
}
}
}
fn normalize_did(did: &str) -> String {
if let Some(rest) = did.strip_prefix("did:web:") {
match split_did_web_domain_path(rest) {
Ok((domain, Some(path))) => {
format!("did:web:{}:{}", domain.to_ascii_lowercase(), path)
}
Ok((domain, None)) => {
format!("did:web:{}", domain.to_ascii_lowercase())
}
Err(_) => did.to_string(),
}
} else {
did.to_string()
}
}
fn canonicalize_did_web_domain(domain: &str) -> Result<String, String> {
if domain.starts_with('[') {
let bracket_end = domain
.find(']')
.ok_or_else(|| "did:web IPv6 host missing closing ']'".to_string())?;
let after = &domain[bracket_end + 1..];
if !after.is_empty() && !after.starts_with(':') {
return Err(format!(
"unexpected characters after IPv6 bracket in did:web: '{domain}'"
));
}
if let Some(port_str) = after.strip_prefix(':')
&& port_str.parse::<u16>().is_err()
{
return Err(format!(
"did:web port '{port_str}' is not a valid 1–65535 u16"
));
}
return Ok(domain.to_string());
}
let (host_part, port_part) = match domain.rsplit_once(':') {
Some((h, p)) if !h.is_empty() && !p.is_empty() => (h, Some(p)),
_ => (domain, None),
};
if let Some(p) = port_part
&& p.parse::<u16>().is_err()
{
return Err(format!("did:web port '{p}' is not a valid 1–65535 u16"));
}
let canonical_host = match url::Host::parse(host_part) {
Ok(url::Host::Domain(d)) => d,
Ok(url::Host::Ipv4(v4)) => v4.to_string(),
Ok(url::Host::Ipv6(v6)) => format!("[{v6}]"),
Err(e) => {
return Err(format!(
"did:web host '{host_part}' is not a valid hostname: {e}"
));
}
};
Ok(match port_part {
Some(p) => format!("{canonical_host}:{p}"),
None => canonical_host,
})
}
fn did_document_url(did: &str, plc_directory: &str) -> Result<String, String> {
if did.starts_with("did:plc:") {
Ok(format!("{}/{did}", plc_directory.trim_end_matches('/')))
} else if did.starts_with("did:web:") {
let raw = did.strip_prefix("did:web:").ok_or("invalid did:web")?;
let (domain_raw, path_raw) = split_did_web_domain_path(raw)?;
let domain = domain_raw.replace("%3A", ":").replace("%3a", ":");
let domain = canonicalize_did_web_domain(&domain)?;
if let Some(path_segments) = path_raw {
let path = path_segments.replace(':', "/");
Ok(format!("https://{domain}/{path}/did.json"))
} else {
Ok(format!("https://{domain}/.well-known/did.json"))
}
} else {
Err(format!("unsupported DID method: {did}"))
}
}
#[derive(Debug)]
#[allow(clippy::type_complexity)]
struct ConcurrencyGuard {
counter: Arc<AtomicU64>,
cleanup: Option<(Arc<DashMap<String, Arc<AtomicU64>>>, String)>,
}
impl ConcurrencyGuard {
fn new(counter: Arc<AtomicU64>) -> Self {
Self {
counter,
cleanup: None,
}
}
fn new_domain(
counter: Arc<AtomicU64>,
map: Arc<DashMap<String, Arc<AtomicU64>>>,
key: String,
) -> Self {
Self {
counter,
cleanup: Some((map, key)),
}
}
}
impl Drop for ConcurrencyGuard {
fn drop(&mut self) {
self.counter.fetch_sub(1, Ordering::Relaxed);
if let Some((map, key)) = &self.cleanup {
let ours = Arc::clone(&self.counter);
map.remove_if(key, |_, v| {
Arc::ptr_eq(v, &ours) && v.load(Ordering::Relaxed) == 0
});
}
}
}
fn split_host_port(domain: &str) -> (String, u16) {
if domain.starts_with('[') {
if let Some(end) = domain.find(']') {
let host = domain[1..end].to_string();
let after = &domain[end + 1..];
if let Some(port_str) = after.strip_prefix(':') {
let port = port_str.parse().unwrap_or(443);
return (host, port);
}
return (host, 443);
}
return (domain.to_string(), 443);
}
match domain.rsplit_once(':') {
Some((host, port_str)) if !host.is_empty() => match port_str.parse() {
Ok(p) => (host.to_string(), p),
Err(_) => (domain.to_string(), 443),
},
_ => (domain.to_string(), 443),
}
}
async fn validate_and_resolve_domain(
domain: &str,
global_counter: Arc<AtomicU64>,
domain_fetch_counts: Arc<DashMap<String, Arc<AtomicU64>>>,
async_resolver: &TokioResolver,
) -> Result<(std::net::SocketAddr, ConcurrencyGuard, ConcurrencyGuard), DidError> {
let global_prev = global_counter.fetch_add(1, Ordering::Relaxed);
let global_guard = ConcurrencyGuard::new(global_counter);
if global_prev >= GLOBAL_DIDWEB_FETCH_CONCURRENCY_LIMIT {
return Err(DidError::Transient(format!(
"global did:web concurrency limit exceeded \
({GLOBAL_DIDWEB_FETCH_CONCURRENCY_LIMIT} concurrent fetches)"
)));
}
let domain_owned = domain.to_string();
let domain_prev;
let domain_counter = {
let entry = domain_fetch_counts
.entry(domain_owned.clone())
.or_insert_with(|| Arc::new(AtomicU64::new(0)));
domain_prev = entry.fetch_add(1, Ordering::Relaxed);
Arc::clone(&*entry)
};
let domain_guard = ConcurrencyGuard::new_domain(
domain_counter,
Arc::clone(&domain_fetch_counts),
domain_owned.clone(),
);
if domain_prev >= DOMAIN_FETCH_CONCURRENCY_LIMIT {
return Err(DidError::Transient(format!(
"concurrency limit exceeded for did:web domain '{domain_owned}' \
({DOMAIN_FETCH_CONCURRENCY_LIMIT} concurrent fetches)"
)));
}
let (host, port) = split_host_port(&domain_owned);
let lookup = async_resolver.lookup_ip(host.as_str()).await.map_err(|e| {
DidError::Transient(format!("failed to resolve domain '{domain_owned}': {e}"))
})?;
let ips: Vec<std::net::IpAddr> = lookup.iter().collect();
if ips.is_empty() {
return Err(DidError::Transient(format!(
"domain '{domain_owned}' resolved to no addresses"
)));
}
for ip in &ips {
let checked = match ip {
std::net::IpAddr::V6(v6) => v6
.to_ipv4()
.map(std::net::IpAddr::V4)
.unwrap_or(std::net::IpAddr::V6(*v6)),
std::net::IpAddr::V4(v4) => std::net::IpAddr::V4(*v4),
};
if checked.is_loopback()
|| checked.is_unspecified()
|| is_private_ip(&checked)
|| is_link_local(&checked)
{
return Err(DidError::Authoritative(format!(
"did:web domain '{domain_owned}' resolves to private/loopback address {checked}, \
request blocked (SSRF protection)"
)));
}
}
Ok((
std::net::SocketAddr::new(ips[0], port),
global_guard,
domain_guard,
))
}
fn is_private_ip(ip: &std::net::IpAddr) -> bool {
match ip {
std::net::IpAddr::V4(v4) => {
let octets = v4.octets();
octets[0] == 0
|| octets[0] == 10
|| (octets[0] == 172 && (16..=31).contains(&octets[1]))
|| (octets[0] == 192 && octets[1] == 168)
|| (octets[0] == 100 && (64..=127).contains(&octets[1]))
|| (octets[0] == 169 && octets[1] == 254)
|| (octets[0] == 192 && octets[1] == 0 && octets[2] == 0)
|| (octets[0] == 192 && octets[1] == 0 && octets[2] == 2)
|| (octets[0] == 198 && (octets[1] == 18 || octets[1] == 19))
|| (octets[0] == 198 && octets[1] == 51 && octets[2] == 100)
|| (octets[0] == 203 && octets[1] == 0 && octets[2] == 113)
|| octets[0] >= 224
}
std::net::IpAddr::V6(v6) => {
let segments = v6.segments();
(segments[0] & 0xfe00) == 0xfc00
|| (segments[0] & 0xff00) == 0xff00
|| (segments[0] == 0x0100
&& segments[1] == 0
&& segments[2] == 0
&& segments[3] == 0)
|| (segments[0] == 0x2001 && segments[1] == 0x0002 && segments[2] == 0x0000)
|| (segments[0] == 0x2001 && segments[1] == 0x0db8)
}
}
}
fn is_link_local(ip: &std::net::IpAddr) -> bool {
match ip {
std::net::IpAddr::V4(v4) => v4.is_link_local(),
std::net::IpAddr::V6(v6) => (v6.segments()[0] & 0xffc0) == 0xfe80,
}
}
fn extract_hostname(domain: &str) -> String {
if domain.starts_with('[') {
match domain.find(']') {
Some(idx) => {
let after = &domain[idx + 1..];
if after.is_empty() || after.starts_with(':') {
domain[..=idx].to_string()
} else {
domain.to_string()
}
}
None => domain.to_string(),
}
} else {
domain.split(':').next().unwrap_or(domain).to_string()
}
}
impl DidResolver {
pub fn new() -> Self {
let client = reqwest::Client::builder()
.timeout(DID_FETCH_TIMEOUT)
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("failed to build reqwest client");
let cache = moka::future::Cache::builder()
.max_capacity(MAX_CACHE_ENTRIES)
.time_to_live(DEFAULT_CACHE_TTL)
.build();
let negative_cache = Cache::builder()
.max_capacity(MAX_CACHE_ENTRIES)
.time_to_live(NEGATIVE_CACHE_TTL)
.build();
let domain_clients = Cache::builder()
.max_capacity(MAX_DOMAIN_CLIENTS)
.time_to_live(DEFAULT_CACHE_TTL)
.build();
let domain_fetch_counts = Arc::new(DashMap::new());
let global_didweb_fetch_count = Arc::new(AtomicU64::new(0));
let mut opts = ResolverOpts::default();
opts.timeout = Duration::from_secs(3);
opts.attempts = 2;
let async_resolver = TokioResolver::builder_tokio()
.expect("failed to initialise DNS resolver (system config)")
.with_options(opts)
.build()
.expect("failed to build DNS resolver");
Self {
client,
cache,
negative_cache,
domain_clients,
domain_fetch_counts,
global_didweb_fetch_count,
async_resolver,
plc_directory: DEFAULT_PLC_DIRECTORY.to_string(),
}
}
#[cfg(test)]
pub fn with_plc_directory(plc_directory: String) -> Self {
let mut resolver = Self::new();
resolver.plc_directory = plc_directory;
resolver
}
pub async fn resolve_key(&self, did: &str) -> Result<ResolvedKey, DidError> {
let normalized = normalize_did(did);
let did = normalized.as_str();
if let Some(cached_err) = self.negative_cache.get(did) {
return Err(DidError::Authoritative(format!(
"DID resolution recently failed (cached): {cached_err}"
)));
}
let negative_cache = self.negative_cache.clone();
let did_owned = did.to_string();
self.cache
.try_get_with(did.to_string(), async {
let doc = self.fetch_did_document(&did_owned).await.inspect_err(|e| {
if !e.is_transient() {
negative_cache.insert(did_owned.clone(), e.message().to_string());
}
})?;
extract_signing_key(&doc)
.map_err(DidError::Authoritative)
.inspect_err(|e| {
negative_cache.insert(did_owned.clone(), e.message().to_string());
})
})
.await
.map_err(|e: Arc<DidError>| (*e).clone())
}
async fn fetch_did_document(&self, did: &str) -> Result<DidDocument, DidError> {
let url = did_document_url(did, &self.plc_directory).map_err(DidError::Authoritative)?;
let mut _global_guard: Option<ConcurrencyGuard> = None;
let mut _domain_guard: Option<ConcurrencyGuard> = None;
let pinned_client = if did.starts_with("did:web:") {
let raw = did
.strip_prefix("did:web:")
.ok_or_else(|| DidError::Authoritative("invalid did:web".to_string()))?;
let (domain_raw, _) =
split_did_web_domain_path(raw).map_err(DidError::Authoritative)?;
let domain = domain_raw.replace("%3A", ":").replace("%3a", ":");
let domain = canonicalize_did_web_domain(&domain).map_err(DidError::Authoritative)?;
let cache_key = domain.clone();
let client = if let Some(cached) = self.domain_clients.get(&cache_key) {
let global_prev = self
.global_didweb_fetch_count
.fetch_add(1, Ordering::Relaxed);
_global_guard = Some(ConcurrencyGuard::new(Arc::clone(
&self.global_didweb_fetch_count,
)));
if global_prev >= GLOBAL_DIDWEB_FETCH_CONCURRENCY_LIMIT {
return Err(DidError::Transient(format!(
"global did:web concurrency limit exceeded \
({GLOBAL_DIDWEB_FETCH_CONCURRENCY_LIMIT} concurrent fetches)"
)));
}
let domain_prev;
let domain_counter = {
let entry = self
.domain_fetch_counts
.entry(domain.clone())
.or_insert_with(|| Arc::new(AtomicU64::new(0)));
domain_prev = entry.fetch_add(1, Ordering::Relaxed);
Arc::clone(&*entry)
};
_domain_guard = Some(ConcurrencyGuard::new_domain(
domain_counter,
Arc::clone(&self.domain_fetch_counts),
domain.clone(),
));
if domain_prev >= DOMAIN_FETCH_CONCURRENCY_LIMIT {
return Err(DidError::Transient(format!(
"concurrency limit exceeded for did:web domain '{domain}' \
({DOMAIN_FETCH_CONCURRENCY_LIMIT} concurrent fetches)"
)));
}
cached
} else {
let (validated_addr, global_guard, domain_guard) = validate_and_resolve_domain(
&domain,
Arc::clone(&self.global_didweb_fetch_count),
Arc::clone(&self.domain_fetch_counts),
&self.async_resolver,
)
.await?;
_global_guard = Some(global_guard);
_domain_guard = Some(domain_guard);
let host_only = extract_hostname(&domain);
let pinned = reqwest::Client::builder()
.timeout(DID_FETCH_TIMEOUT)
.redirect(reqwest::redirect::Policy::none())
.resolve(&host_only, validated_addr)
.build()
.map_err(|e| {
DidError::Transient(format!("failed to build pinned HTTP client: {e}"))
})?;
self.domain_clients.insert(cache_key, pinned.clone());
pinned
};
Some(client)
} else {
None
};
let client = pinned_client.as_ref().unwrap_or(&self.client);
let resp = client
.get(&url)
.header("Accept", "application/json")
.send()
.await
.map_err(|e| DidError::Transient(format!("DID document fetch failed: {e}")))?;
let status = resp.status();
if !status.is_success() {
let msg = format!("DID directory returned {status} for {did}");
if status.is_client_error() {
return Err(DidError::Authoritative(msg));
}
return Err(DidError::Transient(msg));
}
if let Some(len) = resp.content_length()
&& len as usize > MAX_DID_DOCUMENT_SIZE
{
return Err(DidError::Authoritative(format!(
"DID document response too large ({len} bytes, max {MAX_DID_DOCUMENT_SIZE})"
)));
}
use futures_util::StreamExt;
let mut stream = resp.bytes_stream();
let mut body = Vec::new();
while let Some(chunk) = stream.next().await {
let chunk = chunk.map_err(|e| {
DidError::Transient(format!("failed to read DID document body: {e}"))
})?;
if body.len() + chunk.len() > MAX_DID_DOCUMENT_SIZE {
return Err(DidError::Authoritative(format!(
"DID document response too large (>{MAX_DID_DOCUMENT_SIZE} bytes), aborting"
)));
}
body.extend_from_slice(&chunk);
}
serde_json::from_slice::<DidDocument>(&body)
.map_err(|e| DidError::Authoritative(format!("invalid DID document JSON: {e}")))
}
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
struct DidDocument {
verification_method: Option<Vec<VerificationMethod>>,
}
#[derive(Debug, Deserialize)]
#[serde(rename_all = "camelCase")]
struct VerificationMethod {
id: String,
#[serde(rename = "type")]
key_type: String,
public_key_multibase: Option<String>,
}
const MAX_MULTIBASE_ENCODED_LEN: usize = 128;
const P256_MULTICODEC: [u8; 2] = [0x80, 0x24];
const K256_MULTICODEC: [u8; 2] = [0xe7, 0x01];
fn extract_signing_key(doc: &DidDocument) -> Result<ResolvedKey, String> {
let methods = doc
.verification_method
.as_ref()
.ok_or("DID document has no verificationMethod")?;
let method = methods
.iter()
.find(|m| m.id.ends_with("#atproto"))
.ok_or("no #atproto verification method in DID document")?;
if method.key_type != "Multikey" {
return Err(format!(
"unsupported verification method type: {}",
method.key_type
));
}
let multibase = method
.public_key_multibase
.as_ref()
.ok_or("verification method missing publicKeyMultibase")?;
let encoded = multibase.strip_prefix('z').ok_or_else(|| {
let prefix = multibase
.chars()
.next()
.map_or_else(|| "(empty)".to_string(), |c| c.to_string());
format!("unsupported multibase prefix: {prefix}")
})?;
if encoded.len() > MAX_MULTIBASE_ENCODED_LEN {
return Err(format!(
"publicKeyMultibase too long ({} chars, max {MAX_MULTIBASE_ENCODED_LEN}); \
valid P-256 keys encode to ~48 characters",
encoded.len()
));
}
let bytes = bs58::decode(encoded)
.into_vec()
.map_err(|e| format!("base58 decode failed: {e}"))?;
if bytes.len() < 2 {
return Err("multicodec key too short".to_string());
}
let (prefix, key_bytes) = bytes.split_at(2);
if prefix == P256_MULTICODEC {
decode_p256_key(key_bytes)
} else if prefix == K256_MULTICODEC {
decode_k256_key(key_bytes)
} else {
Err(format!(
"unknown multicodec prefix: [{:#04x}, {:#04x}]",
prefix[0], prefix[1]
))
}
}
fn decode_p256_key(compressed: &[u8]) -> Result<ResolvedKey, String> {
let public_key = p256::PublicKey::from_sec1_bytes(compressed)
.map_err(|e| format!("invalid P-256 public key: {e}"))?;
let encoded = public_key.to_encoded_point(false);
let x_bytes = encoded
.x()
.ok_or("P-256 public key missing x coordinate after decompression")?;
let y_bytes = encoded
.y()
.ok_or("P-256 public key missing y coordinate after decompression")?;
let x_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(x_bytes);
let y_b64 = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(y_bytes);
let key = DecodingKey::from_ec_components(&x_b64, &y_b64)
.map_err(|e| format!("failed to create DecodingKey from P-256 components: {e}"))?;
Ok(ResolvedKey::P256(key))
}
fn decode_k256_key(compressed: &[u8]) -> Result<ResolvedKey, String> {
let public_key = k256::PublicKey::from_sec1_bytes(compressed)
.map_err(|e| format!("invalid secp256k1 public key: {e}"))?;
Ok(ResolvedKey::K256(public_key))
}
#[cfg(test)]
mod tests {
use super::*;
fn make_test_multibase_key() -> String {
use p256::elliptic_curve::sec1::ToEncodedPoint;
let secret = p256::SecretKey::from_slice(&[
0x9f, 0x86, 0xd0, 0x81, 0x88, 0x4c, 0x7d, 0x65, 0x9a, 0x2f, 0xea, 0xa0, 0xc5, 0x5a,
0xd0, 0x15, 0xa3, 0xbf, 0x4f, 0x1b, 0x2b, 0x0b, 0x82, 0x2c, 0xd1, 0x5d, 0x6c, 0x15,
0xb0, 0xf0, 0x0a, 0x08,
])
.expect("valid test key");
let public = secret.public_key();
let point = public.to_encoded_point(true); let compressed = point.as_bytes();
let mut prefixed = Vec::with_capacity(2 + compressed.len());
prefixed.extend_from_slice(&P256_MULTICODEC);
prefixed.extend_from_slice(compressed);
let encoded = bs58::encode(&prefixed).into_string();
format!("z{encoded}")
}
fn make_test_k256_multibase_key() -> String {
use k256::elliptic_curve::sec1::ToEncodedPoint;
let secret = k256::SecretKey::from_slice(&[
0x9f, 0x86, 0xd0, 0x81, 0x88, 0x4c, 0x7d, 0x65, 0x9a, 0x2f, 0xea, 0xa0, 0xc5, 0x5a,
0xd0, 0x15, 0xa3, 0xbf, 0x4f, 0x1b, 0x2b, 0x0b, 0x82, 0x2c, 0xd1, 0x5d, 0x6c, 0x15,
0xb0, 0xf0, 0x0a, 0x08,
])
.expect("valid test key");
let public = secret.public_key();
let point = public.to_encoded_point(true);
let compressed = point.as_bytes();
let mut prefixed = Vec::with_capacity(2 + compressed.len());
prefixed.extend_from_slice(&K256_MULTICODEC);
prefixed.extend_from_slice(compressed);
let encoded = bs58::encode(&prefixed).into_string();
format!("z{encoded}")
}
fn assert_err_contains(result: Result<ResolvedKey, String>, needle: &str) {
match result {
Err(e) => assert!(e.contains(needle), "unexpected error: {e}"),
Ok(_) => panic!("expected error containing '{needle}', got Ok"),
}
}
#[test]
fn extract_p256_key_from_did_document() {
let multibase_key = make_test_multibase_key();
let doc = DidDocument {
verification_method: Some(vec![VerificationMethod {
id: "did:plc:test#atproto".to_string(),
key_type: "Multikey".to_string(),
public_key_multibase: Some(multibase_key),
}]),
};
let result = extract_signing_key(&doc);
assert!(result.is_ok(), "failed to extract key: {:?}", result.err());
assert!(matches!(result.unwrap(), ResolvedKey::P256(_)));
}
#[test]
fn rejects_missing_atproto_method() {
let doc = DidDocument {
verification_method: Some(vec![VerificationMethod {
id: "did:plc:test#other".to_string(),
key_type: "Multikey".to_string(),
public_key_multibase: Some("zNotUsed".to_string()),
}]),
};
assert_err_contains(extract_signing_key(&doc), "#atproto");
}
#[test]
fn extracts_k256_key_from_did_document() {
let multibase_key = make_test_k256_multibase_key();
let doc = DidDocument {
verification_method: Some(vec![VerificationMethod {
id: "did:plc:test#atproto".to_string(),
key_type: "Multikey".to_string(),
public_key_multibase: Some(multibase_key),
}]),
};
let result = extract_signing_key(&doc);
assert!(
result.is_ok(),
"failed to extract K-256 key: {:?}",
result.err()
);
assert!(matches!(result.unwrap(), ResolvedKey::K256(_)));
}
#[test]
fn rejects_empty_verification_methods() {
let doc = DidDocument {
verification_method: None,
};
assert_err_contains(extract_signing_key(&doc), "no verificationMethod");
}
#[test]
fn rejects_unsupported_key_type() {
let doc = DidDocument {
verification_method: Some(vec![VerificationMethod {
id: "did:plc:test#atproto".to_string(),
key_type: "Ed25519VerificationKey2020".to_string(),
public_key_multibase: Some("zNotUsed".to_string()),
}]),
};
assert_err_contains(
extract_signing_key(&doc),
"unsupported verification method type",
);
}
#[test]
fn did_web_domain_only_resolves_to_well_known() {
let url = did_document_url("did:web:example.com", DEFAULT_PLC_DIRECTORY).unwrap();
assert_eq!(url, "https://example.com/.well-known/did.json");
}
#[test]
fn did_web_path_based_resolves_without_well_known() {
let url = did_document_url("did:web:example.com:u:alice", DEFAULT_PLC_DIRECTORY).unwrap();
assert_eq!(url, "https://example.com/u/alice/did.json");
}
#[test]
fn did_web_single_path_segment() {
let url = did_document_url("did:web:example.com:users", DEFAULT_PLC_DIRECTORY).unwrap();
assert_eq!(url, "https://example.com/users/did.json");
}
#[test]
fn did_web_percent_encoded_port() {
let url = did_document_url("did:web:example.com%3A8443", DEFAULT_PLC_DIRECTORY).unwrap();
assert_eq!(url, "https://example.com:8443/.well-known/did.json");
}
#[test]
fn did_web_percent_encoded_port_with_path() {
let url =
did_document_url("did:web:example.com%3A8443:u:bob", DEFAULT_PLC_DIRECTORY).unwrap();
assert_eq!(url, "https://example.com:8443/u/bob/did.json");
}
#[test]
fn did_plc_resolves_to_plc_directory() {
let url = did_document_url("did:plc:abc123", "https://plc.directory").unwrap();
assert_eq!(url, "https://plc.directory/did:plc:abc123");
}
#[test]
fn unsupported_did_method_returns_error() {
let result = did_document_url("did:key:z123", DEFAULT_PLC_DIRECTORY);
let err = result.unwrap_err();
assert!(
err.contains("unsupported DID method"),
"unexpected error: {err}"
);
}
#[test]
fn did_web_ipv6_domain_only() {
let url = did_document_url("did:web:[2001:db8::1]", DEFAULT_PLC_DIRECTORY).unwrap();
assert_eq!(url, "https://[2001:db8::1]/.well-known/did.json");
}
#[test]
fn did_web_ipv6_with_path() {
let url = did_document_url("did:web:[2001:db8::1]:u:alice", DEFAULT_PLC_DIRECTORY).unwrap();
assert_eq!(url, "https://[2001:db8::1]/u/alice/did.json");
}
#[test]
fn did_web_ipv6_with_port() {
let url = did_document_url("did:web:[2001:db8::1]%3A8443", DEFAULT_PLC_DIRECTORY).unwrap();
assert_eq!(url, "https://[2001:db8::1]:8443/.well-known/did.json");
}
#[test]
fn did_web_ipv6_with_port_and_path() {
let url =
did_document_url("did:web:[2001:db8::1]%3A8443:u:bob", DEFAULT_PLC_DIRECTORY).unwrap();
assert_eq!(url, "https://[2001:db8::1]:8443/u/bob/did.json");
}
#[test]
fn canonicalize_unicode_hostname_to_punycode() {
let canonical = canonicalize_did_web_domain("tést.attacker.com").unwrap();
assert_eq!(canonical, "xn--tst-bma.attacker.com");
}
#[test]
fn canonicalize_preserves_port_after_idna() {
let canonical = canonicalize_did_web_domain("tést.attacker.com:8443").unwrap();
assert_eq!(canonical, "xn--tst-bma.attacker.com:8443");
}
#[test]
fn canonicalize_ascii_hostname_unchanged() {
let canonical = canonicalize_did_web_domain("example.com").unwrap();
assert_eq!(canonical, "example.com");
}
#[test]
fn canonicalize_lowercases_ascii_hostname() {
let canonical = canonicalize_did_web_domain("EXAMPLE.com").unwrap();
assert_eq!(canonical, "example.com");
}
#[test]
fn canonicalize_ipv6_bracketed_unchanged() {
let canonical = canonicalize_did_web_domain("[2001:db8::1]").unwrap();
assert_eq!(canonical, "[2001:db8::1]");
}
#[test]
fn canonicalize_ipv6_with_port_unchanged() {
let canonical = canonicalize_did_web_domain("[2001:db8::1]:8443").unwrap();
assert_eq!(canonical, "[2001:db8::1]:8443");
}
#[test]
fn canonicalize_rejects_unclosed_ipv6_bracket() {
let result = canonicalize_did_web_domain("[2001:db8::1");
assert!(result.is_err());
assert!(result.unwrap_err().contains("missing closing"));
}
#[test]
fn canonicalize_rejects_non_numeric_port() {
let result = canonicalize_did_web_domain("attacker.com:abcd");
assert!(result.is_err(), "expected non-numeric port to be rejected");
assert!(result.unwrap_err().contains("not a valid"));
}
#[test]
fn canonicalize_rejects_out_of_range_port() {
let result = canonicalize_did_web_domain("example.com:99999");
assert!(result.is_err());
}
#[test]
fn canonicalize_rejects_non_numeric_port_with_ipv6() {
let result = canonicalize_did_web_domain("[2001:db8::1]:abcd");
assert!(result.is_err());
assert!(result.unwrap_err().contains("not a valid"));
}
#[test]
fn did_web_idn_url_uses_punycode_host() {
let url = did_document_url("did:web:tést.attacker.com", DEFAULT_PLC_DIRECTORY).unwrap();
assert_eq!(url, "https://xn--tst-bma.attacker.com/.well-known/did.json");
}
#[test]
fn did_web_idn_with_port_url_uses_punycode_host() {
let url = did_document_url(
"did:web:tést.attacker.com%3A8443:u:bob",
DEFAULT_PLC_DIRECTORY,
)
.unwrap();
assert_eq!(url, "https://xn--tst-bma.attacker.com:8443/u/bob/did.json");
}
#[test]
fn did_web_ipv6_unclosed_bracket_errors() {
let result = did_document_url("did:web:[2001:db8::1", DEFAULT_PLC_DIRECTORY);
assert!(result.is_err());
assert!(result.unwrap_err().contains("no closing ']'"));
}
async fn resolve_domain_for_test(
domain: &str,
) -> Result<(std::net::SocketAddr, ConcurrencyGuard, ConcurrencyGuard), DidError> {
let mut opts = ResolverOpts::default();
opts.timeout = Duration::from_secs(3);
opts.attempts = 2;
let resolver = TokioResolver::builder_tokio()
.expect("build resolver")
.with_options(opts)
.build()
.expect("build resolver");
validate_and_resolve_domain(
domain,
Arc::new(AtomicU64::new(0)),
Arc::new(DashMap::new()),
&resolver,
)
.await
}
#[tokio::test]
async fn ssrf_rejects_loopback() {
let result = resolve_domain_for_test("localhost").await;
let err = result.expect_err("expected loopback to be rejected");
assert!(
err.message().contains("SSRF protection"),
"error should mention SSRF: {err}"
);
assert!(
!err.is_transient(),
"SSRF block must be authoritative so the negative cache short-circuits retries"
);
}
#[tokio::test]
async fn ssrf_rejects_127_0_0_1() {
let result = resolve_domain_for_test("127.0.0.1").await;
assert!(result.is_err(), "expected 127.0.0.1 to be rejected");
}
#[tokio::test]
async fn ssrf_rejects_private_10_network() {
let result = resolve_domain_for_test("10.0.0.1").await;
assert!(result.is_err(), "expected 10.x.x.x to be rejected");
}
#[tokio::test]
async fn ssrf_rejects_private_172_network() {
let result = resolve_domain_for_test("172.16.0.1").await;
assert!(result.is_err(), "expected 172.16.x.x to be rejected");
}
#[tokio::test]
async fn ssrf_rejects_private_192_168_network() {
let result = resolve_domain_for_test("192.168.1.1").await;
assert!(result.is_err(), "expected 192.168.x.x to be rejected");
}
#[tokio::test]
async fn ssrf_rejects_link_local() {
let result = resolve_domain_for_test("169.254.1.1").await;
assert!(result.is_err(), "expected link-local to be rejected");
}
#[tokio::test]
async fn ssrf_rejects_loopback_with_port() {
let result = resolve_domain_for_test("127.0.0.1:10250").await;
assert!(result.is_err(), "expected 127.0.0.1:port to be rejected");
}
#[tokio::test]
async fn ssrf_rejects_this_network_0_0_0_1() {
let result = resolve_domain_for_test("0.0.0.1").await;
assert!(
result.is_err(),
"expected 0.0.0.1 (this-network) to be rejected"
);
}
#[tokio::test]
async fn ssrf_rejects_multicast() {
let result = resolve_domain_for_test("224.0.0.1").await;
assert!(
result.is_err(),
"expected multicast 224.0.0.1 to be rejected"
);
}
#[tokio::test]
async fn ssrf_rejects_broadcast() {
let result = resolve_domain_for_test("255.255.255.255").await;
assert!(result.is_err(), "expected broadcast to be rejected");
}
#[test]
fn rejects_oversized_multibase_key() {
let long_key = format!("z{}", "1".repeat(MAX_MULTIBASE_ENCODED_LEN + 1));
let doc = DidDocument {
verification_method: Some(vec![VerificationMethod {
id: "did:plc:test#atproto".to_string(),
key_type: "Multikey".to_string(),
public_key_multibase: Some(long_key),
}]),
};
assert_err_contains(extract_signing_key(&doc), "too long");
}
#[test]
fn multibyte_multibase_prefix_does_not_panic() {
let doc = DidDocument {
verification_method: Some(vec![VerificationMethod {
id: "did:plc:test#atproto".to_string(),
key_type: "Multikey".to_string(),
public_key_multibase: Some("\u{1F680}rest".to_string()), }]),
};
let result = extract_signing_key(&doc);
let err = match result {
Err(e) => e,
Ok(_) => panic!("should reject non-'z' prefix"),
};
assert!(
err.contains("unsupported multibase prefix"),
"unexpected error: {err}"
);
assert!(
err.contains('\u{1F680}'),
"error should include the actual prefix char"
);
}
#[tokio::test]
async fn transient_resolver_errors_are_not_negative_cached() {
let resolver = DidResolver::with_plc_directory("http://127.0.0.1:1".to_string());
let err = resolver
.resolve_key("did:plc:probe1")
.await
.err()
.expect("unreachable PLC must fail");
assert!(
err.is_transient(),
"connect-refused must map to Transient, got: {err:?}"
);
assert!(
resolver.negative_cache.get("did:plc:probe1").is_none(),
"transient errors must not be negative-cached"
);
}
#[tokio::test]
async fn ssrf_errors_are_authoritative_and_cached() {
let resolver = DidResolver::new();
let did = "did:web:127.0.0.1";
let err = resolver
.resolve_key(did)
.await
.err()
.expect("loopback did:web must be rejected");
assert!(
!err.is_transient(),
"SSRF block must be Authoritative so the negative cache can retain it, got: {err:?}"
);
assert!(
resolver.negative_cache.get(&normalize_did(did)).is_some(),
"authoritative SSRF block should populate the negative cache"
);
}
}