use anyhow::Result;
use atproto_client::{
client::Auth,
com::atproto::repo::{GetRecordResponse, get_record},
};
use atproto_identity::resolve::{DnsResolver, resolve_subject};
use serde_json::Value;
use tracing::instrument;
use crate::{errors::LexiconResolveError, validation};
#[async_trait::async_trait]
pub trait LexiconResolver: Send + Sync {
async fn resolve(&self, nsid: &str) -> Result<Value>;
}
#[derive(Clone)]
pub struct DefaultLexiconResolver<R> {
http_client: reqwest::Client,
dns_resolver: R,
}
impl<R> DefaultLexiconResolver<R> {
pub fn new(http_client: reqwest::Client, dns_resolver: R) -> Self {
Self {
http_client,
dns_resolver,
}
}
}
#[async_trait::async_trait]
impl<R> LexiconResolver for DefaultLexiconResolver<R>
where
R: DnsResolver + Send + Sync,
{
#[instrument(skip(self), err)]
async fn resolve(&self, nsid: &str) -> Result<Value> {
let dns_name = validation::nsid_to_dns_name(nsid)?;
let did = resolve_lexicon_dns(&self.dns_resolver, &dns_name).await?;
let resolved_did = resolve_subject(&self.http_client, &self.dns_resolver, &did).await?;
let pds_endpoint = get_pds_from_did(&self.http_client, &resolved_did).await?;
let lexicon =
fetch_lexicon_from_pds(&self.http_client, &pds_endpoint, &resolved_did, nsid).await?;
Ok(lexicon)
}
}
#[instrument(skip(dns_resolver), err)]
pub async fn resolve_lexicon_dns<R: DnsResolver + ?Sized>(
dns_resolver: &R,
lookup_dns: &str,
) -> Result<String, LexiconResolveError> {
let txt_records = dns_resolver.resolve_txt(lookup_dns).await?;
let dids: Vec<String> = txt_records
.iter()
.filter_map(|record| {
record
.strip_prefix("did=")
.or_else(|| record.strip_prefix("did:"))
.map(|did| {
if did.starts_with("plc:") || did.starts_with("web:") {
format!("did:{}", did)
} else if did.starts_with("did:") {
did.to_string()
} else {
format!("did:{}", did)
}
})
})
.collect();
if dids.is_empty() {
return Err(LexiconResolveError::NoDIDsFound);
}
if dids.len() > 1 {
return Err(LexiconResolveError::MultipleDIDsFound);
}
Ok(dids[0].clone())
}
#[instrument(skip(http_client), err)]
pub async fn get_pds_from_did(http_client: &reqwest::Client, did: &str) -> Result<String> {
use atproto_identity::{
model::Document,
plc,
resolve::{InputType, parse_input},
web,
};
let did_document: Document = match parse_input(did)? {
InputType::Plc(did) => plc::query(http_client, "plc.directory", &did).await?,
InputType::Web(did) => web::query(http_client, &did).await?,
_ => {
return Err(LexiconResolveError::InvalidDIDFormat {
did: did.to_string(),
}
.into());
}
};
for service in &did_document.service {
if service.r#type == "AtprotoPersonalDataServer" {
return Ok(service.service_endpoint.clone());
}
}
Err(LexiconResolveError::NoPDSEndpoint.into())
}
#[instrument(skip(http_client), err)]
pub async fn fetch_lexicon_from_pds(
http_client: &reqwest::Client,
pds_endpoint: &str,
did: &str,
nsid: &str,
) -> Result<Value> {
let collection = "com.atproto.lexicon.schema";
let auth = Auth::None;
let response = get_record(
http_client,
&auth,
pds_endpoint,
did,
collection,
nsid,
None,
)
.await
.map_err(|e| LexiconResolveError::PDSFetchFailed {
details: e.to_string(),
})?;
match response {
GetRecordResponse::Record { value, .. } => Ok(value),
GetRecordResponse::Error(err) => {
let msg = err
.message
.or(err.error_description)
.or(err.error)
.unwrap_or_else(|| "Unknown error".to_string());
Err(LexiconResolveError::PDSErrorResponse {
nsid: nsid.to_string(),
message: msg,
}
.into())
}
}
}
#[instrument(skip(http_client, dns_resolver), err)]
pub async fn get_lexicon<R: DnsResolver + ?Sized>(
http_client: &reqwest::Client,
dns_resolver: &R,
nsid: &str,
repo: Option<&str>,
) -> Result<Value> {
if !validation::is_valid_nsid(nsid) {
return Err(LexiconResolveError::InvalidNsid {
nsid: nsid.to_string(),
}
.into());
}
let resolved_did = match repo {
Some(subject) => resolve_subject(http_client, dns_resolver, subject).await?,
None => {
let dns_name = validation::nsid_to_dns_name(nsid)?;
let did = resolve_lexicon_dns(dns_resolver, &dns_name).await?;
resolve_subject(http_client, dns_resolver, &did).await?
}
};
let pds_endpoint = get_pds_from_did(http_client, &resolved_did).await?;
fetch_lexicon_from_pds(http_client, &pds_endpoint, &resolved_did, nsid).await
}