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 Ok(Self(s.to_string()))
128 }
129
130 pub fn as_str(&self) -> &str {
132 &self.0
133 }
134}
135
136impl fmt::Display for RandomId {
137 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
138 f.write_str(&self.0)
139 }
140}
141
142impl TryFrom<String> for RandomId {
143 type Error = Error;
144
145 fn try_from(s: String) -> Result<Self, Self::Error> {
146 Self::parse(&s)
147 }
148}
149
150impl From<RandomId> for String {
151 fn from(id: RandomId) -> Self {
152 id.0
153 }
154}
155
156pub type GrantId = RandomId;
158
159pub type PopNonce = RandomId;
161
162#[derive(Clone, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
169#[serde(try_from = "String", into = "String")]
170pub struct ClientId(String);
171
172impl ClientId {
173 pub fn new(s: &str) -> Result<Self, Error> {
175 if s.is_empty() {
176 return Err(Error::InvalidFormat("ClientId must not be empty"));
177 }
178 if s.len() > CLIENT_ID_MAX_LENGTH {
179 return Err(Error::InvalidFormat(
180 "ClientId must be at most 253 characters",
181 ));
182 }
183 Ok(Self(s.to_string()))
184 }
185
186 pub fn as_str(&self) -> &str {
188 &self.0
189 }
190}
191
192impl fmt::Display for ClientId {
193 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
194 f.write_str(&self.0)
195 }
196}
197
198impl TryFrom<String> for ClientId {
199 type Error = Error;
200
201 fn try_from(s: String) -> Result<Self, Self::Error> {
202 Self::new(&s)
203 }
204}
205
206impl TryFrom<&str> for ClientId {
207 type Error = Error;
208
209 fn try_from(s: &str) -> Result<Self, Self::Error> {
210 Self::new(s)
211 }
212}
213
214impl From<ClientId> for String {
215 fn from(id: ClientId) -> Self {
216 id.0
217 }
218}
219
220#[derive(thiserror::Error, Debug)]
224pub enum Error {
225 #[error("{0}")]
227 InvalidFormat(&'static str),
228
229 #[error("JSON parse error: {0}")]
231 JsonParse(String),
232}
233
234#[cfg(test)]
235mod tests {
236
237 use super::*;
238
239 #[test]
240 fn try_into_test() {
241 let domain = "example.com";
242 let client_id: ClientId = domain.try_into().unwrap();
243 assert_eq!(client_id.as_str(), domain);
244 }
245
246 #[test]
247 fn random_id_generate_is_valid() {
248 let id = RandomId::generate();
249 assert!(!id.as_str().is_empty());
250 assert!(id.as_str().len() <= RANDOM_ID_MAX_LENGTH);
251 assert_eq!(id.as_str().len(), 22);
253 }
254
255 #[test]
256 fn random_id_uniqueness() {
257 let a = RandomId::generate();
258 let b = RandomId::generate();
259 assert_ne!(a, b);
260 }
261
262 #[test]
263 fn random_id_parse_valid() {
264 RandomId::parse("abc123").unwrap();
265 RandomId::parse("a").unwrap(); }
267
268 #[test]
269 fn random_id_parse_rejects_empty() {
270 assert!(RandomId::parse("").is_err());
271 }
272
273 #[test]
274 fn random_id_parse_rejects_too_long() {
275 let long = "a".repeat(RANDOM_ID_MAX_LENGTH + 1);
276 assert!(RandomId::parse(&long).is_err());
277 }
278
279 #[test]
280 fn random_id_serde_roundtrip() {
281 let id = RandomId::generate();
282 let json = serde_json::to_string(&id).unwrap();
283 let parsed: RandomId = serde_json::from_str(&json).unwrap();
284 assert_eq!(id, parsed);
285 }
286
287 #[test]
288 fn client_id_valid() {
289 ClientId::new("franky.pubky.app").unwrap();
290 ClientId::new("a").unwrap();
291 }
292
293 #[test]
294 fn client_id_rejects_empty() {
295 assert!(ClientId::new("").is_err());
296 }
297
298 #[test]
299 fn client_id_rejects_too_long() {
300 let long = "a".repeat(CLIENT_ID_MAX_LENGTH + 1);
301 assert!(ClientId::new(&long).is_err());
302 }
303
304 #[test]
305 fn client_id_serde_roundtrip() {
306 let id = ClientId::new("test.app").unwrap();
307 let json = serde_json::to_string(&id).unwrap();
308 let parsed: ClientId = serde_json::from_str(&json).unwrap();
309 assert_eq!(id, parsed);
310 }
311
312 #[test]
313 fn sign_jws_round_trips_through_decode_jws_payload() {
314 let kp = Keypair::random();
315 #[derive(Serialize, Deserialize, PartialEq, Debug)]
316 struct Claims {
317 sub: String,
318 iat: u64,
319 }
320 let claims = Claims {
321 sub: "alice".into(),
322 iat: 1_700_000_000,
323 };
324 let compact = sign_jws(&kp, "pubky-test", &claims);
325
326 assert_eq!(compact.matches('.').count(), 2);
328
329 let decoded: Claims = decode_jws_payload(&compact).unwrap();
331 assert_eq!(decoded, claims);
332 }
333
334 #[test]
335 fn sign_jws_signature_verifies_with_raw_ed25519() {
336 let kp = Keypair::random();
337 let claims = serde_json::json!({"foo": "bar"});
338 let compact = sign_jws(&kp, "pubky-test", &claims);
339
340 let mut parts = compact.splitn(3, '.');
341 let header_b64 = parts.next().unwrap();
342 let payload_b64 = parts.next().unwrap();
343 let signature_b64 = parts.next().unwrap();
344 let signing_input = format!("{header_b64}.{payload_b64}");
345
346 let signature_bytes = URL_SAFE_NO_PAD.decode(signature_b64).unwrap();
347 assert_eq!(signature_bytes.len(), 64);
348 let signature_arr: [u8; 64] = signature_bytes.try_into().unwrap();
349 let signature = ed25519_dalek::Signature::from_bytes(&signature_arr);
350 kp.public_key()
351 .verify(signing_input.as_bytes(), &signature)
352 .expect("signature must verify against the keypair's public key");
353 }
354
355 #[test]
356 fn sign_jws_header_contains_alg_and_typ() {
357 let kp = Keypair::random();
358 let compact = sign_jws(&kp, GRANT_JWS_TYP, &serde_json::json!({}));
359 let header_b64 = compact.split('.').next().unwrap();
360 let header_bytes = URL_SAFE_NO_PAD.decode(header_b64).unwrap();
361 let header: serde_json::Value = serde_json::from_slice(&header_bytes).unwrap();
362 assert_eq!(header["alg"], "EdDSA");
363 assert_eq!(header["typ"], GRANT_JWS_TYP);
364 }
365
366 #[test]
367 fn decode_jws_payload_valid() {
368 let payload = URL_SAFE_NO_PAD.encode(b"{\"sub\":\"hello\"}");
371 let header = URL_SAFE_NO_PAD.encode(b"{\"alg\":\"EdDSA\"}");
372 let compact = format!("{}.{}.fakesig", header, payload);
373
374 #[derive(Deserialize)]
375 struct Claims {
376 sub: String,
377 }
378
379 let claims: Claims = decode_jws_payload(&compact).unwrap();
380 assert_eq!(claims.sub, "hello");
381 }
382
383 #[test]
384 fn decode_jws_payload_rejects_malformed() {
385 assert!(decode_jws_payload::<serde_json::Value>("not.a.valid.jws.toomanyparts").is_err());
386 assert!(decode_jws_payload::<serde_json::Value>("only-one-part").is_err());
387 }
388}