cairn-mod 1.7.0

Lightweight, Rust-native ATProto labeler
Documentation
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
//! Client for the four PDS endpoints the CLI depends on (§5.3).
//!
//! - `com.atproto.server.createSession` — exchange handle+password
//!   for `accessJwt`/`refreshJwt` during `cairn login`.
//! - `com.atproto.server.refreshSession` — trade `refreshJwt` for
//!   new tokens when `accessJwt` is rejected (401 on getServiceAuth).
//! - `com.atproto.server.deleteSession` — invalidate the refresh
//!   token on `cairn logout`; §5.3 requires the PDS-side revocation
//!   in addition to local session-file cleanup.
//! - `com.atproto.server.getServiceAuth` — mint a short-lived
//!   service auth JWT for a given `aud`+`lxm`. Called fresh for
//!   every authed CLI command per §5.3 (no client-side caching).
//!
//! Uses a vanilla `reqwest::Client` — SSRF filtering (#11) applies
//! server-side to attacker-influenced URLs, but CLI URLs are
//! user-supplied (Q2 in the criteria confirmation). The module stays
//! decoupled from the server's DNS resolver.

use std::time::Duration;

use reqwest::{Client, StatusCode};
use serde::{Deserialize, Serialize};
use thiserror::Error;
use url::Url;

/// Default connect + request timeout. Login/report are interactive
/// operations; 30s is generous against any reasonable PDS latency
/// while still bounding hangs.
const DEFAULT_TIMEOUT: Duration = Duration::from_secs(30);

/// Successful `createSession` response — the fields the CLI
/// consumes. Additional PDS-returned fields (`email`, `didDoc`,
/// `active`, ...) are ignored by serde `deny_unknown_fields` being
/// absent, and intentionally not surfaced into the session file.
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CreateSessionResponse {
    /// Short-lived access JWT. Used for authed PDS requests until
    /// it expires, then rotated via `refreshSession`.
    pub access_jwt: String,
    /// Long-lived refresh JWT. Used to mint new access tokens.
    pub refresh_jwt: String,
    /// Authoritative DID the PDS authenticated this session as.
    pub did: String,
    /// Handle associated with the authenticated DID.
    pub handle: String,
}

/// `refreshSession` response. Structurally mirrors `createSession`;
/// the CLI only persists the two rotated tokens.
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct RefreshSessionResponse {
    /// New short-lived access JWT.
    pub access_jwt: String,
    /// New long-lived refresh JWT (rotates on each refresh).
    pub refresh_jwt: String,
}

/// `getServiceAuth` response. `token` is an ES256K JWT the PDS
/// signs with the moderator's `#atproto` key — the CLI forwards
/// this verbatim as the `Authorization: Bearer <...>` header on the
/// Cairn request.
#[derive(Debug, Clone, Deserialize)]
struct GetServiceAuthResponse {
    token: String,
}

/// `putRecord` 200 response body.
#[derive(Debug, Clone, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct PutRecordResponse {
    /// AT-URI of the record that was written
    /// (`at://<did>/<collection>/<rkey>`).
    pub uri: String,
    /// Content-addressed ID of the written record. Used as the
    /// `swap_record` guard on subsequent conditional writes.
    pub cid: String,
}

/// `com.atproto.repo.getRecord` response shape. Unauthenticated
/// public endpoint — used by `cairn serve`'s startup verify check
/// (#8) to fetch the published `app.bsky.labeler.service` record
/// for comparison against the local config.
#[derive(Debug, Clone, Deserialize)]
pub struct GetRecordResponse {
    /// AT-URI of the record (`at://<did>/<collection>/<rkey>`).
    pub uri: String,
    /// Content-addressed ID of the record at the time of fetch.
    /// Optional in the lexicon; some PDS implementations omit it
    /// for legacy records.
    #[serde(default)]
    pub cid: Option<String>,
    /// The record body itself, opaque on the PdsClient side. The
    /// caller deserializes into the appropriate per-collection
    /// shape (e.g. `crate::service_record::ServiceRecord`).
    pub value: serde_json::Value,
}

/// Wire shape of an XRPC error body (`{error, message}`), used to
/// surface meaningful CLI error output without echoing the whole
/// PDS response.
#[derive(Debug, Clone, Deserialize, Default)]
struct XrpcErrorBody {
    #[serde(default)]
    error: String,
    #[serde(default)]
    message: String,
}

/// `createSession` request body.
#[derive(Debug, Serialize)]
struct CreateSessionRequest<'a> {
    identifier: &'a str,
    password: &'a str,
}

/// Taxonomy for PDS interaction failures. Carries enough context
/// for the CLI dispatcher to map to an exit code + human-readable
/// message without leaking PDS internals.
#[derive(Debug, Error)]
pub enum PdsError {
    /// `Url::parse` failed on the configured PDS base.
    #[error("invalid PDS URL {url}: {source}")]
    InvalidUrl {
        /// URL string that failed to parse.
        url: String,
        /// Underlying parse error.
        #[source]
        source: url::ParseError,
    },
    /// Transport-level failure (DNS, TLS, connection, timeout).
    #[error("network error contacting {url}: {source}")]
    Network {
        /// URL the request was sent to.
        url: String,
        /// Underlying reqwest error.
        #[source]
        source: reqwest::Error,
    },
    /// 401 from any endpoint. `context` identifies the lexicon
    /// method (`createSession` / `refreshSession` / ...) so the
    /// caller can distinguish "bad app password" from "refresh
    /// token expired" without parsing error strings.
    #[error("PDS rejected credentials on {context}: {error} — {message}")]
    Unauthorized {
        /// Short lexicon method name that returned 401.
        context: &'static str,
        /// `error` field from the XRPC response body.
        error: String,
        /// `message` field from the XRPC response body.
        message: String,
    },
    /// Any non-2xx, non-401 status.
    #[error("PDS {context} failed with status {status}: {error} — {message}")]
    UnexpectedStatus {
        /// Short lexicon method name that failed.
        context: &'static str,
        /// HTTP status code returned.
        status: u16,
        /// `error` field from the XRPC response body.
        error: String,
        /// `message` field from the XRPC response body.
        message: String,
    },
    /// Response deserialization failed (PDS sent 2xx but the body
    /// didn't match the expected shape).
    #[error("PDS {context} returned malformed JSON: {source}")]
    MalformedResponse {
        /// Short lexicon method name whose response didn't parse.
        context: &'static str,
        /// Underlying reqwest/serde error.
        #[source]
        source: reqwest::Error,
    },
    /// §F1 swap-race. Distinct from `UnexpectedStatus` so
    /// `publish-service-record` can exit with a specific message
    /// directing the operator to inspect + reconcile manually
    /// before re-running.
    #[error(
        "another process has modified the service record on the PDS since Cairn's last publish: {message}"
    )]
    SwapRace {
        /// Message body from the PDS's `InvalidSwap` response.
        message: String,
    },
}

/// Thin wrapper over `reqwest::Client` pinned to one PDS base URL.
/// Construct once per `cairn` process; one instance serves the
/// whole command's lifetime.
#[derive(Debug, Clone)]
pub struct PdsClient {
    client: Client,
    base: Url,
}

impl PdsClient {
    /// Construct a client for the given PDS base URL
    /// (e.g., `https://bsky.social`). Trailing slashes are
    /// normalized away.
    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 })
    }

    /// Inject a preconfigured `reqwest::Client`. Used by tests that
    /// need to disable TLS verification (mock PDS at plain HTTP)
    /// without plumbing feature flags through the public API.
    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 {
        // base.join("xrpc/<lxm>") handles missing trailing slash
        // correctly since we push a relative segment; the Url crate
        // does the right concatenation.
        let mut u = self.base.clone();
        // Ensure the path ends with `/` before joining a relative
        // path, else `join("xrpc/...")` replaces the final segment.
        if !u.path().ends_with('/') {
            u.set_path(&format!("{}/", u.path()));
        }
        u.join(&format!("xrpc/{lxm}"))
            .expect("xrpc/{lxm} always joins")
    }

    /// Exchange handle+app-password for a PDS session.
    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
    }

    /// Rotate access+refresh tokens using a valid refresh token.
    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
    }

    /// Invalidate the refresh token server-side. Success on 2xx;
    /// callers typically treat any error here as non-fatal (local
    /// cleanup proceeds regardless, per Q3 in the criteria
    /// confirmation).
    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)
        }
    }

    /// Put a record at (repo, collection, rkey). When `swap_record`
    /// is `Some`, the request is conditional on the PDS's current
    /// record having that CID — §F1 swap-race detection rides on
    /// this. Used by `cairn publish-service-record` to emit the
    /// `app.bsky.labeler.service` record at rkey=self.
    ///
    /// Distinct `context` discriminators for auth failures:
    /// `"putRecord"` generally, mapped upward to a specific
    /// swap-race error via [`PdsError::SwapRace`] when the PDS's
    /// response body carries the `InvalidSwap` shape.
    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 {
            // Surface InvalidSwap as its own variant so callers can
            // branch on §F1 swap-race detection without parsing
            // error strings. Everything else falls through to the
            // generic classifier.
            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,
                }
            })
        }
    }

    /// Fetch a record via `com.atproto.repo.getRecord`. Unauthenticated
    /// public endpoint — used by `cairn serve`'s startup verify check
    /// to read the published `app.bsky.labeler.service` record
    /// without requiring an operator session on the serve host.
    ///
    /// Returns:
    /// - `Ok(Some(_))` — record exists, body deserialized
    /// - `Ok(None)` — record does not exist (HTTP 404 OR XRPC
    ///   `RecordNotFound` error body). Distinct from a transport
    ///   failure: callers that need to differentiate "absent"
    ///   from "unreachable" branch on this distinction.
    /// - `Err(PdsError::Network { .. })` — transport-level
    ///   failure (DNS, TLS, timeout, refused).
    /// - `Err(PdsError::UnexpectedStatus { .. })` — any other
    ///   non-2xx (auth-required-on-private-PDS, server error,
    ///   rate-limit response).
    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,
                });
        }

        // Distinguish "record not found" from any other failure.
        // PDS implementations return either HTTP 400 with
        // `error: "RecordNotFound"` in the XRPC body, or HTTP 404
        // — both should map to Ok(None) so the caller can branch
        // cleanly on absent-vs-unreachable.
        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,
        })
    }

    /// Delete a record at (repo, collection, rkey) via
    /// `com.atproto.repo.deleteRecord`. When `swap_record` is `Some`,
    /// the request is conditional on the PDS's current record having
    /// that CID — same swap-race semantics as `put_record`. Used by
    /// `cairn unpublish-service-record` (#34) to remove the published
    /// `app.bsky.labeler.service` record.
    ///
    /// **Idempotency on the wire:** real PDSes return 200 even when
    /// the target record is already absent, matching the ATProto
    /// spec. Callers that want a "did we actually delete something"
    /// distinction should consult their own state (e.g.,
    /// `labeler_config`) before calling — that's the path the
    /// unpublish flow takes.
    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() {
            // PDS may return `{ "commit": ... }` or an empty body;
            // we don't look at it — a 2xx is the contract.
            return Ok(());
        }
        // Surface InvalidSwap as its own variant — same posture as
        // put_record so concurrent operator workflows can distinguish
        // a swap-race from an auth/transport failure.
        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,
            }
        })
    }

    /// Mint a fresh service auth JWT for calling `aud` with the
    /// given lexicon method. Returns the opaque token string; the
    /// CLI presents it as `Authorization: Bearer <token>` to Cairn.
    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)
    }
}

/// Deserialize a 2xx JSON body into `T`, or classify the response
/// as an error.
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)
    }
}

/// Inspect a non-2xx response and produce the right `PdsError`
/// variant. 401 is called out specifically so callers can branch on
/// "credentials rejected" vs. generic failure.
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() {
        // A PDS on a non-root path (e.g., reverse-proxied).
        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 { .. }));
    }
}