fix(web): harden JDBC uploads and CORS

This commit is contained in:
t8y2 2026-07-11 17:55:29 +08:00
parent 3747f08baa
commit e4563f9659
2 changed files with 39 additions and 16 deletions

View File

@ -17,10 +17,8 @@ use axum::routing::{delete, get, post};
use axum::Router;
use dbx_core::connection::AppState;
use dbx_core::storage::Storage;
use tokio::sync::RwLock;
use tower_http::cors::{Any, CorsLayer};
use state::WebState;
use tokio::sync::RwLock;
fn web_body_limit_bytes() -> usize {
const DEFAULT_MB: usize = 1024;
@ -210,9 +208,6 @@ async fn main() {
export_files: RwLock::new(HashMap::new()),
});
// CORS
let cors = CorsLayer::new().allow_origin(Any).allow_methods(Any).allow_headers(Any);
// API routes
let api = Router::new()
// Auth
@ -594,8 +589,7 @@ async fn main() {
let mut app = Router::new()
.nest("/api", api)
.layer(DefaultBodyLimit::max(web_body_limit_bytes()))
.layer(tower_http::trace::TraceLayer::new_for_http())
.layer(cors);
.layer(tower_http::trace::TraceLayer::new_for_http());
// Static file serving
if let Ok(static_dir) = std::env::var("DBX_STATIC_DIR") {

View File

@ -1,3 +1,4 @@
use std::path::{Component, Path as FsPath};
use std::sync::Arc;
use axum::extract::{Multipart, Path, State};
@ -12,6 +13,20 @@ use tokio::sync::broadcast;
use crate::error::AppError;
use crate::state::WebState;
fn safe_upload_file_name(file_name: Option<&str>, fallback: &str, extension: &str) -> Result<String, AppError> {
let file_name = file_name.unwrap_or(fallback);
let mut components = FsPath::new(file_name).components();
let is_single_normal_component =
matches!(components.next(), Some(Component::Normal(_))) && components.next().is_none();
if !is_single_normal_component
|| file_name.contains(['/', '\\'])
|| !file_name.to_ascii_lowercase().ends_with(extension)
{
return Err(AppError::bad_request("Invalid upload file name"));
}
Ok(file_name.to_string())
}
// ---- JDBC Drivers ----
pub async fn list_jdbc_drivers(State(state): State<Arc<WebState>>) -> Result<Json<Vec<JdbcDriverInfo>>, AppError> {
@ -62,10 +77,7 @@ pub async fn import_jdbc_drivers(
let mut imported = Vec::new();
while let Ok(Some(field)) = multipart.next_field().await {
let file_name = field.file_name().unwrap_or("driver.jar").to_string();
if !file_name.to_lowercase().ends_with(".jar") {
return Err(AppError::bad_request("Only .jar files can be imported"));
}
let file_name = safe_upload_file_name(field.file_name(), "driver.jar", ".jar")?;
let data = field.bytes().await.map_err(|e| AppError::internal(e.to_string()))?;
let target = jdbc::unique_target_path(&upload_dir, &file_name);
std::fs::write(&target, &data).map_err(|e| AppError::internal(e.to_string()))?;
@ -136,10 +148,7 @@ pub async fn install_jdbc_plugin_local(
std::fs::create_dir_all(&temp_dir).map_err(|e| AppError::internal(e.to_string()))?;
if let Ok(Some(field)) = multipart.next_field().await {
let file_name = field.file_name().unwrap_or("plugin.zip").to_string();
if !file_name.to_lowercase().ends_with(".zip") {
return Err(AppError::bad_request("Only .zip files can be imported for JDBC plugin"));
}
let file_name = safe_upload_file_name(field.file_name(), "plugin.zip", ".zip")?;
let data = field.bytes().await.map_err(|e| AppError::internal(e.to_string()))?;
let tmp_path = temp_dir.join(&file_name);
std::fs::write(&tmp_path, &data).map_err(|e| AppError::internal(e.to_string()))?;
@ -163,6 +172,26 @@ pub async fn list_system_fonts() -> Result<Json<Vec<String>>, AppError> {
Ok(Json(jdbc::list_system_fonts()))
}
#[cfg(test)]
mod tests {
use super::safe_upload_file_name;
#[test]
fn upload_file_name_rejects_path_components() {
for value in ["../driver.jar", "sub/driver.jar", r"sub\driver.jar", "/tmp/driver.jar"] {
assert!(safe_upload_file_name(Some(value), "driver.jar", ".jar").is_err(), "{value}");
}
}
#[test]
fn upload_file_name_accepts_plain_expected_extension() {
let Ok(file_name) = safe_upload_file_name(Some("postgresql-42.7.jar"), "driver.jar", ".jar") else {
panic!("plain JAR filename should be accepted");
};
assert_eq!(file_name, "postgresql-42.7.jar");
}
}
async fn progress_sender(state: &WebState, operation_id: &str) -> broadcast::Sender<String> {
let mut channels = state.sse_channels.write().await;
channels