use crate::status::DownloadStats;
use crate::{
peer_connection::PeerConnection,
peer_explorer::{
Peer,
channel::{PeerExplorerChannelMessage, PeerExplorerChannelReceiver},
},
piece_manager::channel::{PieceManagerChannelSender, PieceManagerMessage},
};
use std::net::{Ipv4Addr, SocketAddrV4};
use std::{collections::HashMap, sync::Arc};
use tokio::{sync::oneshot, task, task::JoinSet};
use tracing::{error, warn};
pub mod peer_selection_strategy;
const MAX_PEERS: usize = 50;
struct Connection {
peer: Peer,
outbound: bool,
}
pub struct PeerManager<S, W>
where
S: peer_selection_strategy::PeerSelectionStrategy + Sync + Send + 'static,
W: crate::store::Store + Send + Sync + 'static,
W::Error: std::error::Error + Send + Sync + 'static,
{
peer_slection_strategy: S,
info_hash: [u8; 20],
peer_id: [u8; 20],
stats: Arc<DownloadStats>,
download_completed: bool,
listening_port: u16,
store: Arc<W>,
}
impl<S, W> PeerManager<S, W>
where
S: peer_selection_strategy::PeerSelectionStrategy + Sync + Send + 'static,
W: crate::store::Store + Send + Sync + 'static,
W::Error: std::error::Error + Send + Sync + 'static,
{
pub fn new(
peer_slection_strategy: S,
info_hash: &[u8; 20],
peer_id: &[u8; 20],
stats: Arc<DownloadStats>,
listening_port: u16,
store: Arc<W>,
) -> Self {
Self {
peer_slection_strategy,
info_hash: *info_hash,
peer_id: *peer_id,
stats,
download_completed: false,
listening_port,
store,
}
}
async fn is_download_completed(&mut self, sender: &PieceManagerChannelSender) -> bool {
if self.download_completed {
return true;
}
let (response_sender, response) = oneshot::channel();
if sender
.send(PieceManagerMessage::IsCompleted { response_sender })
.await
.is_err()
{
self.download_completed = true;
return true;
}
self.download_completed = response.await.unwrap_or(true);
self.download_completed
}
pub async fn start(
mut self,
mut peer_explorer_channel_receiver: PeerExplorerChannelReceiver,
piece_manager_channel_sender: PieceManagerChannelSender,
) {
let mut connections: JoinSet<()> = JoinSet::new();
let mut live: HashMap<task::Id, Connection> = HashMap::new();
let listener = tokio::net::TcpListener::bind(SocketAddrV4::new(
Ipv4Addr::UNSPECIFIED,
self.listening_port,
))
.await
.unwrap();
loop {
tokio::select! {
Some(outcome) = connections.join_next_with_id() => {
let (id, panicked) = match outcome {
Ok((id, ())) => (id, false),
Err(e) => (e.id(), e.is_panic()),
};
let Some(connection) = live.remove(&id) else {
continue;
};
if panicked {
error!("{}: connection task panicked", connection.peer.address);
}
self.stats.peer_disconnected();
if connection.outbound {
self.peer_slection_strategy.push(connection.peer, true);
}
}
Some(msg) = peer_explorer_channel_receiver.recv() => {
match msg {
PeerExplorerChannelMessage::PeerFound(peer) => {
self.peer_slection_strategy.push(peer, false);
}
}
}
Ok((stream, address)) = listener.accept() => {
let peer = Peer {
peer_id: None,
address,
};
let stats = self.stats.clone();
let piece_manager_channel_sender = piece_manager_channel_sender.clone();
let info_hash = self.info_hash;
let peer_id = self.peer_id;
let store = self.store.clone();
self.stats.peer_connected();
let handle = connections.spawn(
PeerConnection::from_stream(
stream,
address,
piece_manager_channel_sender,
&info_hash,
&peer_id,
stats,
store,
)
.await.start()
);
live.insert(
handle.id(),
Connection {
peer,
outbound: false,
},
);
}
Some(attempt) = self.peer_slection_strategy.pop(),
if !self.download_completed
&& self.stats.active_peers() < MAX_PEERS
&& self.peer_slection_strategy.peek().is_some() =>
{
if self.is_download_completed(&piece_manager_channel_sender).await {
continue;
}
self.stats.peer_connected();
let peer = attempt.peer;
let stats = self.stats.clone();
let piece_manager_channel_sender = piece_manager_channel_sender.clone();
let info_hash = self.info_hash;
let peer_id = self.peer_id;
let store = self.store.clone();
let handle = connections.spawn(async move {
match PeerConnection::connect(
peer,
piece_manager_channel_sender,
&info_hash,
&peer_id,
stats,
store,
)
.await
{
Ok(peer_connection) => peer_connection.start().await,
Err(e) => warn!("{}", e),
}
});
live.insert(
handle.id(),
Connection {
peer,
outbound: true,
},
);
}
}
}
}
}