Skip to main content

pubky_common/auth/
jws.rs

1//! Shared JWS encoding/decoding utilities and token identifier types.
2//!
3//! This module provides:
4//! - JWS Compact Serialization signing for `EdDSA` (Ed25519)
5//! - Lightweight JWS payload decoding (no signature verification)
6//! - Typed identifiers for grants, tokens, nonces, and client IDs
7
8use std::fmt;
9
10use base64::{engine::general_purpose::URL_SAFE_NO_PAD, Engine};
11use serde::{Deserialize, Serialize};
12
13use crate::crypto::{random_bytes, Keypair};
14
15/// JWS header `typ` for Grant tokens.
16pub const GRANT_JWS_TYP: &str = "pubky-grant";
17
18/// JWS header `typ` for Proof-of-Possession proofs.
19pub const POP_JWS_TYP: &str = "pubky-pop";
20
21/// Maximum length for a [`RandomId`] generated from 128-bit random bytes as base64url.
22const RANDOM_ID_MAX_LENGTH: usize = 22;
23
24/// Maximum length for a [`ClientId`] (DNS domain name limit per RFC 1035).
25const CLIENT_ID_MAX_LENGTH: usize = 253;
26
27// ── JWS Encoding ────────────────────────────────────────────────────────────
28
29/// Sign claims as a JWS Compact Serialization string with Ed25519 (EdDSA).
30///
31/// Implements RFC 7515 (JWS) + RFC 8037 (CFRG curves):
32/// - Header: `{"alg":"EdDSA","typ":"<typ>"}`
33/// - Payload: JSON-encoded `claims`
34/// - Signature: Ed25519 over the ASCII bytes `b64url(header) || "." || b64url(payload)`
35///
36/// Returns the canonical compact form `<header>.<payload>.<signature>` so it can
37/// be passed straight into the homeserver's JSON request body or any RFC-7515
38/// JWS verifier (e.g. `jsonwebtoken::decode`).
39pub fn sign_jws<T: Serialize>(keypair: &Keypair, typ: &str, claims: &T) -> String {
40    let signing_input = jws_signing_input(typ, claims);
41    let signature = keypair.sign(signing_input.as_bytes());
42    finish_jws(signing_input, signature.to_bytes())
43}
44
45/// Build the canonical JWS signing input `base64url(header).base64url(payload)`.
46///
47/// This is useful for runtimes where the private key is held by an external
48/// signer, such as WebCrypto, while keeping the JWS bytes identical to
49/// [`sign_jws`].
50pub fn jws_signing_input<T: Serialize>(typ: &str, claims: &T) -> String {
51    let header = serde_json::json!({ "alg": "EdDSA", "typ": typ });
52    let header_b64 = URL_SAFE_NO_PAD.encode(
53        serde_json::to_vec(&header)
54            .expect("invariant: serde_json serialization of a static header object cannot fail"),
55    );
56    let payload_b64 = URL_SAFE_NO_PAD.encode(
57        serde_json::to_vec(claims).expect("invariant: claims must be serde_json-serializable"),
58    );
59
60    format!("{header_b64}.{payload_b64}")
61}
62
63/// Finish compact JWS serialization from a signing input and raw signature.
64/// Merges the signing input with the base64url-encoded signature to produce the final
65/// JWS string. This is useful for runtimes where the signing step is separate from
66/// the signing input construction, such as when using an external signer.
67#[must_use]
68pub fn finish_jws(signing_input: String, signature: impl AsRef<[u8]>) -> String {
69    let signature_b64 = URL_SAFE_NO_PAD.encode(signature);
70    format!("{signing_input}.{signature_b64}")
71}
72
73// ── JWS Decoding ────────────────────────────────────────────────────────────
74
75/// Decode a JWS Compact Serialization string's payload WITHOUT signature verification.
76///
77/// `compact` is a JWS in Compact Serialization form (RFC 7515 §7.1):
78/// three base64url-encoded segments separated by dots (`header.payload.signature`).
79///
80/// Splits on `.`, base64url-decodes the payload (second) segment,
81/// and deserializes from JSON. Useful for the SDK to inspect token
82/// contents and check expiry without needing the signer's public key.
83pub fn decode_jws_payload<T: serde::de::DeserializeOwned>(compact: &str) -> Result<T, Error> {
84    let parts: Vec<&str> = compact.splitn(3, '.').collect();
85    if parts.len() != 3 {
86        return Err(Error::InvalidFormat(
87            "JWS compact must have 3 dot-separated parts",
88        ));
89    }
90
91    let payload_bytes = URL_SAFE_NO_PAD
92        .decode(parts[1])
93        .map_err(|_| Error::InvalidFormat("invalid base64url in JWS payload"))?;
94
95    serde_json::from_slice(&payload_bytes).map_err(|e| Error::JsonParse(e.to_string()))
96}
97
98// ── RandomId ────────────────────────────────────────────────────────────────
99
100/// A cryptographically random identifier, max 22 characters.
101///
102/// Uses base64url (22 chars for 128-bit).
103/// Serde-transparent: serializes as a plain string in JSON.
104#[derive(Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
105#[serde(try_from = "String", into = "String")]
106pub struct RandomId(String);
107
108impl RandomId {
109    /// Generate a new random ID: 128-bit random bytes → base64url (22 chars).
110    pub fn generate() -> Self {
111        let bytes = random_bytes::<16>();
112        Self(URL_SAFE_NO_PAD.encode(bytes))
113    }
114
115    /// Parse and validate an existing ID string.
116    ///
117    /// Must be non-empty, at most 22 characters, and contain only base64url characters.
118    pub fn parse(s: &str) -> Result<Self, Error> {
119        if s.is_empty() {
120            return Err(Error::InvalidFormat("RandomId must not be empty"));
121        }
122        if s.len() > RANDOM_ID_MAX_LENGTH {
123            return Err(Error::InvalidFormat(
124                "RandomId must be at most 22 characters",
125            ));
126        }
127        if !s
128            .bytes()
129            .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_'))
130        {
131            return Err(Error::InvalidFormat(
132                "RandomId must contain only base64url characters",
133            ));
134        }
135        Ok(Self(s.to_string()))
136    }
137
138    /// Returns the inner string representation.
139    pub fn as_str(&self) -> &str {
140        &self.0
141    }
142}
143
144impl fmt::Display for RandomId {
145    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
146        f.write_str(&self.0)
147    }
148}
149
150impl TryFrom<String> for RandomId {
151    type Error = Error;
152
153    fn try_from(s: String) -> Result<Self, Self::Error> {
154        Self::parse(&s)
155    }
156}
157
158impl From<RandomId> for String {
159    fn from(id: RandomId) -> Self {
160        id.0
161    }
162}
163
164/// Grant identifier — a [`RandomId`] used as the `jti` claim in a Grant JWS.
165pub type GrantId = RandomId;
166
167/// Proof-of-Possession nonce — a [`RandomId`] used to prevent PoP replay.
168pub type PopNonce = RandomId;
169
170// ── ClientId ────────────────────────────────────────────────────────────────
171
172/// An application identifier, typically a domain string (e.g., `franky.pubky.app`).
173///
174/// Max 253 characters (DNS domain name limit per RFC 1035).
175/// Serde-transparent: serializes as a plain string in JSON.
176#[derive(Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
177#[serde(try_from = "String", into = "String")]
178pub struct ClientId(String);
179
180impl ClientId {
181    /// Create a new [`ClientId`], validating that it is non-empty and at most 253 characters.
182    pub fn new(s: &str) -> Result<Self, Error> {
183        if s.is_empty() {
184            return Err(Error::InvalidFormat("ClientId must not be empty"));
185        }
186        if s.len() > CLIENT_ID_MAX_LENGTH {
187            return Err(Error::InvalidFormat(
188                "ClientId must be at most 253 characters",
189            ));
190        }
191        Ok(Self(s.to_string()))
192    }
193
194    /// Returns the inner string representation.
195    pub fn as_str(&self) -> &str {
196        &self.0
197    }
198}
199
200impl fmt::Display for ClientId {
201    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
202        f.write_str(&self.0)
203    }
204}
205
206impl TryFrom<String> for ClientId {
207    type Error = Error;
208
209    fn try_from(s: String) -> Result<Self, Self::Error> {
210        Self::new(&s)
211    }
212}
213
214impl TryFrom<&str> for ClientId {
215    type Error = Error;
216
217    fn try_from(s: &str) -> Result<Self, Self::Error> {
218        Self::new(s)
219    }
220}
221
222impl From<ClientId> for String {
223    fn from(id: ClientId) -> Self {
224        id.0
225    }
226}
227
228// ── Errors ──────────────────────────────────────────────────────────────────
229
230/// Errors from JWS decoding and ID parsing.
231#[derive(thiserror::Error, Debug)]
232pub enum Error {
233    /// The input format is invalid.
234    #[error("{0}")]
235    InvalidFormat(&'static str),
236
237    /// JSON parsing failed.
238    #[error("JSON parse error: {0}")]
239    JsonParse(String),
240}
241
242#[cfg(test)]
243mod tests {
244
245    use super::*;
246
247    #[test]
248    fn try_into_test() {
249        let domain = "example.com";
250        let client_id: ClientId = domain.try_into().unwrap();
251        assert_eq!(client_id.as_str(), domain);
252    }
253
254    #[test]
255    fn random_id_generate_is_valid() {
256        let id = RandomId::generate();
257        assert!(!id.as_str().is_empty());
258        assert!(id.as_str().len() <= RANDOM_ID_MAX_LENGTH);
259        // base64url of 16 bytes = 22 chars
260        assert_eq!(id.as_str().len(), 22);
261    }
262
263    #[test]
264    fn random_id_uniqueness() {
265        let a = RandomId::generate();
266        let b = RandomId::generate();
267        assert_ne!(a, b);
268    }
269
270    #[test]
271    fn random_id_parse_valid() {
272        RandomId::parse("abc123").unwrap();
273        RandomId::parse("AZaz09-_").unwrap();
274        RandomId::parse("a").unwrap(); // min length
275    }
276
277    #[test]
278    fn random_id_parse_rejects_non_base64url_characters() {
279        for value in [
280            ".",
281            "..",
282            "../../../pub/a.txt",
283            "a/b",
284            "a?b",
285            "a#b",
286            "a+b",
287            "a=b",
288            "a b",
289        ] {
290            assert!(RandomId::parse(value).is_err(), "accepted {value:?}");
291        }
292    }
293
294    #[test]
295    fn random_id_parse_rejects_empty() {
296        assert!(RandomId::parse("").is_err());
297    }
298
299    #[test]
300    fn random_id_parse_rejects_too_long() {
301        let long = "a".repeat(RANDOM_ID_MAX_LENGTH + 1);
302        assert!(RandomId::parse(&long).is_err());
303    }
304
305    #[test]
306    fn random_id_serde_roundtrip() {
307        let id = RandomId::generate();
308        let json = serde_json::to_string(&id).unwrap();
309        let parsed: RandomId = serde_json::from_str(&json).unwrap();
310        assert_eq!(id, parsed);
311    }
312
313    #[test]
314    fn client_id_valid() {
315        ClientId::new("franky.pubky.app").unwrap();
316        ClientId::new("a").unwrap();
317    }
318
319    #[test]
320    fn client_id_rejects_empty() {
321        assert!(ClientId::new("").is_err());
322    }
323
324    #[test]
325    fn client_id_rejects_too_long() {
326        let long = "a".repeat(CLIENT_ID_MAX_LENGTH + 1);
327        assert!(ClientId::new(&long).is_err());
328    }
329
330    #[test]
331    fn client_id_serde_roundtrip() {
332        let id = ClientId::new("test.app").unwrap();
333        let json = serde_json::to_string(&id).unwrap();
334        let parsed: ClientId = serde_json::from_str(&json).unwrap();
335        assert_eq!(id, parsed);
336    }
337
338    #[test]
339    fn sign_jws_round_trips_through_decode_jws_payload() {
340        let kp = Keypair::random();
341        #[derive(Serialize, Deserialize, PartialEq, Debug)]
342        struct Claims {
343            sub: String,
344            iat: u64,
345        }
346        let claims = Claims {
347            sub: "alice".into(),
348            iat: 1_700_000_000,
349        };
350        let compact = sign_jws(&kp, "pubky-test", &claims);
351
352        // Three dot-separated parts.
353        assert_eq!(compact.matches('.').count(), 2);
354
355        // Payload survives decode.
356        let decoded: Claims = decode_jws_payload(&compact).unwrap();
357        assert_eq!(decoded, claims);
358    }
359
360    #[test]
361    fn sign_jws_signature_verifies_with_raw_ed25519() {
362        let kp = Keypair::random();
363        let claims = serde_json::json!({"foo": "bar"});
364        let compact = sign_jws(&kp, "pubky-test", &claims);
365
366        let mut parts = compact.splitn(3, '.');
367        let header_b64 = parts.next().unwrap();
368        let payload_b64 = parts.next().unwrap();
369        let signature_b64 = parts.next().unwrap();
370        let signing_input = format!("{header_b64}.{payload_b64}");
371
372        let signature_bytes = URL_SAFE_NO_PAD.decode(signature_b64).unwrap();
373        assert_eq!(signature_bytes.len(), 64);
374        let signature_arr: [u8; 64] = signature_bytes.try_into().unwrap();
375        let signature = ed25519_dalek::Signature::from_bytes(&signature_arr);
376        kp.public_key()
377            .verify(signing_input.as_bytes(), &signature)
378            .expect("signature must verify against the keypair's public key");
379    }
380
381    #[test]
382    fn sign_jws_header_contains_alg_and_typ() {
383        let kp = Keypair::random();
384        let compact = sign_jws(&kp, GRANT_JWS_TYP, &serde_json::json!({}));
385        let header_b64 = compact.split('.').next().unwrap();
386        let header_bytes = URL_SAFE_NO_PAD.decode(header_b64).unwrap();
387        let header: serde_json::Value = serde_json::from_slice(&header_bytes).unwrap();
388        assert_eq!(header["alg"], "EdDSA");
389        assert_eq!(header["typ"], GRANT_JWS_TYP);
390    }
391
392    #[test]
393    fn decode_jws_payload_valid() {
394        // Manually construct a JWS-like string: header.payload.signature
395        // Payload: {"sub":"hello"}
396        let payload = URL_SAFE_NO_PAD.encode(b"{\"sub\":\"hello\"}");
397        let header = URL_SAFE_NO_PAD.encode(b"{\"alg\":\"EdDSA\"}");
398        let compact = format!("{}.{}.fakesig", header, payload);
399
400        #[derive(Deserialize)]
401        struct Claims {
402            sub: String,
403        }
404
405        let claims: Claims = decode_jws_payload(&compact).unwrap();
406        assert_eq!(claims.sub, "hello");
407    }
408
409    #[test]
410    fn decode_jws_payload_rejects_malformed() {
411        assert!(decode_jws_payload::<serde_json::Value>("not.a.valid.jws.toomanyparts").is_err());
412        assert!(decode_jws_payload::<serde_json::Value>("only-one-part").is_err());
413    }
414}