1use std::{collections::HashMap, sync::Arc};
6
7use message_encoding::MessageEncoding;
8use sequenced_broadcast::{SequencedBroadcastSettings, SettingsError};
9use tokio::{
10 sync::{
11 mpsc::{self, error::SendError},
12 Mutex,
13 },
14 task::JoinHandle,
15};
16
17use crate::{
18 cluster::{
19 leader::{LeaderTask, LeaderTiming},
20 node_state::{NodeState, PeerState},
21 peer_connections::PeerConnections,
22 peer_discovery::{PeerDiscoveryTask, PeerDiscoveryTiming},
23 rpc_server::RpcServer,
24 state_sync::{StateSyncTask, StateSyncTiming},
25 },
26 protocol::messages::{ElectionTerm, LeaderMode, LeaderState},
27 state::{
28 deterministic_state::DeterministicState,
29 recoverable_state::RecoverableState,
30 subscribable_state::{StateHandle, SubscribableState},
31 },
32 transport::{channels::NetIoSettings, traits::SyncIOListener},
33 utils::unique_state_id,
34};
35
36pub struct SharedStateConfig<I: SyncIOListener, D: DeterministicState> {
37 pub io: Arc<I>,
38 pub my_address: I::Address,
39 pub can_lead: bool,
40 pub initial_peers: Vec<I::Address>,
41 pub initial_state: D,
42 pub settings: SharedStateSettings,
43}
44
45pub struct SharedStateRecoverableConfig<I: SyncIOListener, D: DeterministicState> {
46 pub io: Arc<I>,
47 pub my_address: I::Address,
48 pub can_lead: bool,
49 pub initial_peers: Vec<I::Address>,
50 pub initial_state: RecoverableState<D>,
51 pub settings: SharedStateSettings,
52}
53
54#[derive(Clone, Debug, Default)]
55pub struct SharedStateSettings {
56 pub net: NetIoSettings,
57 pub broadcast: SequencedBroadcastSettings,
58 pub discovery_timing: PeerDiscoveryTiming,
59 pub leader_timing: LeaderTiming,
60 pub sync_timing: StateSyncTiming,
61}
62
63const ACTION_QUEUE_CAPACITY: usize = 512;
64
65pub struct SharedState<I: SyncIOListener, D: DeterministicState> {
67 node: Arc<NodeState<I::Address, D>>,
68 actions_tx: mpsc::Sender<(I::Address, D::Action)>,
69 tasks: Vec<JoinHandle<()>>,
70}
71
72impl<I, D> SharedState<I, D>
73where
74 I: SyncIOListener,
75 D: DeterministicState + MessageEncoding,
76 D::Action: MessageEncoding,
77 D::AuthorityAction: MessageEncoding,
78{
79 pub fn start(config: SharedStateConfig<I, D>) -> Result<Self, SettingsError> {
80 let SharedStateConfig {
81 io,
82 my_address,
83 can_lead,
84 initial_peers,
85 initial_state,
86 settings,
87 } = config;
88
89 Self::start_recoverable(SharedStateRecoverableConfig {
90 io,
91 my_address,
92 can_lead,
93 initial_peers,
94 initial_state: RecoverableState::new(unique_state_id(&my_address), initial_state),
95 settings,
96 })
97 }
98
99 pub fn start_recoverable(config: SharedStateRecoverableConfig<I, D>) -> Result<Self, SettingsError> {
100 let SharedStateRecoverableConfig {
101 io,
102 my_address,
103 can_lead,
104 initial_peers,
105 initial_state,
106 settings,
107 } = config;
108
109 let peers = initial_peers
110 .into_iter()
111 .filter(|peer| *peer != my_address)
112 .map(|peer| (peer, PeerState::empty(peer)))
113 .collect::<HashMap<_, _>>();
114
115 let node = Arc::new(NodeState {
116 my_address,
117 can_lead,
118 peers: Mutex::new(peers),
119 state: SubscribableState::new(initial_state, settings.broadcast.clone())?,
120 leader_state: Mutex::new(LeaderState {
121 term: ElectionTerm::default(),
122 mode: LeaderMode::NoLeader,
123 }),
124 });
125
126 let peer_connections = Arc::new(PeerConnections::new(io.clone(), settings.net.clone(), node.clone()));
127 let (actions_tx, actions_rx) = mpsc::channel(ACTION_QUEUE_CAPACITY);
128 let rpc_server = Arc::new(RpcServer::new(node.clone(), actions_tx.clone()));
129
130 let tasks = vec![
131 rpc_server.start_listener(io.clone(), settings.net.clone()),
132 tokio::spawn(
133 PeerDiscoveryTask::new(node.clone(), peer_connections.clone(), settings.discovery_timing).run(),
134 ),
135 tokio::spawn(LeaderTask::new(node.clone(), settings.leader_timing).run()),
136 tokio::spawn(
137 StateSyncTask::new(node.clone(), peer_connections, io, settings.net, actions_rx, settings.sync_timing)
138 .run(),
139 ),
140 ];
141
142 Ok(Self {
143 node,
144 actions_tx,
145 tasks,
146 })
147 }
148
149 pub fn my_address(&self) -> I::Address {
150 self.node.my_address
151 }
152
153 pub fn can_lead(&self) -> bool {
154 self.node.can_lead
155 }
156
157 pub fn node(&self) -> &Arc<NodeState<I::Address, D>> {
159 &self.node
160 }
161
162 pub fn state_handle(&self) -> StateHandle<D> {
164 self.node.state.create_handle()
165 }
166
167 pub async fn leader_state(&self) -> LeaderState<I::Address> {
168 self.node.leader_state.lock().await.clone()
169 }
170
171 pub async fn submit_action(&self, action: D::Action) -> Result<(), SendError<(I::Address, D::Action)>> {
174 self.actions_tx.send((self.node.my_address, action)).await
175 }
176
177 pub fn actions_sender(&self) -> mpsc::Sender<(I::Address, D::Action)> {
179 self.actions_tx.clone()
180 }
181}
182
183impl<I: SyncIOListener, D: DeterministicState> Drop for SharedState<I, D> {
184 fn drop(&mut self) {
185 for task in &self.tasks {
186 task.abort();
187 }
188 }
189}
190
191#[cfg(test)]
192mod tests {
193 use std::{
194 collections::BTreeMap,
195 io::Result,
196 time::{Duration, Instant},
197 };
198
199 use super::*;
200 use crate::{
201 cluster::{leader::LeaderTiming, peer_discovery::PeerDiscoveryTiming, state_sync::StateSyncTiming},
202 state::recoverable_state::RecoverableStateAction,
203 transport::simulated::{SimulatedIo, SimulatedNet},
204 };
205
206 #[derive(Clone, Debug, Default, PartialEq, Eq)]
207 struct KvState {
208 seq: u64,
209 values: BTreeMap<u64, u64>,
210 }
211
212 impl DeterministicState for KvState {
213 type Action = (u64, u64);
214 type AuthorityAction = (u64, u64);
215
216 fn accept_seq(&self) -> u64 {
217 self.seq
218 }
219
220 fn authority(&self, action: Self::Action) -> Self::AuthorityAction {
221 action
222 }
223
224 fn update(&mut self, (key, value): &Self::AuthorityAction) {
225 self.values.insert(*key, *value);
226 self.seq += 1;
227 }
228 }
229
230 impl MessageEncoding for KvState {
231 fn write_to<T: std::io::Write>(&self, out: &mut T) -> Result<usize> {
232 let mut sum = self.seq.write_to(out)?;
233 sum += (self.values.len() as u64).write_to(out)?;
234 for (key, value) in &self.values {
235 sum += key.write_to(out)?;
236 sum += value.write_to(out)?;
237 }
238 Ok(sum)
239 }
240
241 fn read_from<T: std::io::Read>(read: &mut T) -> Result<Self> {
242 let seq = MessageEncoding::read_from(read)?;
243 let len = u64::read_from(read)? as usize;
244 let mut values = BTreeMap::new();
245 for _ in 0..len {
246 values.insert(MessageEncoding::read_from(read)?, MessageEncoding::read_from(read)?);
247 }
248 Ok(Self { seq, values })
249 }
250 }
251
252 fn fast_settings() -> SharedStateSettings {
253 SharedStateSettings {
254 net: NetIoSettings {
255 process_timeout: Duration::from_secs(1),
256 message_timeout: Duration::from_secs(2),
257 },
258 broadcast: SequencedBroadcastSettings::default(),
259 discovery_timing: PeerDiscoveryTiming {
260 observation_interval: Duration::from_millis(50),
261 max_concurrent_observations: 8,
262 },
263 leader_timing: LeaderTiming {
264 tick_interval: Duration::from_millis(25),
265 },
266 sync_timing: StateSyncTiming {
267 leader_poll_interval: Duration::from_millis(20),
268 retry_delay: Duration::from_millis(50),
269 },
270 }
271 }
272
273 async fn start_node(
274 net: &SimulatedNet,
275 address: u64,
276 can_lead: bool,
277 peers: &[u64],
278 ) -> SharedState<SimulatedIo, KvState> {
279 let io = net.start_io(address).await;
280 SharedState::start(SharedStateConfig {
281 io,
282 my_address: address,
283 can_lead,
284 initial_peers: peers.to_vec(),
285 initial_state: KvState::default(),
286 settings: fast_settings(),
287 })
288 .unwrap()
289 }
290
291 async fn wait_for<F: FnMut() -> bool>(what: &str, mut check: F) {
292 let deadline = Instant::now() + Duration::from_secs(30);
293 while !check() {
294 assert!(Instant::now() < deadline, "timed out waiting for {what}");
295 tokio::time::sleep(Duration::from_millis(20)).await;
296 }
297 }
298
299 async fn wait_for_value(node: &SharedState<SimulatedIo, KvState>, key: u64, value: u64) {
300 let mut handle = node.state_handle();
301 wait_for(&format!("node {} to see {key}={value}", node.my_address()), || {
302 handle.read_with(|state| state.state().values.get(&key) == Some(&value))
303 })
304 .await;
305 }
306
307 async fn wait_for_state(node: &SharedState<SimulatedIo, KvState>, expected: &BTreeMap<u64, u64>) {
308 let mut handle = node.state_handle();
309 let deadline = Instant::now() + Duration::from_secs(30);
310 loop {
311 let actual = handle.read_with(|state| state.state().clone());
312 if actual.seq == expected.len() as u64 && actual.values == *expected {
313 return;
314 }
315 assert!(
316 Instant::now() < deadline,
317 "timed out waiting for node {} to settle on {expected:?}, actual seq {} values {:?}",
318 node.my_address(),
319 actual.seq,
320 actual.values,
321 );
322 tokio::time::sleep(Duration::from_millis(20)).await;
323 }
324 }
325
326 async fn wait_for_cluster_state(nodes: &[&SharedState<SimulatedIo, KvState>], expected: &BTreeMap<u64, u64>) {
327 for node in nodes {
328 wait_for_state(node, expected).await;
329 }
330 }
331
332 async fn wait_for_common_leader(nodes: &[&SharedState<SimulatedIo, KvState>]) -> u64 {
333 let deadline = Instant::now() + Duration::from_secs(30);
334 loop {
335 let mut observed_leader = None;
336 let mut leader_is_leading = false;
337 let mut unsettled = Vec::new();
338
339 for node in nodes {
340 let state = node.leader_state().await;
341 let leader = match state.mode {
342 LeaderMode::Leading => {
343 leader_is_leading = true;
344 node.my_address()
345 }
346 LeaderMode::Following { leader } => leader,
347 _ => {
348 unsettled.push((node.my_address(), state));
349 continue;
350 }
351 };
352
353 if observed_leader.is_some_and(|observed| observed != leader) {
354 unsettled.push((node.my_address(), state));
355 }
356 observed_leader.get_or_insert(leader);
357 }
358
359 if let Some(leader) = observed_leader {
360 if unsettled.is_empty() && leader_is_leading {
361 return leader;
362 }
363 }
364
365 assert!(
366 Instant::now() < deadline,
367 "nodes never settled on a common leader, unsettled states {unsettled:?}",
368 );
369 tokio::time::sleep(Duration::from_millis(20)).await;
370 }
371 }
372
373 async fn wait_for_leader(nodes: &[&SharedState<SimulatedIo, KvState>], leader: u64) {
374 for node in nodes {
375 let deadline = Instant::now() + Duration::from_secs(30);
376 loop {
377 let state = node.leader_state().await;
378 let settled = match &state.mode {
379 LeaderMode::Leading => node.my_address() == leader,
380 LeaderMode::Following { leader: followed } => *followed == leader,
381 _ => false,
382 };
383 if settled {
384 break;
385 }
386 assert!(
387 Instant::now() < deadline,
388 "node {} never settled on leader {leader}, last state {state:?}",
389 node.my_address(),
390 );
391 tokio::time::sleep(Duration::from_millis(20)).await;
392 }
393 }
394 }
395
396 #[tokio::test]
397 async fn cluster_replicates_actions_from_any_node() {
398 let net = SimulatedNet::new();
399 let node1 = start_node(&net, 1, true, &[2, 3]).await;
400 let node2 = start_node(&net, 2, true, &[1, 3]).await;
401 let node3 = start_node(&net, 3, false, &[1]).await;
402
403 wait_for_leader(&[&node1, &node2, &node3], 1).await;
404
405 node1.submit_action((10, 100)).await.unwrap();
407 wait_for_value(&node1, 10, 100).await;
408 wait_for_value(&node2, 10, 100).await;
409 wait_for_value(&node3, 10, 100).await;
410
411 node2.submit_action((20, 200)).await.unwrap();
413 node3.submit_action((30, 300)).await.unwrap();
414 for node in [&node1, &node2, &node3] {
415 wait_for_value(node, 20, 200).await;
416 wait_for_value(node, 30, 300).await;
417 }
418 }
419
420 #[tokio::test]
421 async fn start_recoverable_preserves_initial_recovery_details() {
422 let net = SimulatedNet::new();
423 let io = net.start_io(1).await;
424
425 let mut initial_state = RecoverableState::new(101, KvState::default());
426 initial_state.update(&RecoverableStateAction::StateAction { action: (1, 10) });
427 initial_state.update(&RecoverableStateAction::BumpGeneration { new_id: 202 });
428 initial_state.update(&RecoverableStateAction::StateAction { action: (2, 20) });
429 let expected_details = initial_state.details().clone();
430
431 let node = SharedState::start_recoverable(SharedStateRecoverableConfig {
432 io,
433 my_address: 1,
434 can_lead: true,
435 initial_peers: Vec::new(),
436 initial_state,
437 settings: fast_settings(),
438 })
439 .unwrap();
440
441 let mut handle = node.state_handle();
442 let actual_details = handle.recover_details();
443
444 assert_eq!(actual_details, expected_details);
445 }
446
447 #[tokio::test]
448 async fn follower_relays_through_peer_when_leader_is_unreachable() {
449 let net = SimulatedNet::new();
450 let node1 = start_node(&net, 1, true, &[2, 3]).await;
451 let node2 = start_node(&net, 2, true, &[1, 3]).await;
452 let node3 = start_node(&net, 3, false, &[1, 2]).await;
453
454 wait_for_leader(&[&node1, &node2, &node3], 1).await;
455
456 node1.submit_action((1, 1)).await.unwrap();
457 wait_for_value(&node3, 1, 1).await;
458
459 net.set_edge_blocked(1, 3, true).await;
462
463 node1.submit_action((2, 2)).await.unwrap();
464 wait_for_value(&node3, 2, 2).await;
465
466 node3.submit_action((3, 3)).await.unwrap();
467 wait_for_value(&node1, 3, 3).await;
468 wait_for_value(&node2, 3, 3).await;
469 wait_for_value(&node3, 3, 3).await;
470 }
471
472 #[tokio::test]
473 async fn follower_recovers_when_leader_link_goes_silent() {
474 let net = SimulatedNet::new();
475 let node1 = start_node(&net, 1, true, &[2, 3]).await;
476 let node2 = start_node(&net, 2, true, &[1, 3]).await;
477 let node3 = start_node(&net, 3, false, &[1, 2]).await;
478
479 wait_for_leader(&[&node1, &node2, &node3], 1).await;
480 node1.submit_action((1, 1)).await.unwrap();
481 wait_for_value(&node3, 1, 1).await;
482
483 net.set_edge_blackholed(1, 3, true).await;
488
489 node1.submit_action((2, 2)).await.unwrap();
490 wait_for_value(&node3, 2, 2).await;
491
492 node3.submit_action((3, 3)).await.unwrap();
494 wait_for_value(&node1, 3, 3).await;
495 wait_for_value(&node2, 3, 3).await;
496 }
497
498 #[tokio::test]
499 async fn old_leader_rejoins_as_follower_and_its_actions_apply() {
500 let net = SimulatedNet::new();
501 let node1 = start_node(&net, 1, true, &[2, 3]).await;
502 let node2 = start_node(&net, 2, true, &[1, 3]).await;
503 let node3 = start_node(&net, 3, true, &[1, 2]).await;
504
505 wait_for_leader(&[&node1, &node2, &node3], 1).await;
506 node1.submit_action((1, 1)).await.unwrap();
507 wait_for_value(&node3, 1, 1).await;
508
509 net.set_node_blocked(1, true).await;
511 wait_for_leader(&[&node2, &node3], 2).await;
512
513 node2.submit_action((2, 2)).await.unwrap();
514 wait_for_value(&node3, 2, 2).await;
515
516 net.set_node_blocked(1, false).await;
518 let deadline = Instant::now() + Duration::from_secs(10);
519 loop {
520 let state = node1.leader_state().await;
521 if matches!(state.mode, LeaderMode::Following { leader: 2 }) {
522 break;
523 }
524 assert!(Instant::now() < deadline, "node 1 never conceded to node 2, last state {state:?}");
525 tokio::time::sleep(Duration::from_millis(20)).await;
526 }
527 wait_for_value(&node1, 2, 2).await;
528
529 node1.submit_action((3, 3)).await.unwrap();
531 wait_for_value(&node1, 3, 3).await;
532 wait_for_value(&node2, 3, 3).await;
533 wait_for_value(&node3, 3, 3).await;
534 }
535
536 #[tokio::test]
537 async fn observer_actions_apply_after_leader_change() {
538 let net = SimulatedNet::new();
539 let node1 = start_node(&net, 1, true, &[2, 3, 4]).await;
540 let node2 = start_node(&net, 2, true, &[1, 3, 4]).await;
541 let node3 = start_node(&net, 3, true, &[1, 2, 4]).await;
542 let node4 = start_node(&net, 4, false, &[1, 2, 3]).await;
543
544 wait_for_leader(&[&node1, &node2, &node3, &node4], 1).await;
545
546 node4.submit_action((1, 1)).await.unwrap();
547 wait_for_value(&node1, 1, 1).await;
548 wait_for_value(&node4, 1, 1).await;
549
550 net.set_node_blocked(1, true).await;
553 net.stop_node(1).await;
554 drop(node1);
555
556 wait_for_leader(&[&node2, &node3, &node4], 2).await;
557
558 node4.submit_action((2, 2)).await.unwrap();
559 wait_for_value(&node2, 2, 2).await;
560 wait_for_value(&node3, 2, 2).await;
561 wait_for_value(&node4, 2, 2).await;
562 }
563
564 #[tokio::test(flavor = "multi_thread")]
565 async fn action_flood_during_failover_does_not_wedge_sync() {
566 let net = SimulatedNet::new();
567 let node1 = Arc::new(start_node(&net, 1, true, &[2, 3]).await);
568 let node2 = Arc::new(start_node(&net, 2, true, &[1, 3]).await);
569 let node3 = Arc::new(start_node(&net, 3, true, &[1, 2]).await);
570
571 let old_leader = wait_for_common_leader(&[&node1, &node2, &node3]).await;
572 let (survivor_a, survivor_b, new_leader, moved_follower) = match old_leader {
573 1 => (node2.clone(), node3.clone(), 2, node3.clone()),
574 2 => (node1.clone(), node3.clone(), 1, node3.clone()),
575 3 => (node1.clone(), node2.clone(), 1, node2.clone()),
576 _ => unreachable!("test only starts nodes 1, 2, and 3"),
577 };
578
579 let flood = {
583 let survivor_a = survivor_a.clone();
584 let survivor_b = survivor_b.clone();
585 tokio::spawn(async move {
586 let mut i = 0u64;
587 loop {
588 let _ = survivor_a.submit_action((1000 + i, i)).await;
589 let _ = survivor_b.submit_action((2000 + i, i)).await;
590 i += 1;
591 tokio::time::sleep(Duration::from_millis(1)).await;
592 }
593 })
594 };
595
596 net.set_node_blocked(old_leader, true).await;
597 net.stop_node(old_leader).await;
598 flood.abort();
599
600 wait_for_leader(&[&survivor_a, &survivor_b], new_leader).await;
601
602 moved_follower.submit_action((1, 1)).await.unwrap();
605 wait_for_value(&survivor_a, 1, 1).await;
606 wait_for_value(&survivor_b, 1, 1).await;
607 }
608
609 #[tokio::test]
610 async fn cluster_recovers_after_leader_failure() {
611 let net = SimulatedNet::new();
612 let node1 = start_node(&net, 1, true, &[2, 3]).await;
613 let node2 = start_node(&net, 2, true, &[1, 3]).await;
614 let node3 = start_node(&net, 3, true, &[1, 2]).await;
615
616 wait_for_leader(&[&node1, &node2, &node3], 1).await;
617
618 node1.submit_action((1, 1)).await.unwrap();
619 wait_for_value(&node2, 1, 1).await;
620 wait_for_value(&node3, 1, 1).await;
621
622 net.set_node_blocked(1, true).await;
625 net.stop_node(1).await;
626 drop(node1);
627
628 wait_for_leader(&[&node2, &node3], 2).await;
629
630 node3.submit_action((2, 2)).await.unwrap();
631 wait_for_value(&node2, 2, 2).await;
632 wait_for_value(&node3, 2, 2).await;
633 }
634
635 #[tokio::test]
636 async fn five_node_cluster_replicates_from_all_nodes_after_leader_failure() {
637 let net = SimulatedNet::new();
638 let node1 = start_node(&net, 1, true, &[2, 3, 4, 5]).await;
639 let node2 = start_node(&net, 2, true, &[1, 3, 4, 5]).await;
640 let node3 = start_node(&net, 3, true, &[1, 2, 4, 5]).await;
641 let node4 = start_node(&net, 4, false, &[1, 2, 3, 5]).await;
642 let node5 = start_node(&net, 5, false, &[1, 2, 3, 4]).await;
643
644 let all_nodes = [&node1, &node2, &node3, &node4, &node5];
645 let first_leader = wait_for_common_leader(&all_nodes).await;
646 assert_eq!(first_leader, 1);
647
648 let mut expected = BTreeMap::new();
649 for node in all_nodes {
650 let key = 100 + node.my_address();
651 let value = key * 10;
652 node.submit_action((key, value)).await.unwrap();
653 expected.insert(key, value);
654 }
655 wait_for_cluster_state(&[&node1, &node2, &node3, &node4, &node5], &expected).await;
656
657 net.set_node_blocked(first_leader, true).await;
658 net.stop_node(first_leader).await;
659 drop(node1);
660
661 let remaining_nodes = [&node2, &node3, &node4, &node5];
662 let second_leader = wait_for_common_leader(&remaining_nodes).await;
663 assert_ne!(second_leader, first_leader);
664 assert_eq!(second_leader, 2);
665
666 for node in remaining_nodes {
667 let key = 200 + node.my_address();
668 let value = key * 10;
669 node.submit_action((key, value)).await.unwrap();
670 expected.insert(key, value);
671 }
672 wait_for_cluster_state(&[&node2, &node3, &node4, &node5], &expected).await;
673 }
674}