Skip to main content

karyon_p2p/peer/
mod.rs

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
45/// A connected peer. Holds a `Wire` that hides the wire shape
46/// (single framed pipe vs. per-protocol streams).
47pub 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}