use std::time::Duration;
use reqwest::{Client, StatusCode};
use serde::{Deserialize, Serialize};
use thiserror::Error;
use url::Url;
const DEFAULT_TIMEOUT: Duration = Duration::from_secs(30);
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CreateSessionResponse {
pub access_jwt: String,
pub refresh_jwt: String,
pub did: String,
pub handle: String,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct RefreshSessionResponse {
pub access_jwt: String,
pub refresh_jwt: String,
}
#[derive(Debug, Clone, Deserialize)]
struct GetServiceAuthResponse {
token: String,
}
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PutRecordResponse {
pub uri: String,
pub cid: String,
}
#[derive(Debug, Clone, Deserialize)]
pub struct GetRecordResponse {
pub uri: String,
#[serde(default)]
pub cid: Option<String>,
pub value: serde_json::Value,
}
#[derive(Debug, Clone, Deserialize, Default)]
struct XrpcErrorBody {
#[serde(default)]
error: String,
#[serde(default)]
message: String,
}
#[derive(Debug, Serialize)]
struct CreateSessionRequest<'a> {
identifier: &'a str,
password: &'a str,
}
#[derive(Debug, Error)]
pub enum PdsError {
#[error("invalid PDS URL {url}: {source}")]
InvalidUrl {
url: String,
#[source]
source: url::ParseError,
},
#[error("network error contacting {url}: {source}")]
Network {
url: String,
#[source]
source: reqwest::Error,
},
#[error("PDS rejected credentials on {context}: {error} — {message}")]
Unauthorized {
context: &'static str,
error: String,
message: String,
},
#[error("PDS {context} failed with status {status}: {error} — {message}")]
UnexpectedStatus {
context: &'static str,
status: u16,
error: String,
message: String,
},
#[error("PDS {context} returned malformed JSON: {source}")]
MalformedResponse {
context: &'static str,
#[source]
source: reqwest::Error,
},
#[error(
"another process has modified the service record on the PDS since Cairn's last publish: {message}"
)]
SwapRace {
message: String,
},
}
#[derive(Debug, Clone)]
pub struct PdsClient {
client: Client,
base: Url,
}
impl PdsClient {
pub fn new(base_url: &str) -> Result<Self, PdsError> {
let base = Url::parse(base_url).map_err(|source| PdsError::InvalidUrl {
url: base_url.to_string(),
source,
})?;
let client = Client::builder()
.timeout(DEFAULT_TIMEOUT)
.build()
.expect("reqwest client build with default tls config should not fail");
Ok(Self { client, base })
}
pub fn with_http_client(base_url: &str, client: Client) -> Result<Self, PdsError> {
let base = Url::parse(base_url).map_err(|source| PdsError::InvalidUrl {
url: base_url.to_string(),
source,
})?;
Ok(Self { client, base })
}
fn endpoint(&self, lxm: &str) -> Url {
let mut u = self.base.clone();
if !u.path().ends_with('/') {
u.set_path(&format!("{}/", u.path()));
}
u.join(&format!("xrpc/{lxm}"))
.expect("xrpc/{lxm} always joins")
}
pub async fn create_session(
&self,
identifier: &str,
password: &str,
) -> Result<CreateSessionResponse, PdsError> {
const CTX: &str = "createSession";
let url = self.endpoint("com.atproto.server.createSession");
let resp = self
.client
.post(url.clone())
.json(&CreateSessionRequest {
identifier,
password,
})
.send()
.await
.map_err(|source| PdsError::Network {
url: url.to_string(),
source,
})?;
deserialize_or_xrpc_error(CTX, resp).await
}
pub async fn refresh_session(
&self,
refresh_jwt: &str,
) -> Result<RefreshSessionResponse, PdsError> {
const CTX: &str = "refreshSession";
let url = self.endpoint("com.atproto.server.refreshSession");
let resp = self
.client
.post(url.clone())
.bearer_auth(refresh_jwt)
.send()
.await
.map_err(|source| PdsError::Network {
url: url.to_string(),
source,
})?;
deserialize_or_xrpc_error(CTX, resp).await
}
pub async fn delete_session(&self, refresh_jwt: &str) -> Result<(), PdsError> {
const CTX: &str = "deleteSession";
let url = self.endpoint("com.atproto.server.deleteSession");
let resp = self
.client
.post(url.clone())
.bearer_auth(refresh_jwt)
.send()
.await
.map_err(|source| PdsError::Network {
url: url.to_string(),
source,
})?;
if resp.status().is_success() {
Ok(())
} else {
Err(classify_error(CTX, resp).await)
}
}
pub async fn put_record(
&self,
access_jwt: &str,
repo: &str,
collection: &str,
rkey: &str,
record: &serde_json::Value,
swap_record: Option<&str>,
) -> Result<PutRecordResponse, PdsError> {
const CTX: &str = "putRecord";
let url = self.endpoint("com.atproto.repo.putRecord");
let mut body = serde_json::json!({
"repo": repo,
"collection": collection,
"rkey": rkey,
"record": record,
});
if let Some(cid) = swap_record {
body.as_object_mut()
.expect("json object")
.insert("swapRecord".into(), serde_json::Value::String(cid.into()));
}
let resp = self
.client
.post(url.clone())
.bearer_auth(access_jwt)
.json(&body)
.send()
.await
.map_err(|source| PdsError::Network {
url: url.to_string(),
source,
})?;
if resp.status().is_success() {
resp.json::<PutRecordResponse>()
.await
.map_err(|source| PdsError::MalformedResponse {
context: CTX,
source,
})
} else {
let status = resp.status();
let body = resp.json::<XrpcErrorBody>().await.unwrap_or_default();
if body.error == "InvalidSwap" {
return Err(PdsError::SwapRace {
message: body.message,
});
}
Err(if status == reqwest::StatusCode::UNAUTHORIZED {
PdsError::Unauthorized {
context: CTX,
error: body.error,
message: body.message,
}
} else {
PdsError::UnexpectedStatus {
context: CTX,
status: status.as_u16(),
error: body.error,
message: body.message,
}
})
}
}
pub async fn get_record(
&self,
repo: &str,
collection: &str,
rkey: &str,
) -> Result<Option<GetRecordResponse>, PdsError> {
const CTX: &str = "getRecord";
let url = self.endpoint("com.atproto.repo.getRecord");
let resp = self
.client
.get(url.clone())
.query(&[("repo", repo), ("collection", collection), ("rkey", rkey)])
.send()
.await
.map_err(|source| PdsError::Network {
url: url.to_string(),
source,
})?;
if resp.status().is_success() {
return resp
.json::<GetRecordResponse>()
.await
.map(Some)
.map_err(|source| PdsError::MalformedResponse {
context: CTX,
source,
});
}
let status = resp.status();
let body = resp.json::<XrpcErrorBody>().await.unwrap_or_default();
if status == reqwest::StatusCode::NOT_FOUND || body.error == "RecordNotFound" {
return Ok(None);
}
Err(PdsError::UnexpectedStatus {
context: CTX,
status: status.as_u16(),
error: body.error,
message: body.message,
})
}
pub async fn delete_record(
&self,
access_jwt: &str,
repo: &str,
collection: &str,
rkey: &str,
swap_record: Option<&str>,
) -> Result<(), PdsError> {
const CTX: &str = "deleteRecord";
let url = self.endpoint("com.atproto.repo.deleteRecord");
let mut body = serde_json::json!({
"repo": repo,
"collection": collection,
"rkey": rkey,
});
if let Some(cid) = swap_record {
body.as_object_mut()
.expect("json object")
.insert("swapRecord".into(), serde_json::Value::String(cid.into()));
}
let resp = self
.client
.post(url.clone())
.bearer_auth(access_jwt)
.json(&body)
.send()
.await
.map_err(|source| PdsError::Network {
url: url.to_string(),
source,
})?;
if resp.status().is_success() {
return Ok(());
}
let status = resp.status();
let body = resp.json::<XrpcErrorBody>().await.unwrap_or_default();
if body.error == "InvalidSwap" {
return Err(PdsError::SwapRace {
message: body.message,
});
}
Err(if status == reqwest::StatusCode::UNAUTHORIZED {
PdsError::Unauthorized {
context: CTX,
error: body.error,
message: body.message,
}
} else {
PdsError::UnexpectedStatus {
context: CTX,
status: status.as_u16(),
error: body.error,
message: body.message,
}
})
}
pub async fn get_service_auth(
&self,
access_jwt: &str,
aud: &str,
lxm: &str,
) -> Result<String, PdsError> {
const CTX: &str = "getServiceAuth";
let mut url = self.endpoint("com.atproto.server.getServiceAuth");
url.query_pairs_mut()
.append_pair("aud", aud)
.append_pair("lxm", lxm);
let resp = self
.client
.get(url.clone())
.bearer_auth(access_jwt)
.send()
.await
.map_err(|source| PdsError::Network {
url: url.to_string(),
source,
})?;
let body: GetServiceAuthResponse = deserialize_or_xrpc_error(CTX, resp).await?;
Ok(body.token)
}
}
async fn deserialize_or_xrpc_error<T: for<'de> Deserialize<'de>>(
context: &'static str,
resp: reqwest::Response,
) -> Result<T, PdsError> {
if resp.status().is_success() {
resp.json::<T>()
.await
.map_err(|source| PdsError::MalformedResponse { context, source })
} else {
Err(classify_error(context, resp).await)
}
}
async fn classify_error(context: &'static str, resp: reqwest::Response) -> PdsError {
let status = resp.status();
let body = resp.json::<XrpcErrorBody>().await.unwrap_or_default();
if status == StatusCode::UNAUTHORIZED {
PdsError::Unauthorized {
context,
error: body.error,
message: body.message,
}
} else {
PdsError::UnexpectedStatus {
context,
status: status.as_u16(),
error: body.error,
message: body.message,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn endpoint_joins_correctly_without_trailing_slash() {
let c = PdsClient::new("https://bsky.social").unwrap();
let u = c.endpoint("com.atproto.server.createSession");
assert_eq!(
u.as_str(),
"https://bsky.social/xrpc/com.atproto.server.createSession"
);
}
#[test]
fn endpoint_joins_correctly_with_trailing_slash() {
let c = PdsClient::new("https://bsky.social/").unwrap();
let u = c.endpoint("com.atproto.server.createSession");
assert_eq!(
u.as_str(),
"https://bsky.social/xrpc/com.atproto.server.createSession"
);
}
#[test]
fn endpoint_preserves_base_path() {
let c = PdsClient::new("https://example.com/pds").unwrap();
let u = c.endpoint("com.atproto.server.createSession");
assert_eq!(
u.as_str(),
"https://example.com/pds/xrpc/com.atproto.server.createSession"
);
}
#[test]
fn invalid_url_returns_structured_error() {
let err = PdsClient::new("not a url").unwrap_err();
assert!(matches!(err, PdsError::InvalidUrl { .. }));
}
}