1mod connection;
2mod peer_id;
3
4use std::{
5 collections::HashSet,
6 fmt,
7 sync::{Arc, Weak},
8};
9
10use async_channel::{Receiver, Sender};
11use log::{error, trace};
12
13use karyon_core::{
14 async_runtime::Executor,
15 async_util::{TaskGroup, TaskResult},
16};
17
18use crate::{
19 conn_queue::QueuedConn,
20 endpoint::Endpoint,
21 peer_pool::PeerPool,
22 protocol::{PeerConn, ProtocolEvent, ProtocolID},
23 Config, Result,
24};
25
26pub use peer_id::PeerID;
27
28use connection::Wire;
29
30#[derive(Clone, Debug)]
31pub enum ConnDirection {
32 Inbound,
33 Outbound,
34}
35
36impl fmt::Display for ConnDirection {
37 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
38 match self {
39 ConnDirection::Inbound => write!(f, "Inbound"),
40 ConnDirection::Outbound => write!(f, "Outbound"),
41 }
42 }
43}
44
45pub struct Peer {
48 id: PeerID,
49 peer_pool: Weak<PeerPool>,
50
51 direction: ConnDirection,
52 remote_endpoint: Endpoint,
53
54 connection: Arc<dyn Wire>,
55 disconnect_signal: Sender<Result<()>>,
56
57 negotiated_protocols: HashSet<ProtocolID>,
58 stop_chan: (Sender<Result<()>>, Receiver<Result<()>>),
59 config: Arc<Config>,
60 executor: Executor,
61 task_group: TaskGroup,
62}
63
64impl Peer {
65 pub(crate) async fn new(
66 peer_pool: Arc<PeerPool>,
67 queued: QueuedConn,
68 id: PeerID,
69 negotiated_protocols: HashSet<ProtocolID>,
70 protocol_ids: impl IntoIterator<Item = ProtocolID> + Clone,
71 ) -> Result<Arc<Self>> {
72 let config = peer_pool.config.clone();
73 let executor = peer_pool.executor.clone();
74 let task_group = TaskGroup::with_executor(executor.clone());
75 let stop_chan = async_channel::bounded::<Result<()>>(1);
76
77 let remote_endpoint = queued.remote_endpoint.clone();
78 let direction = queued.direction.clone();
79 let disconnect_signal = queued.disconnect_signal.clone();
80
81 let connection = connection::from_queued(
82 queued,
83 &negotiated_protocols,
84 protocol_ids,
85 &task_group,
86 stop_chan.0.clone(),
87 )
88 .await?;
89
90 let peer_pool_weak = Arc::downgrade(&peer_pool);
91 Ok(Arc::new(Peer {
92 id,
93 peer_pool: peer_pool_weak,
94 direction,
95 remote_endpoint,
96 connection,
97 disconnect_signal,
98 negotiated_protocols,
99 stop_chan,
100 config,
101 executor,
102 task_group,
103 }))
104 }
105
106 pub async fn send(&self, proto_id: ProtocolID, msg: Vec<u8>) -> Result<()> {
107 self.connection.send(&proto_id, msg).await
108 }
109
110 pub async fn recv(&self, proto_id: &ProtocolID) -> Result<ProtocolEvent> {
111 self.connection.recv(proto_id).await
112 }
113
114 pub async fn broadcast(&self, proto_id: &ProtocolID, msg: Vec<u8>) {
115 self.peer_pool().broadcast(proto_id, msg).await;
116 }
117
118 pub fn id(&self) -> &PeerID {
119 &self.id
120 }
121
122 pub fn config(&self) -> Arc<Config> {
123 self.config.clone()
124 }
125
126 pub fn executor(&self) -> Executor {
127 self.executor.clone()
128 }
129
130 pub fn remote_endpoint(&self) -> &Endpoint {
131 &self.remote_endpoint
132 }
133
134 pub fn is_inbound(&self) -> bool {
135 matches!(self.direction, ConnDirection::Inbound)
136 }
137
138 pub fn direction(&self) -> &ConnDirection {
139 &self.direction
140 }
141
142 pub fn negotiated_protocols(&self) -> &HashSet<ProtocolID> {
143 &self.negotiated_protocols
144 }
145
146 pub(crate) async fn run(self: Arc<Self>) -> Result<()> {
147 self.run_connect_protocols().await;
148 let stop_signal = self.stop_chan.1.recv().await?;
149 stop_signal
150 }
151
152 pub(crate) async fn shutdown(self: &Arc<Self>) -> Result<()> {
153 trace!("peer {} shutting down", self.id);
154
155 let _ = self.connection.shutdown().await;
156 let _ = self.stop_chan.0.try_send(Ok(()));
157
158 let _ = self.disconnect_signal.send(Ok(())).await;
159 self.task_group.cancel().await;
160 Ok(())
161 }
162
163 async fn run_connect_protocols(self: &Arc<Self>) {
164 for (proto_id, constructor) in self.peer_pool().protocols.read().await.iter() {
165 if !self.negotiated_protocols.contains(proto_id) {
166 trace!("peer {} skip protocol {proto_id} (not negotiated)", self.id);
167 continue;
168 }
169 trace!("peer {} run protocol {proto_id}", self.id);
170
171 let peer_conn = PeerConn::new(self.clone(), proto_id.clone());
172 let protocol = constructor(peer_conn);
173
174 let on_failure = {
175 let this = self.clone();
176 let proto_id = proto_id.clone();
177 |result: TaskResult<Result<()>>| async move {
178 if let TaskResult::Completed(res) = result {
179 if res.is_err() {
180 error!("protocol {proto_id} stopped");
181 }
182 let _ = this.stop_chan.0.try_send(res);
183 }
184 }
185 };
186
187 self.task_group.spawn_then(protocol.start(), on_failure);
188 }
189 }
190
191 fn peer_pool(&self) -> Arc<PeerPool> {
192 self.peer_pool.upgrade().unwrap()
193 }
194}