pub mod baidupan;
pub mod localfs;
pub mod webdav;
use async_trait::async_trait;
use futures_util::stream::BoxStream;
use serde::{Deserialize, Serialize};
use crate::error::{ApiError, ApiResult};
use crate::registry::DataSource;
#[derive(Debug, Clone, Serialize)]
#[serde(rename_all = "camelCase")]
pub struct Entry {
#[serde(skip_serializing_if = "Option::is_none")]
pub id: Option<u64>,
pub name: String,
pub is_dir: bool,
pub size: u64,
pub mtime: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CloudShare {
pub url: String,
pub password: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ImportedEntry {
pub source_name: String,
pub name: String,
}
pub type ByteStream = BoxStream<'static, std::io::Result<bytes::Bytes>>;
pub type ProgressFn = std::sync::Arc<dyn Fn(u64) + Send + Sync>;
#[async_trait]
pub trait Storage: Send + Sync {
fn max_range_size(&self) -> Option<u64> {
None
}
async fn list(&self, path: &str) -> ApiResult<Vec<Entry>>;
async fn mkdir(&self, path: &str) -> ApiResult<()>;
async fn delete(&self, path: &str) -> ApiResult<()>;
async fn rename(&self, from: &str, to: &str) -> ApiResult<()>;
async fn share(&self, _paths: &[String]) -> ApiResult<CloudShare> {
Err(ApiError::BadRequest("该数据源不支持分享".into()))
}
async fn import_share(
&self,
_share: &CloudShare,
_dest: &str,
) -> ApiResult<Vec<ImportedEntry>> {
Err(ApiError::BadRequest("该数据源不支持导入分享".into()))
}
async fn get(&self, path: &str) -> ApiResult<(Option<u64>, ByteStream)>;
async fn get_range(&self, path: &str, start: u64, end: u64) -> ApiResult<ByteStream> {
use futures_util::StreamExt;
if end < start {
return Err(ApiError::BadRequest("非法字节区间".into()));
}
let (_, stream) = self.get(path).await?;
let mut skipped = 0u64;
let mut remaining = end - start + 1;
let filtered = stream.filter_map(move |item| {
let out = match item {
Err(e) => Some(Err(e)),
Ok(b) => {
let mut b = b;
if skipped < start {
let drop_n = ((start - skipped).min(b.len() as u64)) as usize;
skipped += drop_n as u64;
b = b.slice(drop_n..);
}
if b.is_empty() || remaining == 0 {
None
} else {
let take = (b.len() as u64).min(remaining) as usize;
remaining -= take as u64;
Some(Ok(b.slice(..take)))
}
}
};
async move { out }
});
Ok(filtered.boxed())
}
async fn put(&self, path: &str, body: ByteStream) -> ApiResult<()>;
async fn put_sized(&self, path: &str, _size: u64, body: ByteStream) -> ApiResult<()> {
self.put(path, body).await
}
async fn put_sized_tracked(
&self,
path: &str,
size: u64,
body: ByteStream,
progress: ProgressFn,
) -> ApiResult<()> {
use futures_util::StreamExt;
let counted = body.map(move |item| {
if let Ok(b) = &item {
progress(b.len() as u64);
}
item
});
self.put_sized(path, size, counted.boxed()).await
}
}
pub fn make_with_token_persister(
ds: &DataSource,
http: reqwest::Client,
persist_tokens: Option<baidupan::TokenPersister>,
) -> ApiResult<Box<dyn Storage>> {
match ds.ds_type.as_str() {
"localfs" => Ok(Box::new(localfs::LocalFs::from_config(&ds.config)?)),
"webdav" => Ok(Box::new(webdav::WebdavFs::from_config(&ds.config, http)?)),
"baidupan" => Ok(Box::new(baidupan::BaiduPanFs::from_config_with_persister(
&ds.config,
http,
persist_tokens,
)?)),
other => Err(ApiError::BadRequest(format!("未知数据源类型: {other}"))),
}
}
pub fn make_arc_with_token_persister(
ds: &DataSource,
http: reqwest::Client,
persist_tokens: Option<baidupan::TokenPersister>,
) -> ApiResult<std::sync::Arc<dyn Storage>> {
Ok(std::sync::Arc::from(make_with_token_persister(
ds,
http,
persist_tokens,
)?))
}
pub(crate) fn log_headers(headers: &reqwest::header::HeaderMap) -> String {
headers
.iter()
.map(|(k, v)| {
let v = if k == reqwest::header::SET_COOKIE {
"…(已脱敏)".into()
} else {
String::from_utf8_lossy(v.as_bytes()).into_owned()
};
format!("{k}: {v}")
})
.collect::<Vec<_>>()
.join("; ")
}
pub fn sanitize(path: &str) -> ApiResult<String> {
let mut parts: Vec<&str> = Vec::new();
for seg in path.split('/') {
if seg.is_empty() || seg == "." {
continue;
}
if seg == ".." || seg.contains('\\') || seg.bytes().any(|b| b < 0x20 || b == 0x7f) {
return Err(ApiError::BadRequest(format!("非法路径: {path}")));
}
parts.push(seg);
}
Ok(parts.join("/"))
}
#[cfg(test)]
mod tests {
use super::sanitize;
#[test]
fn sanitize_normalizes_and_rejects() {
assert_eq!(sanitize("").unwrap(), "");
assert_eq!(sanitize("/").unwrap(), "");
assert_eq!(sanitize("a/b/c").unwrap(), "a/b/c");
assert_eq!(sanitize("/a//b/./c/").unwrap(), "a/b/c");
assert!(sanitize("a/../b").is_err());
assert!(sanitize("..").is_err());
assert!(sanitize("a\\b").is_err());
assert!(sanitize("a/\u{0}b").is_err());
}
}