use super::*;
use chrono::Utc;
use dragonfly_api::common::v2::{Hdfs, HuggingFace, ModelScope, ObjectStorage, Range, TrafficType};
use dragonfly_client_backend::{BackendFactory, GetRequest};
use dragonfly_client_config::dfdaemon::Config;
use dragonfly_client_core::{error::BackendError, Error, Result};
use dragonfly_client_metric::{
collect_backend_request_failure_metrics, collect_backend_request_finished_metrics,
collect_backend_request_started_metrics, collect_download_piece_traffic_metrics,
};
use dragonfly_client_storage::{content::RangeReader, metadata, Storage};
use dragonfly_client_util::net::format_socket_addr;
use leaky_bucket::RateLimiter;
use reqwest::header::HeaderMap;
use std::collections::HashMap;
use std::net::IpAddr;
use std::str::FromStr;
use std::sync::Arc;
use std::time::Instant;
use tokio::io::{AsyncBufRead, AsyncRead, AsyncReadExt};
use tracing::{debug, error, instrument, warn, Span};
pub const MAX_PIECE_COUNT: u64 = 500;
pub const MIN_PIECE_LENGTH: u64 = 4 * 1024 * 1024;
pub const MAX_PIECE_LENGTH: u64 = 64 * 1024 * 1024;
pub enum PieceLengthStrategy {
OptimizeByFileLength(u64),
FixedPieceLength(u64),
}
pub struct Piece {
config: Arc<Config>,
storage: Arc<Storage>,
tcp_downloader: Arc<dyn piece_downloader::Downloader>,
quic_downloader: Arc<dyn piece_downloader::Downloader>,
backend_factory: Arc<BackendFactory>,
download_bandwidth_limiter: Arc<RateLimiter>,
prefetch_bandwidth_limiter: Arc<RateLimiter>,
back_to_source_bandwidth_limiter: Arc<RateLimiter>,
}
impl Piece {
pub fn new(
config: Arc<Config>,
storage: Arc<Storage>,
backend_factory: Arc<BackendFactory>,
download_bandwidth_limiter: Arc<RateLimiter>,
prefetch_bandwidth_limiter: Arc<RateLimiter>,
back_to_source_bandwidth_limiter: Arc<RateLimiter>,
) -> Result<Self> {
Ok(Self {
config: config.clone(),
storage,
tcp_downloader: piece_downloader::DownloaderFactory::new("tcp", config.clone())?
.build(),
quic_downloader: piece_downloader::DownloaderFactory::new("quic", config)?.build(),
backend_factory,
download_bandwidth_limiter,
prefetch_bandwidth_limiter,
back_to_source_bandwidth_limiter,
})
}
#[inline]
pub fn id(&self, task_id: &str, number: u32) -> String {
self.storage.piece_id(task_id, number)
}
pub fn get(&self, piece_id: &str) -> Result<Option<metadata::Piece>> {
self.storage.get_piece(piece_id)
}
pub fn get_all(&self, task_id: &str) -> Result<Vec<metadata::Piece>> {
self.storage.get_pieces(task_id)
}
pub fn calculate_interested(
&self,
piece_length: u64,
content_length: u64,
range: Option<Range>,
) -> Result<Vec<metadata::Piece>> {
if content_length == 0 {
return Ok(Vec::new());
}
if let Some(range) = range {
if range.length == 0 {
error!("range length is 0");
return Err(Error::InvalidParameter);
}
let mut number = 0;
let mut offset = 0;
let mut pieces: Vec<metadata::Piece> = Vec::new();
loop {
if offset >= content_length {
let mut piece = pieces.pop().ok_or_else(|| {
error!("piece not found");
Error::InvalidParameter
})?;
piece.length = piece_length + content_length - offset;
pieces.push(piece);
break;
}
if offset >= range.start + range.length {
break;
}
if offset + piece_length > range.start {
pieces.push(metadata::Piece {
number: number as u32,
offset,
length: piece_length,
digest: "".to_string(),
parent_id: None,
uploading_count: 0,
uploaded_count: 0,
updated_at: Utc::now().naive_utc(),
created_at: Utc::now().naive_utc(),
finished_at: None,
});
}
offset = (number + 1) * piece_length;
number += 1;
}
debug!(
"calculate interested pieces by range {:?}, content length: {:?}, piece length: {:?}, pieces: count {}, range [{}, {}]",
range,
content_length,
piece_length,
pieces.len(),
pieces.first().map(|piece| piece.number).unwrap_or_default(),
pieces.last().map(|piece| piece.number).unwrap_or_default()
);
return Ok(pieces);
}
let mut number = 0;
let mut offset = 0;
let mut pieces: Vec<metadata::Piece> = Vec::new();
loop {
if offset >= content_length {
let mut piece =
pieces
.pop()
.ok_or(Error::InvalidParameter)
.inspect_err(|_err| {
error!("piece not found");
})?;
piece.length = piece_length + content_length - offset;
pieces.push(piece);
break;
}
pieces.push(metadata::Piece {
number: number as u32,
offset,
length: piece_length,
digest: "".to_string(),
parent_id: None,
uploading_count: 0,
uploaded_count: 0,
updated_at: Utc::now().naive_utc(),
created_at: Utc::now().naive_utc(),
finished_at: None,
});
offset = (number + 1) * piece_length;
number += 1;
}
debug!(
"calculate interested pieces by content length: {:?}, piece length: {:?}, pieces: count {}, range [{}, {}]",
content_length,
piece_length,
pieces.len(),
pieces.first().map(|piece| piece.number).unwrap_or_default(),
pieces.last().map(|piece| piece.number).unwrap_or_default()
);
Ok(pieces)
}
#[instrument(level = "debug", skip_all)]
pub fn remove_finished_from_interested(
&self,
finished_pieces: Vec<metadata::Piece>,
interested_pieces: Vec<metadata::Piece>,
) -> Vec<metadata::Piece> {
interested_pieces
.iter()
.filter(|piece| {
!finished_pieces
.iter()
.any(|finished_piece| finished_piece.number == piece.number)
})
.cloned()
.collect::<Vec<metadata::Piece>>()
}
#[instrument(level = "debug", skip_all)]
pub fn merge_finished_pieces(
&self,
finished_pieces: Vec<metadata::Piece>,
old_finished_pieces: Vec<metadata::Piece>,
) -> Vec<metadata::Piece> {
let mut pieces: HashMap<u32, metadata::Piece> = HashMap::new();
for finished_piece in finished_pieces.into_iter() {
pieces.insert(finished_piece.number, finished_piece);
}
for old_finished_piece in old_finished_pieces.into_iter() {
pieces
.entry(old_finished_piece.number)
.or_insert(old_finished_piece);
}
pieces.into_values().collect()
}
pub fn calculate_piece_length(&self, strategy: PieceLengthStrategy) -> u64 {
match strategy {
PieceLengthStrategy::OptimizeByFileLength(content_length) => {
let piece_length = (content_length as f64 / MAX_PIECE_COUNT as f64) as u64;
let actual_piece_length = piece_length.next_power_of_two();
match (
actual_piece_length > MIN_PIECE_LENGTH,
actual_piece_length < MAX_PIECE_LENGTH,
) {
(true, true) => actual_piece_length,
(_, false) => MAX_PIECE_LENGTH,
(false, _) => MIN_PIECE_LENGTH,
}
}
PieceLengthStrategy::FixedPieceLength(piece_length) => piece_length,
}
}
pub fn calculate_piece_count(&self, piece_length: u64, content_length: u64) -> u32 {
(content_length as f64 / piece_length as f64).ceil() as u32
}
#[instrument(level = "debug", skip_all, fields(piece_id))]
pub async fn download_from_local_into_range_reader(
&self,
piece_id: &str,
task_id: &str,
length: u64,
range: Option<Range>,
) -> Result<RangeReader> {
Span::current().record("piece_id", piece_id);
Span::current().record("piece_length", length);
self.storage.upload_piece(piece_id, task_id, range).await
}
#[instrument(level = "debug", skip_all)]
pub fn download_from_local(&self, length: u64) {
collect_download_piece_traffic_metrics(&TrafficType::LocalPeer, length);
}
#[allow(clippy::too_many_arguments)]
#[instrument(skip_all, fields(piece_id))]
pub async fn download_from_parent(
&self,
piece_id: &str,
host_id: &str,
task_id: &str,
number: u32,
length: u64,
parent: piece_collector::CollectedParent,
is_prefetch: bool,
) -> Result<metadata::Piece> {
Span::current().record("piece_id", piece_id);
Span::current().record("piece_length", length);
let piece = self
.storage
.download_piece_started(piece_id, number)
.await?;
if piece.is_finished() {
debug!("finished piece {} from local", piece_id);
return Ok(piece);
}
let guard = scopeguard::guard((), |_| {
if let Some(err) = self.storage.download_piece_failed(piece_id).err() {
error!("set piece metadata failed: {}", err)
};
});
if is_prefetch {
self.prefetch_bandwidth_limiter
.acquire(length as usize)
.await;
}
self.download_bandwidth_limiter
.acquire(length as usize)
.await;
let (mut reader, offset, digest) = match (
self.config.download.protocol.as_str(),
parent.download_ip,
parent.download_tcp_port,
parent.download_quic_port,
) {
("tcp", Some(ip), Some(port), _) => {
self.tcp_downloader
.download_piece(
&format_socket_addr(IpAddr::from_str(&ip)?, port as u16),
number,
host_id,
task_id,
)
.await?
}
("quic", Some(ip), _, Some(port)) => {
self.quic_downloader
.download_piece(
&format_socket_addr(IpAddr::from_str(&ip)?, port as u16),
number,
host_id,
task_id,
)
.await?
}
_ => {
warn!("fall back to grpc downloader");
let host = parent.host.clone().ok_or_else(|| {
error!("parent host is empty");
Error::InvalidPeer(parent.id.clone())
})?;
self.tcp_downloader
.download_piece(
&format_socket_addr(IpAddr::from_str(&host.ip)?, host.port as u16),
number,
host_id,
task_id,
)
.await
.inspect_err(|err| {
error!("download piece failed: {}", err);
})?
}
};
match self
.storage
.download_piece_from_parent_finished(
piece_id,
task_id,
offset,
length,
digest.as_str(),
parent.id.as_str(),
&mut reader,
self.config.storage.write_piece_timeout,
)
.await
{
Ok(piece) => {
collect_download_piece_traffic_metrics(&TrafficType::RemotePeer, length);
scopeguard::ScopeGuard::into_inner(guard);
Ok(piece)
}
Err(err) => {
error!("download piece finished: {}", err);
Err(err)
}
}
}
#[allow(clippy::too_many_arguments)]
#[instrument(skip_all, fields(piece_id))]
pub async fn download_from_source(
&self,
piece_id: &str,
task_id: &str,
number: u32,
url: &str,
offset: u64,
length: u64,
request_header: HeaderMap,
is_prefetch: bool,
object_storage: Option<ObjectStorage>,
hdfs: Option<Hdfs>,
hugging_face: Option<HuggingFace>,
model_scope: Option<ModelScope>,
) -> Result<metadata::Piece> {
Span::current().record("piece_id", piece_id);
Span::current().record("piece_length", length);
let piece = self
.storage
.download_piece_started(piece_id, number)
.await?;
if piece.is_finished() {
debug!("finished piece {} from local", piece_id);
return Ok(piece);
}
let guard = scopeguard::guard((), |_| {
if let Some(err) = self.storage.download_piece_failed(piece_id).err() {
error!("set piece metadata failed: {}", err)
};
});
if is_prefetch {
self.prefetch_bandwidth_limiter
.acquire(length as usize)
.await;
}
self.back_to_source_bandwidth_limiter
.acquire(length as usize)
.await;
self.download_bandwidth_limiter
.acquire(length as usize)
.await;
let backend = self.backend_factory.build(url).inspect_err(|err| {
error!("build backend failed: {}", err);
})?;
let start_time = Instant::now();
collect_backend_request_started_metrics(
backend.scheme().as_str(),
http::Method::GET.as_str(),
);
let mut response = backend
.get(GetRequest {
task_id: task_id.to_string(),
piece_id: piece_id.to_string(),
url: url.to_string(),
range: Some(Range {
start: offset,
length,
}),
http_header: Some(request_header),
timeout: self.config.download.piece_timeout,
client_cert: None,
object_storage,
hdfs,
hugging_face,
model_scope,
})
.await
.inspect_err(|err| {
collect_backend_request_failure_metrics(
backend.scheme().as_str(),
http::Method::GET.as_str(),
);
error!("backend get failed: {}", err);
})?;
if !response.success {
collect_backend_request_failure_metrics(
backend.scheme().as_str(),
http::Method::GET.as_str(),
);
let mut buffer = String::new();
response
.reader
.read_to_string(&mut buffer)
.await
.unwrap_or_default();
let error_message = response.error_message.unwrap_or_default();
error!("backend get failed: {} {}", error_message, buffer.as_str());
return Err(Error::BackendError(Box::new(BackendError {
message: error_message,
status_code: Some(response.http_status_code.unwrap_or_default()),
header: Some(response.http_header.unwrap_or_default()),
})));
}
collect_backend_request_finished_metrics(
backend.scheme().as_str(),
http::Method::GET.as_str(),
start_time.elapsed(),
);
match self
.storage
.download_piece_from_source_finished(
piece_id,
task_id,
offset,
length,
&mut response.reader,
self.config.storage.write_piece_timeout,
)
.await
{
Ok(piece) => {
collect_download_piece_traffic_metrics(&TrafficType::BackToSource, length);
scopeguard::ScopeGuard::into_inner(guard);
Ok(piece)
}
Err(err) => {
error!("download piece finished: {}", err);
Err(err)
}
}
}
#[inline]
pub fn persistent_id(&self, task_id: &str, number: u32) -> String {
self.storage.persistent_piece_id(task_id, number)
}
#[instrument(level = "debug", skip_all)]
pub fn get_persistent(&self, piece_id: &str) -> Result<Option<metadata::Piece>> {
self.storage.get_persistent_piece(piece_id)
}
#[instrument(skip_all)]
pub async fn create_persistent<R: AsyncRead + Unpin + ?Sized>(
&self,
piece_id: &str,
task_id: &str,
number: u32,
offset: u64,
length: u64,
reader: &mut R,
) -> Result<metadata::Piece> {
self.storage
.create_persistent_piece(piece_id, task_id, number, offset, length, reader)
.await
}
#[instrument(skip_all)]
pub fn register_persistent(
&self,
piece_id: &str,
number: u32,
offset: u64,
length: u64,
) -> Result<metadata::Piece> {
self.storage
.register_persistent_piece(piece_id, number, offset, length)
}
#[instrument(level = "debug", skip_all, fields(piece_id))]
pub async fn download_persistent_from_local_into_async_read(
&self,
piece_id: &str,
task_id: &str,
length: u64,
range: Option<Range>,
) -> Result<impl AsyncBufRead> {
Span::current().record("piece_id", piece_id);
Span::current().record("piece_length", length);
self.storage
.upload_persistent_piece(piece_id, task_id, range)
.await
}
#[instrument(level = "debug", skip_all)]
pub fn download_persistent_from_local(&self, length: u64) {
collect_download_piece_traffic_metrics(&TrafficType::LocalPeer, length);
}
#[allow(clippy::too_many_arguments)]
#[instrument(skip_all, fields(piece_id))]
pub async fn download_persistent_from_parent(
&self,
piece_id: &str,
host_id: &str,
task_id: &str,
number: u32,
length: u64,
parent: piece_collector::CollectedParent,
) -> Result<metadata::Piece> {
Span::current().record("piece_id", piece_id);
Span::current().record("piece_length", length);
self.download_bandwidth_limiter
.acquire(length as usize)
.await;
let piece = self
.storage
.download_persistent_piece_started(piece_id, number)
.await?;
if piece.is_finished() {
debug!("finished persistent piece {} from local", piece_id);
return Ok(piece);
}
let guard = scopeguard::guard((), |_| {
if let Some(err) = self
.storage
.download_persistent_piece_failed(piece_id)
.err()
{
error!("set persistent piece metadata failed: {}", err)
};
});
let (mut reader, offset, digest) = match (
self.config.download.protocol.as_str(),
parent.download_ip,
parent.download_tcp_port,
parent.download_quic_port,
) {
("tcp", Some(ip), Some(port), _) => {
self.tcp_downloader
.download_persistent_piece(
&format_socket_addr(IpAddr::from_str(&ip)?, port as u16),
number,
host_id,
task_id,
)
.await?
}
("quic", Some(ip), _, Some(port)) => {
let quic_downloader =
piece_downloader::DownloaderFactory::new("quic", self.config.clone())?.build();
quic_downloader
.download_persistent_piece(
&format_socket_addr(IpAddr::from_str(&ip)?, port as u16),
number,
host_id,
task_id,
)
.await?
}
_ => {
warn!("fall back to grpc downloader");
let host = parent.host.clone().ok_or_else(|| {
error!("parent host is empty");
Error::InvalidPeer(parent.id.clone())
})?;
self.tcp_downloader
.download_persistent_piece(
&format_socket_addr(IpAddr::from_str(&host.ip)?, host.port as u16),
number,
host_id,
task_id,
)
.await
.inspect_err(|err| {
error!("download persistent piece failed: {}", err);
})?
}
};
match self
.storage
.download_persistent_piece_from_parent_finished(
piece_id,
task_id,
offset,
length,
digest.as_str(),
parent.id.as_str(),
&mut reader,
)
.await
{
Ok(piece) => {
collect_download_piece_traffic_metrics(&TrafficType::RemotePeer, length);
scopeguard::ScopeGuard::into_inner(guard);
Ok(piece)
}
Err(err) => {
error!("download persistent piece finished: {}", err);
Err(err)
}
}
}
#[allow(clippy::too_many_arguments)]
#[instrument(skip_all, fields(piece_id))]
pub async fn download_persistent_from_source(
&self,
piece_id: &str,
task_id: &str,
number: u32,
url: &str,
offset: u64,
length: u64,
request_header: HeaderMap,
object_storage: Option<ObjectStorage>,
hdfs: Option<Hdfs>,
hugging_face: Option<HuggingFace>,
model_scope: Option<ModelScope>,
) -> Result<metadata::Piece> {
Span::current().record("piece_id", piece_id);
Span::current().record("piece_length", length);
let piece = self
.storage
.download_persistent_piece_started(piece_id, number)
.await?;
if piece.is_finished() {
debug!("finished piece {} from local", piece_id);
return Ok(piece);
}
let guard = scopeguard::guard((), |_| {
if let Some(err) = self
.storage
.download_persistent_piece_failed(piece_id)
.err()
{
error!("set piece metadata failed: {}", err)
};
});
self.back_to_source_bandwidth_limiter
.acquire(length as usize)
.await;
self.download_bandwidth_limiter
.acquire(length as usize)
.await;
let backend = self.backend_factory.build(url).inspect_err(|err| {
error!("build backend failed: {}", err);
})?;
let start_time = Instant::now();
collect_backend_request_started_metrics(
backend.scheme().as_str(),
http::Method::GET.as_str(),
);
let mut response = backend
.get(GetRequest {
task_id: task_id.to_string(),
piece_id: piece_id.to_string(),
url: url.to_string(),
range: Some(Range {
start: offset,
length,
}),
http_header: Some(request_header),
timeout: self.config.download.piece_timeout,
client_cert: None,
object_storage,
hdfs,
hugging_face,
model_scope,
})
.await
.inspect_err(|err| {
collect_backend_request_failure_metrics(
backend.scheme().as_str(),
http::Method::GET.as_str(),
);
error!("backend get failed: {}", err);
})?;
if !response.success {
collect_backend_request_failure_metrics(
backend.scheme().as_str(),
http::Method::GET.as_str(),
);
let mut buffer = String::new();
response
.reader
.read_to_string(&mut buffer)
.await
.unwrap_or_default();
let error_message = response.error_message.unwrap_or_default();
error!("backend get failed: {} {}", error_message, buffer.as_str());
return Err(Error::BackendError(Box::new(BackendError {
message: error_message,
status_code: Some(response.http_status_code.unwrap_or_default()),
header: Some(response.http_header.unwrap_or_default()),
})));
}
collect_backend_request_finished_metrics(
backend.scheme().as_str(),
http::Method::GET.as_str(),
start_time.elapsed(),
);
match self
.storage
.download_persistent_piece_from_source_finished(
piece_id,
task_id,
offset,
length,
&mut response.reader,
self.config.storage.write_piece_timeout,
)
.await
{
Ok(piece) => {
collect_download_piece_traffic_metrics(&TrafficType::BackToSource, length);
scopeguard::ScopeGuard::into_inner(guard);
Ok(piece)
}
Err(err) => {
error!("download piece finished: {}", err);
Err(err)
}
}
}
#[inline]
pub fn persistent_cache_id(&self, task_id: &str, number: u32) -> String {
self.storage.persistent_cache_piece_id(task_id, number)
}
#[instrument(level = "debug", skip_all)]
pub fn get_persistent_cache(&self, piece_id: &str) -> Result<Option<metadata::Piece>> {
self.storage.get_persistent_cache_piece(piece_id)
}
#[instrument(level = "debug", skip_all)]
pub async fn create_persistent_cache<R: AsyncRead + Unpin + ?Sized>(
&self,
piece_id: &str,
task_id: &str,
number: u32,
offset: u64,
length: u64,
reader: &mut R,
) -> Result<metadata::Piece> {
self.storage
.create_persistent_cache_piece(piece_id, task_id, number, offset, length, reader)
.await
}
#[instrument(level = "debug", skip_all)]
pub fn register_persistent_cache(
&self,
piece_id: &str,
number: u32,
offset: u64,
length: u64,
) -> Result<metadata::Piece> {
self.storage
.register_persistent_cache_piece(piece_id, number, offset, length)
}
#[instrument(level = "debug", skip_all, fields(piece_id))]
pub async fn download_persistent_cache_from_local_into_async_read(
&self,
piece_id: &str,
task_id: &str,
length: u64,
range: Option<Range>,
) -> Result<impl AsyncBufRead> {
Span::current().record("piece_id", piece_id);
Span::current().record("piece_length", length);
self.storage
.upload_persistent_cache_piece(piece_id, task_id, range)
.await
}
#[instrument(level = "debug", skip_all)]
pub fn download_persistent_cache_from_local(&self, length: u64) {
collect_download_piece_traffic_metrics(&TrafficType::LocalPeer, length);
}
#[allow(clippy::too_many_arguments)]
#[instrument(skip_all, fields(piece_id))]
pub async fn download_persistent_cache_from_parent(
&self,
piece_id: &str,
host_id: &str,
task_id: &str,
number: u32,
length: u64,
parent: piece_collector::CollectedParent,
) -> Result<metadata::Piece> {
Span::current().record("piece_id", piece_id);
Span::current().record("piece_length", length);
self.download_bandwidth_limiter
.acquire(length as usize)
.await;
let piece = self
.storage
.download_persistent_cache_piece_started(piece_id, number)
.await?;
if piece.is_finished() {
debug!("finished persistent cache piece {} from local", piece_id);
return Ok(piece);
}
let guard = scopeguard::guard((), |_| {
if let Some(err) = self
.storage
.download_persistent_cache_piece_failed(piece_id)
.err()
{
error!("set persistent cache piece metadata failed: {}", err)
};
});
let (mut reader, offset, digest) = match (
self.config.download.protocol.as_str(),
parent.download_ip,
parent.download_tcp_port,
parent.download_quic_port,
) {
("tcp", Some(ip), Some(port), _) => {
self.tcp_downloader
.download_persistent_cache_piece(
&format_socket_addr(IpAddr::from_str(&ip)?, port as u16),
number,
host_id,
task_id,
)
.await?
}
("quic", Some(ip), _, Some(port)) => {
let quic_downloader =
piece_downloader::DownloaderFactory::new("quic", self.config.clone())?.build();
quic_downloader
.download_persistent_cache_piece(
&format_socket_addr(IpAddr::from_str(&ip)?, port as u16),
number,
host_id,
task_id,
)
.await?
}
_ => {
warn!("fall back to grpc downloader");
let host = parent.host.clone().ok_or_else(|| {
error!("parent host is empty");
Error::InvalidPeer(parent.id.clone())
})?;
self.tcp_downloader
.download_persistent_cache_piece(
&format_socket_addr(IpAddr::from_str(&host.ip)?, host.port as u16),
number,
host_id,
task_id,
)
.await
.inspect_err(|err| {
error!("download persistent cache piece failed: {}", err);
})?
}
};
match self
.storage
.download_persistent_cache_piece_from_parent_finished(
piece_id,
task_id,
offset,
length,
digest.as_str(),
parent.id.as_str(),
&mut reader,
)
.await
{
Ok(piece) => {
collect_download_piece_traffic_metrics(&TrafficType::RemotePeer, length);
scopeguard::ScopeGuard::into_inner(guard);
Ok(piece)
}
Err(err) => {
error!("download persistent cache piece finished: {}", err);
Err(err)
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::tempdir;
#[tokio::test]
async fn test_calculate_interested() {
let temp_dir = tempdir().unwrap();
let config = Config::default();
let config = Arc::new(config);
let storage = Storage::new(
config.clone(),
temp_dir.path(),
temp_dir.path().to_path_buf(),
)
.await
.unwrap();
let storage = Arc::new(storage);
let backend_factory = BackendFactory::new(config.clone(), None).unwrap();
let backend_factory = Arc::new(backend_factory);
let download_bandwidth_limiter = Arc::new(RateLimiter::builder().build());
let prefetch_bandwidth_limiter = Arc::new(RateLimiter::builder().build());
let back_to_source_bandwidth_limiter = Arc::new(RateLimiter::builder().build());
let piece = Piece::new(
config.clone(),
storage.clone(),
backend_factory.clone(),
download_bandwidth_limiter,
prefetch_bandwidth_limiter,
back_to_source_bandwidth_limiter,
)
.unwrap();
let test_cases = vec![
(1000, 1, None, 1, vec![0], 0, 1),
(1000, 5000, None, 5, vec![0, 1, 2, 3, 4], 4000, 1000),
(5000, 1000, None, 1, vec![0], 0, 1000),
(
10,
101,
None,
11,
vec![0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10],
100,
1,
),
(
1000,
5000,
Some(Range {
start: 1500,
length: 2000,
}),
3,
vec![1, 2, 3],
3000,
1000,
),
(
1000,
5000,
Some(Range {
start: 0,
length: 1,
}),
1,
vec![0],
0,
1000,
),
];
for (
piece_length,
content_length,
range,
expected_len,
expected_numbers,
expected_last_piece_offset,
expected_last_piece_length,
) in test_cases
{
let pieces = piece
.calculate_interested(piece_length, content_length, range)
.unwrap();
assert_eq!(pieces.len(), expected_len);
assert_eq!(
pieces
.iter()
.map(|piece| piece.number)
.collect::<Vec<u32>>(),
expected_numbers
);
let last_piece = pieces.last().unwrap();
assert_eq!(last_piece.offset, expected_last_piece_offset);
assert_eq!(last_piece.length, expected_last_piece_length);
}
}
}