1use crate::socket::NoiseSocket;
2use crate::store::persistence_manager::PersistenceManager;
3use crate::transport::{Transport, TransportEvent};
4use log::{debug, info, warn};
5use std::sync::Arc;
6use std::sync::atomic::{AtomicU32, Ordering};
7use std::time::Duration;
8use thiserror::Error;
9use wacore::handshake::{
10 HandshakeError as CoreHandshakeError, IkHandshakeState, IkServerHelloOutcome,
11 VerifiedServerCertChain, XxFallbackHandshakeState, XxHandshakeState, build_handshake_header,
12};
13use wacore::noise::NoiseCipher;
14use wacore::runtime::{Runtime, timeout as rt_timeout};
15use wacore::store::DeviceCommand;
16use wacore_binary::consts::WA_CONN_HEADER;
17
18const NOISE_HANDSHAKE_RESPONSE_TIMEOUT: Duration = Duration::from_secs(20);
19
20const IK_FAILURE_THRESHOLD: u32 = 1;
22
23#[derive(Debug, Error)]
24#[non_exhaustive]
25pub enum HandshakeError {
26 #[error("Transport error: {0}")]
27 Transport(#[from] anyhow::Error),
28 #[error("Core handshake error: {0}")]
29 Core(#[from] CoreHandshakeError),
30 #[error("Timed out waiting for handshake response")]
31 Timeout,
32 #[error("Transport event stream closed before handshake completed")]
38 StreamClosed,
39 #[error("Disconnected during handshake")]
40 Disconnected,
41 #[error("Unexpected event during handshake: {0}")]
42 UnexpectedEvent(String),
43}
44
45impl HandshakeError {
46 pub fn is_timeout(&self) -> bool {
51 match self {
52 HandshakeError::Timeout => true,
53 HandshakeError::Transport(_)
54 | HandshakeError::Core(_)
55 | HandshakeError::StreamClosed
56 | HandshakeError::Disconnected
57 | HandshakeError::UnexpectedEvent(_) => false,
58 }
59 }
60}
61
62impl HandshakeError {
63 pub fn is_transient(&self) -> bool {
66 matches!(
67 self,
68 Self::Transport(_) | Self::Timeout | Self::Disconnected | Self::StreamClosed
69 )
70 }
71
72 pub fn is_crypto_fatal(&self) -> bool {
84 let Self::Core(inner) = self else {
85 return false;
86 };
87 use wacore::handshake::HandshakeError as Core;
88 use wacore::noise::NoiseError;
89 match inner {
90 Core::Noise(NoiseError::Decrypt(_))
94 | Core::Noise(NoiseError::CiphertextTooShort)
95 | Core::Noise(NoiseError::InvalidKeyLength { .. }) => true,
96 Core::CertVerification(_) => true,
99 Core::IncompleteResponse
103 | Core::InvalidLength { .. }
104 | Core::InvalidKeyLength
105 | Core::ProtoDecode(_) => true,
106 Core::Crypto(_)
112 | Core::Noise(NoiseError::Encrypt(_))
113 | Core::Noise(NoiseError::HkdfExpandFailed)
114 | Core::Noise(NoiseError::InvalidPatternLength { .. })
115 | Core::Noise(NoiseError::CounterExhausted) => false,
116 }
117 }
118}
119
120type Result<T> = std::result::Result<T, HandshakeError>;
121
122#[derive(Debug, Clone, Copy, PartialEq, Eq)]
124enum HandshakePattern {
125 Xx,
127 Ik([u8; 32]),
129}
130
131fn select_pattern(
132 device: &wacore::store::Device,
133 ik_failures: u32,
134 now_secs: i64,
135) -> HandshakePattern {
136 if !device.is_registered() {
140 return HandshakePattern::Xx;
141 }
142 if ik_failures >= IK_FAILURE_THRESHOLD {
143 return HandshakePattern::Xx;
144 }
145 let Some(chain) = device.server_cert_chain.as_ref() else {
146 return HandshakePattern::Xx;
147 };
148 if now_secs < chain.leaf.not_before
150 || now_secs < chain.intermediate.not_before
151 || now_secs >= chain.leaf.not_after
152 || now_secs >= chain.intermediate.not_after
153 {
154 return HandshakePattern::Xx;
155 }
156 HandshakePattern::Ik(chain.leaf.key)
157}
158
159struct HandshakeSuccess {
162 write_cipher: NoiseCipher,
163 read_cipher: NoiseCipher,
164 server_cert_chain: Option<VerifiedServerCertChain>,
165}
166
167fn should_persist_cert_chain(device: &wacore::store::Device) -> bool {
168 device.is_registered()
169}
170
171#[cfg_attr(
172 feature = "tracing",
173 tracing::instrument(name = "wa.conn.handshake", level = "debug", skip_all, err(Debug))
174)]
175pub async fn do_handshake(
176 runtime: Arc<dyn Runtime>,
177 persistence_manager: &PersistenceManager,
178 ik_handshake_failures: &AtomicU32,
179 transport: Arc<dyn Transport>,
180 transport_events: &mut async_channel::Receiver<TransportEvent>,
181 stats: Option<Arc<wacore::stats::SessionStats>>,
182) -> Result<Arc<NoiseSocket>> {
183 let device_snapshot = persistence_manager.get_device_snapshot();
184 let now_secs = wacore::time::now_secs();
185 let pattern = select_pattern(
186 &device_snapshot,
187 ik_handshake_failures.load(Ordering::Acquire),
188 now_secs,
189 );
190
191 let mut fallback_taken = false;
192
193 let result = match pattern {
194 HandshakePattern::Xx => {
195 debug!("[socket] doFullHandshake: openChatSocket send hello");
196 run_xx_handshake(
197 &runtime,
198 &device_snapshot,
199 transport.clone(),
200 transport_events,
201 )
202 .await
203 }
204 HandshakePattern::Ik(server_static_pub) => {
205 debug!("[socket] resumeNoiseHandshake started");
206 run_ik_handshake(
207 &runtime,
208 &device_snapshot,
209 server_static_pub,
210 transport.clone(),
211 transport_events,
212 &mut fallback_taken,
213 )
214 .await
215 }
216 };
217
218 match result {
219 Ok(success) => {
220 if let Some(chain) = success.server_cert_chain
221 && should_persist_cert_chain(&device_snapshot)
222 {
223 persistence_manager
224 .process_command(DeviceCommand::SetServerCertChain(chain.into()))
225 .await;
226 }
227 ik_handshake_failures.store(0, Ordering::Release);
228 Ok(Arc::new(NoiseSocket::with_stats(
229 runtime,
230 transport,
231 success.write_cipher,
232 success.read_cipher,
233 stats,
234 )))
235 }
236 Err(e) => {
237 if matches!(pattern, HandshakePattern::Ik(_)) && !fallback_taken && e.is_crypto_fatal()
241 {
242 warn!(
243 "[socket] resumeNoiseHandshake failed crypto-fatally; \
244 clearing cached server cert chain and forcing XX next connect: {e}"
245 );
246 ik_handshake_failures.fetch_add(1, Ordering::AcqRel);
247 persistence_manager
248 .process_command(DeviceCommand::ClearServerCertChain)
249 .await;
250 }
251 Err(e)
252 }
253 }
254}
255
256#[cfg_attr(
257 feature = "tracing",
258 tracing::instrument(name = "wa.conn.handshake.xx", level = "debug", skip_all, err(Debug))
259)]
260async fn run_xx_handshake(
261 runtime: &Arc<dyn Runtime>,
262 device: &wacore::store::Device,
263 transport: Arc<dyn Transport>,
264 transport_events: &mut async_channel::Receiver<TransportEvent>,
265) -> Result<HandshakeSuccess> {
266 let client_payload = waproto::codec::client_payload_to_vec(&device.get_client_payload());
267 let mut handshake_state =
268 XxHandshakeState::new(device.noise_key.clone(), client_payload, &WA_CONN_HEADER)?;
269 let mut frame_decoder = wacore::framing::FrameDecoder::new();
270
271 let client_hello_bytes = handshake_state.build_client_hello()?;
272 send_first_handshake_message(&transport, device, &client_hello_bytes).await?;
273
274 let resp_frame = recv_frame(runtime, transport_events, &mut frame_decoder).await?;
275 debug!("[socket] openChatSocket rcv hello");
276
277 let client_finish_bytes =
278 handshake_state.read_server_hello_and_build_client_finish(&resp_frame)?;
279
280 debug!("[socket] continueFullHandshakeCore client finish and deriving secrets");
281 let framed = wacore::framing::encode_frame(&client_finish_bytes, None)
282 .map_err(HandshakeError::Transport)?;
283 transport.send(bytes::Bytes::from(framed)).await?;
284
285 let outcome = handshake_state.finish()?;
286 info!("Handshake complete (XX), switching to encrypted communication");
287
288 Ok(HandshakeSuccess {
289 write_cipher: outcome.write_cipher,
290 read_cipher: outcome.read_cipher,
291 server_cert_chain: Some(outcome.server_cert_chain),
292 })
293}
294
295#[cfg_attr(
298 feature = "tracing",
299 tracing::instrument(name = "wa.conn.handshake.ik", level = "debug", skip_all, err(Debug))
300)]
301async fn run_ik_handshake(
302 runtime: &Arc<dyn Runtime>,
303 device: &wacore::store::Device,
304 server_static_pub: [u8; 32],
305 transport: Arc<dyn Transport>,
306 transport_events: &mut async_channel::Receiver<TransportEvent>,
307 fallback_taken: &mut bool,
308) -> Result<HandshakeSuccess> {
309 let client_payload = waproto::codec::client_payload_to_vec(&device.get_client_payload());
310 let mut ik = IkHandshakeState::new(
311 device.noise_key.clone(),
312 server_static_pub,
313 client_payload,
314 &WA_CONN_HEADER,
315 )?;
316 let mut frame_decoder = wacore::framing::FrameDecoder::new();
317
318 debug!("[socket] resumeNoiseHandshake send hello");
319 let client_hello_bytes = ik.build_client_hello()?;
320 send_first_handshake_message(&transport, device, &client_hello_bytes).await?;
321
322 let resp_frame = recv_frame(runtime, transport_events, &mut frame_decoder).await?;
323 debug!("[socket] resumeNoiseHandshake rcv hello");
324
325 match ik.read_server_hello(&resp_frame)? {
326 IkServerHelloOutcome::Continue(out) => {
327 debug!("[socket] resumeNoiseHandshake deriving secrets");
328 info!("Handshake complete (IK), switching to encrypted communication");
329 Ok(HandshakeSuccess {
330 write_cipher: out.write_cipher,
331 read_cipher: out.read_cipher,
332 server_cert_chain: None,
333 })
334 }
335 IkServerHelloOutcome::Fallback(inputs) => {
336 *fallback_taken = true;
337 warn!(
338 "[socket] resumeNoiseHandshake failed: serverStaticCiphertext not null — \
339 doFallbackHandshake continuing handshake with given server hello"
340 );
341 let mut fb = XxFallbackHandshakeState::from_ik_failure(*inputs, &WA_CONN_HEADER)?;
342 let client_finish_bytes = fb.build_client_finish()?;
343 debug!(
344 "[socket] continueFullHandshakeCore client finish and deriving secrets (XXfallback)"
345 );
346 let framed = wacore::framing::encode_frame(&client_finish_bytes, None)
347 .map_err(HandshakeError::Transport)?;
348 transport.send(bytes::Bytes::from(framed)).await?;
349 let outcome = fb.finish()?;
350 info!("Handshake complete (XXfallback), switching to encrypted communication");
351 Ok(HandshakeSuccess {
352 write_cipher: outcome.write_cipher,
353 read_cipher: outcome.read_cipher,
354 server_cert_chain: Some(outcome.server_cert_chain),
355 })
356 }
357 }
358}
359
360async fn send_first_handshake_message(
361 transport: &Arc<dyn Transport>,
362 device: &wacore::store::Device,
363 payload_bytes: &[u8],
364) -> Result<()> {
365 let (header, used_edge_routing) = build_handshake_header(device.edge_routing_info.as_deref());
366 if used_edge_routing {
367 debug!("Sending edge routing pre-intro for optimized reconnection");
368 } else if device.edge_routing_info.is_some() {
369 warn!("Edge routing info provided but not used (possibly too large)");
370 }
371 let framed = wacore::framing::encode_frame(payload_bytes, Some(&header))
372 .map_err(HandshakeError::Transport)?;
373 transport.send(bytes::Bytes::from(framed)).await?;
374 Ok(())
375}
376
377async fn recv_frame(
378 runtime: &Arc<dyn Runtime>,
379 transport_events: &mut async_channel::Receiver<TransportEvent>,
380 frame_decoder: &mut wacore::framing::FrameDecoder,
381) -> Result<bytes::BytesMut> {
382 loop {
383 match rt_timeout(
384 &**runtime,
385 NOISE_HANDSHAKE_RESPONSE_TIMEOUT,
386 transport_events.recv(),
387 )
388 .await
389 {
390 Ok(Ok(TransportEvent::DataReceived(data))) => {
391 frame_decoder.feed(&data);
392 if let Some(frame) = frame_decoder.decode_frame() {
393 return Ok(frame);
394 }
395 continue;
396 }
397 Ok(Ok(TransportEvent::Connected)) => continue,
398 Ok(Ok(TransportEvent::Disconnected(reason))) => {
399 debug!("Transport disconnected during handshake: {reason}");
400 return Err(HandshakeError::Disconnected);
401 }
402 Ok(Err(_)) => return Err(HandshakeError::StreamClosed),
404 Err(_) => return Err(HandshakeError::Timeout),
405 }
406 }
407}
408
409#[cfg(test)]
410mod tests {
411 use super::*;
412 use wacore::store::CachedNoiseCert;
413 use wacore::store::CachedServerCertChain;
414
415 fn cached_chain(
416 leaf_key: [u8; 32],
417 leaf_not_after: i64,
418 intermediate_not_after: i64,
419 ) -> CachedServerCertChain {
420 CachedServerCertChain {
421 intermediate: CachedNoiseCert {
422 key: [0xCC; 32],
423 not_before: 1_700_000_000,
424 not_after: intermediate_not_after,
425 },
426 leaf: CachedNoiseCert {
427 key: leaf_key,
428 not_before: 1_700_000_000,
429 not_after: leaf_not_after,
430 },
431 }
432 }
433
434 fn paired_device() -> wacore::store::Device {
435 let mut device = wacore::store::Device::new();
436 device.pn = Some("12345@s.whatsapp.net".parse().unwrap());
437 device
438 }
439
440 #[test]
441 fn select_pattern_no_cache_returns_xx() {
442 let device = paired_device();
443 assert_eq!(
444 select_pattern(&device, 0, 1_800_000_000),
445 HandshakePattern::Xx
446 );
447 }
448
449 #[test]
450 fn select_pattern_with_valid_cache_returns_ik() {
451 let mut device = paired_device();
452 let pub_key = [0xAA; 32];
453 device.server_cert_chain = Some(cached_chain(pub_key, 1_900_000_000, 1_900_000_000));
454 assert_eq!(
455 select_pattern(&device, 0, 1_800_000_000),
456 HandshakePattern::Ik(pub_key)
457 );
458 }
459
460 #[test]
461 fn select_pattern_after_one_failure_returns_xx() {
462 let mut device = paired_device();
463 device.server_cert_chain = Some(cached_chain([0xAA; 32], 1_900_000_000, 1_900_000_000));
464 assert_eq!(
465 select_pattern(&device, IK_FAILURE_THRESHOLD, 1_800_000_000),
466 HandshakePattern::Xx
467 );
468 }
469
470 #[test]
471 fn select_pattern_with_expired_leaf_returns_xx() {
472 let mut device = paired_device();
473 device.server_cert_chain = Some(cached_chain([0xAA; 32], 1_700_000_500, 1_900_000_000));
474 assert_eq!(
475 select_pattern(&device, 0, 1_800_000_000),
476 HandshakePattern::Xx
477 );
478 }
479
480 #[test]
481 fn select_pattern_with_expired_intermediate_returns_xx() {
482 let mut device = paired_device();
483 device.server_cert_chain = Some(cached_chain([0xAA; 32], 1_900_000_000, 1_700_000_500));
484 assert_eq!(
485 select_pattern(&device, 0, 1_800_000_000),
486 HandshakePattern::Xx
487 );
488 }
489
490 #[test]
491 fn select_pattern_with_clock_before_leaf_not_before_returns_xx() {
492 let mut device = paired_device();
493 device.server_cert_chain = Some(cached_chain([0xAA; 32], 1_900_000_000, 1_900_000_000));
494 assert_eq!(
495 select_pattern(&device, 0, 1_699_999_999),
496 HandshakePattern::Xx
497 );
498 }
499
500 #[test]
501 fn select_pattern_with_clock_before_intermediate_not_before_returns_xx() {
502 let mut device = paired_device();
503 let mut chain = cached_chain([0xAA; 32], 1_900_000_000, 1_900_000_000);
504 chain.intermediate.not_before = 1_800_000_001;
505 device.server_cert_chain = Some(chain);
506 assert_eq!(
507 select_pattern(&device, 0, 1_800_000_000),
508 HandshakePattern::Xx
509 );
510 }
511
512 #[test]
513 fn select_pattern_unregistered_device_returns_xx_even_with_valid_cache() {
514 let mut device = wacore::store::Device::new();
515 assert!(
516 !device.is_registered(),
517 "fresh Device::new() must be unpaired"
518 );
519 device.server_cert_chain = Some(cached_chain([0xAA; 32], 1_900_000_000, 1_900_000_000));
520 assert_eq!(
521 select_pattern(&device, 0, 1_800_000_000),
522 HandshakePattern::Xx
523 );
524 }
525
526 #[test]
527 fn should_persist_cert_chain_unregistered_returns_false() {
528 let device = wacore::store::Device::new();
529 assert!(!device.is_registered());
530 assert!(!should_persist_cert_chain(&device));
531 }
532
533 #[test]
534 fn should_persist_cert_chain_registered_returns_true() {
535 let device = paired_device();
536 assert!(device.is_registered());
537 assert!(should_persist_cert_chain(&device));
538 }
539
540 #[test]
541 fn handshake_error_classification() {
542 assert!(HandshakeError::Timeout.is_transient());
544 assert!(HandshakeError::Disconnected.is_transient());
545 assert!(HandshakeError::StreamClosed.is_transient());
546 assert!(!HandshakeError::Timeout.is_crypto_fatal());
547 assert!(!HandshakeError::Disconnected.is_crypto_fatal());
548 assert!(!HandshakeError::StreamClosed.is_crypto_fatal());
549
550 for err in [
552 HandshakeError::Core(CoreHandshakeError::IncompleteResponse),
553 HandshakeError::Core(CoreHandshakeError::CertVerification("x".into())),
554 HandshakeError::Core(CoreHandshakeError::InvalidKeyLength),
555 ] {
556 assert!(err.is_crypto_fatal(), "{err:?} should be crypto-fatal");
557 assert!(!err.is_transient(), "{err:?} should not be transient");
558 }
559
560 let bug = HandshakeError::Core(CoreHandshakeError::Crypto("bug".into()));
563 assert!(
564 !bug.is_crypto_fatal(),
565 "generic Crypto(String) errors must not invalidate the cache"
566 );
567 assert!(!bug.is_transient());
568 }
569
570 #[test]
580 fn xx_and_ik_share_same_first_frame_prologue() {
581 let (xx_header, xx_used) = build_handshake_header(None);
583 let (ik_header, ik_used) = build_handshake_header(None);
584 assert_eq!(xx_header, ik_header);
585 assert_eq!(xx_used, ik_used);
586 assert!(xx_header.starts_with(b"WA"));
587
588 let routing = vec![0xDE, 0xAD, 0xBE, 0xEF];
590 let (xx_h2, xx_used2) = build_handshake_header(Some(&routing));
591 let (ik_h2, ik_used2) = build_handshake_header(Some(&routing));
592 assert_eq!(xx_h2, ik_h2);
593 assert_eq!(xx_used2, ik_used2);
594 assert!(xx_used2);
595 assert!(xx_h2.starts_with(b"ED\x00\x01"));
596 assert!(xx_h2.ends_with(b"WA\x06\x03") || xx_h2.ends_with(b"WA\x06\x04"));
597 }
598}