1use 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
15pub const GRANT_JWS_TYP: &str = "pubky-grant";
17
18pub const POP_JWS_TYP: &str = "pubky-pop";
20
21const RANDOM_ID_MAX_LENGTH: usize = 22;
23
24const CLIENT_ID_MAX_LENGTH: usize = 253;
26
27pub 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
45pub 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#[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
73pub 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#[derive(Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
105#[serde(try_from = "String", into = "String")]
106pub struct RandomId(String);
107
108impl RandomId {
109 pub fn generate() -> Self {
111 let bytes = random_bytes::<16>();
112 Self(URL_SAFE_NO_PAD.encode(bytes))
113 }
114
115 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 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
164pub type GrantId = RandomId;
166
167pub type PopNonce = RandomId;
169
170#[derive(Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
177#[serde(try_from = "String", into = "String")]
178pub struct ClientId(String);
179
180impl ClientId {
181 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 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#[derive(thiserror::Error, Debug)]
232pub enum Error {
233 #[error("{0}")]
235 InvalidFormat(&'static str),
236
237 #[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 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(); }
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 assert_eq!(compact.matches('.').count(), 2);
354
355 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 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}