Files
aria2-rust-pro/crates/aria2-rust-pro-protocol/src/torrent/peer_wire/framing.rs
T

507 lines
19 KiB
Rust

use super::{
PEER_WIRE_BITFIELD_ID, PEER_WIRE_CANCEL_ID, PEER_WIRE_CHOKE_ID, PEER_WIRE_EXTENSION_ID,
PEER_WIRE_HAVE_ID, PEER_WIRE_INTERESTED_ID, PEER_WIRE_NOT_INTERESTED_ID, PEER_WIRE_PIECE_ID,
PEER_WIRE_PORT_ID, PEER_WIRE_REQUEST_ID, PEER_WIRE_UNCHOKE_ID, PeerWireBitfieldModel,
PeerWireBlockRequestModel, PeerWireExtensionMessageModel, PeerWireFrameHeaderModel,
PeerWireMessageKind, PeerWireMessageModel, PeerWirePieceBlockModel,
PeerWireUnknownMessageModel, TorrentMessageModel,
};
impl PeerWireMessageModel {
#[must_use]
/// Wraps one parsed torrent message with an optional peer id.
pub fn new(peer_id: Option<[u8; 20]>, message: TorrentMessageModel) -> Self {
Self { peer_id, message }
}
/// Parses one framed peer-wire message and reports the number of bytes consumed.
///
/// # Errors
///
/// Returns an error when the frame header or payload is malformed.
pub fn parse_frame(input: &[u8]) -> Result<(Self, usize), String> {
let (message, consumed) = TorrentMessageModel::parse_peer_wire_frame(input)?;
Ok((
Self {
peer_id: None,
message,
},
consumed,
))
}
/// Parses one complete framed peer-wire message.
///
/// # Errors
///
/// Returns an error when the frame is malformed or contains trailing bytes.
pub fn parse_frame_exact(input: &[u8]) -> Result<Self, String> {
let (message, consumed) = Self::parse_frame(input)?;
if consumed != input.len() {
return Err("trailing bytes after peer-wire frame".to_owned());
}
Ok(message)
}
/// Serializes the wrapped message as a framed peer-wire payload.
///
/// # Errors
///
/// Returns an error when the inner message cannot be represented as peer-wire bytes.
pub fn serialize_frame(&self) -> Result<Vec<u8>, String> {
self.message.serialize_peer_wire_frame()
}
}
impl PeerWireBitfieldModel {
#[must_use]
/// Builds a bitfield from per-piece completion flags.
pub fn from_piece_flags(flags: &[bool]) -> Self {
let mut bytes = vec![0_u8; flags.len().div_ceil(8)];
for (index, &present) in flags.iter().enumerate() {
if present {
bytes[index / 8] |= 1 << (7 - (index % 8));
}
}
Self { bytes }
}
#[must_use]
/// Returns the maximum number of pieces represented by this bitfield.
pub fn piece_capacity(&self) -> usize {
self.bytes.len() * 8
}
#[must_use]
/// Returns whether the bitfield marks `piece_index` as present.
pub fn has_piece(&self, piece_index: usize) -> bool {
let byte = piece_index / 8;
let bit = piece_index % 8;
self.bytes
.get(byte)
.is_some_and(|value| value & (1 << (7 - bit)) != 0)
}
#[must_use]
/// Expands the bitfield into per-piece completion flags.
pub fn to_piece_flags(&self, piece_count: usize) -> Vec<bool> {
(0..piece_count)
.map(|index| self.has_piece(index))
.collect()
}
}
impl TorrentMessageModel {
#[must_use]
/// Builds an internal torrent message from a peer-wire message variant.
pub fn from_peer_wire_kind(kind: PeerWireMessageKind) -> Self {
match kind {
PeerWireMessageKind::KeepAlive => Self {
message_type: "keepalive".to_owned(),
payload: Vec::new(),
},
PeerWireMessageKind::Choke => Self {
message_type: "choke".to_owned(),
payload: Vec::new(),
},
PeerWireMessageKind::Unchoke => Self {
message_type: "unchoke".to_owned(),
payload: Vec::new(),
},
PeerWireMessageKind::Interested => Self {
message_type: "interested".to_owned(),
payload: Vec::new(),
},
PeerWireMessageKind::NotInterested => Self {
message_type: "not_interested".to_owned(),
payload: Vec::new(),
},
PeerWireMessageKind::Have(piece_index) => Self {
message_type: "have".to_owned(),
payload: piece_index.to_be_bytes().to_vec(),
},
PeerWireMessageKind::Bitfield(bitfield) => Self {
message_type: "bitfield".to_owned(),
payload: bitfield.bytes,
},
PeerWireMessageKind::Request(request) => Self {
message_type: "request".to_owned(),
payload: encode_block_request_payload(&request),
},
PeerWireMessageKind::Piece(piece) => Self {
message_type: "piece".to_owned(),
payload: encode_piece_payload(&piece),
},
PeerWireMessageKind::Cancel(request) => Self {
message_type: "cancel".to_owned(),
payload: encode_block_request_payload(&request),
},
PeerWireMessageKind::Port(port) => Self {
message_type: "port".to_owned(),
payload: port.to_be_bytes().to_vec(),
},
PeerWireMessageKind::Extension(extension) => {
let mut payload = Vec::with_capacity(1 + extension.payload.len());
payload.push(extension.extension_message_id);
payload.extend_from_slice(&extension.payload);
Self {
message_type: "extension".to_owned(),
payload,
}
}
PeerWireMessageKind::Unknown(message) => Self {
message_type: format!("unknown:{}", message.message_id),
payload: message.payload,
},
}
}
/// Reconstructs the typed peer-wire message kind from the internal message payload.
///
/// # Errors
///
/// Returns an error when the message type or payload shape is unsupported.
pub fn peer_wire_kind(&self) -> Result<PeerWireMessageKind, String> {
peer_wire_kind_from_raw(&self.message_type, &self.payload)
}
/// Inspects a framed peer-wire message header without fully decoding the payload.
///
/// # Errors
///
/// Returns an error when the frame is truncated or malformed.
pub fn inspect_peer_wire_frame(
input: &[u8],
) -> Result<(PeerWireFrameHeaderModel, usize), String> {
if input.len() < 4 {
return Err("truncated peer-wire frame: missing length prefix".to_owned());
}
let frame_len =
usize::try_from(u32::from_be_bytes([input[0], input[1], input[2], input[3]]))
.map_err(|_| "peer-wire frame length does not fit usize".to_owned())?;
let total_len = 4_usize
.checked_add(frame_len)
.ok_or_else(|| "peer-wire frame length overflow".to_owned())?;
if input.len() < total_len {
return Err(format!(
"truncated peer-wire frame: expected {total_len} bytes, got {}",
input.len()
));
}
if frame_len == 0 {
return Ok((
PeerWireFrameHeaderModel {
message_id: None,
payload_len: 0,
},
total_len,
));
}
let message_id = input[4];
Ok((
PeerWireFrameHeaderModel {
message_id: Some(message_id),
payload_len: frame_len - 1,
},
total_len,
))
}
/// Parses a framed peer-wire message and reports the number of consumed bytes.
///
/// # Errors
///
/// Returns an error when the frame is truncated or malformed.
pub fn parse_peer_wire_frame(input: &[u8]) -> Result<(Self, usize), String> {
let (header, consumed) = Self::inspect_peer_wire_frame(input)?;
let Some(message_id) = header.message_id else {
return Ok((
Self::from_peer_wire_kind(PeerWireMessageKind::KeepAlive),
consumed,
));
};
let payload = &input[5..consumed];
let kind = peer_wire_kind_from_message_id(message_id, payload)?;
Ok((Self::from_peer_wire_kind(kind), consumed))
}
/// Parses one complete framed peer-wire message.
///
/// # Errors
///
/// Returns an error when the frame is truncated, malformed, or has trailing bytes.
pub fn parse_peer_wire_frame_exact(input: &[u8]) -> Result<Self, String> {
let (message, consumed) = Self::parse_peer_wire_frame(input)?;
if consumed != input.len() {
return Err("trailing bytes after peer-wire frame".to_owned());
}
Ok(message)
}
/// Serializes the message as a framed peer-wire payload.
///
/// # Errors
///
/// Returns an error when the message cannot be represented as a supported peer-wire frame.
pub fn serialize_peer_wire_frame(&self) -> Result<Vec<u8>, String> {
let kind = self.peer_wire_kind()?;
serialize_peer_wire_kind(&kind)
}
}
/// Serializes one peer-wire message kind into a framed peer-wire payload.
fn serialize_peer_wire_kind(kind: &PeerWireMessageKind) -> Result<Vec<u8>, String> {
let mut payload = Vec::new();
let message_id = match kind {
PeerWireMessageKind::KeepAlive => None,
PeerWireMessageKind::Choke => Some(PEER_WIRE_CHOKE_ID),
PeerWireMessageKind::Unchoke => Some(PEER_WIRE_UNCHOKE_ID),
PeerWireMessageKind::Interested => Some(PEER_WIRE_INTERESTED_ID),
PeerWireMessageKind::NotInterested => Some(PEER_WIRE_NOT_INTERESTED_ID),
PeerWireMessageKind::Have(piece_index) => {
payload.extend_from_slice(&piece_index.to_be_bytes());
Some(PEER_WIRE_HAVE_ID)
}
PeerWireMessageKind::Bitfield(bitfield) => {
payload.extend_from_slice(&bitfield.bytes);
Some(PEER_WIRE_BITFIELD_ID)
}
PeerWireMessageKind::Request(request) => {
payload.extend_from_slice(&encode_block_request_payload(request));
Some(PEER_WIRE_REQUEST_ID)
}
PeerWireMessageKind::Piece(piece) => {
payload.extend_from_slice(&encode_piece_payload(piece));
Some(PEER_WIRE_PIECE_ID)
}
PeerWireMessageKind::Cancel(request) => {
payload.extend_from_slice(&encode_block_request_payload(request));
Some(PEER_WIRE_CANCEL_ID)
}
PeerWireMessageKind::Port(port) => {
payload.extend_from_slice(&port.to_be_bytes());
Some(PEER_WIRE_PORT_ID)
}
PeerWireMessageKind::Extension(extension) => {
payload.push(extension.extension_message_id);
payload.extend_from_slice(&extension.payload);
Some(PEER_WIRE_EXTENSION_ID)
}
PeerWireMessageKind::Unknown(message) => {
payload.extend_from_slice(&message.payload);
Some(message.message_id)
}
};
let Some(message_id) = message_id else {
return Ok(vec![0, 0, 0, 0]);
};
let frame_len = 1 + payload.len();
let frame_len_u32 = u32::try_from(frame_len)
.map_err(|_| "peer-wire frame exceeds u32 length prefix".to_owned())?;
let mut bytes = Vec::with_capacity(4 + frame_len);
bytes.extend_from_slice(&frame_len_u32.to_be_bytes());
bytes.push(message_id);
bytes.extend_from_slice(&payload);
Ok(bytes)
}
/// Interprets one internal message-type label and payload as a peer-wire message kind.
fn peer_wire_kind_from_raw(
message_type: &str,
payload: &[u8],
) -> Result<PeerWireMessageKind, String> {
match message_type {
"keepalive" => {
expect_empty_payload("keepalive", payload).map(|()| PeerWireMessageKind::KeepAlive)
}
"choke" => expect_empty_payload("choke", payload).map(|()| PeerWireMessageKind::Choke),
"unchoke" => {
expect_empty_payload("unchoke", payload).map(|()| PeerWireMessageKind::Unchoke)
}
"interested" => {
expect_empty_payload("interested", payload).map(|()| PeerWireMessageKind::Interested)
}
"not_interested" | "not-interested" => expect_empty_payload("not_interested", payload)
.map(|()| PeerWireMessageKind::NotInterested),
"have" => parse_have_payload(payload),
"bitfield" => Ok(PeerWireMessageKind::Bitfield(PeerWireBitfieldModel {
bytes: payload.to_vec(),
})),
"request" => {
parse_block_request_payload("request", payload).map(PeerWireMessageKind::Request)
}
"piece" => parse_piece_payload(payload).map(PeerWireMessageKind::Piece),
"cancel" => parse_block_request_payload("cancel", payload).map(PeerWireMessageKind::Cancel),
"port" => parse_port_payload(payload),
"extension" => parse_extension_payload(payload),
_ => parse_unknown_message_id(message_type).map_or_else(
|| {
Err(format!(
"unsupported peer-wire message type: {message_type}"
))
},
|message_id| {
Ok(PeerWireMessageKind::Unknown(PeerWireUnknownMessageModel {
message_id,
payload: payload.to_vec(),
}))
},
),
}
}
/// Interprets a peer-wire message id and payload as a typed peer-wire message kind.
fn peer_wire_kind_from_message_id(
message_id: u8,
payload: &[u8],
) -> Result<PeerWireMessageKind, String> {
match message_id {
PEER_WIRE_CHOKE_ID => {
expect_empty_payload("choke", payload).map(|()| PeerWireMessageKind::Choke)
}
PEER_WIRE_UNCHOKE_ID => {
expect_empty_payload("unchoke", payload).map(|()| PeerWireMessageKind::Unchoke)
}
PEER_WIRE_INTERESTED_ID => {
expect_empty_payload("interested", payload).map(|()| PeerWireMessageKind::Interested)
}
PEER_WIRE_NOT_INTERESTED_ID => expect_empty_payload("not_interested", payload)
.map(|()| PeerWireMessageKind::NotInterested),
PEER_WIRE_HAVE_ID => parse_have_payload(payload),
PEER_WIRE_BITFIELD_ID => Ok(PeerWireMessageKind::Bitfield(PeerWireBitfieldModel {
bytes: payload.to_vec(),
})),
PEER_WIRE_REQUEST_ID => {
parse_block_request_payload("request", payload).map(PeerWireMessageKind::Request)
}
PEER_WIRE_PIECE_ID => parse_piece_payload(payload).map(PeerWireMessageKind::Piece),
PEER_WIRE_CANCEL_ID => {
parse_block_request_payload("cancel", payload).map(PeerWireMessageKind::Cancel)
}
PEER_WIRE_PORT_ID => parse_port_payload(payload),
PEER_WIRE_EXTENSION_ID => parse_extension_payload(payload),
_ => Ok(PeerWireMessageKind::Unknown(PeerWireUnknownMessageModel {
message_id,
payload: payload.to_vec(),
})),
}
}
/// Verifies that a peer-wire control payload is empty.
fn expect_empty_payload(name: &str, payload: &[u8]) -> Result<(), String> {
if payload.is_empty() {
Ok(())
} else {
Err(format!(
"peer-wire {name} payload must be empty, got {} bytes",
payload.len()
))
}
}
/// Parses a `have` payload into its piece index variant.
fn parse_have_payload(payload: &[u8]) -> Result<PeerWireMessageKind, String> {
let piece_index = read_u32(payload, "have", 0)?;
Ok(PeerWireMessageKind::Have(piece_index))
}
/// Parses a `request` or `cancel` payload into block coordinates.
fn parse_block_request_payload(
name: &str,
payload: &[u8],
) -> Result<PeerWireBlockRequestModel, String> {
if payload.len() != 12 {
return Err(format!(
"peer-wire {name} payload must be 12 bytes, got {}",
payload.len()
));
}
Ok(PeerWireBlockRequestModel {
piece_index: read_u32(payload, name, 0)?,
block_offset: read_u32(payload, name, 4)?,
block_length: read_u32(payload, name, 8)?,
})
}
/// Parses a `piece` payload into block coordinates plus data.
fn parse_piece_payload(payload: &[u8]) -> Result<PeerWirePieceBlockModel, String> {
if payload.len() < 8 {
return Err(format!(
"peer-wire piece payload must be at least 8 bytes, got {}",
payload.len()
));
}
Ok(PeerWirePieceBlockModel {
piece_index: read_u32(payload, "piece", 0)?,
block_offset: read_u32(payload, "piece", 4)?,
block: payload[8..].to_vec(),
})
}
/// Parses a `port` payload into the corresponding peer-wire message variant.
fn parse_port_payload(payload: &[u8]) -> Result<PeerWireMessageKind, String> {
if payload.len() != 2 {
return Err(format!(
"peer-wire port payload must be 2 bytes, got {}",
payload.len()
));
}
Ok(PeerWireMessageKind::Port(u16::from_be_bytes([
payload[0], payload[1],
])))
}
/// Parses an extension-protocol payload into the typed extension message variant.
fn parse_extension_payload(payload: &[u8]) -> Result<PeerWireMessageKind, String> {
let Some((&extension_message_id, rest)) = payload.split_first() else {
return Err("peer-wire extension payload must include extension message id".to_owned());
};
Ok(PeerWireMessageKind::Extension(
PeerWireExtensionMessageModel {
extension_message_id,
payload: rest.to_vec(),
},
))
}
/// Parses a synthetic `unknown:<id>` message-type label into a raw peer-wire id.
fn parse_unknown_message_id(message_type: &str) -> Option<u8> {
message_type
.strip_prefix("unknown:")
.and_then(|value| value.parse::<u8>().ok())
}
/// Encodes request or cancel block coordinates into peer-wire payload bytes.
fn encode_block_request_payload(request: &PeerWireBlockRequestModel) -> Vec<u8> {
let mut payload = Vec::with_capacity(12);
payload.extend_from_slice(&request.piece_index.to_be_bytes());
payload.extend_from_slice(&request.block_offset.to_be_bytes());
payload.extend_from_slice(&request.block_length.to_be_bytes());
payload
}
/// Encodes a piece block into peer-wire payload bytes.
fn encode_piece_payload(piece: &PeerWirePieceBlockModel) -> Vec<u8> {
let mut payload = Vec::with_capacity(8 + piece.block.len());
payload.extend_from_slice(&piece.piece_index.to_be_bytes());
payload.extend_from_slice(&piece.block_offset.to_be_bytes());
payload.extend_from_slice(&piece.block);
payload
}
/// Reads one big-endian `u32` from a peer-wire payload.
fn read_u32(payload: &[u8], name: &str, start: usize) -> Result<u32, String> {
let end = start + 4;
let bytes = payload
.get(start..end)
.ok_or_else(|| format!("peer-wire {name} payload truncated at byte offset {start}"))?;
Ok(u32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]))
}