use std::net::SocketAddr;
use std::sync::Arc;
use crate::{
peer_connection::error::{PeerConnectionError, PeerConnectionResult},
peer_explorer::Peer,
piece_manager::channel::{PieceManagerChannelSender, PieceManagerMessage},
status::DownloadStats,
wire_protocol::{Bitfield, Handshake, Message, WireCodec, WireItem},
};
use futures::{SinkExt, StreamExt};
use tokio::{net::TcpStream, select, sync::oneshot, task::JoinSet, time};
use tokio_util::codec::Framed;
use tracing::{debug, warn};
pub mod channels;
pub mod error;
pub mod request_manager;
pub struct PeerConnection<W>
where
W: crate::store::Store + Send + Sync + 'static,
W::Error: std::error::Error + Send + Sync + 'static,
{
pub stats: Arc<DownloadStats>,
pub peer: Peer,
pub is_outbound: bool,
pub piece_manager_channel_sender: PieceManagerChannelSender,
pub stream: Option<TcpStream>,
pub info_hash: [u8; 20],
pub peer_id: [u8; 20],
pub store: Arc<W>,
}
impl<W> PeerConnection<W>
where
W: crate::store::Store + Send + Sync + 'static,
W::Error: std::error::Error + Send + Sync + 'static,
{
pub async fn connect(
peer: Peer,
piece_manager_channel_sender: PieceManagerChannelSender,
info_hash: &[u8; 20],
peer_id: &[u8; 20],
stats: Arc<DownloadStats>,
store: Arc<W>,
) -> PeerConnectionResult<Self> {
let stream = TcpStream::connect(peer.address).await.map_err(|source| {
PeerConnectionError::ConnectFailed {
peer: Box::new(peer),
source,
}
})?;
Ok(PeerConnection {
stats,
peer,
is_outbound: true,
piece_manager_channel_sender,
stream: Some(stream),
info_hash: *info_hash,
peer_id: *peer_id,
store,
})
}
pub async fn from_stream(
stream: TcpStream,
address: SocketAddr,
piece_manager_channel_sender: PieceManagerChannelSender,
info_hash: &[u8; 20],
peer_id: &[u8; 20],
stats: Arc<DownloadStats>,
store: Arc<W>,
) -> Self {
PeerConnection {
stats,
peer: Peer {
peer_id: None,
address,
},
is_outbound: false,
piece_manager_channel_sender,
stream: Some(stream),
info_hash: *info_hash,
peer_id: *peer_id,
store,
}
}
pub async fn start(mut self) {
match self.run().await {
Ok(()) | Err(PeerConnectionError::PeerDisconnected) => {
debug!("{}: connection ended", peer_addr_direct(&self.peer));
}
Err(e) => warn!("{}: connection ended: {}", peer_addr_direct(&self.peer), e),
}
close(&Some(self.peer));
}
async fn run(&mut self) -> PeerConnectionResult<()> {
let stream = self.stream.take().unwrap();
stream.set_nodelay(true)?;
let mut framed = Framed::new(stream, WireCodec::new());
self.handshake(&mut framed).await?;
let snapshot = piece_manager_request(&mut self.piece_manager_channel_sender, |tx| {
PieceManagerMessage::GetBitfield {
response_sender: tx,
}
})
.await?;
framed
.send(WireItem::Message(Message::Bitfield(snapshot.bitfield)))
.await?;
let peer_bitfield: Bitfield;
select! {
_ = time::sleep(time::Duration::from_secs(30)) => {
return Err(PeerConnectionError::Timeout);
},
item = framed.next() => {
match item {
Some(Ok(WireItem::Message(Message::Bitfield(bitfield)))) => {
peer_bitfield = bitfield;
}
_ => {
return Err(PeerConnectionError::PeerDisconnected);
}
}
}
}
let (incoming_sender, incoming_receiver) = channels::new_incoming_channel();
let (outgoing_sender, mut outgoing_receiver) = channels::new_outgoing_channel();
let (mut sink, mut stream) = framed.split();
let mut joinset: JoinSet<()> = JoinSet::new();
joinset.spawn(async move {
while let Some(item) = outgoing_receiver.recv().await {
if let Err(e) = sink.send(item).await {
warn!("wire write failed: {e}");
break;
}
}
});
joinset.spawn(async move {
loop {
match stream.next().await {
Some(Ok(item)) => {
if incoming_sender.send(item).await.is_err() {
break;
}
}
Some(Err(e)) => {
warn!("wire decode failed: {e}");
break;
}
None => break,
}
}
});
let request_manager = request_manager::RequestManager::new(
Some(self.peer),
self.info_hash,
self.peer_id,
peer_bitfield,
self.piece_manager_channel_sender.clone(),
incoming_receiver,
outgoing_sender,
snapshot.events,
self.stats.clone(),
self.store.clone(),
);
joinset.spawn(async move {
request_manager.start().await;
});
joinset.join_all().await;
Ok(())
}
fn our_handshake(&self) -> Handshake {
Handshake {
pstrlen: 19,
pstr: "BitTorrent protocol".to_string(),
reserved: [0; 8],
info_hash: self.info_hash,
peer_id: self.peer_id,
}
}
pub async fn handshake(
&mut self,
framed: &mut Framed<TcpStream, WireCodec>,
) -> error::PeerConnectionResult<()> {
if self.is_outbound {
self.handshake_outbound(framed).await?;
} else {
self.handshake_inbound(framed).await?;
}
Ok(())
}
async fn handshake_outbound(
&mut self,
framed: &mut Framed<TcpStream, WireCodec>,
) -> error::PeerConnectionResult<()> {
framed
.send(WireItem::Handshake(self.our_handshake()))
.await?;
loop {
select! {
_ = time::sleep(time::Duration::from_secs(30)) => {
return Err(PeerConnectionError::Timeout);
},
item = framed.next() => {
match item {
Some(Ok(WireItem::HandshakePartial { info_hash })) => {
if info_hash != self.info_hash {
warn!(
"{}: handshake failed, info_hash mismatch",
peer_addr_direct(&self.peer)
);
return Err(error::PeerConnectionError::InfoHashMismatch);
}
}
Some(Ok(WireItem::Handshake(handshake))) => {
self.peer.peer_id = Some(handshake.peer_id);
debug!("Outbound handshake complete with {}", self.peer.address);
return Ok(());
}
_ => {
warn!(
"{}: handshake failed, unexpected message or connection closed",
peer_addr_direct(&self.peer)
);
return Err(error::PeerConnectionError::UnexpectedMessage);
}
}
},
}
}
}
async fn handshake_inbound(
&mut self,
framed: &mut Framed<TcpStream, WireCodec>,
) -> error::PeerConnectionResult<()> {
loop {
select! {
_ = time::sleep(time::Duration::from_secs(30)) => {
return Err(PeerConnectionError::Timeout);
},
item = framed.next() => {
match item {
Some(Ok(WireItem::HandshakePartial { info_hash })) => {
if info_hash != self.info_hash {
warn!(
"{}: handshake failed, unknown info_hash",
peer_addr_direct(&self.peer)
);
return Err(error::PeerConnectionError::InfoHashMismatch);
}
framed
.send(WireItem::Handshake(self.our_handshake()))
.await?;
debug!("Inbound handshake info_hash verified, replied with our handshake");
}
Some(Ok(WireItem::Handshake(handshake))) => {
let addr = framed.get_ref().peer_addr()?;
self.peer.peer_id = Some(handshake.peer_id);
self.peer.address = SocketAddr::new(addr.ip(), addr.port());
debug!("Inbound handshake complete with {}", addr);
return Ok(());
}
_ => {
warn!(
"{}: handshake failed, unexpected message or connection closed",
peer_addr_direct(&self.peer)
);
return Err(error::PeerConnectionError::UnexpectedMessage);
}
}
}
}
}
}
}
pub fn close(peer: &Option<Peer>) {
debug!("{}: closing connection", peer_addr(peer));
}
fn peer_addr(peer: &Option<Peer>) -> String {
match peer.as_ref() {
Some(peer) => format!("{}", peer.address),
None => "unknown".to_string(),
}
}
fn peer_addr_direct(peer: &Peer) -> String {
format!("{}", peer.address)
}
async fn piece_manager_request<T>(
piece_manager_channel_sender: &mut PieceManagerChannelSender,
build: impl FnOnce(oneshot::Sender<T>) -> PieceManagerMessage,
) -> error::PeerConnectionResult<T> {
let (tx, rx) = oneshot::channel();
piece_manager_channel_sender.send(build(tx)).await?;
Ok(rx.await?)
}