Skip to main content

cairn_mod/cli/
pds.rs

1//! Client for the four PDS endpoints the CLI depends on (§5.3).
2//!
3//! - `com.atproto.server.createSession` — exchange handle+password
4//!   for `accessJwt`/`refreshJwt` during `cairn login`.
5//! - `com.atproto.server.refreshSession` — trade `refreshJwt` for
6//!   new tokens when `accessJwt` is rejected (401 on getServiceAuth).
7//! - `com.atproto.server.deleteSession` — invalidate the refresh
8//!   token on `cairn logout`; §5.3 requires the PDS-side revocation
9//!   in addition to local session-file cleanup.
10//! - `com.atproto.server.getServiceAuth` — mint a short-lived
11//!   service auth JWT for a given `aud`+`lxm`. Called fresh for
12//!   every authed CLI command per §5.3 (no client-side caching).
13//!
14//! Uses a vanilla `reqwest::Client` — SSRF filtering (#11) applies
15//! server-side to attacker-influenced URLs, but CLI URLs are
16//! user-supplied (Q2 in the criteria confirmation). The module stays
17//! decoupled from the server's DNS resolver.
18
19use std::time::Duration;
20
21use reqwest::{Client, StatusCode};
22use serde::{Deserialize, Serialize};
23use thiserror::Error;
24use url::Url;
25
26/// Default connect + request timeout. Login/report are interactive
27/// operations; 30s is generous against any reasonable PDS latency
28/// while still bounding hangs.
29const DEFAULT_TIMEOUT: Duration = Duration::from_secs(30);
30
31/// Successful `createSession` response — the fields the CLI
32/// consumes. Additional PDS-returned fields (`email`, `didDoc`,
33/// `active`, ...) are ignored by serde `deny_unknown_fields` being
34/// absent, and intentionally not surfaced into the session file.
35#[derive(Debug, Clone, Deserialize)]
36#[serde(rename_all = "camelCase")]
37pub struct CreateSessionResponse {
38    /// Short-lived access JWT. Used for authed PDS requests until
39    /// it expires, then rotated via `refreshSession`.
40    pub access_jwt: String,
41    /// Long-lived refresh JWT. Used to mint new access tokens.
42    pub refresh_jwt: String,
43    /// Authoritative DID the PDS authenticated this session as.
44    pub did: String,
45    /// Handle associated with the authenticated DID.
46    pub handle: String,
47}
48
49/// `refreshSession` response. Structurally mirrors `createSession`;
50/// the CLI only persists the two rotated tokens.
51#[derive(Debug, Clone, Deserialize)]
52#[serde(rename_all = "camelCase")]
53pub struct RefreshSessionResponse {
54    /// New short-lived access JWT.
55    pub access_jwt: String,
56    /// New long-lived refresh JWT (rotates on each refresh).
57    pub refresh_jwt: String,
58}
59
60/// `getServiceAuth` response. `token` is an ES256K JWT the PDS
61/// signs with the moderator's `#atproto` key — the CLI forwards
62/// this verbatim as the `Authorization: Bearer <...>` header on the
63/// Cairn request.
64#[derive(Debug, Clone, Deserialize)]
65struct GetServiceAuthResponse {
66    token: String,
67}
68
69/// `putRecord` 200 response body.
70#[derive(Debug, Clone, Deserialize)]
71#[serde(rename_all = "camelCase")]
72pub struct PutRecordResponse {
73    /// AT-URI of the record that was written
74    /// (`at://<did>/<collection>/<rkey>`).
75    pub uri: String,
76    /// Content-addressed ID of the written record. Used as the
77    /// `swap_record` guard on subsequent conditional writes.
78    pub cid: String,
79}
80
81/// `com.atproto.repo.getRecord` response shape. Unauthenticated
82/// public endpoint — used by `cairn serve`'s startup verify check
83/// (#8) to fetch the published `app.bsky.labeler.service` record
84/// for comparison against the local config.
85#[derive(Debug, Clone, Deserialize)]
86pub struct GetRecordResponse {
87    /// AT-URI of the record (`at://<did>/<collection>/<rkey>`).
88    pub uri: String,
89    /// Content-addressed ID of the record at the time of fetch.
90    /// Optional in the lexicon; some PDS implementations omit it
91    /// for legacy records.
92    #[serde(default)]
93    pub cid: Option<String>,
94    /// The record body itself, opaque on the PdsClient side. The
95    /// caller deserializes into the appropriate per-collection
96    /// shape (e.g. `crate::service_record::ServiceRecord`).
97    pub value: serde_json::Value,
98}
99
100/// Wire shape of an XRPC error body (`{error, message}`), used to
101/// surface meaningful CLI error output without echoing the whole
102/// PDS response.
103#[derive(Debug, Clone, Deserialize, Default)]
104struct XrpcErrorBody {
105    #[serde(default)]
106    error: String,
107    #[serde(default)]
108    message: String,
109}
110
111/// `createSession` request body.
112#[derive(Debug, Serialize)]
113struct CreateSessionRequest<'a> {
114    identifier: &'a str,
115    password: &'a str,
116}
117
118/// Taxonomy for PDS interaction failures. Carries enough context
119/// for the CLI dispatcher to map to an exit code + human-readable
120/// message without leaking PDS internals.
121#[derive(Debug, Error)]
122pub enum PdsError {
123    /// `Url::parse` failed on the configured PDS base.
124    #[error("invalid PDS URL {url}: {source}")]
125    InvalidUrl {
126        /// URL string that failed to parse.
127        url: String,
128        /// Underlying parse error.
129        #[source]
130        source: url::ParseError,
131    },
132    /// Transport-level failure (DNS, TLS, connection, timeout).
133    #[error("network error contacting {url}: {source}")]
134    Network {
135        /// URL the request was sent to.
136        url: String,
137        /// Underlying reqwest error.
138        #[source]
139        source: reqwest::Error,
140    },
141    /// 401 from any endpoint. `context` identifies the lexicon
142    /// method (`createSession` / `refreshSession` / ...) so the
143    /// caller can distinguish "bad app password" from "refresh
144    /// token expired" without parsing error strings.
145    #[error("PDS rejected credentials on {context}: {error} — {message}")]
146    Unauthorized {
147        /// Short lexicon method name that returned 401.
148        context: &'static str,
149        /// `error` field from the XRPC response body.
150        error: String,
151        /// `message` field from the XRPC response body.
152        message: String,
153    },
154    /// Any non-2xx, non-401 status.
155    #[error("PDS {context} failed with status {status}: {error} — {message}")]
156    UnexpectedStatus {
157        /// Short lexicon method name that failed.
158        context: &'static str,
159        /// HTTP status code returned.
160        status: u16,
161        /// `error` field from the XRPC response body.
162        error: String,
163        /// `message` field from the XRPC response body.
164        message: String,
165    },
166    /// Response deserialization failed (PDS sent 2xx but the body
167    /// didn't match the expected shape).
168    #[error("PDS {context} returned malformed JSON: {source}")]
169    MalformedResponse {
170        /// Short lexicon method name whose response didn't parse.
171        context: &'static str,
172        /// Underlying reqwest/serde error.
173        #[source]
174        source: reqwest::Error,
175    },
176    /// §F1 swap-race. Distinct from `UnexpectedStatus` so
177    /// `publish-service-record` can exit with a specific message
178    /// directing the operator to inspect + reconcile manually
179    /// before re-running.
180    #[error(
181        "another process has modified the service record on the PDS since Cairn's last publish: {message}"
182    )]
183    SwapRace {
184        /// Message body from the PDS's `InvalidSwap` response.
185        message: String,
186    },
187}
188
189/// Thin wrapper over `reqwest::Client` pinned to one PDS base URL.
190/// Construct once per `cairn` process; one instance serves the
191/// whole command's lifetime.
192#[derive(Debug, Clone)]
193pub struct PdsClient {
194    client: Client,
195    base: Url,
196}
197
198impl PdsClient {
199    /// Construct a client for the given PDS base URL
200    /// (e.g., `https://bsky.social`). Trailing slashes are
201    /// normalized away.
202    pub fn new(base_url: &str) -> Result<Self, PdsError> {
203        let base = Url::parse(base_url).map_err(|source| PdsError::InvalidUrl {
204            url: base_url.to_string(),
205            source,
206        })?;
207        let client = Client::builder()
208            .timeout(DEFAULT_TIMEOUT)
209            .build()
210            .expect("reqwest client build with default tls config should not fail");
211        Ok(Self { client, base })
212    }
213
214    /// Inject a preconfigured `reqwest::Client`. Used by tests that
215    /// need to disable TLS verification (mock PDS at plain HTTP)
216    /// without plumbing feature flags through the public API.
217    pub fn with_http_client(base_url: &str, client: Client) -> Result<Self, PdsError> {
218        let base = Url::parse(base_url).map_err(|source| PdsError::InvalidUrl {
219            url: base_url.to_string(),
220            source,
221        })?;
222        Ok(Self { client, base })
223    }
224
225    fn endpoint(&self, lxm: &str) -> Url {
226        // base.join("xrpc/<lxm>") handles missing trailing slash
227        // correctly since we push a relative segment; the Url crate
228        // does the right concatenation.
229        let mut u = self.base.clone();
230        // Ensure the path ends with `/` before joining a relative
231        // path, else `join("xrpc/...")` replaces the final segment.
232        if !u.path().ends_with('/') {
233            u.set_path(&format!("{}/", u.path()));
234        }
235        u.join(&format!("xrpc/{lxm}"))
236            .expect("xrpc/{lxm} always joins")
237    }
238
239    /// Exchange handle+app-password for a PDS session.
240    pub async fn create_session(
241        &self,
242        identifier: &str,
243        password: &str,
244    ) -> Result<CreateSessionResponse, PdsError> {
245        const CTX: &str = "createSession";
246        let url = self.endpoint("com.atproto.server.createSession");
247        let resp = self
248            .client
249            .post(url.clone())
250            .json(&CreateSessionRequest {
251                identifier,
252                password,
253            })
254            .send()
255            .await
256            .map_err(|source| PdsError::Network {
257                url: url.to_string(),
258                source,
259            })?;
260        deserialize_or_xrpc_error(CTX, resp).await
261    }
262
263    /// Rotate access+refresh tokens using a valid refresh token.
264    pub async fn refresh_session(
265        &self,
266        refresh_jwt: &str,
267    ) -> Result<RefreshSessionResponse, PdsError> {
268        const CTX: &str = "refreshSession";
269        let url = self.endpoint("com.atproto.server.refreshSession");
270        let resp = self
271            .client
272            .post(url.clone())
273            .bearer_auth(refresh_jwt)
274            .send()
275            .await
276            .map_err(|source| PdsError::Network {
277                url: url.to_string(),
278                source,
279            })?;
280        deserialize_or_xrpc_error(CTX, resp).await
281    }
282
283    /// Invalidate the refresh token server-side. Success on 2xx;
284    /// callers typically treat any error here as non-fatal (local
285    /// cleanup proceeds regardless, per Q3 in the criteria
286    /// confirmation).
287    pub async fn delete_session(&self, refresh_jwt: &str) -> Result<(), PdsError> {
288        const CTX: &str = "deleteSession";
289        let url = self.endpoint("com.atproto.server.deleteSession");
290        let resp = self
291            .client
292            .post(url.clone())
293            .bearer_auth(refresh_jwt)
294            .send()
295            .await
296            .map_err(|source| PdsError::Network {
297                url: url.to_string(),
298                source,
299            })?;
300        if resp.status().is_success() {
301            Ok(())
302        } else {
303            Err(classify_error(CTX, resp).await)
304        }
305    }
306
307    /// Put a record at (repo, collection, rkey). When `swap_record`
308    /// is `Some`, the request is conditional on the PDS's current
309    /// record having that CID — §F1 swap-race detection rides on
310    /// this. Used by `cairn publish-service-record` to emit the
311    /// `app.bsky.labeler.service` record at rkey=self.
312    ///
313    /// Distinct `context` discriminators for auth failures:
314    /// `"putRecord"` generally, mapped upward to a specific
315    /// swap-race error via [`PdsError::SwapRace`] when the PDS's
316    /// response body carries the `InvalidSwap` shape.
317    pub async fn put_record(
318        &self,
319        access_jwt: &str,
320        repo: &str,
321        collection: &str,
322        rkey: &str,
323        record: &serde_json::Value,
324        swap_record: Option<&str>,
325    ) -> Result<PutRecordResponse, PdsError> {
326        const CTX: &str = "putRecord";
327        let url = self.endpoint("com.atproto.repo.putRecord");
328
329        let mut body = serde_json::json!({
330            "repo": repo,
331            "collection": collection,
332            "rkey": rkey,
333            "record": record,
334        });
335        if let Some(cid) = swap_record {
336            body.as_object_mut()
337                .expect("json object")
338                .insert("swapRecord".into(), serde_json::Value::String(cid.into()));
339        }
340
341        let resp = self
342            .client
343            .post(url.clone())
344            .bearer_auth(access_jwt)
345            .json(&body)
346            .send()
347            .await
348            .map_err(|source| PdsError::Network {
349                url: url.to_string(),
350                source,
351            })?;
352        if resp.status().is_success() {
353            resp.json::<PutRecordResponse>()
354                .await
355                .map_err(|source| PdsError::MalformedResponse {
356                    context: CTX,
357                    source,
358                })
359        } else {
360            // Surface InvalidSwap as its own variant so callers can
361            // branch on §F1 swap-race detection without parsing
362            // error strings. Everything else falls through to the
363            // generic classifier.
364            let status = resp.status();
365            let body = resp.json::<XrpcErrorBody>().await.unwrap_or_default();
366            if body.error == "InvalidSwap" {
367                return Err(PdsError::SwapRace {
368                    message: body.message,
369                });
370            }
371            Err(if status == reqwest::StatusCode::UNAUTHORIZED {
372                PdsError::Unauthorized {
373                    context: CTX,
374                    error: body.error,
375                    message: body.message,
376                }
377            } else {
378                PdsError::UnexpectedStatus {
379                    context: CTX,
380                    status: status.as_u16(),
381                    error: body.error,
382                    message: body.message,
383                }
384            })
385        }
386    }
387
388    /// Fetch a record via `com.atproto.repo.getRecord`. Unauthenticated
389    /// public endpoint — used by `cairn serve`'s startup verify check
390    /// to read the published `app.bsky.labeler.service` record
391    /// without requiring an operator session on the serve host.
392    ///
393    /// Returns:
394    /// - `Ok(Some(_))` — record exists, body deserialized
395    /// - `Ok(None)` — record does not exist (HTTP 404 OR XRPC
396    ///   `RecordNotFound` error body). Distinct from a transport
397    ///   failure: callers that need to differentiate "absent"
398    ///   from "unreachable" branch on this distinction.
399    /// - `Err(PdsError::Network { .. })` — transport-level
400    ///   failure (DNS, TLS, timeout, refused).
401    /// - `Err(PdsError::UnexpectedStatus { .. })` — any other
402    ///   non-2xx (auth-required-on-private-PDS, server error,
403    ///   rate-limit response).
404    pub async fn get_record(
405        &self,
406        repo: &str,
407        collection: &str,
408        rkey: &str,
409    ) -> Result<Option<GetRecordResponse>, PdsError> {
410        const CTX: &str = "getRecord";
411        let url = self.endpoint("com.atproto.repo.getRecord");
412
413        let resp = self
414            .client
415            .get(url.clone())
416            .query(&[("repo", repo), ("collection", collection), ("rkey", rkey)])
417            .send()
418            .await
419            .map_err(|source| PdsError::Network {
420                url: url.to_string(),
421                source,
422            })?;
423
424        if resp.status().is_success() {
425            return resp
426                .json::<GetRecordResponse>()
427                .await
428                .map(Some)
429                .map_err(|source| PdsError::MalformedResponse {
430                    context: CTX,
431                    source,
432                });
433        }
434
435        // Distinguish "record not found" from any other failure.
436        // PDS implementations return either HTTP 400 with
437        // `error: "RecordNotFound"` in the XRPC body, or HTTP 404
438        // — both should map to Ok(None) so the caller can branch
439        // cleanly on absent-vs-unreachable.
440        let status = resp.status();
441        let body = resp.json::<XrpcErrorBody>().await.unwrap_or_default();
442        if status == reqwest::StatusCode::NOT_FOUND || body.error == "RecordNotFound" {
443            return Ok(None);
444        }
445        Err(PdsError::UnexpectedStatus {
446            context: CTX,
447            status: status.as_u16(),
448            error: body.error,
449            message: body.message,
450        })
451    }
452
453    /// Delete a record at (repo, collection, rkey) via
454    /// `com.atproto.repo.deleteRecord`. When `swap_record` is `Some`,
455    /// the request is conditional on the PDS's current record having
456    /// that CID — same swap-race semantics as `put_record`. Used by
457    /// `cairn unpublish-service-record` (#34) to remove the published
458    /// `app.bsky.labeler.service` record.
459    ///
460    /// **Idempotency on the wire:** real PDSes return 200 even when
461    /// the target record is already absent, matching the ATProto
462    /// spec. Callers that want a "did we actually delete something"
463    /// distinction should consult their own state (e.g.,
464    /// `labeler_config`) before calling — that's the path the
465    /// unpublish flow takes.
466    pub async fn delete_record(
467        &self,
468        access_jwt: &str,
469        repo: &str,
470        collection: &str,
471        rkey: &str,
472        swap_record: Option<&str>,
473    ) -> Result<(), PdsError> {
474        const CTX: &str = "deleteRecord";
475        let url = self.endpoint("com.atproto.repo.deleteRecord");
476
477        let mut body = serde_json::json!({
478            "repo": repo,
479            "collection": collection,
480            "rkey": rkey,
481        });
482        if let Some(cid) = swap_record {
483            body.as_object_mut()
484                .expect("json object")
485                .insert("swapRecord".into(), serde_json::Value::String(cid.into()));
486        }
487
488        let resp = self
489            .client
490            .post(url.clone())
491            .bearer_auth(access_jwt)
492            .json(&body)
493            .send()
494            .await
495            .map_err(|source| PdsError::Network {
496                url: url.to_string(),
497                source,
498            })?;
499        if resp.status().is_success() {
500            // PDS may return `{ "commit": ... }` or an empty body;
501            // we don't look at it — a 2xx is the contract.
502            return Ok(());
503        }
504        // Surface InvalidSwap as its own variant — same posture as
505        // put_record so concurrent operator workflows can distinguish
506        // a swap-race from an auth/transport failure.
507        let status = resp.status();
508        let body = resp.json::<XrpcErrorBody>().await.unwrap_or_default();
509        if body.error == "InvalidSwap" {
510            return Err(PdsError::SwapRace {
511                message: body.message,
512            });
513        }
514        Err(if status == reqwest::StatusCode::UNAUTHORIZED {
515            PdsError::Unauthorized {
516                context: CTX,
517                error: body.error,
518                message: body.message,
519            }
520        } else {
521            PdsError::UnexpectedStatus {
522                context: CTX,
523                status: status.as_u16(),
524                error: body.error,
525                message: body.message,
526            }
527        })
528    }
529
530    /// Mint a fresh service auth JWT for calling `aud` with the
531    /// given lexicon method. Returns the opaque token string; the
532    /// CLI presents it as `Authorization: Bearer <token>` to Cairn.
533    pub async fn get_service_auth(
534        &self,
535        access_jwt: &str,
536        aud: &str,
537        lxm: &str,
538    ) -> Result<String, PdsError> {
539        const CTX: &str = "getServiceAuth";
540        let mut url = self.endpoint("com.atproto.server.getServiceAuth");
541        url.query_pairs_mut()
542            .append_pair("aud", aud)
543            .append_pair("lxm", lxm);
544        let resp = self
545            .client
546            .get(url.clone())
547            .bearer_auth(access_jwt)
548            .send()
549            .await
550            .map_err(|source| PdsError::Network {
551                url: url.to_string(),
552                source,
553            })?;
554        let body: GetServiceAuthResponse = deserialize_or_xrpc_error(CTX, resp).await?;
555        Ok(body.token)
556    }
557}
558
559/// Deserialize a 2xx JSON body into `T`, or classify the response
560/// as an error.
561async fn deserialize_or_xrpc_error<T: for<'de> Deserialize<'de>>(
562    context: &'static str,
563    resp: reqwest::Response,
564) -> Result<T, PdsError> {
565    if resp.status().is_success() {
566        resp.json::<T>()
567            .await
568            .map_err(|source| PdsError::MalformedResponse { context, source })
569    } else {
570        Err(classify_error(context, resp).await)
571    }
572}
573
574/// Inspect a non-2xx response and produce the right `PdsError`
575/// variant. 401 is called out specifically so callers can branch on
576/// "credentials rejected" vs. generic failure.
577async fn classify_error(context: &'static str, resp: reqwest::Response) -> PdsError {
578    let status = resp.status();
579    let body = resp.json::<XrpcErrorBody>().await.unwrap_or_default();
580    if status == StatusCode::UNAUTHORIZED {
581        PdsError::Unauthorized {
582            context,
583            error: body.error,
584            message: body.message,
585        }
586    } else {
587        PdsError::UnexpectedStatus {
588            context,
589            status: status.as_u16(),
590            error: body.error,
591            message: body.message,
592        }
593    }
594}
595
596#[cfg(test)]
597mod tests {
598    use super::*;
599
600    #[test]
601    fn endpoint_joins_correctly_without_trailing_slash() {
602        let c = PdsClient::new("https://bsky.social").unwrap();
603        let u = c.endpoint("com.atproto.server.createSession");
604        assert_eq!(
605            u.as_str(),
606            "https://bsky.social/xrpc/com.atproto.server.createSession"
607        );
608    }
609
610    #[test]
611    fn endpoint_joins_correctly_with_trailing_slash() {
612        let c = PdsClient::new("https://bsky.social/").unwrap();
613        let u = c.endpoint("com.atproto.server.createSession");
614        assert_eq!(
615            u.as_str(),
616            "https://bsky.social/xrpc/com.atproto.server.createSession"
617        );
618    }
619
620    #[test]
621    fn endpoint_preserves_base_path() {
622        // A PDS on a non-root path (e.g., reverse-proxied).
623        let c = PdsClient::new("https://example.com/pds").unwrap();
624        let u = c.endpoint("com.atproto.server.createSession");
625        assert_eq!(
626            u.as_str(),
627            "https://example.com/pds/xrpc/com.atproto.server.createSession"
628        );
629    }
630
631    #[test]
632    fn invalid_url_returns_structured_error() {
633        let err = PdsClient::new("not a url").unwrap_err();
634        assert!(matches!(err, PdsError::InvalidUrl { .. }));
635    }
636}