1use std::fmt;
19
20use sha2::{Digest, Sha384};
21
22use crate::binding::{
23 verify_connect_binding, verify_status, verify_tls_binding, BindingError, SignedTbs,
24};
25use crate::cbor::{self, Value};
26use crate::node_key::{
27 carried_key_well_formed, node_id_of, puzzle_solved, signature_size, verify, KeyError, NodeKey,
28};
29use crate::profile::Profile;
30
31pub const VERSION: i64 = 4;
34
35const NONCE_SIZE: usize = 32;
36const MAX_PROTOCOL_INT: i64 = 1 << 53;
37const MLDSA_KEY_SIZE: usize = 2592;
38const CONNECT_PROOF_LABEL: &[u8] = b"MACULA-PQ-CONNECT-PROOF-V1";
39
40const OPENER_KEYS: &[&str] = &["frame_type", "version"];
41const CHALLENGE_KEYS: &[&str] = &[
42 "frame_type",
43 "identity_key",
44 "nonce",
45 "profile",
46 "tls_binding",
47 "tls_status",
48 "version",
49];
50const CONNECT_KEYS: &[&str] = &[
54 "capabilities",
55 "connect_binding",
56 "connect_key",
57 "connect_status",
58 "frame_type",
59 "identity_key",
60 "member_endorsement",
61 "proof",
62 "version",
63];
64const HELLO_ACCEPTED_KEYS: &[&str] = &["accepted", "capabilities", "frame_type", "version"];
65const HELLO_REFUSED_KEYS: &[&str] = &[
66 "accepted",
67 "capabilities",
68 "frame_type",
69 "refusal_code",
70 "version",
71];
72const STATUS_KEYS: &[&str] = &["frame_type", "statement", "version"];
73
74#[derive(Debug, Clone, Copy, PartialEq, Eq)]
76pub enum RefusalCode {
77 UnsupportedVersion,
79 PuzzleInvalid,
81 NotAccepted,
83}
84
85impl RefusalCode {
86 fn name(self) -> &'static str {
87 match self {
88 RefusalCode::UnsupportedVersion => "unsupported_version",
89 RefusalCode::PuzzleInvalid => "puzzle_invalid",
90 RefusalCode::NotAccepted => "not_accepted",
91 }
92 }
93
94 fn parse(name: &str) -> Option<RefusalCode> {
95 match name {
96 "unsupported_version" => Some(RefusalCode::UnsupportedVersion),
97 "puzzle_invalid" => Some(RefusalCode::PuzzleInvalid),
98 "not_accepted" => Some(RefusalCode::NotAccepted),
99 _ => None,
100 }
101 }
102}
103
104#[derive(Debug, Clone, PartialEq, Eq)]
107pub enum HandshakeError {
108 UnexpectedFrame,
110 UnsupportedVersion,
112 Malformed,
115 ProfileMismatch,
117 KeyPurposeReuse,
120 PeerIdentityMismatch {
122 expected: [u8; 32],
123 derived: [u8; 32],
124 },
125 PuzzleInvalid,
127 ProofInvalid,
129 Refused(RefusalCode),
131 InvalidStationSession,
133 Binding(BindingError),
135 Key(KeyError),
137}
138
139impl fmt::Display for HandshakeError {
140 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
141 match self {
142 HandshakeError::UnexpectedFrame => f.write_str("unexpected frame"),
143 HandshakeError::UnsupportedVersion => f.write_str("unsupported frame version"),
144 HandshakeError::Malformed => f.write_str("malformed frame"),
145 HandshakeError::ProfileMismatch => f.write_str("the peer names another profile"),
146 HandshakeError::KeyPurposeReuse => f.write_str("a key would serve two purposes"),
147 HandshakeError::PeerIdentityMismatch { expected, derived } => write!(
148 f,
149 "dialed node_id {}, but the station's key derives {}",
150 hex_of(expected),
151 hex_of(derived)
152 ),
153 HandshakeError::PuzzleInvalid => f.write_str("the node_id does not meet the puzzle"),
154 HandshakeError::ProofInvalid => f.write_str("the CONNECT proof does not verify"),
155 HandshakeError::Refused(code) => {
156 write!(f, "the station refused the connection: {}", code.name())
157 }
158 HandshakeError::InvalidStationSession => {
159 f.write_str("the station session has an unknown puzzle difficulty")
160 }
161 HandshakeError::Binding(e) => write!(f, "{e}"),
162 HandshakeError::Key(e) => write!(f, "{e}"),
163 }
164 }
165}
166
167impl std::error::Error for HandshakeError {}
168
169impl From<BindingError> for HandshakeError {
170 fn from(e: BindingError) -> Self {
171 HandshakeError::Binding(e)
172 }
173}
174
175#[derive(Debug, Clone, Copy, PartialEq, Eq)]
177pub enum PuzzleMode {
178 Off,
180 LogOnly,
182 Enforce,
184}
185
186#[derive(Debug, Clone, Copy, PartialEq, Eq)]
188pub enum PuzzleResult {
189 Solved,
190 Unsolved,
191 NotChecked,
192}
193
194pub fn opener() -> Vec<u8> {
197 encode_frame("opener", Vec::new())
198}
199
200pub fn read_opener(frame: &[u8]) -> Result<(), HandshakeError> {
202 decode(frame, "opener", &[OPENER_KEYS]).map(|_| ())
203}
204
205#[derive(Debug, Clone)]
208pub struct StationMaterial {
209 pub profile: Profile,
210 pub identity_key: Vec<u8>,
211 pub tls_binding: SignedTbs,
212 pub tls_status: SignedTbs,
213}
214
215pub fn challenge(m: &StationMaterial) -> Result<Vec<u8>, HandshakeError> {
218 let mut nonce = [0u8; NONCE_SIZE];
219 aws_lc_rs::rand::fill(&mut nonce)
220 .map_err(|_| HandshakeError::Key(KeyError::RandomnessUnavailable))?;
221 Ok(encode_frame(
222 "challenge",
223 vec![
224 entry("nonce", Value::Bytes(nonce.to_vec())),
225 entry("profile", Value::text(m.profile.name())),
226 entry("identity_key", Value::Bytes(m.identity_key.clone())),
227 entry("tls_binding", m.tls_binding.to_value()),
228 entry("tls_status", m.tls_status.to_value()),
229 ],
230 ))
231}
232
233pub struct ClientSession<'a> {
239 pub profile: Profile,
240 pub expected_node_id: [u8; 32],
241 pub leaf: &'a [u8],
242 pub identity_key: Vec<u8>,
243 pub connect_key: &'a NodeKey,
244 pub connect_binding: &'a SignedTbs,
245 pub connect_status: &'a SignedTbs,
246 pub capabilities: u64,
247 pub now_ms: i64,
248 pub member_endorsement: Vec<u8>,
249}
250
251#[derive(Debug, Clone, PartialEq, Eq)]
253pub struct Station {
254 pub node_id: [u8; 32],
255 pub identity_key: Vec<u8>,
256 pub tls_binding: SignedTbs,
257 pub status_expires_at: i64,
258 pub binding_not_after: i64,
259}
260
261pub fn answer_challenge(
267 challenge: &[u8],
268 s: &ClientSession<'_>,
269) -> Result<(Vec<u8>, Station), HandshakeError> {
270 let f = decode(challenge, "challenge", &[CHALLENGE_KEYS])?;
271 let station_key = f.bytes("identity_key");
272 let connect_key = s.connect_key.public_key();
273 if f.text("profile") != s.profile.name() {
274 return Err(HandshakeError::ProfileMismatch);
275 }
276 if !carried_key_well_formed(station_key, s.profile) {
277 return Err(HandshakeError::Malformed);
278 }
279 if shares_a_half(&s.identity_key, &connect_key)
280 || in_leaf(station_key, s.leaf)
281 || in_leaf(&connect_key, s.leaf)
282 {
283 return Err(HandshakeError::KeyPurposeReuse);
284 }
285 let station_node_id = node_id_of(station_key, s.profile);
286 if station_node_id != s.expected_node_id {
287 return Err(HandshakeError::PeerIdentityMismatch {
288 expected: s.expected_node_id,
289 derived: station_node_id,
290 });
291 }
292 let tls_binding = f.signed("tls_binding");
293 let binding = verify_tls_binding(&tls_binding, station_key, s.profile, s.leaf, s.now_ms)?;
294 let expires_at = verify_status(
295 &f.signed("tls_status"),
296 &tls_binding,
297 station_key,
298 s.profile,
299 s.now_ms,
300 )?;
301 let client_node_id = node_id_of(&s.identity_key, s.profile);
302 let proof = s
303 .connect_key
304 .sign(&proof_message(
305 f.bytes("nonce"),
306 &station_node_id,
307 &client_node_id,
308 s.leaf,
309 challenge,
310 ))
311 .map_err(HandshakeError::Key)?;
312 let connect = encode_frame(
313 "connect",
314 vec![
315 entry("identity_key", Value::Bytes(s.identity_key.clone())),
316 entry("connect_key", Value::Bytes(connect_key)),
317 entry("connect_binding", s.connect_binding.to_value()),
318 entry("connect_status", s.connect_status.to_value()),
319 entry("proof", Value::Bytes(proof)),
320 entry("capabilities", Value::Int(i128::from(s.capabilities))),
321 entry(
322 "member_endorsement",
323 Value::Bytes(s.member_endorsement.clone()),
324 ),
325 ],
326 );
327 Ok((
328 connect,
329 Station {
330 node_id: station_node_id,
331 identity_key: station_key.to_vec(),
332 tls_binding,
333 status_expires_at: expires_at,
334 binding_not_after: binding.not_after,
335 },
336 ))
337}
338
339#[derive(Debug, Clone)]
343pub struct StationSession {
344 pub profile: Profile,
345 pub challenge: Vec<u8>,
346 pub leaf: Vec<u8>,
347 pub puzzle_difficulty: u32,
348 pub puzzle_mode: PuzzleMode,
349 pub capabilities: u64,
350 pub now_ms: i64,
351}
352
353#[derive(Debug, Clone, PartialEq, Eq)]
355pub struct Client {
356 pub node_id: [u8; 32],
357 pub identity_key: Vec<u8>,
358 pub connect_key: Vec<u8>,
359 pub connect_binding: SignedTbs,
360 pub capabilities: u64,
361 pub status_expires_at: i64,
362 pub binding_not_after: i64,
363 pub puzzle: PuzzleResult,
364 pub member_endorsement: Vec<u8>,
367}
368
369pub fn accept_connect(
377 connect: &[u8],
378 s: &StationSession,
379) -> (Result<Client, HandshakeError>, Vec<u8>) {
380 match check_connect(connect, s) {
381 Ok(client) => (Ok(client), hello(None, s.capabilities)),
382 Err(e) => {
383 let code = match e {
384 HandshakeError::UnsupportedVersion => RefusalCode::UnsupportedVersion,
385 HandshakeError::PuzzleInvalid => RefusalCode::PuzzleInvalid,
386 _ => RefusalCode::NotAccepted,
387 };
388 (Err(e), hello(Some(code), s.capabilities))
389 }
390 }
391}
392
393fn check_connect(connect: &[u8], s: &StationSession) -> Result<Client, HandshakeError> {
394 if s.puzzle_difficulty > 256 {
395 return Err(HandshakeError::InvalidStationSession);
396 }
397 let f = decode(connect, "connect", &[CONNECT_KEYS])?;
398 let (identity_key, connect_key, proof) = (
399 f.bytes("identity_key"),
400 f.bytes("connect_key"),
401 f.bytes("proof"),
402 );
403 if !carried_key_well_formed(identity_key, s.profile)
404 || !carried_key_well_formed(connect_key, s.profile)
405 || proof.len() != signature_size(s.profile)
406 {
407 return Err(HandshakeError::Malformed);
408 }
409 if shares_a_half(identity_key, connect_key) || in_leaf(connect_key, &s.leaf) {
410 return Err(HandshakeError::KeyPurposeReuse);
411 }
412 let node_id = node_id_of(identity_key, s.profile);
413 let puzzle = match s.puzzle_mode {
414 PuzzleMode::Off => PuzzleResult::NotChecked,
415 _ if puzzle_solved(&node_id, s.puzzle_difficulty) => PuzzleResult::Solved,
416 _ => PuzzleResult::Unsolved,
417 };
418 if puzzle == PuzzleResult::Unsolved && s.puzzle_mode == PuzzleMode::Enforce {
419 return Err(HandshakeError::PuzzleInvalid);
420 }
421 let connect_binding = f.signed("connect_binding");
422 let binding = verify_connect_binding(
423 &connect_binding,
424 identity_key,
425 s.profile,
426 connect_key,
427 s.now_ms,
428 )?;
429 let expires_at = verify_status(
430 &f.signed("connect_status"),
431 &connect_binding,
432 identity_key,
433 s.profile,
434 s.now_ms,
435 )?;
436 if !proof_verifies(s, &node_id, connect_key, proof) {
437 return Err(HandshakeError::ProofInvalid);
438 }
439 Ok(Client {
440 node_id,
441 identity_key: identity_key.to_vec(),
442 connect_key: connect_key.to_vec(),
443 connect_binding,
444 capabilities: f.uint("capabilities"),
445 status_expires_at: expires_at,
446 binding_not_after: binding.not_after,
447 puzzle,
448 member_endorsement: f.bytes("member_endorsement").to_vec(),
449 })
450}
451
452fn proof_verifies(
455 s: &StationSession,
456 client_node_id: &[u8; 32],
457 connect_key: &[u8],
458 proof: &[u8],
459) -> bool {
460 let Ok(challenge) = decode(&s.challenge, "challenge", &[CHALLENGE_KEYS]) else {
461 return false;
462 };
463 let station_node_id = node_id_of(challenge.bytes("identity_key"), s.profile);
464 let message = proof_message(
465 challenge.bytes("nonce"),
466 &station_node_id,
467 client_node_id,
468 &s.leaf,
469 &s.challenge,
470 );
471 verify(&message, proof, connect_key, s.profile)
472}
473
474fn proof_message(
478 nonce: &[u8],
479 station_node_id: &[u8; 32],
480 client_node_id: &[u8; 32],
481 leaf: &[u8],
482 challenge: &[u8],
483) -> Vec<u8> {
484 let mut out = Vec::with_capacity(CONNECT_PROOF_LABEL.len() + 1 + nonce.len() + 64 + 96);
485 out.extend_from_slice(CONNECT_PROOF_LABEL);
486 out.push(0);
487 out.extend_from_slice(nonce);
488 out.extend_from_slice(station_node_id);
489 out.extend_from_slice(client_node_id);
490 out.extend_from_slice(&Sha384::digest(leaf));
491 out.extend_from_slice(&Sha384::digest(challenge));
492 out
493}
494
495fn hello(refusal: Option<RefusalCode>, capabilities: u64) -> Vec<u8> {
496 let mut entries = vec![
497 entry("accepted", Value::Int(i128::from(refusal.is_none()))),
498 entry("capabilities", Value::Int(i128::from(capabilities))),
499 ];
500 if let Some(code) = refusal {
501 entries.push(entry("refusal_code", Value::text(code.name())));
502 }
503 encode_frame("hello", entries)
504}
505
506pub fn read_hello(frame: &[u8]) -> Result<u64, HandshakeError> {
509 let f = decode(frame, "hello", &[HELLO_ACCEPTED_KEYS, HELLO_REFUSED_KEYS])?;
510 let accepted = f.int("accepted");
511 match (accepted, f.0.get("refusal_code")) {
512 (1, None) => Ok(f.uint("capabilities")),
513 (0, Some(Value::Text(code))) => Err(HandshakeError::Refused(
514 RefusalCode::parse(code).ok_or(HandshakeError::Malformed)?,
515 )),
516 _ => Err(HandshakeError::Malformed),
517 }
518}
519
520pub fn status_frame(statement: &SignedTbs) -> Vec<u8> {
522 encode_frame("status", vec![entry("statement", statement.to_value())])
523}
524
525#[derive(Debug, Clone)]
529pub struct Peer {
530 pub profile: Profile,
531 pub identity_key: Vec<u8>,
532 pub binding: SignedTbs,
533 pub now_ms: i64,
534}
535
536pub fn read_status(frame: &[u8], p: &Peer) -> Result<i64, HandshakeError> {
538 let f = decode(frame, "status", &[STATUS_KEYS])?;
539 Ok(verify_status(
540 &f.signed("statement"),
541 &p.binding,
542 &p.identity_key,
543 p.profile,
544 p.now_ms,
545 )?)
546}
547
548struct Fields(std::collections::HashMap<String, Value>);
550
551impl Fields {
552 fn bytes(&self, key: &str) -> &[u8] {
553 match self.0.get(key) {
554 Some(Value::Bytes(b)) => b,
555 _ => &[],
556 }
557 }
558
559 fn text(&self, key: &str) -> &str {
560 match self.0.get(key) {
561 Some(Value::Text(t)) => t,
562 _ => "",
563 }
564 }
565
566 fn int(&self, key: &str) -> i128 {
567 match self.0.get(key) {
568 Some(Value::Int(n)) => *n,
569 _ => -1,
570 }
571 }
572
573 fn uint(&self, key: &str) -> u64 {
574 u64::try_from(self.int(key)).unwrap_or(0)
575 }
576
577 fn signed(&self, key: &str) -> SignedTbs {
578 self.0
579 .get(key)
580 .and_then(|v| SignedTbs::from_value(v).ok())
581 .unwrap_or(SignedTbs {
582 tbs: Vec::new(),
583 signature: Vec::new(),
584 })
585 }
586}
587
588fn decode(frame: &[u8], frame_type: &str, layouts: &[&[&str]]) -> Result<Fields, HandshakeError> {
592 let Ok(Value::Map(pairs)) = cbor::decode(frame) else {
593 return Err(HandshakeError::Malformed);
594 };
595 let mut fields = std::collections::HashMap::with_capacity(pairs.len());
596 let mut non_text = 0;
597 for (key, value) in pairs {
598 match key {
599 Value::Text(name) => {
600 fields.insert(name, value);
601 }
602 _ => non_text += 1,
603 }
604 }
605 match fields.get("version") {
606 Some(Value::Int(v)) if *v == i128::from(VERSION) => {}
607 Some(Value::Int(_)) => return Err(HandshakeError::UnsupportedVersion),
608 _ => return Err(HandshakeError::Malformed),
609 }
610 match fields.get("frame_type") {
611 Some(Value::Text(t)) if t == frame_type => {}
612 Some(Value::Text(_)) => return Err(HandshakeError::UnexpectedFrame),
613 _ => return Err(HandshakeError::Malformed),
614 }
615 let mut keys: Vec<&str> = fields.keys().map(String::as_str).collect();
616 keys.sort_unstable();
617 let has_layout = layouts.contains(&keys.as_slice());
618 if non_text > 0 || !has_layout || !fields.iter().all(|(k, v)| field_typed(k, v)) {
619 return Err(HandshakeError::Malformed);
620 }
621 Ok(Fields(fields))
622}
623
624fn field_typed(key: &str, v: &Value) -> bool {
625 match key {
626 "version" | "frame_type" => true,
627 "profile" => matches!(v, Value::Text(_)),
628 "nonce" => matches!(v, Value::Bytes(b) if b.len() == NONCE_SIZE),
629 "identity_key" | "connect_key" | "proof" | "member_endorsement" => {
630 matches!(v, Value::Bytes(_))
631 }
632 "tls_binding" | "tls_status" | "connect_binding" | "connect_status" | "statement" => {
633 SignedTbs::from_value(v).is_ok()
634 }
635 "capabilities" => {
636 matches!(v, Value::Int(n) if *n >= 0 && *n < i128::from(MAX_PROTOCOL_INT))
637 }
638 "accepted" => matches!(v, Value::Int(0 | 1)),
639 "refusal_code" => matches!(v, Value::Text(t) if RefusalCode::parse(t).is_some()),
640 _ => false,
641 }
642}
643
644fn in_leaf(key: &[u8], leaf: &[u8]) -> bool {
646 key.len() >= MLDSA_KEY_SIZE
647 && leaf
648 .windows(MLDSA_KEY_SIZE)
649 .any(|w| w == &key[..MLDSA_KEY_SIZE])
650}
651
652fn shares_a_half(a: &[u8], b: &[u8]) -> bool {
654 if a.len() < MLDSA_KEY_SIZE || b.len() < MLDSA_KEY_SIZE {
655 return a == b;
656 }
657 let (classical_a, classical_b) = (&a[MLDSA_KEY_SIZE..], &b[MLDSA_KEY_SIZE..]);
658 a[..MLDSA_KEY_SIZE] == b[..MLDSA_KEY_SIZE]
659 || (!classical_a.is_empty() && classical_a == classical_b)
660}
661
662fn encode_frame(frame_type: &str, entries: Vec<(Value, Value)>) -> Vec<u8> {
663 let mut all = vec![
664 entry("version", Value::Int(i128::from(VERSION))),
665 entry("frame_type", Value::text(frame_type)),
666 ];
667 all.extend(entries);
668 cbor::encode(&Value::Map(all)).expect("a handshake frame's integers are all below 2^53")
669}
670
671fn entry(key: &str, value: Value) -> (Value, Value) {
672 (Value::text(key), value)
673}
674
675fn hex_of(bytes: &[u8]) -> String {
676 bytes.iter().map(|b| format!("{b:02x}")).collect()
677}