use dashmap::DashMap;
use jsonwebtoken::DecodingKey;
use moka::sync::Cache;
use p256::pkcs8::{EncodePublicKey, LineEnding};
use serde::Deserialize;
use std::net::ToSocketAddrs;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
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>,
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 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", ":");
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)]
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
});
}
}
}
async fn validate_and_resolve_domain(
domain: &str,
global_counter: Arc<AtomicU64>,
domain_fetch_counts: Arc<DashMap<String, Arc<AtomicU64>>>,
) -> Result<(std::net::SocketAddr, ConcurrencyGuard, ConcurrencyGuard), String> {
let host_port = if domain.starts_with('[') {
if domain.contains("]:") {
domain.to_string()
} else {
format!("{domain}:443")
}
} else if domain.contains(':') {
domain.to_string()
} else {
format!("{domain}:443")
};
let domain_owned = domain.to_string();
tokio::task::spawn_blocking(move || {
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(format!(
"global did:web concurrency limit exceeded \
({GLOBAL_DIDWEB_FETCH_CONCURRENCY_LIMIT} concurrent fetches)"
));
}
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, domain_fetch_counts, domain_owned.clone());
if domain_prev >= DOMAIN_FETCH_CONCURRENCY_LIMIT {
return Err(format!(
"concurrency limit exceeded for did:web domain '{domain_owned}' \
({DOMAIN_FETCH_CONCURRENCY_LIMIT} concurrent fetches)"
));
}
let addrs: Vec<std::net::SocketAddr> = host_port
.to_socket_addrs()
.map(|iter| iter.collect::<Vec<_>>())
.map_err(|e| format!("failed to resolve domain '{domain_owned}': {e}"))?;
if addrs.is_empty() {
return Err(format!("domain '{domain_owned}' resolved to no addresses"));
}
for addr in &addrs {
let ip = match addr.ip() {
std::net::IpAddr::V6(v6) => v6
.to_ipv4_mapped()
.map(std::net::IpAddr::V4)
.unwrap_or(std::net::IpAddr::V6(v6)),
v4 => v4,
};
if ip.is_loopback() || ip.is_unspecified() || is_private_ip(&ip) || is_link_local(&ip) {
return Err(format!(
"did:web domain '{domain_owned}' resolves to private/loopback address {ip}, \
request blocked (SSRF protection)"
));
}
}
Ok((addrs[0], global_guard, domain_guard))
})
.await
.map_err(|e| format!("DNS resolution task panicked: {e}"))?
}
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) => domain[..=idx].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));
Self {
client,
cache,
negative_cache,
domain_clients,
domain_fetch_counts,
global_didweb_fetch_count,
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, String> {
let normalized = normalize_did(did);
let did = normalized.as_str();
if let Some(cached_err) = self.negative_cache.get(did) {
return Err(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.map_err(|e| {
negative_cache.insert(did_owned.clone(), e.clone());
e
})?;
extract_signing_key(&doc).map_err(|e| {
negative_cache.insert(did_owned.clone(), e.clone());
e
})
})
.await
.map_err(|e: Arc<String>| (*e).clone())
}
async fn fetch_did_document(&self, did: &str) -> Result<DidDocument, String> {
let url = did_document_url(did, &self.plc_directory)?;
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("invalid did:web")?;
let (domain_raw, _) = split_did_web_domain_path(raw)?;
let domain = domain_raw.replace("%3A", ":").replace("%3a", ":");
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(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(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),
)
.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| 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| format!("DID document fetch failed: {e}"))?;
let status = resp.status();
if !status.is_success() {
return Err(format!("DID directory returned {status} for {did}"));
}
if let Some(len) = resp.content_length() {
if len as usize > MAX_DID_DOCUMENT_SIZE {
return Err(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| format!("failed to read DID document body: {e}"))?;
if body.len() + chunk.len() > MAX_DID_DOCUMENT_SIZE {
return Err(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| 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 pem = public_key
.to_public_key_pem(LineEnding::LF)
.map_err(|e| format!("P-256 PEM encoding failed: {e}"))?;
let key = DecodingKey::from_ec_pem(pem.as_bytes())
.map_err(|e| format!("failed to create DecodingKey from P-256 PEM: {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 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), String> {
validate_and_resolve_domain(
domain,
Arc::new(AtomicU64::new(0)),
Arc::new(DashMap::new()),
)
.await
}
#[tokio::test]
async fn ssrf_rejects_loopback() {
let result = resolve_domain_for_test("localhost").await;
assert!(result.is_err(), "expected loopback to be rejected");
assert!(
result.unwrap_err().contains("SSRF protection"),
"error should mention SSRF"
);
}
#[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"
);
}
}