1use std::collections::VecDeque;
10use std::fmt;
11use std::future::Future;
12use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
13use std::sync::mpsc::{SyncSender, sync_channel};
14use std::sync::{Arc, Condvar, Mutex, MutexGuard};
15use std::thread::{self, JoinHandle};
16use std::time::{Duration, Instant as StdInstant, SystemTime, UNIX_EPOCH};
17
18use futures_util::{SinkExt, StreamExt};
19use rustls::ClientConfig;
20#[cfg(target_os = "android")]
21use rustls::RootCertStore;
22use rustls::pki_types::ServerName;
23#[cfg(not(target_os = "android"))]
24use rustls_platform_verifier::ConfigVerifierExt;
25use tokio::net::TcpStream;
26use tokio::runtime::Builder as RuntimeBuilder;
27use tokio::sync::{Notify, mpsc as async_mpsc};
28use tokio::time::{Instant, sleep_until, timeout};
29use tokio_rustls::TlsConnector;
30use tokio_rustls::client::TlsStream;
31use tokio_tungstenite::tungstenite::client::IntoClientRequest;
32use tokio_tungstenite::tungstenite::error::Error as WebSocketError;
33use tokio_tungstenite::tungstenite::http::{HeaderValue, header::SEC_WEBSOCKET_PROTOCOL};
34use tokio_tungstenite::tungstenite::protocol::{Message, WebSocketConfig};
35use tokio_tungstenite::{WebSocketStream, client_async_with_config};
36use url::Url;
37
38use crate::runtime_facade::AccountRuntimeHandle;
39use crate::{
40 AccountTransportAction, AccountTransportActionKind, AccountTransportBatch,
41 AccountTransportResult, AccountTransportResultKind, AccountTypingIndicator,
42 MAX_NOISE_FRAME_CIPHERTEXT, MAX_RELAY_FRAME_BYTES, NoiseClientHandshake, NoiseTransport,
43 RelayConnectionState, SoftchatError, SyncEngineSnapshot,
44};
45
46const REVISION_WAKE_INTERVAL: Duration = Duration::from_millis(250);
47const CONNECT_TIMEOUT: Duration = Duration::from_secs(20);
48const SEND_TIMEOUT: Duration = Duration::from_secs(15);
49const CLOSE_TIMEOUT: Duration = Duration::from_secs(2);
50const NOISE_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
51const START_TIMEOUT: Duration = Duration::from_secs(30);
52const COMMAND_TIMEOUT: Duration = Duration::from_secs(30);
53const CONTROL_CHANNEL_CAPACITY: usize = 16;
54const MAX_WAIT_MILLIS: u64 = 300_000;
55const READ_BUFFER_BYTES: usize = 64 * 1024;
56const MAX_WRITE_BUFFER_BYTES: usize = MAX_RELAY_FRAME_BYTES * 2;
57const NOISE_WEBSOCKET_SUBPROTOCOL: &str = "nostr-noise-nk.v1.25519-chachapoly-sha256";
58
59static NEXT_RUN_SEQUENCE: AtomicU64 = AtomicU64::new(1);
60
61type NativeWebSocket = WebSocketStream<TlsStream<TcpStream>>;
62
63#[derive(Clone, Copy, Debug, Eq, PartialEq)]
65#[cfg_attr(feature = "native-bindings", derive(uniffi::Enum))]
66pub enum ManagedTransportFailure {
67 NetworkUnavailable,
69 Connection,
71 Tls,
73 WebSocket,
75 Noise,
77 Protocol,
79 Account,
81}
82
83#[derive(Clone, Debug, Eq, PartialEq)]
85#[cfg_attr(feature = "native-bindings", derive(uniffi::Record))]
86pub struct ManagedTransportSnapshot {
87 pub sequence: u64,
89 pub revision: i64,
91 pub connection_state: RelayConnectionState,
93 pub idle: bool,
95 pub typing_indicators: Vec<AccountTypingIndicator>,
97 pub synchronization: Option<SyncEngineSnapshot>,
99 pub last_failure: Option<ManagedTransportFailure>,
101 pub network_available: bool,
103 pub stopped: bool,
105}
106
107impl ManagedTransportSnapshot {
108 fn initial(revision: i64) -> Self {
109 Self {
110 sequence: 0,
111 revision,
112 connection_state: RelayConnectionState::Disconnected,
113 idle: false,
114 typing_indicators: Vec::new(),
115 synchronization: None,
116 last_failure: None,
117 network_available: true,
118 stopped: false,
119 }
120 }
121
122 fn same_content(&self, other: &Self) -> bool {
123 self.revision == other.revision
124 && self.connection_state == other.connection_state
125 && self.idle == other.idle
126 && self.typing_indicators == other.typing_indicators
127 && self.synchronization == other.synchronization
128 && self.last_failure == other.last_failure
129 && self.network_available == other.network_available
130 && self.stopped == other.stopped
131 }
132}
133
134#[derive(Debug)]
135struct SnapshotState {
136 value: Mutex<ManagedTransportSnapshot>,
137 changed: Condvar,
138}
139
140impl SnapshotState {
141 fn new(revision: i64) -> Self {
142 Self {
143 value: Mutex::new(ManagedTransportSnapshot::initial(revision)),
144 changed: Condvar::new(),
145 }
146 }
147
148 fn current(&self) -> ManagedTransportSnapshot {
149 lock(&self.value).clone()
150 }
151
152 fn publish(&self, mut next: ManagedTransportSnapshot) {
153 let mut current = lock(&self.value);
154 if current.same_content(&next) {
155 return;
156 }
157 next.sequence = current.sequence.saturating_add(1);
158 *current = next;
159 self.changed.notify_all();
160 }
161
162 fn wait_after(&self, sequence: u64, wait: Duration) -> ManagedTransportSnapshot {
163 let deadline = StdInstant::now().checked_add(wait);
164 let mut current = lock(&self.value);
165 while current.sequence <= sequence && !current.stopped {
166 let Some(deadline) = deadline else {
167 break;
168 };
169 let remaining = deadline.saturating_duration_since(StdInstant::now());
170 if remaining.is_zero() {
171 break;
172 }
173 match self.changed.wait_timeout(current, remaining) {
174 Ok((guard, result)) => {
175 current = guard;
176 if result.timed_out() {
177 break;
178 }
179 }
180 Err(poisoned) => {
181 let (guard, _) = poisoned.into_inner();
182 current = guard;
183 break;
184 }
185 }
186 }
187 current.clone()
188 }
189}
190
191enum ManagerCommand {
192 ForceResync(SyncSender<Result<(), SoftchatError>>),
193}
194
195#[derive(Debug)]
196struct ManagerSignals {
197 wake_pending: AtomicBool,
198 network_available: AtomicBool,
199 network_changed: AtomicBool,
200 notify: Notify,
201}
202
203impl ManagerSignals {
204 fn new() -> Self {
205 Self {
206 wake_pending: AtomicBool::new(false),
207 network_available: AtomicBool::new(true),
208 network_changed: AtomicBool::new(false),
209 notify: Notify::new(),
210 }
211 }
212
213 fn request_wake(&self) {
214 if !self.wake_pending.swap(true, Ordering::AcqRel) {
215 self.notify.notify_one();
216 }
217 }
218
219 fn set_network_available(&self, available: bool) {
220 if self.network_available.swap(available, Ordering::AcqRel) != available {
221 self.network_changed.store(true, Ordering::Release);
222 self.notify.notify_one();
223 }
224 }
225}
226
227#[cfg_attr(feature = "native-bindings", derive(uniffi::Object))]
234pub struct ManagedAccountTransport {
235 account: Arc<AccountRuntimeHandle>,
236 commands: async_mpsc::Sender<ManagerCommand>,
237 signals: Arc<ManagerSignals>,
238 snapshots: Arc<SnapshotState>,
239 shutdown_requested: Arc<AtomicBool>,
240 shutdown_notify: Arc<Notify>,
241 worker: Mutex<Option<JoinHandle<()>>>,
242}
243
244impl fmt::Debug for ManagedAccountTransport {
245 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
246 formatter
247 .debug_struct("ManagedAccountTransport")
248 .field("account_id", &self.account.account_id())
249 .field("snapshot", &self.snapshots.current())
250 .finish_non_exhaustive()
251 }
252}
253
254#[cfg_attr(feature = "native-bindings", uniffi::export)]
255impl ManagedAccountTransport {
256 #[cfg_attr(feature = "native-bindings", uniffi::constructor)]
262 pub fn new(account: Arc<AccountRuntimeHandle>) -> Result<Self, SoftchatError> {
263 let revision = account.database_info()?.revision;
264 let snapshots = Arc::new(SnapshotState::new(revision));
265 let shutdown_requested = Arc::new(AtomicBool::new(false));
266 let shutdown_notify = Arc::new(Notify::new());
267 let signals = Arc::new(ManagerSignals::new());
268 let (commands, command_receiver) = async_mpsc::channel(CONTROL_CHANNEL_CAPACITY);
269 let (started_sender, started_receiver) = sync_channel(1);
270 let run_sequence = NEXT_RUN_SEQUENCE.fetch_add(1, Ordering::Relaxed);
271 let run_id = format!("managed-{}-{run_sequence}", std::process::id());
272
273 let thread_account = Arc::clone(&account);
274 let thread_snapshots = Arc::clone(&snapshots);
275 let thread_shutdown = Arc::clone(&shutdown_requested);
276 let thread_notify = Arc::clone(&shutdown_notify);
277 let thread_signals = Arc::clone(&signals);
278 let worker = thread::Builder::new()
279 .name(format!("softchat-transport-{run_sequence}"))
280 .spawn(move || {
281 let runtime = RuntimeBuilder::new_current_thread()
282 .enable_io()
283 .enable_time()
284 .build()
285 .map_err(|_| SoftchatError::InternalFailure);
286 let runtime = match runtime {
287 Ok(runtime) => runtime,
288 Err(error) => {
289 let _ = started_sender.send(Err(error));
290 return;
291 }
292 };
293 let initial = thread_account.start_account_transport(run_id);
294 let initial = match initial {
295 Ok(initial) => initial,
296 Err(error) => {
297 let _ = started_sender.send(Err(error));
298 return;
299 }
300 };
301 let manager = TransportManager::new(
302 thread_account,
303 command_receiver,
304 thread_signals,
305 thread_snapshots,
306 thread_shutdown,
307 thread_notify,
308 initial,
309 );
310 let _ = started_sender.send(Ok(()));
311 runtime.block_on(manager.run());
312 })
313 .map_err(|_| SoftchatError::InternalFailure)?;
314
315 match started_receiver.recv_timeout(START_TIMEOUT) {
316 Ok(Ok(())) => Ok(Self {
317 account,
318 commands,
319 signals,
320 snapshots,
321 shutdown_requested,
322 shutdown_notify,
323 worker: Mutex::new(Some(worker)),
324 }),
325 Ok(Err(error)) => {
326 let _ = worker.join();
327 Err(error)
328 }
329 Err(_) => {
330 shutdown_requested.store(true, Ordering::Release);
331 shutdown_notify.notify_one();
332 let _ = worker.join();
333 Err(SoftchatError::InternalFailure)
334 }
335 }
336 }
337
338 #[must_use]
340 pub fn snapshot(&self) -> ManagedTransportSnapshot {
341 self.snapshots.current()
342 }
343
344 #[must_use]
349 pub fn wait_for_update(
350 &self,
351 after_sequence: u64,
352 timeout_millis: u64,
353 ) -> ManagedTransportSnapshot {
354 self.snapshots.wait_after(
355 after_sequence,
356 Duration::from_millis(timeout_millis.min(MAX_WAIT_MILLIS)),
357 )
358 }
359
360 pub fn wake(&self) {
365 self.signals.request_wake();
366 }
367
368 pub fn set_network_available(&self, available: bool) {
370 self.signals.set_network_available(available);
371 }
372
373 pub fn force_resync(&self) -> Result<(), SoftchatError> {
380 let (sender, receiver) = sync_channel(1);
381 self.commands
382 .try_send(ManagerCommand::ForceResync(sender))
383 .map_err(|error| match error {
384 async_mpsc::error::TrySendError::Full(_) => SoftchatError::InternalFailure,
385 async_mpsc::error::TrySendError::Closed(_) => SoftchatError::AccountClosed,
386 })?;
387 receiver
388 .recv_timeout(COMMAND_TIMEOUT)
389 .map_err(|_| SoftchatError::InternalFailure)?
390 }
391
392 pub fn flush(&self, timeout_millis: u64) -> Result<bool, SoftchatError> {
400 let target_revision = self.account.database_info()?.revision;
401 self.wake();
402 let timeout_duration = Duration::from_millis(timeout_millis.min(MAX_WAIT_MILLIS));
403 let deadline = StdInstant::now()
404 .checked_add(timeout_duration)
405 .ok_or(SoftchatError::InternalFailure)?;
406 let mut snapshot = self.snapshots.current();
407 loop {
408 if snapshot.stopped {
409 return Ok(false);
410 }
411 if snapshot.connection_state == RelayConnectionState::Ready
412 && snapshot.idle
413 && snapshot.revision >= target_revision
414 {
415 return Ok(true);
416 }
417 let remaining = deadline.saturating_duration_since(StdInstant::now());
418 if remaining.is_zero() {
419 return Ok(false);
420 }
421 snapshot = self.snapshots.wait_after(snapshot.sequence, remaining);
422 }
423 }
424
425 pub fn shutdown(&self) {
429 self.request_shutdown();
430 }
431}
432
433impl ManagedAccountTransport {
434 fn request_shutdown(&self) {
435 if !self.shutdown_requested.swap(true, Ordering::AcqRel) {
436 self.shutdown_notify.notify_one();
440 }
441 if let Some(worker) = lock(&self.worker).take()
442 && worker.join().is_err()
443 {
444 let mut snapshot = self.snapshots.current();
445 snapshot.last_failure = Some(ManagedTransportFailure::Account);
446 snapshot.stopped = true;
447 self.snapshots.publish(snapshot);
448 }
449 }
450}
451
452impl Drop for ManagedAccountTransport {
453 fn drop(&mut self) {
454 self.request_shutdown();
455 }
456}
457
458struct PendingConnect {
459 action: AccountTransportAction,
460 ready_at: Instant,
461}
462
463struct PendingReceive {
464 action: AccountTransportAction,
465 deadline: Instant,
466}
467
468struct RelaySocket {
469 stream: NativeWebSocket,
470 noise: Option<Arc<NoiseTransport>>,
471}
472
473impl RelaySocket {
474 fn discard(self) {
475 if let Some(noise) = self.noise {
476 noise.shutdown();
477 }
478 }
479
480 async fn close(mut self) {
481 if let Some(noise) = self.noise.take() {
482 noise.shutdown();
483 }
484 let _ = timeout(CLOSE_TIMEOUT, self.stream.close(None)).await;
485 }
486
487 async fn send_text(&mut self, text: String) -> Result<(), ManagedTransportFailure> {
488 if text.len() > MAX_RELAY_FRAME_BYTES {
489 return Err(ManagedTransportFailure::Protocol);
490 }
491 if let Some(noise) = &self.noise {
492 let frames = noise
493 .encrypt(text.into_bytes())
494 .map_err(|_| ManagedTransportFailure::Noise)?;
495 for frame in frames {
496 if frame.len() > MAX_NOISE_FRAME_CIPHERTEXT {
497 return Err(ManagedTransportFailure::Noise);
498 }
499 self.stream
500 .send(Message::Binary(frame.into()))
501 .await
502 .map_err(map_write_error)?;
503 }
504 } else {
505 self.stream
506 .send(Message::Text(text.into()))
507 .await
508 .map_err(map_write_error)?;
509 }
510 self.stream.flush().await.map_err(map_write_error)
511 }
512
513 async fn receive(&mut self) -> WireRead {
514 let next = self.stream.next().await;
515 match next {
516 None => WireRead::Disconnected,
517 Some(Err(error)) => WireRead::Failure(map_read_error(&error)),
518 Some(Ok(Message::Text(text))) => {
519 if self.noise.is_some() || text.len() > MAX_RELAY_FRAME_BYTES {
520 WireRead::Failure(ManagedTransportFailure::Protocol)
521 } else {
522 WireRead::Text(text.to_string())
523 }
524 }
525 Some(Ok(Message::Binary(frame))) => {
526 let Some(noise) = &self.noise else {
527 return WireRead::Failure(ManagedTransportFailure::Protocol);
528 };
529 if frame.len() > MAX_NOISE_FRAME_CIPHERTEXT {
530 return WireRead::Failure(ManagedTransportFailure::Noise);
531 }
532 match noise.decrypt(frame.to_vec()) {
533 Ok(None) => WireRead::Continue,
534 Ok(Some(plaintext)) => {
535 if plaintext.len() > MAX_RELAY_FRAME_BYTES {
536 return WireRead::Failure(ManagedTransportFailure::Protocol);
537 }
538 match String::from_utf8(plaintext) {
539 Ok(text) => WireRead::Text(text),
540 Err(_) => WireRead::Failure(ManagedTransportFailure::Noise),
541 }
542 }
543 Err(_) => WireRead::Failure(ManagedTransportFailure::Noise),
544 }
545 }
546 Some(Ok(Message::Ping(_))) => match self.stream.flush().await {
547 Ok(()) => WireRead::Continue,
548 Err(error) => WireRead::Failure(map_read_error(&error)),
549 },
550 Some(Ok(Message::Pong(_))) => WireRead::Continue,
551 Some(Ok(Message::Close(_))) => WireRead::Disconnected,
552 Some(Ok(Message::Frame(_))) => WireRead::Failure(ManagedTransportFailure::Protocol),
553 }
554 }
555}
556
557enum WireRead {
558 Text(String),
559 Continue,
560 Disconnected,
561 Failure(ManagedTransportFailure),
562}
563
564struct TransportManager {
565 account: Arc<AccountRuntimeHandle>,
566 commands: async_mpsc::Receiver<ManagerCommand>,
567 signals: Arc<ManagerSignals>,
568 snapshots: Arc<SnapshotState>,
569 shutdown_requested: Arc<AtomicBool>,
570 shutdown_notify: Arc<Notify>,
571 socket: Option<RelaySocket>,
572 pending_connect: Option<PendingConnect>,
573 pending_receive: Option<PendingReceive>,
574 sends: VecDeque<AccountTransportAction>,
575 closes: VecDeque<AccountTransportAction>,
576 generation: u64,
577 batch: AccountTransportBatch,
578 last_failure: Option<ManagedTransportFailure>,
579 network_available: bool,
580 next_revision_wake: Instant,
581}
582
583impl TransportManager {
584 fn new(
585 account: Arc<AccountRuntimeHandle>,
586 commands: async_mpsc::Receiver<ManagerCommand>,
587 signals: Arc<ManagerSignals>,
588 snapshots: Arc<SnapshotState>,
589 shutdown_requested: Arc<AtomicBool>,
590 shutdown_notify: Arc<Notify>,
591 initial: AccountTransportBatch,
592 ) -> Self {
593 let mut manager = Self {
594 account,
595 commands,
596 signals,
597 snapshots,
598 shutdown_requested,
599 shutdown_notify,
600 socket: None,
601 pending_connect: None,
602 pending_receive: None,
603 sends: VecDeque::new(),
604 closes: VecDeque::new(),
605 generation: 0,
606 batch: initial.clone(),
607 last_failure: None,
608 network_available: true,
609 next_revision_wake: Instant::now() + REVISION_WAKE_INTERVAL,
610 };
611 manager.install_batch(initial);
612 manager
613 }
614
615 async fn run(mut self) {
616 let result = self.run_loop().await;
617 if result.is_err() && self.last_failure.is_none() {
618 self.last_failure = Some(ManagedTransportFailure::Account);
619 }
620 if let Some(socket) = self.socket.take() {
621 socket.close().await;
622 }
623 if let Ok(batch) = self.account.cancel_account_transport() {
624 self.install_batch(batch);
625 }
626 let mut snapshot = self.snapshot();
627 snapshot.stopped = true;
628 self.snapshots.publish(snapshot);
629 }
630
631 async fn run_loop(&mut self) -> Result<(), SoftchatError> {
632 loop {
633 if self.shutdown_requested.load(Ordering::Acquire) {
634 return Ok(());
635 }
636 while let Ok(command) = self.commands.try_recv() {
637 self.handle_command(command);
638 }
639 self.handle_pending_signals()?;
640
641 if let Some(action) = self.sends.pop_front() {
642 if !self.execute_send(action).await? {
643 return Ok(());
644 }
645 continue;
646 }
647 if let Some(action) = self.closes.pop_front() {
648 if let Some(socket) = self.socket.take() {
649 socket.close().await;
650 }
651 self.apply_result(
652 &action,
653 AccountTransportResultKind::Disconnected,
654 String::new(),
655 None,
656 )?;
657 continue;
658 }
659 if self
660 .pending_connect
661 .as_ref()
662 .is_some_and(|pending| pending.ready_at <= Instant::now())
663 {
664 let pending = self
665 .pending_connect
666 .take()
667 .ok_or(SoftchatError::InternalFailure)?;
668 self.execute_connect(pending.action).await?;
669 continue;
670 }
671
672 let event = self.next_event().await;
673 match event {
674 LoopEvent::Shutdown => return Ok(()),
675 LoopEvent::Command(Some(command)) => {
676 self.handle_command(command);
677 }
678 LoopEvent::Command(None) => return Ok(()),
679 LoopEvent::Signal => self.handle_pending_signals()?,
680 LoopEvent::RevisionWake => self.wake_transport()?,
681 LoopEvent::ConnectReady => {}
682 LoopEvent::ReceiveTimedOut => {
683 let pending = self
684 .pending_receive
685 .take()
686 .ok_or(SoftchatError::InternalFailure)?;
687 self.apply_result(
688 &pending.action,
689 AccountTransportResultKind::TimedOut,
690 String::new(),
691 None,
692 )?;
693 }
694 LoopEvent::Wire(read) => self.handle_wire_read(read)?,
695 }
696 }
697 }
698
699 async fn next_event(&mut self) -> LoopEvent {
700 let revision_wake = self.next_revision_wake;
701 if let Some(pending) = &self.pending_connect {
702 let connect_ready = pending.ready_at;
703 return tokio::select! {
704 _ = self.shutdown_notify.notified() => LoopEvent::Shutdown,
705 command = self.commands.recv() => LoopEvent::Command(command),
706 _ = self.signals.notify.notified() => LoopEvent::Signal,
707 _ = sleep_until(connect_ready) => LoopEvent::ConnectReady,
708 _ = sleep_until(revision_wake) => LoopEvent::RevisionWake,
709 };
710 }
711
712 let Some(pending_receive) = &self.pending_receive else {
713 return tokio::select! {
714 _ = self.shutdown_notify.notified() => LoopEvent::Shutdown,
715 command = self.commands.recv() => LoopEvent::Command(command),
716 _ = self.signals.notify.notified() => LoopEvent::Signal,
717 _ = sleep_until(revision_wake) => LoopEvent::RevisionWake,
718 };
719 };
720 let receive_deadline = pending_receive.deadline;
721 let Some(socket) = self.socket.as_mut() else {
722 return LoopEvent::Wire(WireRead::Disconnected);
723 };
724 tokio::select! {
725 _ = self.shutdown_notify.notified() => LoopEvent::Shutdown,
726 command = self.commands.recv() => LoopEvent::Command(command),
727 _ = self.signals.notify.notified() => LoopEvent::Signal,
728 _ = sleep_until(receive_deadline) => LoopEvent::ReceiveTimedOut,
729 _ = sleep_until(revision_wake) => LoopEvent::RevisionWake,
730 read = socket.receive() => LoopEvent::Wire(read),
731 }
732 }
733
734 fn handle_command(&mut self, command: ManagerCommand) {
735 match command {
736 ManagerCommand::ForceResync(response) => {
737 let result = self.account.reset_account_sync(now_seconds()).map(|batch| {
738 self.reset_io();
739 self.install_batch(batch);
740 });
741 let _ = response.send(result);
742 }
743 }
744 }
745
746 fn handle_pending_signals(&mut self) -> Result<(), SoftchatError> {
747 let network_changed = self.signals.network_changed.swap(false, Ordering::AcqRel);
748 if network_changed {
749 let available = self.signals.network_available.load(Ordering::Acquire);
750 if self.network_available != available {
751 self.network_available = available;
752 self.publish_snapshot();
753 }
754 if !available {
755 self.signals.wake_pending.store(false, Ordering::Release);
756 return self.fail_for_network_unavailable();
757 }
758 }
759 let wake_pending = self.signals.wake_pending.swap(false, Ordering::AcqRel);
760 if network_changed || wake_pending {
761 self.wake_transport()?;
762 }
763 Ok(())
764 }
765
766 fn fail_for_network_unavailable(&mut self) -> Result<(), SoftchatError> {
767 if let Some(socket) = self.socket.take() {
768 socket.discard();
769 }
770 let action = self
771 .pending_connect
772 .take()
773 .map(|pending| pending.action)
774 .or_else(|| self.sends.pop_front())
775 .or_else(|| self.pending_receive.take().map(|pending| pending.action));
776 if let Some(action) = action {
777 self.apply_result(
778 &action,
779 AccountTransportResultKind::NetworkUnavailable,
780 String::new(),
781 Some(ManagedTransportFailure::NetworkUnavailable),
782 )?;
783 }
784 Ok(())
785 }
786
787 fn wake_transport(&mut self) -> Result<(), SoftchatError> {
788 self.next_revision_wake = Instant::now() + REVISION_WAKE_INTERVAL;
789 let batch = self.account.wake_account_transport(now_seconds())?;
790 self.install_batch(batch);
791 Ok(())
792 }
793
794 async fn execute_connect(
795 &mut self,
796 action: AccountTransportAction,
797 ) -> Result<(), SoftchatError> {
798 if !self.network_available {
799 return self.apply_result(
800 &action,
801 AccountTransportResultKind::NetworkUnavailable,
802 String::new(),
803 Some(ManagedTransportFailure::NetworkUnavailable),
804 );
805 }
806 if let Some(socket) = self.socket.take() {
807 socket.discard();
808 }
809 let connect = open_socket(&action);
810 let result = tokio::select! {
811 _ = self.shutdown_notify.notified() => return Ok(()),
812 result = timeout(CONNECT_TIMEOUT, connect) => result,
813 };
814 match result {
815 Ok(Ok(socket)) => {
816 self.socket = Some(socket);
817 self.apply_result(
818 &action,
819 AccountTransportResultKind::Connected,
820 String::new(),
821 None,
822 )
823 }
824 Ok(Err(failure)) => self.apply_result(
825 &action,
826 failure.result_kind(),
827 String::new(),
828 Some(failure.snapshot_kind()),
829 ),
830 Err(_) => self.apply_result(
831 &action,
832 AccountTransportResultKind::Disconnected,
833 String::new(),
834 Some(ManagedTransportFailure::Connection),
835 ),
836 }
837 }
838
839 async fn execute_send(
840 &mut self,
841 action: AccountTransportAction,
842 ) -> Result<bool, SoftchatError> {
843 if !self.network_available {
844 self.apply_result(
845 &action,
846 AccountTransportResultKind::NetworkUnavailable,
847 String::new(),
848 Some(ManagedTransportFailure::NetworkUnavailable),
849 )?;
850 return Ok(true);
851 }
852 let Some(socket) = self.socket.as_mut() else {
853 self.apply_result(
854 &action,
855 AccountTransportResultKind::Disconnected,
856 String::new(),
857 Some(ManagedTransportFailure::Connection),
858 )?;
859 return Ok(true);
860 };
861 let send = socket.send_text(action.text_frame.clone());
862 let result = await_send_or_control(
863 send,
864 &mut self.commands,
865 &self.signals,
866 &self.shutdown_notify,
867 )
868 .await;
869 match result {
870 SendWait::Shutdown => return Ok(false),
871 SendWait::NetworkUnavailable => {
872 if let Some(socket) = self.socket.take() {
873 socket.discard();
874 }
875 self.sends.push_front(action);
876 self.handle_pending_signals()?;
877 }
878 SendWait::Command(command) => {
879 if let Some(socket) = self.socket.take() {
880 socket.discard();
881 }
882 self.sends.push_front(action);
883 return match command {
884 Some(command) => {
885 self.handle_command(command);
886 Ok(true)
887 }
888 None => Ok(false),
889 };
890 }
891 SendWait::Written => self.apply_result(
892 &action,
893 AccountTransportResultKind::FrameWritten,
894 String::new(),
895 None,
896 )?,
897 SendWait::Failed(failure) => {
898 if let Some(socket) = self.socket.take() {
899 socket.discard();
900 }
901 self.apply_result(
902 &action,
903 failure_result_kind(failure),
904 String::new(),
905 Some(failure),
906 )?;
907 }
908 SendWait::TimedOut => {
909 if let Some(socket) = self.socket.take() {
910 socket.discard();
911 }
912 self.apply_result(
913 &action,
914 AccountTransportResultKind::Disconnected,
915 String::new(),
916 Some(ManagedTransportFailure::Connection),
917 )?;
918 }
919 }
920 Ok(true)
921 }
922
923 fn handle_wire_read(&mut self, read: WireRead) -> Result<(), SoftchatError> {
924 match read {
925 WireRead::Continue => Ok(()),
926 WireRead::Text(text) => {
927 let pending = self
928 .pending_receive
929 .take()
930 .ok_or(SoftchatError::InternalFailure)?;
931 self.apply_result(
932 &pending.action,
933 AccountTransportResultKind::FrameReceived,
934 text,
935 None,
936 )
937 }
938 WireRead::Disconnected => {
939 if let Some(socket) = self.socket.take() {
940 socket.discard();
941 }
942 let pending = self
943 .pending_receive
944 .take()
945 .ok_or(SoftchatError::InternalFailure)?;
946 self.apply_result(
947 &pending.action,
948 AccountTransportResultKind::Disconnected,
949 String::new(),
950 Some(ManagedTransportFailure::Connection),
951 )
952 }
953 WireRead::Failure(failure) => {
954 if let Some(socket) = self.socket.take() {
955 socket.discard();
956 }
957 let pending = self
958 .pending_receive
959 .take()
960 .ok_or(SoftchatError::InternalFailure)?;
961 self.apply_result(
962 &pending.action,
963 failure_result_kind(failure),
964 String::new(),
965 Some(failure),
966 )
967 }
968 }
969 }
970
971 fn apply_result(
972 &mut self,
973 action: &AccountTransportAction,
974 kind: AccountTransportResultKind,
975 text_frame: String,
976 failure: Option<ManagedTransportFailure>,
977 ) -> Result<(), SoftchatError> {
978 if let Some(failure) = failure {
979 self.last_failure = Some(failure);
980 self.account
981 .logs
982 .emit(crate::account_diagnostics::managed_failure(failure));
983 }
984 let batch = self.account.apply_account_transport_result(
985 AccountTransportResult {
986 run_id: action.run_id.clone(),
987 action_id: action.action_id.clone(),
988 generation: action.generation,
989 kind,
990 text_frame,
991 },
992 now_seconds(),
993 )?;
994 self.install_batch(batch);
995 Ok(())
996 }
997
998 fn install_batch(&mut self, batch: AccountTransportBatch) {
999 self.batch = batch.clone();
1000 if batch.connection_state == RelayConnectionState::Ready {
1001 self.last_failure = None;
1002 }
1003 for action in batch.actions {
1004 if action.generation < self.generation {
1005 continue;
1006 }
1007 if action.generation > self.generation {
1008 self.generation = action.generation;
1009 self.reset_io();
1010 }
1011 match action.kind {
1012 AccountTransportActionKind::Connect => {
1013 self.pending_connect = Some(PendingConnect {
1014 ready_at: Instant::now() + Duration::from_millis(action.delay_ms),
1015 action,
1016 });
1017 }
1018 AccountTransportActionKind::SendText => self.sends.push_back(action),
1019 AccountTransportActionKind::ReceiveText => {
1020 if self.pending_receive.is_none() {
1021 self.pending_receive = Some(PendingReceive {
1022 deadline: Instant::now() + Duration::from_millis(action.delay_ms),
1023 action,
1024 });
1025 }
1026 }
1027 AccountTransportActionKind::Close => self.closes.push_back(action),
1028 }
1029 }
1030 self.publish_snapshot();
1031 }
1032
1033 fn reset_io(&mut self) {
1034 if let Some(socket) = self.socket.take() {
1035 socket.discard();
1036 }
1037 self.pending_connect = None;
1038 self.pending_receive = None;
1039 self.sends.clear();
1040 self.closes.clear();
1041 }
1042
1043 fn snapshot(&self) -> ManagedTransportSnapshot {
1044 ManagedTransportSnapshot {
1045 sequence: 0,
1046 revision: self.batch.revision,
1047 connection_state: self.batch.connection_state,
1048 idle: self.batch.idle,
1049 typing_indicators: self.batch.typing_indicators.clone(),
1050 synchronization: self.batch.sync.clone(),
1051 last_failure: self.last_failure,
1052 network_available: self.network_available,
1053 stopped: false,
1054 }
1055 }
1056
1057 fn publish_snapshot(&self) {
1058 self.snapshots.publish(self.snapshot());
1059 }
1060}
1061
1062enum LoopEvent {
1063 Shutdown,
1064 Command(Option<ManagerCommand>),
1065 Signal,
1066 RevisionWake,
1067 ConnectReady,
1068 ReceiveTimedOut,
1069 Wire(WireRead),
1070}
1071
1072enum SendWait {
1073 Shutdown,
1074 Command(Option<ManagerCommand>),
1075 NetworkUnavailable,
1076 Written,
1077 Failed(ManagedTransportFailure),
1078 TimedOut,
1079}
1080
1081async fn await_send_or_control<F>(
1082 send: F,
1083 commands: &mut async_mpsc::Receiver<ManagerCommand>,
1084 signals: &ManagerSignals,
1085 shutdown_notify: &Notify,
1086) -> SendWait
1087where
1088 F: Future<Output = Result<(), ManagedTransportFailure>>,
1089{
1090 let send = timeout(SEND_TIMEOUT, send);
1091 tokio::pin!(send);
1092 loop {
1093 tokio::select! {
1094 _ = shutdown_notify.notified() => return SendWait::Shutdown,
1095 command = commands.recv() => return SendWait::Command(command),
1096 _ = signals.notify.notified() => {
1097 if signals.network_changed.load(Ordering::Acquire)
1098 && !signals.network_available.load(Ordering::Acquire)
1099 {
1100 return SendWait::NetworkUnavailable;
1101 }
1102 }
1103 result = &mut send => return match result {
1104 Ok(Ok(())) => SendWait::Written,
1105 Ok(Err(failure)) => SendWait::Failed(failure),
1106 Err(_) => SendWait::TimedOut,
1107 },
1108 }
1109 }
1110}
1111
1112#[derive(Clone, Copy, Debug)]
1113enum ConnectFailure {
1114 Connection,
1115 Tls,
1116 WebSocket,
1117 Noise,
1118 Protocol,
1119 Authentication,
1120 RateLimited,
1121}
1122
1123impl ConnectFailure {
1124 const fn snapshot_kind(self) -> ManagedTransportFailure {
1125 match self {
1126 Self::Connection => ManagedTransportFailure::Connection,
1127 Self::Tls | Self::Authentication => ManagedTransportFailure::Tls,
1128 Self::WebSocket | Self::RateLimited => ManagedTransportFailure::WebSocket,
1129 Self::Noise => ManagedTransportFailure::Noise,
1130 Self::Protocol => ManagedTransportFailure::Protocol,
1131 }
1132 }
1133
1134 const fn result_kind(self) -> AccountTransportResultKind {
1135 match self {
1136 Self::Authentication | Self::Tls => AccountTransportResultKind::AuthenticationFailed,
1137 Self::RateLimited => AccountTransportResultKind::RateLimited,
1138 Self::Noise | Self::Protocol | Self::WebSocket => {
1139 AccountTransportResultKind::ProtocolFailed
1140 }
1141 Self::Connection => AccountTransportResultKind::Disconnected,
1142 }
1143 }
1144}
1145
1146async fn open_socket(action: &AccountTransportAction) -> Result<RelaySocket, ConnectFailure> {
1147 let url = Url::parse(&action.network_url).map_err(|_| ConnectFailure::Protocol)?;
1148 if url.scheme() != "wss" || url.username() != "" || url.password().is_some() {
1149 return Err(ConnectFailure::Protocol);
1150 }
1151 let host = url.host_str().ok_or(ConnectFailure::Protocol)?.to_owned();
1152 let port = url
1153 .port_or_known_default()
1154 .ok_or(ConnectFailure::Protocol)?;
1155 let request = websocket_request(action)?;
1156 let tcp = TcpStream::connect((host.as_str(), port))
1157 .await
1158 .map_err(|_| ConnectFailure::Connection)?;
1159 let server_name = ServerName::try_from(host).map_err(|_| ConnectFailure::Protocol)?;
1160 let tls_config = managed_tls_config()?;
1161 let tls = TlsConnector::from(Arc::new(tls_config))
1162 .connect(server_name, tcp)
1163 .await
1164 .map_err(|_| ConnectFailure::Tls)?;
1165 let config = WebSocketConfig::default()
1166 .read_buffer_size(READ_BUFFER_BYTES)
1167 .write_buffer_size(0)
1168 .max_write_buffer_size(MAX_WRITE_BUFFER_BYTES)
1169 .max_message_size(Some(MAX_RELAY_FRAME_BYTES))
1170 .max_frame_size(Some(MAX_RELAY_FRAME_BYTES));
1171 let (mut stream, response) = client_async_with_config(request, tls, Some(config))
1172 .await
1173 .map_err(|error| map_handshake_error(&error))?;
1174
1175 if !action.noise_remote_static_key.is_empty()
1176 && response
1177 .headers()
1178 .get(SEC_WEBSOCKET_PROTOCOL)
1179 .and_then(|value| value.to_str().ok())
1180 != Some(NOISE_WEBSOCKET_SUBPROTOCOL)
1181 {
1182 return Err(ConnectFailure::Noise);
1183 }
1184
1185 let noise = if action.noise_remote_static_key.is_empty() {
1186 None
1187 } else {
1188 let handshake = NoiseClientHandshake::new(action.noise_remote_static_key.clone(), true)
1189 .map_err(|_| ConnectFailure::Noise)?;
1190 stream
1191 .send(Message::Binary(handshake.message_one().into()))
1192 .await
1193 .map_err(|_| ConnectFailure::Noise)?;
1194 let message_two = timeout(
1195 NOISE_HANDSHAKE_TIMEOUT,
1196 receive_noise_message_two(&mut stream),
1197 )
1198 .await
1199 .map_err(|_| ConnectFailure::Noise)??;
1200 Some(
1201 handshake
1202 .complete(message_two)
1203 .map_err(|_| ConnectFailure::Noise)?,
1204 )
1205 };
1206 Ok(RelaySocket { stream, noise })
1207}
1208
1209#[cfg(not(target_os = "android"))]
1210fn managed_tls_config() -> Result<ClientConfig, ConnectFailure> {
1211 ClientConfig::with_platform_verifier().map_err(|_| ConnectFailure::Tls)
1212}
1213
1214#[cfg(target_os = "android")]
1215fn managed_tls_config() -> Result<ClientConfig, ConnectFailure> {
1216 let roots = RootCertStore::from_iter(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
1222 Ok(ClientConfig::builder()
1223 .with_root_certificates(roots)
1224 .with_no_client_auth())
1225}
1226
1227fn websocket_request(
1228 action: &AccountTransportAction,
1229) -> Result<tokio_tungstenite::tungstenite::http::Request<()>, ConnectFailure> {
1230 let mut request = action
1231 .network_url
1232 .as_str()
1233 .into_client_request()
1234 .map_err(|_| ConnectFailure::Protocol)?;
1235 if !action.noise_remote_static_key.is_empty() {
1236 request.headers_mut().insert(
1237 SEC_WEBSOCKET_PROTOCOL,
1238 HeaderValue::from_static(NOISE_WEBSOCKET_SUBPROTOCOL),
1239 );
1240 }
1241 Ok(request)
1242}
1243
1244async fn receive_noise_message_two(
1245 stream: &mut NativeWebSocket,
1246) -> Result<Vec<u8>, ConnectFailure> {
1247 loop {
1248 match stream.next().await {
1249 Some(Ok(Message::Binary(message))) => {
1250 if message.len() > MAX_NOISE_FRAME_CIPHERTEXT {
1251 return Err(ConnectFailure::Noise);
1252 }
1253 return Ok(message.to_vec());
1254 }
1255 Some(Ok(Message::Ping(_))) => {
1256 stream.flush().await.map_err(|_| ConnectFailure::Noise)?;
1257 }
1258 Some(Ok(Message::Pong(_))) => {}
1259 Some(Ok(Message::Close(_))) | None => return Err(ConnectFailure::Connection),
1260 Some(Ok(Message::Text(_) | Message::Frame(_))) => {
1261 return Err(ConnectFailure::Noise);
1262 }
1263 Some(Err(error)) => return Err(map_handshake_error(&error)),
1264 }
1265 }
1266}
1267
1268fn map_handshake_error(error: &WebSocketError) -> ConnectFailure {
1269 match error {
1270 WebSocketError::Http(response) if response.status().as_u16() == 429 => {
1271 ConnectFailure::RateLimited
1272 }
1273 WebSocketError::Http(response) if matches!(response.status().as_u16(), 401 | 403) => {
1274 ConnectFailure::Authentication
1275 }
1276 WebSocketError::Io(_)
1277 | WebSocketError::ConnectionClosed
1278 | WebSocketError::AlreadyClosed => ConnectFailure::Connection,
1279 WebSocketError::Tls(_) => ConnectFailure::Tls,
1280 WebSocketError::Http(_)
1281 | WebSocketError::HttpFormat(_)
1282 | WebSocketError::Protocol(_)
1283 | WebSocketError::Capacity(_)
1284 | WebSocketError::WriteBufferFull(_)
1285 | WebSocketError::Utf8(_)
1286 | WebSocketError::AttackAttempt
1287 | WebSocketError::Url(_) => ConnectFailure::WebSocket,
1288 }
1289}
1290
1291fn map_write_error(error: WebSocketError) -> ManagedTransportFailure {
1292 match error {
1293 WebSocketError::Io(_)
1294 | WebSocketError::ConnectionClosed
1295 | WebSocketError::AlreadyClosed => ManagedTransportFailure::Connection,
1296 WebSocketError::Tls(_) => ManagedTransportFailure::Tls,
1297 _ => ManagedTransportFailure::WebSocket,
1298 }
1299}
1300
1301fn map_read_error(error: &WebSocketError) -> ManagedTransportFailure {
1302 match error {
1303 WebSocketError::Io(_)
1304 | WebSocketError::ConnectionClosed
1305 | WebSocketError::AlreadyClosed => ManagedTransportFailure::Connection,
1306 WebSocketError::Tls(_) => ManagedTransportFailure::Tls,
1307 _ => ManagedTransportFailure::WebSocket,
1308 }
1309}
1310
1311const fn failure_result_kind(failure: ManagedTransportFailure) -> AccountTransportResultKind {
1312 match failure {
1313 ManagedTransportFailure::NetworkUnavailable => {
1314 AccountTransportResultKind::NetworkUnavailable
1315 }
1316 ManagedTransportFailure::Connection => AccountTransportResultKind::Disconnected,
1317 ManagedTransportFailure::Tls => AccountTransportResultKind::AuthenticationFailed,
1318 ManagedTransportFailure::WebSocket
1319 | ManagedTransportFailure::Noise
1320 | ManagedTransportFailure::Protocol
1321 | ManagedTransportFailure::Account => AccountTransportResultKind::ProtocolFailed,
1322 }
1323}
1324
1325fn now_seconds() -> i64 {
1326 SystemTime::now()
1327 .duration_since(UNIX_EPOCH)
1328 .ok()
1329 .and_then(|duration| i64::try_from(duration.as_secs()).ok())
1330 .unwrap_or(0)
1331}
1332
1333fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
1334 match mutex.lock() {
1335 Ok(guard) => guard,
1336 Err(poisoned) => poisoned.into_inner(),
1337 }
1338}
1339
1340#[cfg(test)]
1341mod tests {
1342 use super::*;
1343
1344 fn temporary_directory(label: &str) -> Result<std::path::PathBuf, Box<dyn std::error::Error>> {
1345 let nonce = SystemTime::now().duration_since(UNIX_EPOCH)?.as_nanos();
1346 let path =
1347 std::env::temp_dir().join(format!("softchat-{label}-{}-{nonce}", std::process::id()));
1348 std::fs::create_dir_all(&path)?;
1349 Ok(path)
1350 }
1351
1352 #[test]
1353 fn snapshot_wait_is_bounded_and_sequence_based() {
1354 let state = SnapshotState::new(7);
1355 let initial = state.current();
1356 let timed_out = state.wait_after(initial.sequence, Duration::from_millis(1));
1357 assert_eq!(timed_out, initial);
1358
1359 let mut updated = initial.clone();
1360 updated.network_available = false;
1361 state.publish(updated);
1362 let observed = state.wait_after(initial.sequence, Duration::ZERO);
1363 assert!(observed.sequence > initial.sequence);
1364 assert!(!observed.network_available);
1365 }
1366
1367 #[test]
1368 fn managed_feature_rejects_insecure_or_credentialed_urls_before_io() {
1369 for value in [
1370 "ws://relay.example",
1371 "https://relay.example",
1372 "wss://[email protected]",
1373 "not a url",
1374 ] {
1375 let parsed = Url::parse(value);
1376 let valid = parsed.as_ref().is_ok_and(|url| {
1377 url.scheme() == "wss" && url.username().is_empty() && url.password().is_none()
1378 });
1379 assert!(!valid);
1380 }
1381 }
1382
1383 #[test]
1384 fn noise_websocket_request_advertises_the_deployed_subprotocol() -> Result<(), ConnectFailure> {
1385 let mut action = AccountTransportAction {
1386 run_id: "test-run".to_owned(),
1387 action_id: "test-action".to_owned(),
1388 generation: 1,
1389 kind: AccountTransportActionKind::Connect,
1390 configured_url: "wss+noise://relay.example".to_owned(),
1391 network_url: "wss://relay.example/".to_owned(),
1392 noise_remote_static_key: vec![7_u8; 32],
1393 text_frame: String::new(),
1394 delay_ms: 0,
1395 };
1396 let noise = websocket_request(&action)?;
1397 assert_eq!(
1398 noise
1399 .headers()
1400 .get(SEC_WEBSOCKET_PROTOCOL)
1401 .and_then(|value| value.to_str().ok()),
1402 Some(NOISE_WEBSOCKET_SUBPROTOCOL)
1403 );
1404
1405 action.noise_remote_static_key.clear();
1406 let plain = websocket_request(&action)?;
1407 assert!(plain.headers().get(SEC_WEBSOCKET_PROTOCOL).is_none());
1408 Ok(())
1409 }
1410
1411 #[test]
1412 fn force_resync_control_preempts_a_stalled_send() -> Result<(), Box<dyn std::error::Error>> {
1413 let (sender, mut commands) = async_mpsc::channel(CONTROL_CHANNEL_CAPACITY);
1414 let (response_sender, _response_receiver) = sync_channel(1);
1415 sender.try_send(ManagerCommand::ForceResync(response_sender))?;
1416 let signals = ManagerSignals::new();
1417 let shutdown = Notify::new();
1418 let runtime = RuntimeBuilder::new_current_thread().enable_time().build()?;
1419
1420 let selected = runtime.block_on(await_send_or_control(
1421 std::future::pending::<Result<(), ManagedTransportFailure>>(),
1422 &mut commands,
1423 &signals,
1424 &shutdown,
1425 ));
1426
1427 assert!(matches!(
1428 selected,
1429 SendWait::Command(Some(ManagerCommand::ForceResync(_)))
1430 ));
1431 Ok(())
1432 }
1433
1434 #[test]
1435 fn wake_and_reachability_hints_are_coalesced_to_latest_state() {
1436 let signals = ManagerSignals::new();
1437 for _ in 0..10_000 {
1438 signals.request_wake();
1439 }
1440 assert!(signals.wake_pending.swap(false, Ordering::AcqRel));
1441 assert!(!signals.wake_pending.swap(false, Ordering::AcqRel));
1442
1443 for available in [false, true].into_iter().cycle().take(10_000) {
1444 signals.set_network_available(available);
1445 }
1446 assert!(signals.network_changed.swap(false, Ordering::AcqRel));
1447 assert!(signals.network_available.load(Ordering::Acquire));
1448 assert!(!signals.network_changed.swap(false, Ordering::AcqRel));
1449 }
1450
1451 #[test]
1452 fn control_queue_is_bounded() {
1453 let (sender, _commands) = async_mpsc::channel(CONTROL_CHANNEL_CAPACITY);
1454 for _ in 0..CONTROL_CHANNEL_CAPACITY {
1455 let (response_sender, _response_receiver) = sync_channel(1);
1456 assert!(
1457 sender
1458 .try_send(ManagerCommand::ForceResync(response_sender))
1459 .is_ok()
1460 );
1461 }
1462 let (response_sender, _response_receiver) = sync_channel(1);
1463 assert!(matches!(
1464 sender.try_send(ManagerCommand::ForceResync(response_sender)),
1465 Err(async_mpsc::error::TrySendError::Full(_))
1466 ));
1467 }
1468
1469 #[test]
1470 fn wake_does_not_preempt_a_valid_send() -> Result<(), Box<dyn std::error::Error>> {
1471 let (_sender, mut commands) = async_mpsc::channel(CONTROL_CHANNEL_CAPACITY);
1472 let signals = ManagerSignals::new();
1473 signals.request_wake();
1474 let shutdown = Notify::new();
1475 let runtime = RuntimeBuilder::new_current_thread().enable_time().build()?;
1476
1477 let selected = runtime.block_on(await_send_or_control(
1478 async { Ok(()) },
1479 &mut commands,
1480 &signals,
1481 &shutdown,
1482 ));
1483
1484 assert!(matches!(selected, SendWait::Written));
1485 assert!(signals.wake_pending.load(Ordering::Acquire));
1486 Ok(())
1487 }
1488
1489 #[test]
1490 fn network_down_preempts_a_stalled_send() -> Result<(), Box<dyn std::error::Error>> {
1491 let (_sender, mut commands) = async_mpsc::channel(CONTROL_CHANNEL_CAPACITY);
1492 let signals = ManagerSignals::new();
1493 signals.set_network_available(false);
1494 let shutdown = Notify::new();
1495 let runtime = RuntimeBuilder::new_current_thread().enable_time().build()?;
1496
1497 let selected = runtime.block_on(await_send_or_control(
1498 std::future::pending::<Result<(), ManagedTransportFailure>>(),
1499 &mut commands,
1500 &signals,
1501 &shutdown,
1502 ));
1503
1504 assert!(matches!(selected, SendWait::NetworkUnavailable));
1505 Ok(())
1506 }
1507
1508 #[test]
1509 fn failed_force_resync_preserves_existing_io_state() -> Result<(), Box<dyn std::error::Error>> {
1510 let directory = temporary_directory("managed-resync-order")?;
1511 let account = Arc::new(AccountRuntimeHandle::open(
1512 vec![1_u8; 32],
1513 directory.display().to_string(),
1514 )?);
1515 let initial = account.start_account_transport("resync-order".to_owned())?;
1516 let (_command_sender, command_receiver) = async_mpsc::channel(CONTROL_CHANNEL_CAPACITY);
1517 let signals = Arc::new(ManagerSignals::new());
1518 let snapshots = Arc::new(SnapshotState::new(initial.revision));
1519 let mut manager = TransportManager::new(
1520 Arc::clone(&account),
1521 command_receiver,
1522 signals,
1523 snapshots,
1524 Arc::new(AtomicBool::new(false)),
1525 Arc::new(Notify::new()),
1526 initial,
1527 );
1528 assert!(manager.pending_connect.is_some());
1529 account.close_account();
1530 let (response_sender, response_receiver) = sync_channel(1);
1531
1532 manager.handle_command(ManagerCommand::ForceResync(response_sender));
1533 assert_eq!(response_receiver.recv()?, Err(SoftchatError::AccountClosed),);
1534 assert!(manager.pending_connect.is_some());
1535
1536 drop(manager);
1537 drop(account);
1538 std::fs::remove_dir_all(directory)?;
1539 Ok(())
1540 }
1541
1542 #[test]
1543 fn shutdown_interrupts_a_stalled_tls_handshake_and_joins_the_worker()
1544 -> Result<(), Box<dyn std::error::Error>> {
1545 let listener = std::net::TcpListener::bind(("127.0.0.1", 0))?;
1546 listener.set_nonblocking(true)?;
1547 let relay_url = format!("wss://127.0.0.1:{}/", listener.local_addr()?.port());
1548 let directory = temporary_directory("managed-shutdown")?;
1549 let account = Arc::new(AccountRuntimeHandle::open(
1550 vec![1_u8; 32],
1551 directory.display().to_string(),
1552 )?);
1553 account.add_account_relay(relay_url.clone(), 1_700_000_000)?;
1554 account.set_account_active_relay(relay_url, 1_700_000_001)?;
1555
1556 let transport = ManagedAccountTransport::new(Arc::clone(&account))?;
1557 let accept_deadline = StdInstant::now() + Duration::from_secs(5);
1558 let accepted = loop {
1559 match listener.accept() {
1560 Ok((stream, _)) => break stream,
1561 Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
1562 if StdInstant::now() >= accept_deadline {
1563 return Err("managed transport did not open the test socket".into());
1564 }
1565 thread::sleep(Duration::from_millis(10));
1566 }
1567 Err(error) => return Err(error.into()),
1568 }
1569 };
1570
1571 let shutdown_started = StdInstant::now();
1572 transport.shutdown();
1573 assert!(shutdown_started.elapsed() < Duration::from_secs(5));
1574 assert!(transport.snapshot().stopped);
1575
1576 drop(accepted);
1577 drop(transport);
1578 account.close_account();
1579 drop(account);
1580 std::fs::remove_dir_all(directory)?;
1581 Ok(())
1582 }
1583}