use bytes::{Buf, BufMut, BytesMut};
use speedy::{BigEndian, Readable, Writable};
use std::io::Cursor;
use tokio::io;
use tokio_util::codec::{Decoder, Encoder};
use tracing::warn;
use crate::{bitfield::Bitfield, error::Error};
use super::{Block, BlockInfo, PSTR};
#[derive(Debug, Clone, PartialEq)]
pub enum Message {
KeepAlive,
Bitfield(Bitfield),
Choke,
Unchoke,
Interested,
NotInterested,
Have(usize),
Request(BlockInfo),
Piece(Block),
Cancel(BlockInfo),
Extended((u8, Vec<u8>)),
}
#[repr(u8)]
#[derive(Copy, Clone, Debug, PartialEq)]
pub enum MessageId {
Choke = 0,
Unchoke = 1,
Interested = 2,
NotInterested = 3,
Have = 4,
Bitfield = 5,
Request = 6,
Piece = 7,
Cancel = 8,
Extended = 20,
}
impl TryFrom<u8> for MessageId {
type Error = io::Error;
fn try_from(k: u8) -> Result<Self, Self::Error> {
use MessageId::*;
match k {
k if k == Choke as u8 => Ok(Choke),
k if k == Unchoke as u8 => Ok(Unchoke),
k if k == Interested as u8 => Ok(Interested),
k if k == NotInterested as u8 => Ok(NotInterested),
k if k == Have as u8 => Ok(Have),
k if k == Bitfield as u8 => Ok(Bitfield),
k if k == Request as u8 => Ok(Request),
k if k == Piece as u8 => Ok(Piece),
k if k == Cancel as u8 => Ok(Cancel),
k if k == Extended as u8 => Ok(Extended),
_ => Err(io::Error::new(
io::ErrorKind::InvalidInput,
"Unknown message id",
)),
}
}
}
#[derive(Debug)]
pub struct PeerCodec;
impl Encoder<Message> for PeerCodec {
type Error = io::Error;
fn encode(
&mut self,
item: Message,
buf: &mut BytesMut,
) -> Result<(), Self::Error> {
match item {
Message::KeepAlive => {
buf.put_u32(0);
}
Message::Bitfield(bitfield) => {
let v = bitfield.into_vec();
buf.put_u32(1 + v.len() as u32);
buf.put_u8(MessageId::Bitfield as u8);
buf.extend_from_slice(&v);
}
Message::Choke => {
buf.put_u32(1);
buf.put_u8(MessageId::Choke as u8);
}
Message::Unchoke => {
buf.put_u32(1);
buf.put_u8(MessageId::Unchoke as u8);
}
Message::Interested => {
buf.put_u32(1);
buf.put_u8(MessageId::Interested as u8);
}
Message::NotInterested => {
buf.put_u32(1);
buf.put_u8(MessageId::NotInterested as u8);
}
Message::Have(piece_index) => {
let msg_len = 1 + 4;
buf.put_u32(msg_len);
buf.put_u8(MessageId::Have as u8);
let piece_index = piece_index.try_into().map_err(|e| {
io::Error::new(io::ErrorKind::InvalidInput, e)
})?;
buf.put_u32(piece_index);
}
Message::Request(block) => {
let msg_len = 1 + 4 + 4 + 4;
buf.put_u32(msg_len);
buf.put_u8(MessageId::Request as u8);
block.encode(buf)?;
}
Message::Piece(block) => {
let Block { index, begin, block } = block;
let msg_len = 1 + 4 + 4 + block.len() as u32;
buf.put_u32(msg_len);
buf.put_u8(MessageId::Piece as u8);
let index = index.try_into().map_err(|e| {
io::Error::new(io::ErrorKind::InvalidInput, e)
})?;
buf.put_u32(index);
buf.put_u32(begin);
buf.put(&block[..]);
}
Message::Cancel(block) => {
let msg_len = 1 + 4 + 4 + 4;
buf.put_u32(msg_len);
buf.put_u8(MessageId::Cancel as u8);
block.encode(buf)?;
}
Message::Extended((ext_id, payload)) => {
let msg_len = payload.len() as u32 + 2;
buf.put_u32(msg_len);
buf.put_u8(MessageId::Extended as u8);
buf.put_u8(ext_id);
if !payload.is_empty() {
buf.extend_from_slice(&payload);
}
}
}
Ok(())
}
}
impl Decoder for PeerCodec {
type Item = Message;
type Error = io::Error;
fn decode(
&mut self,
buf: &mut BytesMut,
) -> Result<Option<Self::Item>, Self::Error> {
if buf.remaining() < 4 {
return Ok(None);
}
let mut tmp_buf = Cursor::new(&buf);
let msg_len = tmp_buf.get_u32() as usize;
tmp_buf.set_position(0);
if buf.remaining() >= 4 + msg_len {
buf.advance(4);
if msg_len == 0 {
return Ok(Some(Message::KeepAlive));
}
} else {
tracing::trace!(
"Read buffer is {} bytes long but message is {} bytes long",
buf.remaining(),
msg_len
);
return Ok(None);
}
let msg_id = MessageId::try_from(buf.get_u8())?;
let msg = match msg_id {
MessageId::Choke => Message::Choke,
MessageId::Unchoke => Message::Unchoke,
MessageId::Interested => Message::Interested,
MessageId::NotInterested => Message::NotInterested,
MessageId::Have => {
let piece_index = buf.get_u32();
Message::Have(piece_index as usize)
}
MessageId::Bitfield => {
let mut bitfield = vec![0; msg_len - 1];
buf.copy_to_slice(&mut bitfield);
Message::Bitfield(Bitfield::from_vec(bitfield))
}
MessageId::Request => {
let index = buf.get_u32();
let begin = buf.get_u32();
let len = buf.get_u32();
Message::Request(BlockInfo { index, begin, len })
}
MessageId::Piece => {
let index = buf.get_u32() as usize;
let begin = buf.get_u32();
let mut block = vec![0; msg_len - 9];
buf.copy_to_slice(&mut block);
Message::Piece(Block { index, begin, block })
}
MessageId::Cancel => {
let index = buf.get_u32();
let begin = buf.get_u32();
let len = buf.get_u32();
Message::Cancel(BlockInfo { index, begin, len })
}
MessageId::Extended => {
let ext_id = buf.get_u8();
let mut payload = vec![0u8; msg_len - 2];
buf.copy_to_slice(&mut payload);
Message::Extended((ext_id, payload))
}
};
Ok(Some(msg))
}
}
pub const PROTOCOL_STRING: &str = "BitTorrent protocol";
#[derive(Debug)]
pub struct HandshakeCodec;
impl Encoder<Handshake> for HandshakeCodec {
type Error = io::Error;
fn encode(
&mut self,
handshake: Handshake,
buf: &mut BytesMut,
) -> io::Result<()> {
let Handshake { pstr_len, pstr, reserved, info_hash, peer_id } =
handshake;
debug_assert_eq!(pstr_len, 19);
buf.put_u8(pstr.len() as u8);
debug_assert_eq!(pstr, PROTOCOL_STRING.as_bytes());
buf.extend_from_slice(&pstr);
buf.extend_from_slice(&reserved);
buf.extend_from_slice(&info_hash);
buf.extend_from_slice(&peer_id);
Ok(())
}
}
impl Decoder for HandshakeCodec {
type Item = Handshake;
type Error = io::Error;
fn decode(&mut self, buf: &mut BytesMut) -> io::Result<Option<Handshake>> {
if buf.is_empty() {
return Ok(None);
}
let mut tmp_buf = Cursor::new(&buf);
let prot_len = tmp_buf.get_u8() as usize;
if prot_len != PROTOCOL_STRING.as_bytes().len() {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"Handshake must have the string \"BitTorrent protocol\"",
));
}
let payload_len = prot_len + 8 + 20 + 20;
if buf.remaining() > payload_len {
buf.advance(1);
} else {
return Ok(None);
}
let mut pstr = [0; 19];
buf.copy_to_slice(&mut pstr);
let mut reserved = [0; 8];
buf.copy_to_slice(&mut reserved);
let mut info_hash = [0; 20];
buf.copy_to_slice(&mut info_hash);
let mut peer_id = [0; 20];
buf.copy_to_slice(&mut peer_id);
Ok(Some(Handshake {
pstr,
pstr_len: pstr.len() as u8,
reserved,
info_hash,
peer_id,
}))
}
}
#[derive(Clone, Debug, Writable, Readable)]
pub struct Handshake {
pub pstr_len: u8,
pub pstr: [u8; 19],
pub reserved: [u8; 8],
pub info_hash: [u8; 20],
pub peer_id: [u8; 20],
}
impl Handshake {
pub fn new(info_hash: [u8; 20], peer_id: [u8; 20]) -> Self {
let mut reserved = [0u8; 8];
reserved[5] |= 0x10;
Self {
pstr_len: u8::to_be(19),
pstr: PSTR,
reserved,
info_hash,
peer_id,
}
}
pub fn serialize(&self) -> Result<[u8; 68], Error> {
let mut buf: [u8; 68] = [0u8; 68];
let temp = self
.write_to_vec_with_ctx(BigEndian {})
.map_err(Error::SpeedyError)?;
buf.copy_from_slice(&temp[..]);
Ok(buf)
}
pub fn deserialize(buf: &[u8]) -> Result<Self, Error> {
Self::read_from_buffer_with_ctx(BigEndian {}, buf)
.map_err(Error::SpeedyError)
}
pub fn validate(&self, target: &Self) -> bool {
if target.peer_id.len() != 20 {
warn!("! invalid peer_id from receiving handshake");
return false;
}
if self.info_hash != target.info_hash {
warn!("! info_hash from receiving handshake does not match ours");
return false;
}
if target.pstr_len != 19 {
warn!("! handshake with wrong pstr_len, dropping connection");
return false;
}
if target.pstr != PSTR {
warn!("! handshake with wrong pstr, dropping connection");
return false;
}
true
}
}
#[cfg(test)]
mod tests {
use crate::tcp_wire::BLOCK_LEN;
use super::*;
use bitvec::{bitvec, prelude::Msb0};
use bytes::{Buf, BytesMut};
use tokio_util::codec::{Decoder, Encoder};
#[test]
fn extended() {
let mut buf = BytesMut::new();
let msg = Message::Extended((0, vec![]));
PeerCodec.encode(msg.clone(), &mut buf).unwrap();
assert_eq!(buf.len(), 6);
assert_eq!(buf.get_u32(), 2);
assert_eq!(buf.get_u8(), MessageId::Extended as u8);
assert_eq!(buf.get_u8(), 0);
let mut buf = BytesMut::new();
PeerCodec.encode(msg, &mut buf).unwrap();
let msg = PeerCodec.decode(&mut buf).unwrap().unwrap();
match msg {
Message::Extended((ext_id, _payload)) => {
assert_eq!(ext_id, 0);
}
_ => panic!(),
}
}
#[test]
fn bitfield() {
let mut buf = BytesMut::new();
let mut original = bitvec![u8, Msb0; 0; 10];
original.set(8, true);
original.set(9, true);
let msg = Message::Bitfield(original.clone());
PeerCodec.encode(msg.clone(), &mut buf).unwrap();
assert_eq!(buf.get_u32(), 1 + original.clone().into_vec().len() as u32);
assert_eq!(buf.get_u8(), MessageId::Bitfield as u8);
let mut buf = BytesMut::new();
PeerCodec.encode(msg, &mut buf).unwrap();
let msg = PeerCodec.decode(&mut buf).unwrap().unwrap();
match msg {
Message::Bitfield(mut bitfield) => {
unsafe {
bitfield.set_len(original.len());
}
assert_eq!(bitfield, original);
}
_ => panic!(),
}
}
#[test]
fn request() {
let mut buf = BytesMut::new();
let msg = Message::Request(BlockInfo::default());
PeerCodec.encode(msg.clone(), &mut buf).unwrap();
assert_eq!(buf.len(), 17);
assert_eq!(buf.get_u32(), 13);
assert_eq!(buf.get_u8(), MessageId::Request as u8);
assert_eq!(buf.get_u32(), 0);
assert_eq!(buf.get_u32(), 0);
assert_eq!(buf.get_u32(), BLOCK_LEN);
let mut buf = BytesMut::new();
PeerCodec.encode(msg, &mut buf).unwrap();
let msg = PeerCodec.decode(&mut buf).unwrap().unwrap();
match msg {
Message::Request(block_info) => {
assert_eq!(block_info.index, 0);
assert_eq!(block_info.begin, 0);
assert_eq!(block_info.len, BLOCK_LEN);
}
_ => panic!(),
}
}
#[test]
fn piece() {
let mut buf = BytesMut::new();
let msg = Message::Piece(Block { index: 0, begin: 0, block: vec![0] });
PeerCodec.encode(msg.clone(), &mut buf).unwrap();
assert_eq!(buf.get_u32(), 9 + 1);
assert_eq!(buf.get_u8(), MessageId::Piece as u8);
assert_eq!(buf.get_u32(), 0);
assert_eq!(buf.get_u32(), 0);
let mut block = BytesMut::new();
buf.copy_to_slice(&mut block);
assert_eq!(block.len(), 0);
let mut buf = BytesMut::new();
PeerCodec.encode(msg.clone(), &mut buf).unwrap();
let msg = PeerCodec.decode(&mut buf).unwrap().unwrap();
match msg {
Message::Piece(block) => {
assert_eq!(block.index, 0);
assert_eq!(block.begin, 0);
assert_eq!(block.block.len(), 1);
}
_ => panic!(),
}
}
#[test]
fn handshake() {
let info_hash = [5u8; 20];
let peer_id = [7u8; 20];
let our_handshake = Handshake::new(info_hash, peer_id);
assert_eq!(our_handshake.pstr_len, 19);
assert_eq!(our_handshake.pstr, PSTR);
assert_eq!(our_handshake.peer_id, peer_id);
assert_eq!(our_handshake.info_hash, info_hash);
let our_handshake =
Handshake::new(info_hash, peer_id).serialize().unwrap();
assert_eq!(
our_handshake,
[
19, 66, 105, 116, 84, 111, 114, 114, 101, 110, 116, 32, 112,
114, 111, 116, 111, 99, 111, 108, 0, 0, 0, 0, 0, 16, 0, 0, 5,
5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 7, 7,
7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7
]
);
}
#[test]
fn reserved_bytes() {
let reserved = Bitfield::from_vec(vec![0, 0, 0, 0, 0, 16, 0, 0]);
assert_eq!(reserved.clone().into_vec(), [0, 0, 0, 0, 0, 16, 0, 0]);
let support_extension_protocol = reserved[43];
assert!(support_extension_protocol)
}
}