use std::io;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use bytes::Bytes;
use chacha20::ChaCha20;
use chacha20::cipher::{KeyIvInit, StreamCipher, StreamCipherSeek};
use futures_util::stream::FuturesUnordered;
use futures_util::{Stream, StreamExt};
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use tokio_stream::wrappers::ReceiverStream;
use crate::adapters::Storage;
use crate::crypto::{ChunkPrp, content_cipher_params};
use crate::error::{ApiError, ApiResult};
const HEAD_SMALL_SPLIT: u64 = 256 * 1024;
const HEAD_SMALL_COUNT: usize = 4;
#[derive(Debug, Clone)]
pub struct VolumeMeta {
pub name: String,
pub size: u64,
pub offset: u64,
}
#[derive(Debug, Clone)]
pub struct FileLayout {
pub volumes: Vec<VolumeMeta>,
pub total: u64,
}
pub async fn load_layout(
storage: &dyn Storage,
enc_folder: &str,
pw: &[u8],
) -> ApiResult<FileLayout> {
let entries = storage.list(enc_folder).await?;
let prp = ChunkPrp::new(pw);
let mut indexed: Vec<(usize, String, u64)> = entries
.into_iter()
.into_iter()
.filter(|e| !e.is_dir)
.filter_map(|e| prp.index_of(&e.name).map(|i| (i, e.name, e.size)))
.collect();
indexed.sort_by_key(|(i, ..)| *i);
for (pos, (i, ..)) in indexed.iter().enumerate() {
if *i != pos {
return Err(ApiError::Upstream(format!(
"云端分卷不完整:缺第 {pos} 卷(共发现 {} 卷)",
indexed.len()
)));
}
}
let mut volumes = Vec::with_capacity(indexed.len());
let mut offset = 0u64;
for (_, name, size) in indexed {
volumes.push(VolumeMeta { name, size, offset });
offset += size;
}
Ok(FileLayout {
volumes,
total: offset,
})
}
#[derive(Debug, PartialEq, Eq)]
pub enum RangeSpec {
Full,
Slice { start: u64, end: u64 },
Unsatisfiable,
}
pub fn parse_range(header: Option<&str>, total: u64) -> (RangeSpec, bool) {
let Some(h) = header else {
return (RangeSpec::Full, true);
};
let h = h.trim();
let Some(spec) = h.strip_prefix("bytes=") else {
return (RangeSpec::Full, true);
};
if spec.contains(',') {
return (RangeSpec::Full, true);
}
let Some((a, b)) = spec.split_once('-') else {
return (RangeSpec::Full, true);
};
let (a, b) = (a.trim(), b.trim());
if total == 0 {
return (RangeSpec::Full, false);
}
match (a.is_empty(), b.is_empty()) {
(false, false) => {
let (Ok(s), Ok(e)) = (a.parse::<u64>(), b.parse::<u64>()) else {
return (RangeSpec::Full, true);
};
if s >= total || s > e {
return (RangeSpec::Unsatisfiable, false);
}
(
RangeSpec::Slice {
start: s,
end: e.min(total - 1),
},
false,
)
}
(false, true) => {
let Ok(s) = a.parse::<u64>() else {
return (RangeSpec::Full, true);
};
if s >= total {
return (RangeSpec::Unsatisfiable, false);
}
(
RangeSpec::Slice {
start: s,
end: total - 1,
},
true,
)
}
(true, false) => {
let Ok(n) = b.parse::<u64>() else {
return (RangeSpec::Full, true);
};
if n == 0 {
return (RangeSpec::Unsatisfiable, false);
}
let start = total.saturating_sub(n);
(
RangeSpec::Slice {
start,
end: total - 1,
},
false,
)
}
(true, true) => (RangeSpec::Full, true),
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PlannedChunk {
pub merged_start: u64,
pub len: u64,
pub vol: usize,
pub vol_off: u64,
}
#[cfg(test)]
fn plan_chunks(
layout: &FileLayout,
start: u64,
end: u64,
max_split: u64,
open_ended: bool,
) -> Vec<PlannedChunk> {
plan_chunks_with_head_count(layout, start, end, max_split, open_ended, HEAD_SMALL_COUNT)
}
fn plan_chunks_with_head_count(
layout: &FileLayout,
start: u64,
end: u64,
max_split: u64,
open_ended: bool,
head_count: usize,
) -> Vec<PlannedChunk> {
let split = max_split.max(1); let head: usize = if open_ended && split > HEAD_SMALL_SPLIT {
head_count
} else {
0
};
let mut plan = Vec::new();
let mut cur = start;
let mut vol_idx = 0usize;
while cur <= end && vol_idx < layout.volumes.len() {
let v = &layout.volumes[vol_idx];
if v.size == 0 || cur >= v.offset + v.size {
vol_idx += 1;
continue;
}
let this_split = if plan.len() < head {
HEAD_SMALL_SPLIT
} else {
split
};
let vol_last = v.offset + v.size - 1;
let chunk_end = (cur + this_split - 1).min(vol_last).min(end);
plan.push(PlannedChunk {
merged_start: cur,
len: chunk_end - cur + 1,
vol: vol_idx,
vol_off: cur - v.offset,
});
cur = chunk_end + 1;
}
plan
}
pub struct StreamParams {
pub max_split: u64,
pub max_threads: usize,
pub max_per_volume: usize,
}
#[allow(clippy::too_many_arguments)]
#[cfg_attr(not(test), allow(dead_code))]
pub fn stream_range_cached(
storage: Arc<dyn Storage>,
enc_folder: String,
pw: [u8; crate::crypto::SECRET_LEN],
layout: Arc<FileLayout>,
start: u64,
end: u64,
open_ended: bool,
params: &StreamParams,
cache: Option<Arc<crate::cache::CacheEntry>>,
) -> mpsc::Receiver<io::Result<Bytes>> {
stream_range_cached_mode(
storage, enc_folder, pw, true, layout, start, end, open_ended, params, cache, None,
)
}
pub async fn load_layout_ordered(storage: &dyn Storage, folder: &str) -> ApiResult<FileLayout> {
let mut entries: Vec<_> = storage
.list(folder)
.await?
.into_iter()
.filter(|entry| !entry.is_dir)
.collect();
entries.sort_by(|a, b| a.name.cmp(&b.name));
let mut offset = 0u64;
let volumes = entries
.into_iter()
.map(|entry| {
let volume = VolumeMeta {
name: entry.name,
size: entry.size,
offset,
};
offset += entry.size;
volume
})
.collect();
Ok(FileLayout {
volumes,
total: offset,
})
}
#[allow(clippy::too_many_arguments)]
pub fn stream_range_cached_mode(
storage: Arc<dyn Storage>,
enc_folder: String,
pw: [u8; crate::crypto::SECRET_LEN],
encrypted: bool,
layout: Arc<FileLayout>,
start: u64,
end: u64,
open_ended: bool,
params: &StreamParams,
cache: Option<Arc<crate::cache::CacheEntry>>,
network_progress: Option<crate::adapters::ProgressFn>,
) -> mpsc::Receiver<io::Result<Bytes>> {
let max_split = storage
.max_range_size()
.map_or(params.max_split, |limit| params.max_split.min(limit));
let max_threads = params.max_threads.max(1);
let max_per_volume = params.max_per_volume.max(1);
let plan = plan_chunks_with_head_count(
&layout,
start,
end,
max_split,
open_ended,
max_threads.max(HEAD_SMALL_COUNT),
);
let total_chunks = plan.len();
tracing::debug!(
"stream_range [{start},{end}] chunks={total_chunks} split={} threads={max_threads} per_vol={max_per_volume} open_ended={open_ended} window={max_threads}",
max_split,
);
let item_estimate = 16 * 1024u64;
let chan_buffer = ((max_split / item_estimate) as usize).clamp(8, 512);
let mut senders: Vec<Option<mpsc::Sender<io::Result<Bytes>>>> =
Vec::with_capacity(total_chunks);
let mut receivers: Vec<mpsc::Receiver<io::Result<Bytes>>> = Vec::with_capacity(total_chunks);
for _ in 0..total_chunks {
let (tx, rx) = mpsc::channel(chan_buffer);
senders.push(Some(tx));
receivers.push(rx);
}
let (out_tx, out_rx) = mpsc::channel::<io::Result<Bytes>>(8);
tokio::spawn(async move {
let plan = Arc::new(plan);
let mut handles: Vec<JoinHandle<()>> = Vec::with_capacity(total_chunks);
let mut next_to_spawn = 0usize;
let spawn_one = |idx: usize, tx: mpsc::Sender<io::Result<Bytes>>| -> JoinHandle<()> {
let c = plan[idx].clone();
let vol_name = layout.volumes[c.vol].name.clone();
let obj_path = if enc_folder.is_empty() {
vol_name
} else {
format!("{enc_folder}/{vol_name}")
};
let st = Arc::clone(&storage);
let chunk_cache = cache.clone();
let progress = network_progress.clone();
tokio::spawn(async move {
fetch_chunk(st, obj_path, pw, encrypted, c, tx, chunk_cache, progress).await;
})
};
let abort_all = |handles: &[JoinHandle<()>]| {
for handle in handles {
handle.abort();
}
};
if total_chunks > 0 {
let tx = senders[0].take().expect("首块 sender 仅使用一次");
handles.push(spawn_one(0, tx));
next_to_spawn = 1;
}
let mut initial_window_opened = false;
'outer: for (i, mut rx) in receivers.into_iter().enumerate() {
let expect = plan[i].len;
let mut got = 0u64;
while got < expect {
let item = tokio::select! {
biased;
_ = out_tx.closed() => {
tracing::debug!(
"客户端在 chunk {i} 等待期间断开,abort {} 个 fetcher",
handles.len()
);
abort_all(&handles);
break 'outer;
}
item = rx.recv() => item,
};
match item {
Some(Ok(b)) => {
got += b.len() as u64;
if out_tx.send(Ok(b)).await.is_err() {
tracing::debug!(
"客户端在 chunk {i} 输出期间断开,abort {} 个 fetcher",
handles.len()
);
abort_all(&handles);
break 'outer;
}
if !initial_window_opened {
let target = total_chunks.min(max_threads);
while next_to_spawn < target {
let idx = next_to_spawn;
let tx = senders[idx].take().expect("chunk sender 仅使用一次");
handles.push(spawn_one(idx, tx));
next_to_spawn += 1;
}
initial_window_opened = true;
}
}
Some(Err(e)) => {
let _ = out_tx.send(Err(e)).await;
abort_all(&handles);
break 'outer;
}
None => {
let _ = out_tx
.send(Err(io::Error::other(format!(
"分片 {i} 提前结束({got}/{expect} 字节)"
))))
.await;
abort_all(&handles);
break 'outer;
}
}
}
let target = total_chunks.min(i.saturating_add(1).saturating_add(max_threads));
while next_to_spawn < target {
let idx = next_to_spawn;
let tx = senders[idx].take().expect("chunk sender 仅使用一次");
handles.push(spawn_one(idx, tx));
next_to_spawn += 1;
}
}
});
out_rx
}
async fn fetch_chunk(
storage: Arc<dyn Storage>,
obj_path: String,
pw: [u8; crate::crypto::SECRET_LEN],
encrypted: bool,
c: PlannedChunk,
tx: mpsc::Sender<io::Result<Bytes>>,
cache: Option<Arc<crate::cache::CacheEntry>>,
network_progress: Option<crate::adapters::ProgressFn>,
) {
let merged_end = c.merged_start + c.len - 1;
if let Some(cache) = &cache
&& cache.has_range(c.merged_start, merged_end)
{
let hit_cache = Arc::clone(cache);
let hit_start = c.merged_start;
let hit = tokio::task::spawn_blocking(move || {
let cached_bytes = hit_cache.read_range(hit_start, merged_end)?;
let mut buf = cached_bytes.to_vec();
if encrypted {
crate::crypto::apply_content_keystream(&pw, hit_start, &mut buf);
}
Ok::<Bytes, io::Error>(Bytes::from(buf))
})
.await;
match hit {
Ok(Ok(bytes)) => {
let _ = tx.send(Ok(bytes)).await;
return;
}
Ok(Err(e)) => tracing::warn!("读取密文缓存失败,回源: {e}"),
Err(e) => tracing::warn!("缓存读取任务异常,回源: {e}"),
}
}
if let Some(cache) = &cache {
cache.record_miss();
}
let (key, nonce) = content_cipher_params(&pw);
let mut cipher = ChaCha20::new(&key.into(), &nonce.into());
if encrypted && cipher.try_seek(c.merged_start).is_err() {
let _ = tx.send(Err(io::Error::other("keystream 偏移越界"))).await;
return;
}
let mut remaining = c.len;
let mut attempts = 0usize;
let mut last_error = String::new();
while remaining > 0 && attempts < 4 {
attempts += 1;
let done = c.len - remaining;
let range_start = c.vol_off + done;
let range_end = c.vol_off + c.len - 1;
let mut stream = match storage.get_range(&obj_path, range_start, range_end).await {
Ok(stream) => stream,
Err(e) => {
last_error = e.to_string();
tokio::task::yield_now().await;
continue;
}
};
let before = remaining;
while let Some(item) = stream.next().await {
let item = match item {
Ok(bytes) => bytes,
Err(e) => {
last_error = e.to_string();
break;
}
};
if item.is_empty() {
continue;
}
let take = (item.len() as u64).min(remaining) as usize;
if let Some(progress) = &network_progress {
progress(take as u64);
}
let cache_offset = c.merged_start + (c.len - remaining);
if let Some(cache) = &cache
&& let Err(e) = cache.write_range(cache_offset, &item[..take])
{
tracing::warn!("写入密文缓存失败(不影响本次下载): {e}");
}
let mut buf = item[..take].to_vec();
if encrypted {
cipher.apply_keystream(&mut buf);
}
remaining -= take as u64;
if tx.send(Ok(Bytes::from(buf))).await.is_err() {
return; }
if remaining == 0 {
return;
}
}
if remaining == before && last_error.is_empty() {
last_error = format!("上游未返回 range {range_start}-{range_end}");
}
tracing::debug!(
"分片重试 {attempts}/4: path={obj_path} 已完成={} 剩余={remaining} err={last_error}",
c.len - remaining,
);
tokio::task::yield_now().await;
}
let _ = tx
.send(Err(io::Error::other(format!(
"上游重试 {attempts} 次后仍少 {remaining} 字节: {last_error}"
))))
.await;
}
pub struct UploadProgress {
pub total: u64,
pub encrypted: AtomicU64,
pub uploaded: AtomicU64,
network: Option<Arc<crate::transfer::TransferTracker>>,
}
impl UploadProgress {
#[cfg_attr(not(test), allow(dead_code))]
pub fn new(total: u64) -> Self {
Self {
total,
encrypted: AtomicU64::new(0),
uploaded: AtomicU64::new(0),
network: None,
}
}
pub fn tracked(total: u64, network: Arc<crate::transfer::TransferTracker>) -> Self {
Self {
total,
encrypted: AtomicU64::new(0),
uploaded: AtomicU64::new(0),
network: Some(network),
}
}
}
const MAX_PENDING_UPLOADS: usize = 4;
type UploadTask = (mpsc::Sender<io::Result<Bytes>>, JoinHandle<ApiResult<()>>);
type PendingUploads = FuturesUnordered<JoinHandle<ApiResult<()>>>;
#[allow(clippy::too_many_arguments)]
#[cfg_attr(not(test), allow(dead_code))]
pub async fn upload_stream<S>(
storage: Arc<dyn Storage>,
enc_folder: &str,
pw: &[u8],
total: u64,
volume_size: u64,
names: &[String],
body: S,
progress: Arc<UploadProgress>,
) -> ApiResult<()>
where
S: Stream<Item = io::Result<Bytes>> + Unpin,
{
let sizes = (0..names.len())
.map(|idx| {
let start = idx as u64 * volume_size;
volume_size.min(total.saturating_sub(start))
})
.collect::<Vec<_>>();
upload_stream_planned(
storage, enc_folder, pw, true, total, &sizes, names, body, progress,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub async fn upload_stream_planned<S>(
storage: Arc<dyn Storage>,
enc_folder: &str,
pw: &[u8],
encrypted: bool,
total: u64,
volume_sizes: &[u64],
names: &[String],
mut body: S,
progress: Arc<UploadProgress>,
) -> ApiResult<()>
where
S: Stream<Item = io::Result<Bytes>> + Unpin,
{
if names.len() != volume_sizes.len() || volume_sizes.iter().sum::<u64>() != total {
return Err(ApiError::BadRequest("分卷计划与文件大小不一致".into()));
}
if total == 0 {
if let Some(name) = names.first() {
let path = if enc_folder.is_empty() {
name.clone()
} else {
format!("{enc_folder}/{name}")
};
return storage
.put_sized(&path, 0, futures_util::stream::empty().boxed())
.await;
}
return Ok(());
}
let (key, nonce) = content_cipher_params(pw);
let mut cipher = ChaCha20::new(&key.into(), &nonce.into());
let vol_cap = |idx: usize| -> u64 { volume_sizes[idx] };
let mut vol_idx = 0usize;
let mut sent_in_vol = 0u64;
let mut received = 0u64;
let mut current: Option<UploadTask> = None;
let mut pending = PendingUploads::new();
async fn close_current(cur: &mut Option<UploadTask>, pending: &mut PendingUploads) {
let Some((tx, handle)) = cur.take() else {
return;
};
drop(tx); pending.push(handle);
}
async fn wait_one(pending: &mut PendingUploads) -> ApiResult<()> {
let Some(handle) = pending.next().await else {
return Ok(());
};
handle.map_err(|e| ApiError::Internal(anyhow::anyhow!("上传任务 panic: {e}")))?
}
async fn wait_all(pending: &mut PendingUploads) -> ApiResult<()> {
let mut first_error = None;
while let Some(handle) = pending.next().await {
let result = handle
.map_err(|e| ApiError::Internal(anyhow::anyhow!("上传任务 panic: {e}")))
.and_then(|result| result);
if first_error.is_none() {
first_error = result.err();
}
}
first_error.map_or(Ok(()), Err)
}
while let Some(item) = body.next().await {
let item = match item {
Ok(item) => item,
Err(e) => {
close_current(&mut current, &mut pending).await;
let _ = wait_all(&mut pending).await;
return Err(ApiError::BadRequest(format!("请求体读取失败: {e}")));
}
};
if item.is_empty() {
continue;
}
received += item.len() as u64;
if received > total {
close_current(&mut current, &mut pending).await;
let _ = wait_all(&mut pending).await;
return Err(ApiError::BadRequest("实际字节数超过声明大小".into()));
}
let mut buf = item
.try_into_mut()
.unwrap_or_else(|shared| bytes::BytesMut::from(shared.as_ref()));
if encrypted {
cipher.apply_keystream(&mut buf);
}
progress
.encrypted
.fetch_add(buf.len() as u64, Ordering::Relaxed);
let mut b = buf.freeze();
while !b.is_empty() {
if current.is_none() {
let cap = vol_cap(vol_idx);
let name = names
.get(vol_idx)
.ok_or_else(|| ApiError::BadRequest("分卷数超出计划".into()))?;
let obj_path = if enc_folder.is_empty() {
name.clone()
} else {
format!("{enc_folder}/{name}")
};
let (tx, rx) = mpsc::channel::<io::Result<Bytes>>(8);
let st = Arc::clone(&storage);
let uploaded = Arc::clone(&progress);
let on_upload: crate::adapters::ProgressFn = Arc::new(move |n| {
uploaded.uploaded.fetch_add(n, Ordering::Relaxed);
if let Some(network) = &uploaded.network {
network.upload(n);
}
});
let handle = tokio::spawn(async move {
st.put_sized_tracked(&obj_path, cap, ReceiverStream::new(rx).boxed(), on_upload)
.await
});
current = Some((tx, handle));
sent_in_vol = 0;
}
let cap = vol_cap(vol_idx);
let take = (cap - sent_in_vol).min(b.len() as u64) as usize;
let piece = b.split_to(take);
let send_ok = {
let (tx, _) = current.as_ref().expect("上面刚建立");
tx.send(Ok(piece)).await.is_ok()
};
if !send_ok {
close_current(&mut current, &mut pending).await;
wait_all(&mut pending).await?;
return Err(ApiError::Upstream("分卷写入提前中断".into()));
}
sent_in_vol += take as u64;
if sent_in_vol == cap {
close_current(&mut current, &mut pending).await;
vol_idx += 1;
if pending.len() >= MAX_PENDING_UPLOADS
&& let Err(e) = wait_one(&mut pending).await
{
let _ = wait_all(&mut pending).await;
return Err(e);
}
}
}
}
if received != total {
close_current(&mut current, &mut pending).await;
let _ = wait_all(&mut pending).await;
return Err(ApiError::BadRequest(format!(
"实际字节数 {received} 与声明大小 {total} 不符"
)));
}
close_current(&mut current, &mut pending).await;
wait_all(&mut pending).await
}
#[cfg(test)]
mod tests {
use super::*;
use crate::adapters::localfs::LocalFs;
use crate::adapters::{ByteStream, Entry};
use crate::crypto::{chunk_count, gen_chunk_names, gen_secret};
use futures_util::stream;
use std::sync::atomic::{AtomicUsize, Ordering};
struct FlakyRangeStorage {
encrypted: Bytes,
calls: AtomicUsize,
}
struct SlowFirstFinalizeStorage;
#[async_trait::async_trait]
impl Storage for SlowFirstFinalizeStorage {
async fn list(&self, _: &str) -> ApiResult<Vec<Entry>> {
unreachable!()
}
async fn mkdir(&self, _: &str) -> ApiResult<()> {
unreachable!()
}
async fn delete(&self, _: &str) -> ApiResult<()> {
unreachable!()
}
async fn rename(&self, _: &str, _: &str) -> ApiResult<()> {
unreachable!()
}
async fn get(&self, _: &str) -> ApiResult<(Option<u64>, ByteStream)> {
unreachable!()
}
async fn put(&self, path: &str, mut body: ByteStream) -> ApiResult<()> {
while let Some(item) = body.next().await {
item?;
}
if path.ends_with("/v0") {
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
}
Ok(())
}
}
#[async_trait::async_trait]
impl Storage for FlakyRangeStorage {
async fn list(&self, _: &str) -> ApiResult<Vec<Entry>> {
unreachable!()
}
async fn mkdir(&self, _: &str) -> ApiResult<()> {
unreachable!()
}
async fn delete(&self, _: &str) -> ApiResult<()> {
unreachable!()
}
async fn rename(&self, _: &str, _: &str) -> ApiResult<()> {
unreachable!()
}
async fn get(&self, _: &str) -> ApiResult<(Option<u64>, ByteStream)> {
unreachable!()
}
async fn get_range(&self, _: &str, start: u64, end: u64) -> ApiResult<ByteStream> {
let call = self.calls.fetch_add(1, Ordering::SeqCst);
let bytes = self.encrypted.slice(start as usize..=end as usize);
if call == 0 {
let half = bytes.len() / 2;
Ok(stream::iter(vec![
Ok(bytes.slice(..half)),
Err(io::Error::other("模拟中途断流")),
])
.boxed())
} else {
Ok(stream::iter(vec![Ok(bytes)]).boxed())
}
}
async fn put(&self, _: &str, _: ByteStream) -> ApiResult<()> {
unreachable!()
}
}
fn layout(sizes: &[u64]) -> FileLayout {
let mut offset = 0;
let volumes = sizes
.iter()
.enumerate()
.map(|(i, &size)| {
let v = VolumeMeta {
name: format!("vol{i:02}.bin"),
size,
offset,
};
offset += size;
v
})
.collect();
FileLayout {
volumes,
total: offset,
}
}
struct RecordingRangeStorage {
volumes: std::collections::HashMap<String, Bytes>,
requests: std::sync::Mutex<Vec<(String, u64, u64)>>,
}
#[async_trait::async_trait]
impl Storage for RecordingRangeStorage {
async fn list(&self, _: &str) -> ApiResult<Vec<Entry>> {
unreachable!()
}
async fn mkdir(&self, _: &str) -> ApiResult<()> {
unreachable!()
}
async fn delete(&self, _: &str) -> ApiResult<()> {
unreachable!()
}
async fn rename(&self, _: &str, _: &str) -> ApiResult<()> {
unreachable!()
}
async fn get(&self, _: &str) -> ApiResult<(Option<u64>, ByteStream)> {
unreachable!()
}
async fn get_range(&self, path: &str, start: u64, end: u64) -> ApiResult<ByteStream> {
self.requests
.lock()
.unwrap()
.push((path.to_string(), start, end));
let bytes = self.volumes[path].slice(start as usize..=end as usize);
Ok(stream::iter(vec![Ok(bytes)]).boxed())
}
async fn put(&self, _: &str, _: ByteStream) -> ApiResult<()> {
unreachable!()
}
}
struct SlowEndlessStorage {
calls: AtomicUsize,
}
#[async_trait::async_trait]
impl Storage for SlowEndlessStorage {
async fn list(&self, _: &str) -> ApiResult<Vec<Entry>> {
unreachable!()
}
async fn mkdir(&self, _: &str) -> ApiResult<()> {
unreachable!()
}
async fn delete(&self, _: &str) -> ApiResult<()> {
unreachable!()
}
async fn rename(&self, _: &str, _: &str) -> ApiResult<()> {
unreachable!()
}
async fn get(&self, _: &str) -> ApiResult<(Option<u64>, ByteStream)> {
unreachable!()
}
async fn get_range(&self, _: &str, _: u64, _: u64) -> ApiResult<ByteStream> {
self.calls.fetch_add(1, Ordering::SeqCst);
Ok(stream::unfold((), |()| async {
tokio::time::sleep(std::time::Duration::from_millis(20)).await;
Some((Ok(Bytes::from(vec![0u8; 4096])), ()))
})
.boxed())
}
async fn put(&self, _: &str, _: ByteStream) -> ApiResult<()> {
unreachable!()
}
}
#[tokio::test]
async fn seek_starts_at_target_without_fetching_gap() {
let vol_size = 1_000_000u64;
let all: Vec<u8> = (0..2 * vol_size).map(|i| (i % 251) as u8).collect();
let mut volumes = std::collections::HashMap::new();
volumes.insert(
"vol00.bin".to_string(),
Bytes::copy_from_slice(&all[..vol_size as usize]),
);
volumes.insert(
"vol01.bin".to_string(),
Bytes::copy_from_slice(&all[vol_size as usize..]),
);
let storage = Arc::new(RecordingRangeStorage {
volumes,
requests: std::sync::Mutex::new(Vec::new()),
});
let seek = 1_200_000u64; let mut rx = stream_range_cached_mode(
Arc::clone(&storage) as Arc<dyn Storage>,
String::new(),
[0u8; crate::crypto::SECRET_LEN],
false,
Arc::new(layout(&[vol_size, vol_size])),
seek,
2 * vol_size - 1,
true, &StreamParams {
max_split: 256 * 1024,
max_threads: 4,
max_per_volume: 2,
},
None,
None,
);
let mut out = Vec::new();
while let Some(item) = rx.recv().await {
out.extend_from_slice(&item.unwrap());
}
assert_eq!(out, &all[seek as usize..]);
let requests = storage.requests.lock().unwrap();
assert!(!requests.is_empty());
for (path, start, _) in requests.iter() {
assert_eq!(path, "vol01.bin", "seek 点之前的分卷不应被请求");
assert!(
*start >= seek - vol_size,
"不应回头补 seek 点之前的空档: 请求了卷内偏移 {start}"
);
}
let min_start = requests.iter().map(|(_, start, _)| *start).min().unwrap();
assert_eq!(min_start, seek - vol_size);
}
#[tokio::test]
async fn dropping_receiver_aborts_inflight_fetchers() {
let storage = Arc::new(SlowEndlessStorage {
calls: AtomicUsize::new(0),
});
let total = 64 * 1024 * 1024u64;
let mut rx = stream_range_cached_mode(
Arc::clone(&storage) as Arc<dyn Storage>,
String::new(),
[0u8; crate::crypto::SECRET_LEN],
false,
Arc::new(layout(&[total])),
0,
total - 1,
true,
&StreamParams {
max_split: 1024 * 1024,
max_threads: 4,
max_per_volume: 4,
},
None,
None,
);
let first = rx.recv().await.unwrap().unwrap();
assert!(!first.is_empty());
drop(rx);
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
let after_drop = storage.calls.load(Ordering::SeqCst);
assert!(
after_drop <= 4,
"断开前的在途请求数不应超过并发上限: {after_drop}"
);
tokio::time::sleep(std::time::Duration::from_millis(300)).await;
assert_eq!(
storage.calls.load(Ordering::SeqCst),
after_drop,
"断开后不应再发起新的区间请求"
);
}
#[test]
fn parse_range_cases() {
assert_eq!(parse_range(None, 100), (RangeSpec::Full, true));
assert_eq!(
parse_range(Some("bytes=0-49"), 100),
(RangeSpec::Slice { start: 0, end: 49 }, false)
);
assert_eq!(
parse_range(Some("bytes=10-"), 100),
(RangeSpec::Slice { start: 10, end: 99 }, true),
"开区间 → open_ended"
);
assert_eq!(
parse_range(Some("bytes=-30"), 100),
(RangeSpec::Slice { start: 70, end: 99 }, false)
);
assert_eq!(
parse_range(Some("bytes=0-999"), 100),
(RangeSpec::Slice { start: 0, end: 99 }, false),
"end 截断"
);
assert_eq!(
parse_range(Some("bytes=100-"), 100),
(RangeSpec::Unsatisfiable, false)
);
assert_eq!(
parse_range(Some("bytes=5-2"), 100),
(RangeSpec::Unsatisfiable, false)
);
assert_eq!(parse_range(Some("bytes=abc"), 100), (RangeSpec::Full, true));
assert_eq!(
parse_range(Some("bytes=0-1,5-6"), 100),
(RangeSpec::Full, true)
);
}
#[test]
fn plan_respects_volume_boundaries_and_split() {
let l = layout(&[1000, 1000, 500]);
let plan = plan_chunks(&l, 0, l.total - 1, 400, false);
assert_eq!(plan.len(), 8);
for c in &plan {
let v = &l.volumes[c.vol];
assert!(c.vol_off + c.len <= v.size, "chunk 不跨卷");
assert_eq!(c.merged_start, v.offset + c.vol_off);
}
let mut cur = 0;
for c in &plan {
assert_eq!(c.merged_start, cur);
cur += c.len;
}
assert_eq!(cur, 2500);
}
#[test]
fn plan_head_zone_for_open_ended() {
let l = layout(&[10_000_000]);
let plan = plan_chunks(&l, 0, l.total - 1, 5_000_000, true);
assert_eq!(plan[0].len, HEAD_SMALL_SPLIT);
assert_eq!(plan[3].len, HEAD_SMALL_SPLIT);
assert!(plan[4].len > HEAD_SMALL_SPLIT);
let plan2 = plan_chunks(&l, 0, l.total - 1, 5_000_000, false);
assert_eq!(plan2[0].len, 5_000_000);
}
#[test]
fn open_ended_initial_window_stays_close_to_playback_point() {
let threads = 16;
let l = layout(&[256 * 1024 * 1024]);
let plan = plan_chunks_with_head_count(&l, 0, l.total - 1, 5 * 1024 * 1024, true, threads);
assert!(plan.len() > threads);
assert!(
plan[..threads]
.iter()
.all(|chunk| chunk.len == HEAD_SMALL_SPLIT),
"初始线程窗口必须全部使用小分片"
);
assert_eq!(
plan[threads - 1].merged_start + plan[threads - 1].len,
threads as u64 * HEAD_SMALL_SPLIT,
"默认 16 线程只覆盖播放点附近 4 MiB,而非散到远端"
);
}
#[test]
fn plan_mid_range_starts_in_right_volume() {
let l = layout(&[1000, 1000, 500]);
let plan = plan_chunks(&l, 1500, 2200, 10_000, false);
assert_eq!(plan.len(), 2);
assert_eq!(
plan[0],
PlannedChunk {
merged_start: 1500,
len: 500,
vol: 1,
vol_off: 500
}
);
assert_eq!(
plan[1],
PlannedChunk {
merged_start: 2000,
len: 201,
vol: 2,
vol_off: 0
}
);
}
#[tokio::test]
async fn upload_waits_for_any_finished_volume_not_oldest() {
let storage: Arc<dyn Storage> = Arc::new(SlowFirstFinalizeStorage);
let progress = Arc::new(UploadProgress::new(5));
let sizes = vec![1; MAX_PENDING_UPLOADS + 1];
let names = (0..sizes.len())
.map(|index| format!("v{index}"))
.collect::<Vec<_>>();
let task_progress = Arc::clone(&progress);
let task = tokio::spawn(async move {
upload_stream_planned(
storage,
"folder",
&[0; crate::crypto::SECRET_LEN],
false,
5,
&sizes,
&names,
stream::iter([Ok(Bytes::from_static(b"12345"))]),
task_progress,
)
.await
});
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
assert_eq!(
progress.encrypted.load(Ordering::Relaxed),
5,
"后续卷已完成时,不应被最早卷的收尾响应阻塞前置处理"
);
task.await.unwrap().unwrap();
}
#[tokio::test]
async fn upload_then_stream_roundtrip() {
let dir = tempfile::tempdir().unwrap();
let storage: Arc<dyn Storage> = Arc::from(Box::new(
LocalFs::from_config(&serde_json::json!({"root": dir.path().to_str().unwrap()}))
.unwrap(),
) as Box<dyn Storage>);
let pw = gen_secret();
let plain: Vec<u8> = (0..700_000u32).map(|i| (i * 31 % 256) as u8).collect();
let total = plain.len() as u64;
let volume_size = 256 * 1024u64;
let names = gen_chunk_names(&pw, chunk_count(total, volume_size));
storage.mkdir("ENCFOLDER").await.unwrap();
let body = stream::iter(
plain
.chunks(17_000)
.map(|c| Ok(Bytes::copy_from_slice(c)))
.collect::<Vec<_>>(),
);
let progress = Arc::new(UploadProgress::new(total));
upload_stream(
Arc::clone(&storage),
"ENCFOLDER",
&pw,
total,
volume_size,
&names,
body,
Arc::clone(&progress),
)
.await
.unwrap();
assert_eq!(progress.encrypted.load(Ordering::Relaxed), total);
assert_eq!(
progress.uploaded.load(Ordering::Relaxed),
total,
"localfs 直写:消费即上传"
);
let entries = storage.list("ENCFOLDER").await.unwrap();
assert_eq!(entries.len(), 3);
let disk_total: u64 = entries.iter().map(|e| e.size).sum();
assert_eq!(disk_total, total);
for e in &entries {
assert!(names.contains(&e.name), "{}", e.name);
assert_eq!(e.name.len(), 2);
}
let raw = std::fs::read(dir.path().join("ENCFOLDER").join(&names[0])).unwrap();
assert_ne!(&raw[..], &plain[..raw.len()], "磁盘上必须是密文");
let l = Arc::new(
load_layout(storage.as_ref(), "ENCFOLDER", &pw)
.await
.unwrap(),
);
assert_eq!(
l.volumes.iter().map(|v| v.name.clone()).collect::<Vec<_>>(),
names
);
assert_eq!(l.total, total);
assert_eq!(l.volumes.len(), 3);
let params = StreamParams {
max_split: 100_000,
max_threads: 8,
max_per_volume: 2,
};
let cache_store = crate::cache::CacheStore::new(dir.path().join(".cache")).unwrap();
let cache = cache_store.open("roundtrip", total).unwrap();
let mut rx = stream_range_cached(
Arc::clone(&storage),
"ENCFOLDER".into(),
pw,
Arc::clone(&l),
0,
total - 1,
false,
¶ms,
Some(Arc::clone(&cache)),
);
let mut out = Vec::new();
while let Some(item) = rx.recv().await {
out.extend_from_slice(&item.unwrap());
}
assert_eq!(out, plain, "全量下载解密一致");
assert_eq!(cache_store.stats().bytes_cached, total);
storage.delete("ENCFOLDER").await.unwrap();
for (s, e) in [
(0u64, 0u64),
(262_143, 262_144),
(100_000, 550_000),
(699_999, 699_999),
] {
let mut rx = stream_range_cached(
Arc::clone(&storage),
"ENCFOLDER".into(),
pw,
Arc::clone(&l),
s,
e,
true,
¶ms,
Some(Arc::clone(&cache)),
);
let mut out = Vec::new();
while let Some(item) = rx.recv().await {
out.extend_from_slice(&item.unwrap());
}
assert_eq!(out, &plain[s as usize..=e as usize], "区间 [{s},{e}]");
}
}
#[tokio::test]
async fn upload_size_mismatch_rejected() {
let dir = tempfile::tempdir().unwrap();
let storage: Arc<dyn Storage> = Arc::from(Box::new(
LocalFs::from_config(&serde_json::json!({"root": dir.path().to_str().unwrap()}))
.unwrap(),
) as Box<dyn Storage>);
storage.mkdir("F").await.unwrap();
let pw = gen_secret();
let names = gen_chunk_names(&pw, 1);
let body = stream::iter(vec![Ok(Bytes::from_static(b"short"))]);
let progress = || Arc::new(UploadProgress::new(100));
assert!(
upload_stream(
Arc::clone(&storage),
"F",
&pw,
100,
1024,
&names,
body,
progress()
)
.await
.is_err()
);
let body = stream::iter(vec![Ok(Bytes::from(vec![0u8; 200]))]);
assert!(
upload_stream(
Arc::clone(&storage),
"F",
&pw,
100,
1024,
&names,
body,
progress()
)
.await
.is_err()
);
}
#[tokio::test]
async fn empty_file_uploads_no_volumes() {
let dir = tempfile::tempdir().unwrap();
let storage: Arc<dyn Storage> = Arc::from(Box::new(
LocalFs::from_config(&serde_json::json!({"root": dir.path().to_str().unwrap()}))
.unwrap(),
) as Box<dyn Storage>);
storage.mkdir("E").await.unwrap();
let pw = gen_secret();
let body = stream::iter(Vec::<io::Result<Bytes>>::new());
upload_stream(
Arc::clone(&storage),
"E",
&pw,
0,
1024,
&[],
body,
Arc::new(UploadProgress::new(0)),
)
.await
.unwrap();
let l = load_layout(storage.as_ref(), "E", &pw).await.unwrap();
assert_eq!(l.total, 0);
assert!(l.volumes.is_empty());
}
#[tokio::test]
async fn fetch_chunk_resumes_after_midstream_failure() {
let pw = gen_secret();
let plain = Bytes::from((0..100_000u32).map(|i| (i % 251) as u8).collect::<Vec<_>>());
let mut encrypted = plain.to_vec();
crate::crypto::apply_content_keystream(&pw, 0, &mut encrypted);
let storage = Arc::new(FlakyRangeStorage {
encrypted: Bytes::from(encrypted),
calls: AtomicUsize::new(0),
});
let (tx, mut rx) = mpsc::channel(16);
fetch_chunk(
storage.clone(),
"v".into(),
pw,
true,
PlannedChunk {
merged_start: 0,
len: plain.len() as u64,
vol: 0,
vol_off: 0,
},
tx,
None,
None,
)
.await;
let mut out = Vec::new();
while let Some(item) = rx.recv().await {
out.extend_from_slice(&item.unwrap());
}
assert_eq!(out, plain);
assert_eq!(storage.calls.load(Ordering::SeqCst), 2);
}
}