diff --git a/src-tauri/src/services/webdav.rs b/src-tauri/src/services/webdav.rs index c9d376e70..a6d5be1ab 100644 --- a/src-tauri/src/services/webdav.rs +++ b/src-tauri/src/services/webdav.rs @@ -8,6 +8,7 @@ use std::time::Duration; use crate::error::AppError; use crate::proxy::http_client; +use futures::StreamExt; const DEFAULT_TIMEOUT_SECS: u64 = 30; /// Timeout for large file transfers (PUT/GET of db.sql, skills.zip). @@ -237,15 +238,7 @@ pub async fn put_bytes( ) .send() .await - .map_err(|e| { - webdav_transport_error( - "webdav.put_failed", - "PUT 请求", - "PUT request", - url, - &e, - ) - })?; + .map_err(|e| webdav_transport_error("webdav.put_failed", "PUT 请求", "PUT request", url, &e))?; if resp.status().is_success() { return Ok(()); @@ -259,6 +252,7 @@ pub async fn put_bytes( pub async fn get_bytes( url: &str, auth: &WebDavAuth, + max_bytes: usize, ) -> Result, Option)>, AppError> { let client = http_client::get(); let resp = apply_auth( @@ -269,15 +263,7 @@ pub async fn get_bytes( ) .send() .await - .map_err(|e| { - webdav_transport_error( - "webdav.get_failed", - "GET 请求", - "GET request", - url, - &e, - ) - })?; + .map_err(|e| webdav_transport_error("webdav.get_failed", "GET 请求", "GET request", url, &e))?; if resp.status() == StatusCode::NOT_FOUND { return Ok(None); @@ -285,22 +271,29 @@ pub async fn get_bytes( if !resp.status().is_success() { return Err(webdav_status_error("GET", resp.status(), url)); } + ensure_content_length_within_limit(resp.headers(), max_bytes, url)?; + let etag = resp .headers() .get("etag") .and_then(|v| v.to_str().ok()) .map(|s| s.to_string()); - let bytes = resp - .bytes() - .await - .map_err(|e| { + let mut bytes = Vec::new(); + let mut stream = resp.bytes_stream(); + while let Some(chunk) = stream.next().await { + let chunk = chunk.map_err(|e| { AppError::localized( "webdav.response_read_failed", format!("读取 WebDAV 响应失败: {e}"), format!("Failed to read WebDAV response: {e}"), ) })?; - Ok(Some((bytes.to_vec(), etag))) + if bytes.len().saturating_add(chunk.len()) > max_bytes { + return Err(response_too_large_error(url, max_bytes)); + } + bytes.extend_from_slice(&chunk); + } + Ok(Some((bytes, etag))) } /// HEAD request to retrieve the ETag. Returns `None` on 404. @@ -315,13 +308,7 @@ pub async fn head_etag(url: &str, auth: &WebDavAuth) -> Result, A .send() .await .map_err(|e| { - webdav_transport_error( - "webdav.head_failed", - "HEAD 请求", - "HEAD request", - url, - &e, - ) + webdav_transport_error("webdav.head_failed", "HEAD 请求", "HEAD request", url, &e) })?; if resp.status() == StatusCode::NOT_FOUND { @@ -386,9 +373,7 @@ pub fn webdav_status_error(op: &str, status: StatusCode, url: &str) -> AppError if matches!(status, StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN) { if jgy { - zh.push_str( - "。坚果云请使用「第三方应用密码」,并确认地址指向 /dav/ 下的目录。", - ); + zh.push_str("。坚果云请使用「第三方应用密码」,并确认地址指向 /dav/ 下的目录。"); en.push_str( ". For Jianguoyun, use an app-specific password and ensure the URL points under /dav/.", ); @@ -401,9 +386,7 @@ pub fn webdav_status_error(op: &str, status: StatusCode, url: &str) -> AppError en.push_str(". Common Jianguoyun cause: URL is outside a writable /dav/ directory."); } else if op == "MKCOL" && status == StatusCode::CONFLICT { if jgy { - zh.push_str( - "。坚果云不允许自动创建顶层文件夹,请先在网页端手动创建后重试。", - ); + zh.push_str("。坚果云不允许自动创建顶层文件夹,请先在网页端手动创建后重试。"); en.push_str( ". Jianguoyun does not allow creating top-level folders automatically; create it manually first.", ); @@ -446,9 +429,47 @@ fn redact_url(raw: &str) -> String { } } +fn response_too_large_error(url: &str, max_bytes: usize) -> AppError { + let max_mb = max_bytes / 1024 / 1024; + AppError::localized( + "webdav.response_too_large", + format!( + "WebDAV 响应体超过上限({} MB): {}", + max_mb, + redact_url(url) + ), + format!( + "WebDAV response body exceeds limit ({} MB): {}", + max_mb, + redact_url(url) + ), + ) +} + +fn ensure_content_length_within_limit( + headers: &reqwest::header::HeaderMap, + max_bytes: usize, + url: &str, +) -> Result<(), AppError> { + let Some(content_length) = headers.get(reqwest::header::CONTENT_LENGTH) else { + return Ok(()); + }; + let Ok(raw) = content_length.to_str() else { + return Ok(()); + }; + let Ok(value) = raw.parse::() else { + return Ok(()); + }; + if value > max_bytes as u64 { + return Err(response_too_large_error(url, max_bytes)); + } + Ok(()) +} + #[cfg(test)] mod tests { use super::*; + use reqwest::header::{HeaderMap, HeaderValue, CONTENT_LENGTH}; #[test] fn build_remote_url_encodes_path_segments() { @@ -498,10 +519,34 @@ mod tests { #[test] fn redact_url_hides_credentials_and_query_values() { let redacted = redact_url("https://alice:secret@example.com:8443/dav?token=abc&foo=1"); - assert_eq!( - redacted, - "https://example.com:8443/dav?[keys:foo,token]" - ); + assert_eq!(redacted, "https://example.com:8443/dav?[keys:foo,token]"); assert!(!redacted.contains("secret")); } + + #[test] + fn ensure_content_length_within_limit_accepts_missing_or_small_values() { + let empty = HeaderMap::new(); + assert!( + ensure_content_length_within_limit(&empty, 1024, "https://dav.example.com").is_ok() + ); + + let mut small = HeaderMap::new(); + small.insert(CONTENT_LENGTH, HeaderValue::from_static("1024")); + assert!( + ensure_content_length_within_limit(&small, 1024, "https://dav.example.com").is_ok() + ); + } + + #[test] + fn ensure_content_length_within_limit_rejects_oversized_values() { + let mut large = HeaderMap::new(); + large.insert(CONTENT_LENGTH, HeaderValue::from_static("2048")); + + let err = ensure_content_length_within_limit(&large, 1024, "https://dav.example.com") + .expect_err("oversized response should be rejected"); + assert!( + err.to_string().contains("too large") || err.to_string().contains("超过"), + "unexpected error: {err}" + ); + } } diff --git a/src-tauri/src/services/webdav_sync.rs b/src-tauri/src/services/webdav_sync.rs index 2144e4e80..6bdd6f4b0 100644 --- a/src-tauri/src/services/webdav_sync.rs +++ b/src-tauri/src/services/webdav_sync.rs @@ -35,6 +35,8 @@ const REMOTE_DB_SQL: &str = "db.sql"; const REMOTE_SKILLS_ZIP: &str = "skills.zip"; const REMOTE_MANIFEST: &str = "manifest.json"; const MAX_DEVICE_NAME_LEN: usize = 64; +const MAX_MANIFEST_BYTES: usize = 1024 * 1024; +pub(super) const MAX_SYNC_ARTIFACT_BYTES: u64 = 512 * 1024 * 1024; pub fn sync_mutex() -> &'static tokio::sync::Mutex<()> { static LOCK: OnceLock> = OnceLock::new(); @@ -160,13 +162,15 @@ pub async fn download( let auth = auth_for(settings); let manifest_url = remote_file_url(settings, REMOTE_MANIFEST)?; - let (manifest_bytes, etag) = get_bytes(&manifest_url, &auth).await?.ok_or_else(|| { - localized( - "webdav.sync.remote_empty", - "远端没有可下载的同步数据", - "No downloadable sync data found on the remote.", - ) - })?; + let (manifest_bytes, etag) = get_bytes(&manifest_url, &auth, MAX_MANIFEST_BYTES) + .await? + .ok_or_else(|| { + localized( + "webdav.sync.remote_empty", + "远端没有可下载的同步数据", + "No downloadable sync data found on the remote.", + ) + })?; let manifest: SyncManifest = serde_json::from_slice(&manifest_bytes).map_err(|e| AppError::Json { @@ -196,7 +200,7 @@ pub async fn fetch_remote_info(settings: &WebDavSyncSettings) -> Result String { } fn detect_system_device_name() -> Option { - let env_name = [ - "CC_SWITCH_DEVICE_NAME", - "COMPUTERNAME", - "HOSTNAME", - ] - .iter() - .filter_map(|key| std::env::var(key).ok()) - .find_map(|value| normalize_device_name(&value)); + let env_name = ["CC_SWITCH_DEVICE_NAME", "COMPUTERNAME", "HOSTNAME"] + .iter() + .filter_map(|key| std::env::var(key).ok()) + .find_map(|value| normalize_device_name(&value)); if env_name.is_some() { return env_name; @@ -357,21 +357,26 @@ fn detect_system_device_name() -> Option { } fn normalize_device_name(raw: &str) -> Option { - let compact = raw.chars().fold(String::with_capacity(raw.len()), |mut acc, ch| { - if ch.is_whitespace() { - acc.push(' '); - } else if !ch.is_control() { - acc.push(ch); - } - acc - }); + let compact = raw + .chars() + .fold(String::with_capacity(raw.len()), |mut acc, ch| { + if ch.is_whitespace() { + acc.push(' '); + } else if !ch.is_control() { + acc.push(ch); + } + acc + }); let normalized = compact.split_whitespace().collect::>().join(" "); let trimmed = normalized.trim(); if trimmed.is_empty() { return None; } - let limited = trimmed.chars().take(MAX_DEVICE_NAME_LEN).collect::(); + let limited = trimmed + .chars() + .take(MAX_DEVICE_NAME_LEN) + .collect::(); if limited.is_empty() { None } else { @@ -421,14 +426,18 @@ async fn download_and_verify( format!("Manifest missing artifact: {artifact_name}"), ) })?; + validate_artifact_size_limit(artifact_name, meta.size)?; + let url = remote_file_url(settings, artifact_name)?; - let (bytes, _) = get_bytes(&url, auth).await?.ok_or_else(|| { - localized( - "webdav.sync.remote_missing_artifact", - format!("远端缺少 artifact 文件: {artifact_name}"), - format!("Remote artifact file missing: {artifact_name}"), - ) - })?; + let (bytes, _) = get_bytes(&url, auth, MAX_SYNC_ARTIFACT_BYTES as usize) + .await? + .ok_or_else(|| { + localized( + "webdav.sync.remote_missing_artifact", + format!("远端缺少 artifact 文件: {artifact_name}"), + format!("Remote artifact file missing: {artifact_name}"), + ) + })?; // Quick size check before expensive hash if bytes.len() as u64 != meta.size { @@ -519,6 +528,21 @@ fn auth_for(settings: &WebDavSyncSettings) -> WebDavAuth { auth_from_credentials(&settings.username, &settings.password) } +fn validate_artifact_size_limit(artifact_name: &str, size: u64) -> Result<(), AppError> { + if size > MAX_SYNC_ARTIFACT_BYTES { + let max_mb = MAX_SYNC_ARTIFACT_BYTES / 1024 / 1024; + return Err(localized( + "webdav.sync.artifact_too_large", + format!("artifact {artifact_name} 超过下载上限({} MB)", max_mb), + format!( + "Artifact {artifact_name} exceeds download limit ({} MB)", + max_mb + ), + )); + } + Ok(()) +} + // ─── Tests ─────────────────────────────────────────────────── #[cfg(test)] @@ -662,4 +686,19 @@ mod tests { "manifest should not contain deviceId" ); } + + #[test] + fn validate_artifact_size_limit_rejects_oversized_artifacts() { + let err = validate_artifact_size_limit("skills.zip", MAX_SYNC_ARTIFACT_BYTES + 1) + .expect_err("artifact larger than limit should be rejected"); + assert!( + err.to_string().contains("too large") || err.to_string().contains("超过"), + "unexpected error: {err}" + ); + } + + #[test] + fn validate_artifact_size_limit_accepts_limit_boundary() { + assert!(validate_artifact_size_limit("skills.zip", MAX_SYNC_ARTIFACT_BYTES).is_ok()); + } } diff --git a/src-tauri/src/services/webdav_sync/archive.rs b/src-tauri/src/services/webdav_sync/archive.rs index 2058429e1..aae5de1a8 100644 --- a/src-tauri/src/services/webdav_sync/archive.rs +++ b/src-tauri/src/services/webdav_sync/archive.rs @@ -10,10 +10,8 @@ use zip::DateTime; use crate::error::AppError; use crate::services::skill::SkillService; -use super::{io_context_localized, localized, REMOTE_SKILLS_ZIP}; +use super::{io_context_localized, localized, MAX_SYNC_ARTIFACT_BYTES, REMOTE_SKILLS_ZIP}; -/// Maximum total bytes allowed during zip extraction (512 MB). -const MAX_EXTRACT_BYTES: u64 = 512 * 1024 * 1024; /// Maximum number of entries allowed in a zip archive. const MAX_EXTRACT_ENTRIES: usize = 10_000; @@ -92,8 +90,14 @@ pub(super) fn restore_skills_zip(raw: &[u8]) -> Result<(), AppError> { if archive.len() > MAX_EXTRACT_ENTRIES { return Err(localized( "webdav.sync.skills_zip_too_many_entries", - format!("skills.zip 条目数过多({}),上限 {MAX_EXTRACT_ENTRIES}", archive.len()), - format!("skills.zip has too many entries ({}), limit is {MAX_EXTRACT_ENTRIES}", archive.len()), + format!( + "skills.zip 条目数过多({}),上限 {MAX_EXTRACT_ENTRIES}", + archive.len() + ), + format!( + "skills.zip has too many entries ({}), limit is {MAX_EXTRACT_ENTRIES}", + archive.len() + ), )); } @@ -118,15 +122,13 @@ pub(super) fn restore_skills_zip(raw: &[u8]) -> Result<(), AppError> { fs::create_dir_all(parent).map_err(|e| AppError::io(parent, e))?; } let mut out = fs::File::create(&out_path).map_err(|e| AppError::io(&out_path, e))?; - let written = std::io::copy(&mut entry, &mut out).map_err(|e| AppError::io(&out_path, e))?; - total_bytes += written; - if total_bytes > MAX_EXTRACT_BYTES { - return Err(localized( - "webdav.sync.skills_zip_too_large", - format!("skills.zip 解压后体积超过上限({} MB)", MAX_EXTRACT_BYTES / 1024 / 1024), - format!("skills.zip extracted size exceeds limit ({} MB)", MAX_EXTRACT_BYTES / 1024 / 1024), - )); - } + let _written = copy_entry_with_total_limit( + &mut entry, + &mut out, + &mut total_bytes, + MAX_SYNC_ARTIFACT_BYTES, + &out_path, + )?; } let ssot = SkillService::get_ssot_dir().map_err(|e| { @@ -327,10 +329,47 @@ fn mark_visited_dir(path: &Path, visited: &mut HashSet) -> Result( + reader: &mut R, + writer: &mut W, + total_bytes: &mut u64, + max_total_bytes: u64, + out_path: &Path, +) -> Result { + let mut buffer = [0u8; 16 * 1024]; + let mut written = 0u64; + loop { + let n = reader + .read(&mut buffer) + .map_err(|e| AppError::io(out_path, e))?; + if n == 0 { + break; + } + + if total_bytes.saturating_add(n as u64) > max_total_bytes { + let max_mb = max_total_bytes / 1024 / 1024; + return Err(localized( + "webdav.sync.skills_zip_too_large", + format!("skills.zip 解压后体积超过上限({} MB)", max_mb), + format!("skills.zip extracted size exceeds limit ({} MB)", max_mb), + )); + } + + writer + .write_all(&buffer[..n]) + .map_err(|e| AppError::io(out_path, e))?; + *total_bytes += n as u64; + written += n as u64; + } + Ok(written) +} + #[cfg(test)] mod tests { - use super::mark_visited_dir; + use super::{copy_entry_with_total_limit, mark_visited_dir}; use std::collections::HashSet; + use std::io::Cursor; + use std::path::Path; use tempfile::tempdir; #[test] @@ -343,4 +382,29 @@ mod tests { assert!(mark_visited_dir(&dir, &mut visited).expect("first visit")); assert!(!mark_visited_dir(&dir, &mut visited).expect("second visit")); } + + #[test] + fn copy_entry_with_total_limit_rejects_oversized_stream_before_write() { + let mut reader = Cursor::new(vec![1u8; 16]); + let mut writer = Vec::new(); + let mut total_bytes = 0u64; + + let err = copy_entry_with_total_limit( + &mut reader, + &mut writer, + &mut total_bytes, + 8, + Path::new("skills-extracted/file.bin"), + ) + .expect_err("stream larger than limit should be rejected"); + assert!( + err.to_string().contains("too large") || err.to_string().contains("超过"), + "unexpected error: {err}" + ); + assert_eq!( + writer.len(), + 0, + "should not write when the first chunk exceeds limit" + ); + } }