1use std::{collections::HashMap, mem, net::IpAddr, sync::Arc, time::Duration};
2
3use bytes::{Bytes, BytesMut};
4use tokio::{
5 net::{TcpListener, TcpStream},
6 select, spawn,
7 sync::{
8 mpsc::{UnboundedReceiver, UnboundedSender},
9 watch,
10 },
11 task::{JoinHandle, JoinSet},
12};
13use tokio_util::{sync::CancellationToken, task::JoinMap};
14use tracing::{debug, error, info, trace, warn};
15
16use crate::{
17 Config, Metrics, NetAddr, PublicKey, Role,
18 connection::Connection,
19 delay::DelayQueue,
20 error::NetworkError,
21 msg::{MsgId, Slot, Trailer, hello::Hello},
22 net::{Command, PeerCommand, PeerMessage, RetryPolicy, SendAction, peer::Peer},
23 queue::Queue,
24 util::until,
25};
26
27pub struct Server {
28 key: PublicKey,
29 conf: Arc<Config>,
30 role: Role,
31 msgid: MsgId,
32 lower_bound: Slot,
33 parties: HashMap<PublicKey, Party>,
34 ibound: UnboundedSender<PeerMessage>,
35 obound: UnboundedReceiver<Command>,
36 next_slot: watch::Receiver<Slot>,
37 accept_tasks: JoinSet<Result<Connection, NetworkError>>,
38 hello_tasks: JoinMap<PublicKey, Result<(Hello, Connection, Hello), NetworkError>>,
39 connect_tasks: JoinMap<PublicKey, Connection>,
40 peer_tasks: JoinMap<PublicKey, Peer>,
41 metrics: Arc<dyn Metrics>,
42}
43
44struct Party {
45 role: Role,
46 addr: NetAddr,
47 outbox: Queue<(RetryPolicy, Bytes)>,
48 retry: DelayQueue,
49 peer: PeerState,
50}
51
52enum PeerState {
68 None,
70 Connected(CancellationToken),
75 Reconnect(Peer),
79 Replace(Connection),
85}
86
87impl Server {
88 pub(super) fn spawn(
89 conf: Arc<Config>,
90 listener: TcpListener,
91 role: Role,
92 tx: UnboundedSender<PeerMessage>,
93 rx: UnboundedReceiver<Command>,
94 sx: watch::Receiver<Slot>,
95 metrics: Arc<dyn Metrics>,
96 ) -> JoinHandle<()> {
97 let our_key = conf.keypair.public_key();
98 let parties = conf
99 .parties
100 .iter()
101 .filter(|&(k, _)| *k != our_key)
102 .map(|(k, a)| {
103 let p = Party::new(conf.clone(), Role::Active, a.clone());
104 (*k, p)
105 })
106 .collect();
107
108 let this = Self {
109 key: our_key,
110 conf,
111 role,
112 ibound: tx,
113 obound: rx,
114 parties,
115 accept_tasks: JoinSet::new(),
116 connect_tasks: JoinMap::new(),
117 hello_tasks: JoinMap::new(),
118 peer_tasks: JoinMap::new(),
119 msgid: MsgId::new(0),
120 next_slot: sx,
121 lower_bound: Slot::MIN,
122 metrics,
123 };
124
125 spawn(this.run(listener))
126 }
127
128 async fn run(mut self, listener: TcpListener) {
129 for (k, a) in self
131 .parties
132 .iter()
133 .map(|(k, p)| (*k, p.addr.clone()))
134 .collect::<Vec<_>>()
135 {
136 self.spawn_connect(k, a)
137 }
138
139 loop {
140 select! {
141 x = listener.accept() => match x {
142 Ok((stream, addr)) => {
143 debug!(
144 name = %self.conf.name,
145 node = %self.key,
146 %addr,
147 "accepted new tcp connection"
148 );
149 self.spawn_accept(stream)
150 }
151 Err(err) => {
152 warn!(
153 name = %self.conf.name,
154 node = %self.key,
155 %err,
156 "error accepting tcp connection"
157 )
158 }
159 },
160
161 Some(h) = self.accept_tasks.join_next() => match h {
162 Ok(Ok(conn)) => {
163 self.metrics.set(&self.key, ACCEPT_TASKS, self.accept_tasks.len());
164 if conn.key == self.key {
165 warn!(
166 name = %self.conf.name,
167 node = %self.key,
168 peer = %conn.key,
169 addr = %conn.addr,
170 "rejecting connection with the same key"
171 );
172 self.spawn_hello(conn, Hello::BackOff(Duration::MAX));
173 continue
174 }
175 let Some(party) = self.parties.get_mut(&conn.key) else {
176 info!(
177 name = %self.conf.name,
178 node = %self.key,
179 peer = %conn.key,
180 addr = %conn.addr,
181 "unknown party"
182 );
183 self.spawn_hello(conn, Hello::BackOff(self.conf.backoff_duration));
184 continue
185 };
186 if party.ip_addr_mismatch(conn.addr.ip()) {
187 warn!(
188 name = %self.conf.name,
189 node = %self.key,
190 peer = %conn.key,
191 addr = %conn.addr,
192 "party has invalid ip addr"
193 );
194 self.spawn_hello(conn, Hello::BackOff(self.conf.backoff_duration));
195 continue
196 }
197 self.spawn_hello(conn, Hello::Ok);
198 }
199 Ok(Err(err)) => {
200 self.metrics.set(&self.key, ACCEPT_TASKS, self.accept_tasks.len());
201 warn!(name = %self.conf.name, node = %self.key, %err, "handshake failed")
202 }
203 Err(err) => {
204 self.metrics.set(&self.key, ACCEPT_TASKS, self.accept_tasks.len());
205 if err.is_panic() {
206 error!(
207 name = %self.conf.name,
208 node = %self.key,
209 %err,
210 "handshake task panic"
211 )
212 }
213 }
214 },
215
216 Some(r) = self.hello_tasks.join_next() => match r {
217 (key, Ok(Ok((our_hello, conn, their_hello)))) => {
218 self.metrics.set(&self.key, HELLO_TASKS, self.hello_tasks.len());
219 if key == self.key {
220 continue
223 }
224 let Some(party) = self.parties.get_mut(&key) else {
225 info!(
226 name = %self.conf.name,
227 node = %self.key,
228 peer = %key,
229 addr = %conn.addr,
230 "unknown party"
231 );
232 continue
233 };
234 if !(our_hello.is_ok() && their_hello.is_ok()) {
235 warn!(
236 name = %self.conf.name,
237 node = %self.key,
238 peer = %key,
239 addr = %conn.addr,
240 ours = ?our_hello,
241 theirs = ?their_hello,
242 "hello failed"
243 );
244 continue
245 }
246 match party.peer.take() {
247 PeerState::None => {
248 self.connect_tasks.abort(&key);
249 let peer = Peer::builder()
250 .config(self.conf.clone())
251 .budget(self.conf.peer_budget)
252 .inbound(self.ibound.clone())
253 .messages(party.outbox.clone())
254 .retry(party.retry.clone())
255 .metrics(self.metrics.clone())
256 .build();
257 let cancel = CancellationToken::new();
258 party.peer = PeerState::Connected(cancel.clone());
259 self.spawn_peer(key, peer, conn, cancel);
260 }
261 PeerState::Reconnect(peer) => {
262 self.connect_tasks.abort(&key);
263 let cancel = CancellationToken::new();
264 party.peer = PeerState::Connected(cancel.clone());
265 self.spawn_peer(key, peer, conn, cancel);
266 }
267 PeerState::Connected(cancel) => {
268 if key > self.key {
269 info!(
270 name = %self.conf.name,
271 node = %self.key,
272 peer = %key,
273 addr = %conn.addr,
274 "replacing connection with accepted one"
275 );
276 cancel.cancel();
277 party.peer = PeerState::Replace(conn);
278 } else {
279 party.peer = PeerState::Connected(cancel);
280 }
281 }
282 PeerState::Replace(_) => {
283 party.peer = PeerState::Replace(conn);
284 }
285 }
286 }
287 (key, Ok(Err(err))) => {
288 self.metrics.set(&self.key, HELLO_TASKS, self.hello_tasks.len());
289 warn!(
290 name = %self.conf.name,
291 node = %self.key,
292 peer = %key,
293 %err,
294 "hello task error"
295 )
296 }
297 (key, Err(err)) => {
298 self.metrics.set(&self.key, HELLO_TASKS, self.hello_tasks.len());
299 if err.is_panic() {
300 error!(
301 name = %self.conf.name,
302 node = %self.key,
303 peer = %key,
304 %err,
305 "hello task panic"
306 )
307 }
308 }
309 },
310
311 Some(x) = self.connect_tasks.join_next() => match x {
312 (key, Ok(conn)) => {
313 self.metrics.set(&self.key, CONNECT_TASKS, self.connect_tasks.len());
314 let Some(party) = self.parties.get_mut(&key) else {
315 debug!(
316 name = %self.conf.name,
317 node = %self.key,
318 peer = %key,
319 addr = %conn.addr,
320 "party has been removed"
321 );
322 continue
323 };
324 match party.peer.take() {
325 PeerState::None => {
326 let peer = Peer::builder()
327 .config(self.conf.clone())
328 .budget(self.conf.peer_budget)
329 .inbound(self.ibound.clone())
330 .messages(party.outbox.clone())
331 .retry(party.retry.clone())
332 .metrics(self.metrics.clone())
333 .build();
334 let cancel = CancellationToken::new();
335 party.peer = PeerState::Connected(cancel.clone());
336 self.spawn_peer(key, peer, conn, cancel);
337 }
338 PeerState::Reconnect(peer) => {
339 let cancel = CancellationToken::new();
340 party.peer = PeerState::Connected(cancel.clone());
341 self.spawn_peer(key, peer, conn, cancel);
342 }
343 PeerState::Connected(cancel) => {
344 if key < self.key {
345 info!(
346 name = %self.conf.name,
347 node = %self.key,
348 peer = %key,
349 addr = %conn.addr,
350 "replacing connection with outgoing one"
351 );
352 cancel.cancel();
353 party.peer = PeerState::Replace(conn);
354 } else {
355 party.peer = PeerState::Connected(cancel);
356 }
357 }
358 PeerState::Replace(_) => {
359 party.peer = PeerState::Replace(conn);
360 }
361 }
362 }
363 (key, Err(err)) => {
364 self.metrics.set(&self.key, CONNECT_TASKS, self.connect_tasks.len());
365 if err.is_panic() {
366 error!(
367 name = %self.conf.name,
368 node = %self.key,
369 peer = %key,
370 %err,
371 "connect task panic"
372 );
373 if let Some(party) = self.parties.get_mut(&key)
374 && matches!(party.peer, PeerState::None | PeerState::Reconnect(_))
375 {
376 let addr = party.addr.clone();
377 self.spawn_connect(key, addr);
378 }
379 }
380 }
381 },
382
383 Some(p) = self.peer_tasks.join_next() => match p {
384 (key, Ok(peer)) => {
385 self.metrics.set(&self.key, PEER_TASKS, self.peer_tasks.len());
386 if self.ibound.is_closed() {
387 return
388 }
389 let Some(party) = self.parties.get_mut(&key) else {
390 debug!(
391 name = %self.conf.name,
392 node = %self.key,
393 peer = %key,
394 "party has been removed"
395 );
396 continue
397 };
398 if let PeerState::Replace(conn) = party.peer.take() {
399 let cancel = CancellationToken::new();
400 party.peer = PeerState::Connected(cancel.clone());
401 self.spawn_peer(key, peer, conn, cancel);
402 } else {
403 let addr = party.addr.clone();
404 party.peer = PeerState::Reconnect(peer);
405 self.spawn_connect(key, addr);
406 }
407 }
408 (key, Err(err)) => {
409 self.metrics.set(&self.key, PEER_TASKS, self.peer_tasks.len());
410 if err.is_panic() {
411 error!(
412 name = %self.conf.name,
413 node = %self.key,
414 peer = %key,
415 %err,
416 "peer task panic"
417 );
418 if self.ibound.is_closed() {
419 return
420 }
421 if let Some(party) = self.parties.get_mut(&key) {
422 let addr = party.addr.clone();
423 party.peer = PeerState::None;
424 self.spawn_connect(key, addr);
425 }
426 }
427 }
428 },
429
430 r = self.next_slot.changed() => {
431 if r.is_err() {
432 return
433 }
434 let s = *self.next_slot.borrow_and_update();
435 debug_assert!(s > self.lower_bound); self.lower_bound = s;
437 self.metrics.set(&self.key, LOWER_BOUND, u64::from(s) as usize);
438 for party in self.parties.values() {
439 party.outbox.gc(s);
440 party.retry.gc(s);
441 }
442 }
443
444 cmd = self.obound.recv() => {
445 self.metrics.set(&self.key, CHANNEL_SIZE, self.obound.len());
446 match cmd {
447 Some(Command::Peer(PeerCommand::Add(role, parties))) => {
448 for (k, a) in parties {
449 if k == self.key {
450 self.role = role;
451 continue
452 }
453 if let Some(p) = self.parties.get_mut(&k) {
454 if p.addr == a {
455 p.role = role;
456 } else {
457 info!(
458 name = %self.conf.name,
459 node = %self.key,
460 peer = %k,
461 addr = %a,
462 "updating party address"
463 );
464 p.addr = a.clone();
465 p.role = role;
466 self.connect_tasks.abort(&k);
467 if let PeerState::Connected(cancel) = &p.peer {
468 cancel.cancel()
469 } else {
470 self.spawn_connect(k, a)
471 }
472 }
473 continue
474 }
475 info!(
476 name = %self.conf.name,
477 node = %self.key,
478 peer = %k,
479 addr = %a,
480 "adding new peer"
481 );
482 self.parties.insert(k, Party::new(self.conf.clone(), role, a.clone()));
483 self.spawn_connect(k, a)
484 }
485 }
486 Some(Command::Peer(PeerCommand::Remove(peers))) => {
487 for k in &peers {
488 if *k == self.key {
489 info!(
490 name = %self.conf.name,
491 node = %self.key,
492 "removing self sets role to passive"
493 );
494 self.role = Role::Passive;
495 continue
496 }
497 info!(
498 name = %self.conf.name,
499 node = %self.key,
500 peer = %k,
501 "removing peer"
502 );
503 self.parties.remove(k);
504 self.connect_tasks.abort(k);
505 self.peer_tasks.abort(k);
506 }
507 }
508 Some(Command::Peer(PeerCommand::Assign(role, peers))) => {
509 for k in &peers {
510 if *k == self.key {
511 self.role = role;
512 continue
513 }
514 if let Some(p) = self.parties.get_mut(k) {
515 info!(
516 name = %self.conf.name,
517 node = %self.key,
518 peer = %k,
519 %role,
520 "assigning role to peer"
521 );
522 p.role = role
523 } else {
524 warn!(
525 name = %self.conf.name,
526 node = %self.key,
527 peer = %k,
528 role = %role,
529 "peer to assign role to not found"
530 );
531 }
532 }
533 }
534 Some(Command::Send(cmd)) => match cmd.action {
535 SendAction::Unicast(to, m) => {
536 if cmd.slot < self.lower_bound {
537 continue
538 }
539
540 if to == self.key {
541 trace!(name = %self.conf.name, node = %self.key, "sending message");
542 if let Err(err) = self.ibound.send((self.key, m.into(), None)) {
543 warn!(
544 name = %self.conf.name,
545 node = %self.key,
546 err = %err,
547 "channel closed"
548 );
549 return
550 }
551 trace!(name = %self.conf.name, node = %self.key, "message delivered");
552 continue
553 }
554
555 let msgid = self.next_msgid();
556 let bytes = append_trailer(cmd.retry, cmd.slot, msgid, m);
557
558 if let Some(party) = self.parties.get(&to) {
559 party.outbox.enqueue(cmd.slot, msgid, (cmd.retry, bytes));
560 } else {
561 warn!(
562 name = %self.conf.name,
563 node = %self.key,
564 peer = %to,
565 "unicast target not found"
566 );
567 }
568 }
569 SendAction::Multicast(parties, m) => {
570 if cmd.slot < self.lower_bound {
571 continue
572 }
573
574 let msgid = self.next_msgid();
575 let bytes = append_trailer(cmd.retry, cmd.slot, msgid, m);
576
577 if parties.contains(&self.key) {
578 let bytes = remove_trailer(bytes.clone());
579 trace!(name = %self.conf.name, node = %self.key, "sending message");
580 if let Err(err) = self.ibound.send((self.key, bytes, None)) {
581 warn!(
582 name = %self.conf.name,
583 node = %self.key,
584 err = %err,
585 "channel closed"
586 );
587 return
588 }
589 trace!(name = %self.conf.name, node = %self.key, "message delivered");
590 }
591
592 for (to, party) in &self.parties {
593 if !parties.contains(to) {
594 continue
595 }
596 trace!(name = %self.conf.name, node = %self.key, %to, "sending message");
597 party.outbox.enqueue(cmd.slot, msgid, (cmd.retry, bytes.clone()));
598 }
599 }
600 SendAction::Broadcast(m) => {
601 if cmd.slot < self.lower_bound {
602 continue
603 }
604
605 let msgid = self.next_msgid();
606 let bytes = append_trailer(cmd.retry, cmd.slot, msgid, m);
607
608 if self.role.is_active() {
609 let bytes = remove_trailer(bytes.clone());
610 trace!(name = %self.conf.name, node = %self.key, "sending message");
611 if let Err(err) = self.ibound.send((self.key, bytes, None)) {
612 warn!(
613 name = %self.conf.name,
614 node = %self.key,
615 err = %err,
616 "channel closed"
617 );
618 return
619 }
620 trace!(name = %self.conf.name, node = %self.key, "message delivered");
621 }
622 for (key, party) in &self.parties {
623 if party.role.is_active() {
624 trace!(
625 name = %self.conf.name,
626 node = %self.key,
627 to = %key,
628 "sending message"
629 );
630 party.outbox.enqueue(cmd.slot, msgid, (cmd.retry, bytes.clone()));
631 }
632 }
633 }
634 }
635 Some(Command::Shutdown(tx)) => {
636 debug!(name = %self.conf.name, node = %self.key, "shutting down");
637 let _ = tx.send(());
638 return
639 }
640 None => return
641 }
642 }
643 }
644 }
645 }
646
647 fn spawn_connect(&mut self, key: PublicKey, addr: NetAddr) {
648 if self.key == key {
649 return;
650 }
651 debug!(
652 name = %self.conf.name,
653 node = %self.key,
654 peer = %key,
655 addr = %addr,
656 "spawning connect task"
657 );
658 let conn = Connection::connect(self.conf.clone(), key, addr);
659 self.connect_tasks.spawn(key, conn);
660 self.metrics.add(&key, CONNECT_ATTEMPTS, 1);
661 self.metrics
662 .set(&self.key, CONNECT_TASKS, self.connect_tasks.len());
663 }
664
665 fn spawn_accept(&mut self, stream: TcpStream) {
666 debug!(name = %self.conf.name, node = %self.key, "spawning accept task");
667 let conn = Connection::accept(self.conf.clone(), stream);
668 self.accept_tasks.spawn(conn);
669 self.metrics
670 .set(&self.key, ACCEPT_TASKS, self.accept_tasks.len());
671 }
672
673 fn spawn_hello(&mut self, mut conn: Connection, ours: Hello) {
674 debug!(
675 name = %self.conf.name,
676 node = %self.key,
677 peer = %conn.key,
678 addr = %conn.addr,
679 "spawning hello task"
680 );
681
682 self.metrics.add(&conn.key, HELLOS, 1);
683
684 self.hello_tasks.abort(&conn.key);
685 self.hello_tasks.spawn(
686 conn.key,
687 until(self.conf.handshake_timeout, async move {
688 let theirs = conn.recv_hello().await?;
689 conn.send_hello(ours.clone()).await?;
690 Ok::<_, NetworkError>((ours, conn, theirs))
691 }),
692 );
693
694 self.metrics
695 .set(&self.key, HELLO_TASKS, self.hello_tasks.len());
696 }
697
698 fn spawn_peer(
699 &mut self,
700 key: PublicKey,
701 mut peer: Peer,
702 conn: Connection,
703 cancel: CancellationToken,
704 ) {
705 debug!(
706 name = %self.conf.name,
707 node = %self.key,
708 peer = %key,
709 addr = %conn.addr,
710 "spawning peer task"
711 );
712 let node = self.key;
713 let name = self.conf.name.clone();
714 let metrics = self.metrics.clone();
715 let addr = conn.addr;
716 self.peer_tasks.spawn(key, async move {
717 let Err(err) = peer.start(conn, cancel).await;
718 if !matches!(err, NetworkError::PeerInterrupt) {
719 warn!(
720 %name,
721 %node,
722 peer = %key,
723 %addr,
724 %err,
725 "peer failure"
726 );
727 metrics.add(&key, ERRORS, 1)
728 }
729 peer
730 });
731 self.metrics
732 .set(&self.key, PEER_TASKS, self.peer_tasks.len());
733 }
734
735 fn next_msgid(&mut self) -> MsgId {
736 let current = self.msgid;
737 self.msgid = MsgId::new(self.msgid.0.wrapping_add(1));
738 current
739 }
740}
741
742impl Party {
743 fn new(c: Arc<Config>, r: Role, a: NetAddr) -> Self {
744 Self {
745 addr: a,
746 role: r,
747 outbox: Queue::new(),
748 retry: DelayQueue::new(c),
749 peer: PeerState::None,
750 }
751 }
752
753 fn ip_addr_mismatch(&self, addr: IpAddr) -> bool {
754 let NetAddr::Inet(ip, _) = &self.addr else {
755 return false;
756 };
757 *ip != addr
758 }
759}
760
761impl PeerState {
762 fn take(&mut self) -> Self {
763 mem::replace(self, Self::None)
764 }
765}
766
767fn append_trailer(pol: RetryPolicy, slot: Slot, id: MsgId, bytes: Vec<u8>) -> Bytes {
768 let t = match pol {
769 RetryPolicy::Default => Trailer::Std { slot, id },
770 RetryPolicy::NoRetry => Trailer::NoAck { slot },
771 };
772 let mut msg = BytesMut::from(Bytes::from(bytes));
773 msg.extend_from_slice(t.to_bytes().as_ref());
774 msg.freeze()
775}
776
777fn remove_trailer(mut bytes: Bytes) -> Bytes {
778 let _t = Trailer::from_bytes(&mut bytes);
779 debug_assert!(_t.is_some());
780 bytes
781}
782
783const ACCEPT_TASKS: &str = "accept_tasks";
787
788const CHANNEL_SIZE: &str = "channel_size";
790
791const CONNECT_ATTEMPTS: &str = "connect_attempts";
793
794const CONNECT_TASKS: &str = "connect_tasks";
796
797const ERRORS: &str = "errors";
799
800const HELLOS: &str = "hellos";
802
803const HELLO_TASKS: &str = "hello_tasks";
805
806const LOWER_BOUND: &str = "lower_bound";
808
809const PEER_TASKS: &str = "peer_tasks";