use crate::config::{ServerConfig, read_config}; use axum::body::{Body, HttpBody}; use axum::http::{HeaderMap, Request, Response, StatusCode, header}; use axum::middleware::Next; use axum::routing::{get, post}; use axum::{Router, middleware}; use base64::Engine; use base64::prelude::BASE64_STANDARD; use chrono::Utc; use ed25519_dalek::{Signature, VerifyingKey}; use sha2::{Digest, Sha256}; use std::fs::{File, ReadDir, rename}; use std::future::poll_fn; use std::io::{BufWriter, Read, Write}; use std::path::Path; use std::pin::Pin; use std::sync::OnceLock; use tokio::net::TcpListener; static PUBLIC_KEY: OnceLock = OnceLock::new(); static SERVER_CONFIG: OnceLock = OnceLock::new(); enum ChallengeStatus { BadTimestamp, BadNonce, BadSignature, Success, } fn parse_header(headers: &HeaderMap, key: &str) -> Result { headers .get(key) .and_then(|value| value.to_str().ok()) .and_then(|string| string.parse().ok()) .ok_or(StatusCode::BAD_REQUEST) } fn log_challenge(signature: &str, status: ChallengeStatus) { let short_signature = signature.chars().take(10).collect::(); match status { ChallengeStatus::BadTimestamp => println!( "[*] Challenge {}... was invalid due to bad timestamp", short_signature ), ChallengeStatus::BadNonce => println!( "[*] Challenge {}... was invalid due to malformed nonce", short_signature ), ChallengeStatus::BadSignature => println!( "[*] Challenge {}... was invalid due to an invalid signature", short_signature ), ChallengeStatus::Success => println!("[*] Challenge {}... succeeded", short_signature), } } fn create_symlink(original: &Path, link: &Path, _is_dir: bool) -> std::io::Result<()> { #[cfg(unix)] { let full_target = std::path::absolute(original)?; std::os::unix::fs::symlink(full_target, link) } #[cfg(windows)] { if _is_dir { std::os::windows::fs::symlink_dir(original, link) } else { std::os::windows::fs::symlink_file(original, link) } } } fn perform_download(file: String) -> Result, StatusCode> { let file_bytes = std::fs::read(file).map_err(|_| StatusCode::NOT_FOUND)?; let body = Body::from(file_bytes); let response = Response::builder() .status(StatusCode::OK) .header(header::CONTENT_TYPE, "application/x-keepass2") .body(body) .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; Ok(response) } fn find_short_id_candidates(path: &str, vaults: ReadDir) -> Vec { let candidates: Vec<_> = vaults .filter_map(|entry| entry.ok()) .map(|entry| entry.path()) .filter(|entry| entry.is_file()) .map(|entry| entry.to_string_lossy().into_owned()) .filter(|entry| entry[7..].starts_with(path)) .collect(); candidates } async fn auth_middleware(req: Request, next: Next) -> Result, StatusCode> { let headers = req.headers(); let signature: String = parse_header(headers, "X-Signature")?; let nonce: String = parse_header(headers, "X-Nonce")?; let timestamp: i64 = parse_header(headers, "X-Timestamp")?; if !(0..=60).contains(&(Utc::now().timestamp() - timestamp)) { log_challenge(&signature, ChallengeStatus::BadTimestamp); return Err(StatusCode::UNAUTHORIZED); } if nonce.len() != 32 { log_challenge(&signature, ChallengeStatus::BadNonce); println!("[*] A challenge was failed for a malformed nonce"); return Err(StatusCode::UNAUTHORIZED); } let message = format!("{}{}", nonce, timestamp); let verifier_signature = Signature::from_bytes( &BASE64_STANDARD .decode(signature.as_bytes()) .map_err(|_| StatusCode::UNAUTHORIZED)? .try_into() .map_err(|_| StatusCode::UNAUTHORIZED)?, ); match PUBLIC_KEY .get() .unwrap() .verify_strict(message.as_bytes(), &verifier_signature) { Ok(_) => { log_challenge(&signature, ChallengeStatus::Success); Ok(next.run(req).await) } Err(_) => { log_challenge(&signature, ChallengeStatus::BadSignature); Err(StatusCode::UNAUTHORIZED) } } } async fn upload(mut body: Body) -> Result { let temp_filename = "vaults/temporary.kdbx"; let file = File::create(temp_filename).map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; let mut writer = BufWriter::new(file); let mut hasher = Sha256::new(); loop { let frame_result = poll_fn(|cx| Pin::new(&mut body).poll_frame(cx)).await; match frame_result { Some(Ok(chunk)) => { if let Some(chunk) = chunk.data_ref() { Digest::update(&mut hasher, chunk); writer .write_all(chunk) .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; } } Some(Err(_)) => return Err(StatusCode::INTERNAL_SERVER_ERROR), None => break, } } writer .flush() .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; let file_hash: String = hasher .finalize() .iter() .map(|b| format!("{:02x}", b)) .collect(); println!( "[*] Uploading file with hash {}...", file_hash.chars().take(8).collect::() ); let vaults = std::fs::read_dir("vaults").map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; for vault in vaults { let entry = vault.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; let path = entry.path(); if path.is_file() && let Some(file_name) = path.file_name() && file_name.to_string_lossy().starts_with(file_hash.as_str()) { println!("[*] File already exists, doing nothing..."); std::fs::remove_file(temp_filename).map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; return Ok(StatusCode::NOT_MODIFIED); } } let final_filename = format!("vaults/{}_{}.kdbx", file_hash, Utc::now().timestamp()); rename(temp_filename, &final_filename).map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; // A temporary symlink is created, because creating a symlink is not atomic whereas // renaming a file is atomic in Unix-like systems. This pattern is to prevent data races // from concurrent calls to upload or concurrent calls to get the latest vault. let tmp_symlink_name = format!("vaults/latest-{}.kdbx", file_hash); let target = Path::new(&final_filename); let tmp_symlink = Path::new(&tmp_symlink_name); create_symlink(target, tmp_symlink, false).map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; rename(tmp_symlink, "vaults/latest.kdbx").map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; println!("[*] File uploaded and latest symlink reassigned"); Ok(StatusCode::OK) } async fn download_vault( axum::extract::Path(path): axum::extract::Path, ) -> Result, StatusCode> { if path == "latest" { println!("[*] Downloading latest vault"); return perform_download(String::from("vaults/latest.kdbx")); } if path.len() != 8 && path.len() != 64 { return Err(StatusCode::BAD_REQUEST); } println!("[*] Received download request for path: {}", path); let vaults = std::fs::read_dir("vaults").map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; if path.len() == 8 { let candidates = find_short_id_candidates(&path, vaults); if candidates.is_empty() { println!("[*] No candidates found for the given short ID"); return Err(StatusCode::NOT_FOUND); } if candidates.len() != 1 { println!("[*] Multiple candidates found for the given short ID"); return Err(StatusCode::MULTIPLE_CHOICES); } println!("[*] File found found for the given short ID"); perform_download(candidates[0].clone()) } else { for vault in vaults { let entry = match vault { Ok(v) => v, Err(_) => continue, }; let file_name = entry.path(); if file_name.to_string_lossy().contains(path.as_str()) { println!("[*] File found for the given long ID"); return perform_download(file_name.to_string_lossy().into_owned()); } } Err(StatusCode::NOT_FOUND) } } async fn list_vaults() -> Result { let mut s = String::new(); let path = "vaults"; let mut counter = 0; for entry in std::fs::read_dir(path).map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)? { let entry = entry.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?.path(); if entry.is_file() && !entry.is_symlink() && let Some(path) = entry.as_path().file_stem() && let Some(path) = path.to_str() { s += path; s += "\n"; counter += 1; } } println!("[*] Listed {} vaults", counter); Ok(s) } async fn delete_vault( axum::extract::Path(path): axum::extract::Path, ) -> Result { if path.len() != 8 && path.len() != 64 { return Err(StatusCode::BAD_REQUEST); } println!("[*] Received deletion request for path: {}", path); let vaults = std::fs::read_dir("vaults").map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; if path.len() == 8 { let candidates = find_short_id_candidates(&path, vaults); if candidates.is_empty() { println!("[*] No candidates found for the given short ID"); return Err(StatusCode::NOT_FOUND); } if candidates.len() != 1 { println!("[*] Multiple candidates found for the given short ID"); return Err(StatusCode::MULTIPLE_CHOICES); } println!("[*] File found found for the given short ID"); std::fs::remove_file(candidates[0].clone()) .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; Ok(StatusCode::OK) } else { for vault in vaults { let entry = match vault { Ok(v) => v, Err(_) => continue, }; let file_name = entry.path(); if file_name.to_string_lossy().contains(path.as_str()) { println!("[*] File found for the given long ID"); std::fs::remove_file(file_name).map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; return Ok(StatusCode::OK); } } Err(StatusCode::NOT_FOUND) } } pub async fn server_main() { SERVER_CONFIG.set(read_config().unwrap()).unwrap(); let mut buffer: [u8; 44] = [0u8; 44]; let pubkey_file = SERVER_CONFIG.get().unwrap().public_key.clone(); let file = File::open(pubkey_file); if let Ok(mut file) = file { file.read_exact(&mut buffer).expect("Failed to read file"); } else { eprintln!("Failed to read public key!"); return; } let decoded = BASE64_STANDARD .decode(buffer) .expect("Malformed private key"); let verifying_key = VerifyingKey::from_bytes(&decoded.try_into().expect("Malformed public key")) .expect("Malformed public key"); PUBLIC_KEY.set(verifying_key).unwrap(); let app = Router::new() .route("/vaults", post(upload).get(list_vaults)) .route("/vaults/{*path}", get(download_vault).delete(delete_vault)) .layer(middleware::from_fn(auth_middleware)); let listener = TcpListener::bind("0.0.0.0:3000").await.unwrap(); axum::serve(listener, app).await.unwrap(); }