Skip to main content

whatsapp_rust/
handshake.rs

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
20/// One IK failure per process before falling back to XX (matches WA Web).
21const 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    /// Producer side of `transport_events` was dropped — distinct from a
33    /// timeout because nothing more will ever arrive on the channel,
34    /// regardless of how long we wait. Surfaced separately so callers can
35    /// log it accurately and so retry policies that pace themselves on
36    /// timeout don't silently swallow a teardown.
37    #[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    /// The handshake ran out of time, as opposed to being torn down.
47    ///
48    /// Matched exhaustively so a new variant has to be classified here rather
49    /// than defaulting to "not a timeout" unnoticed.
50    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    /// Transient errors that are expected during reconnect and will resolve
64    /// on retry. These never invalidate the cached server static.
65    pub fn is_transient(&self) -> bool {
66        matches!(
67            self,
68            Self::Transport(_) | Self::Timeout | Self::Disconnected | Self::StreamClosed
69        )
70    }
71
72    /// Crypto-fatal: a cached server static or cert chain is no longer
73    /// trustworthy. The orchestration layer must clear the IK cache and
74    /// fall back to XX on the next attempt.
75    ///
76    /// Narrowed to the `Core` variants that actually point at a stale or
77    /// poisoned cache. Programmer-side bugs (`Proto` encode failure, our
78    /// own crypto provider misuse, HKDF impossible failure, counter
79    /// exhaustion in a single handshake) are NOT crypto-fatal — they
80    /// indicate a code defect, and clearing the cache would mask it. A
81    /// stream-closed event during recv is treated as transient by
82    /// `is_transient`, not here.
83    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            // Server-supplied bytes failed AEAD authentication or had the
91            // wrong shape — canonical "the static we used to derive ee/se
92            // doesn't actually belong to this server" signal.
93            Core::Noise(NoiseError::Decrypt(_))
94            | Core::Noise(NoiseError::CiphertextTooShort)
95            | Core::Noise(NoiseError::InvalidKeyLength { .. }) => true,
96            // Cert content didn't match the static we just decrypted, or
97            // the chain was structurally invalid.
98            Core::CertVerification(_) => true,
99            // Server sent a structurally invalid response. Either it's
100            // out of sync with our cached static or it has a real bug;
101            // either way IK won't recover, so fall back.
102            Core::IncompleteResponse
103            | Core::InvalidLength { .. }
104            | Core::InvalidKeyLength
105            | Core::ProtoDecode(_) => true,
106            // Programmer-side: our encode shouldn't fail with a valid
107            // Device, our own crypto provider shouldn't reject our own
108            // inputs, HKDF can't reasonably fail, and a single handshake
109            // can't exhaust the counter. None of these mean the cache
110            // is bad.
111            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/// Pattern picked at the start of a handshake based on cached state.
123#[derive(Debug, Clone, Copy, PartialEq, Eq)]
124enum HandshakePattern {
125    /// Cold start / pairing / forced fallback after an earlier IK failure.
126    Xx,
127    /// Cached server static + valid cert chain available; attempt IK.
128    Ik([u8; 32]),
129}
130
131fn select_pattern(
132    device: &wacore::store::Device,
133    ik_failures: u32,
134    now_secs: i64,
135) -> HandshakePattern {
136    // Unregistered + cached chain is a signal of a legacy DB written before
137    // the registration gate; `do_handshake` no longer creates that state but
138    // we still need to refuse IK against it.
139    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    // `not_before` covers backwards clock skew, `not_after` is normal expiry.
149    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
159/// `server_cert_chain` is `Some` for XX / XX-fallback (fresh chain to persist)
160/// and `None` for IK Continue (on-disk cache stays authoritative).
161struct 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            // Skip invalidation past the XXfallback pivot: by that point the
238            // server has already accepted our IK ClientHello and the cache
239            // is no longer the implicated party.
240            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/// `fallback_taken` is set to `true` once we pivot from IK to XXfallback,
296/// before any operation that could fail.
297#[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            // Channel closed (no more producers) — distinct from a real timeout.
403            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        // Transient — never invalidate the cache.
543        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        // Stale-cache-indicating Core variants.
551        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        // Programmer-side bug: Crypto(String) wraps generic crypto-provider
561        // misuse; not a server-side cache problem.
562        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    /// Both the XX and IK initial messages must travel inside a frame whose
571    /// prologue is `WA_CONN_HEADER` (optionally preceded by an edge-routing
572    /// pre-intro). The wire-side server validates this prologue when it
573    /// re-derives `h0` for transcript MAC checks, so any divergence between
574    /// the two paths would surface only as a generic AEAD failure.
575    ///
576    /// We compare by fingerprinting the header bytes returned by the shared
577    /// helper for the two relevant scenarios — IK and XX both must hit the
578    /// same builder, with edge-routing applied identically when present.
579    #[test]
580    fn xx_and_ik_share_same_first_frame_prologue() {
581        // No edge routing: pure WA_CONN_HEADER.
582        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        // With edge routing: pre-intro applied identically.
589        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}