use async_trait::async_trait;
use serde::Deserialize;
use super::ssrf::SafeDnsResolver;
#[derive(Debug, thiserror::Error)]
pub enum ResolveError {
#[error("unsupported DID method: {0}")]
UnsupportedMethod(String),
#[error("malformed DID: {0}")]
Malformed(String),
#[error("network error: {0}")]
Network(String),
#[error("non-success HTTP status: {0}")]
BadStatus(u16),
#[error("body parse error: {0}")]
Parse(String),
#[error("SSRF protection rejected host: {0}")]
SsrfBlocked(String),
}
#[derive(Debug, Clone, Deserialize)]
pub struct DidDocument {
pub id: String,
#[serde(rename = "verificationMethod", default)]
pub verification_method: Vec<VerificationMethod>,
}
#[derive(Debug, Clone, Deserialize)]
pub struct VerificationMethod {
pub id: String,
#[serde(rename = "type")]
pub r#type: String,
#[serde(rename = "publicKeyMultibase")]
pub public_key_multibase: String,
}
impl DidDocument {
pub fn find_verification_method(&self, fragment: &str) -> Option<&VerificationMethod> {
self.verification_method
.iter()
.find(|vm| vm.id.ends_with(fragment))
}
}
#[async_trait]
pub trait DidResolver: Send + Sync {
async fn resolve(&self, did: &str) -> Result<DidDocument, ResolveError>;
}
#[derive(Debug, Clone)]
pub struct HttpDidResolver {
client: reqwest::Client,
plc_directory_url: String,
}
impl HttpDidResolver {
pub fn new(plc_directory_url: String, timeout: std::time::Duration) -> Self {
Self::with_dns_resolver(plc_directory_url, timeout, SafeDnsResolver::arc())
}
pub fn with_dns_resolver<R: reqwest::dns::Resolve + 'static>(
plc_directory_url: String,
timeout: std::time::Duration,
dns: std::sync::Arc<R>,
) -> Self {
let client = reqwest::Client::builder()
.timeout(timeout)
.connect_timeout(timeout)
.dns_resolver(dns)
.redirect(reqwest::redirect::Policy::limited(3))
.build()
.expect("reqwest client build");
Self {
client,
plc_directory_url,
}
}
pub fn did_web_url(did: &str) -> Result<url::Url, ResolveError> {
let rest = did
.strip_prefix("did:web:")
.ok_or_else(|| ResolveError::Malformed("not a did:web".into()))?;
if rest.is_empty() {
return Err(ResolveError::Malformed("empty did:web body".into()));
}
let mut parts = rest.splitn(2, ':');
let host_encoded = parts.next().unwrap_or("");
let path_segment = parts.next();
let url_str = match path_segment {
None => format!("https://{host_encoded}/.well-known/did.json"),
Some(rest) => {
let path = rest.replace(':', "/");
format!("https://{host_encoded}/{path}/did.json")
}
};
url::Url::parse(&url_str).map_err(|e| ResolveError::Malformed(format!("bad url: {e}")))
}
async fn resolve_plc(&self, did: &str) -> Result<DidDocument, ResolveError> {
let url = format!("{}/{}", self.plc_directory_url.trim_end_matches('/'), did);
self.fetch_did_doc(&url).await
}
async fn resolve_web(&self, did: &str) -> Result<DidDocument, ResolveError> {
let url = Self::did_web_url(did)?;
self.fetch_did_doc(url.as_str()).await
}
async fn fetch_did_doc(&self, url: &str) -> Result<DidDocument, ResolveError> {
let resp = self.client.get(url).send().await.map_err(|e| {
if e.to_string().contains("SSRF") {
ResolveError::SsrfBlocked(url.to_string())
} else {
ResolveError::Network(e.to_string())
}
})?;
if !resp.status().is_success() {
return Err(ResolveError::BadStatus(resp.status().as_u16()));
}
resp.json::<DidDocument>()
.await
.map_err(|e| ResolveError::Parse(e.to_string()))
}
}
#[async_trait]
impl DidResolver for HttpDidResolver {
async fn resolve(&self, did: &str) -> Result<DidDocument, ResolveError> {
if did.starts_with("did:plc:") {
self.resolve_plc(did).await
} else if did.starts_with("did:web:") {
self.resolve_web(did).await
} else {
let method = did
.strip_prefix("did:")
.and_then(|s| s.split(':').next())
.unwrap_or(did);
Err(ResolveError::UnsupportedMethod(method.to_string()))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn did_web_no_path_component_uses_well_known() {
let url = HttpDidResolver::did_web_url("did:web:example.com").expect("ok");
assert_eq!(url.as_str(), "https://example.com/.well-known/did.json");
}
#[test]
fn did_web_path_components_drops_well_known() {
let url = HttpDidResolver::did_web_url("did:web:example.com:users:alice").expect("ok");
assert_eq!(url.as_str(), "https://example.com/users/alice/did.json");
}
#[test]
fn did_web_single_path_component() {
let url = HttpDidResolver::did_web_url("did:web:example.com:user").expect("ok");
assert_eq!(url.as_str(), "https://example.com/user/did.json");
}
#[test]
fn did_web_malformed_without_method_prefix() {
let err = HttpDidResolver::did_web_url("did:plc:xyz").expect_err("must fail");
assert!(matches!(err, ResolveError::Malformed(_)));
}
#[test]
fn did_doc_finds_method_by_fragment() {
let doc = DidDocument {
id: "did:plc:abc".into(),
verification_method: vec![
VerificationMethod {
id: "did:plc:abc#atproto_label".into(),
r#type: "Multikey".into(),
public_key_multibase: "zOne".into(),
},
VerificationMethod {
id: "did:plc:abc#atproto".into(),
r#type: "Multikey".into(),
public_key_multibase: "zTwo".into(),
},
],
};
let vm = doc.find_verification_method("#atproto").expect("found");
assert_eq!(vm.public_key_multibase, "zTwo");
}
#[test]
fn did_doc_missing_method_returns_none() {
let doc = DidDocument {
id: "did:plc:abc".into(),
verification_method: vec![],
};
assert!(doc.find_verification_method("#atproto").is_none());
}
}