diff --git a/src/api/client.rs b/src/api/client.rs index add7ed9..e327968 100644 --- a/src/api/client.rs +++ b/src/api/client.rs @@ -1,5 +1,7 @@ use reqwest::blocking::{Client, Response}; -use reqwest::header::{ACCEPT, AUTHORIZATION, CONTENT_TYPE, HeaderMap, HeaderValue}; +use reqwest::header::{ + ACCEPT, AUTHORIZATION, CONTENT_LENGTH, CONTENT_TYPE, HeaderMap, HeaderValue, +}; use serde::Serialize; use serde::de::DeserializeOwned; @@ -148,6 +150,33 @@ impl ApiClient { self.handle_response(response) } + pub fn put_file( + &self, + url: &str, + file: std::fs::File, + content_length: u64, + ) -> Result<(), ApiError> { + let response = self + .client + .put(url) + .header(CONTENT_TYPE, "application/gzip") + .header(CONTENT_LENGTH, content_length) + .body(reqwest::blocking::Body::from(file)) + .send() + .map_err(ApiError::NetworkError)?; + + if response.status().is_success() { + Ok(()) + } else { + let status = response.status(); + let body = response.text().map_err(ApiError::NetworkError)?; + Err(ApiError::Other(format!( + "Upload failed ({}): {}", + status, body + ))) + } + } + pub fn delete(&self, path: &str) -> Result { let url = format!("{}{}", self.base_url, path); let response = self diff --git a/src/cli.rs b/src/cli.rs index 0325b5d..1791abb 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -49,6 +49,11 @@ pub enum Commands { #[command(subcommand)] command: DbCommands, }, + /// Manage archives + Archive { + #[command(subcommand)] + command: ArchiveCommands, + }, /// Manage WAF rules and blocklists Waf { #[command(subcommand)] @@ -491,6 +496,35 @@ pub enum SslCommands { }, } +#[derive(Subcommand)] +pub enum ArchiveCommands { + /// Import an archive to a site + Import { + /// Site ID + site_id: String, + /// Path to archive file (.tar.gz) + file: String, + /// Drop all existing tables before import + #[arg(long)] + drop_tables: bool, + /// Disable foreign key checks during import + #[arg(long)] + disable_foreign_keys: bool, + /// Search string for search-and-replace during import + #[arg(long)] + search_replace_from: Option, + /// Replace string for search-and-replace during import + #[arg(long)] + search_replace_to: Option, + /// Wait for import to complete + #[arg(long)] + wait: bool, + /// Seconds between status polls (default: 5) + #[arg(long, default_value = "5")] + poll_interval: u64, + }, +} + #[derive(Subcommand)] pub enum DbCommands { /// Manage archive import sessions diff --git a/src/commands/archive.rs b/src/commands/archive.rs new file mode 100644 index 0000000..9cde764 --- /dev/null +++ b/src/commands/archive.rs @@ -0,0 +1,206 @@ +use std::path::Path; +use std::thread; +use std::time::Duration; + +use serde::Serialize; +use serde_json::Value; + +use crate::api::{ApiClient, ApiError}; +use crate::output::{OutputFormat, format_option, print_json, print_key_value, print_message}; + +#[derive(Debug, Serialize)] +struct CreateImportSessionRequest { + #[serde(skip_serializing_if = "Option::is_none")] + filename: Option, + #[serde(skip_serializing_if = "Option::is_none")] + content_length: Option, + #[serde(skip_serializing_if = "Option::is_none")] + options: Option, +} + +#[derive(Debug, Serialize)] +struct ImportOptions { + #[serde(skip_serializing_if = "std::ops::Not::not")] + drop_tables: bool, + #[serde(skip_serializing_if = "std::ops::Not::not")] + disable_foreign_keys: bool, + #[serde(skip_serializing_if = "Option::is_none")] + search_replace: Option, +} + +#[derive(Debug, Serialize)] +struct SearchReplace { + from: String, + to: String, +} + +#[allow(clippy::too_many_arguments)] +pub fn import( + client: &ApiClient, + site_id: &str, + file: &str, + drop_tables: bool, + disable_foreign_keys: bool, + search_replace_from: Option, + search_replace_to: Option, + wait: bool, + poll_interval: u64, + format: OutputFormat, +) -> Result<(), ApiError> { + let path = Path::new(file); + + if !path.exists() { + return Err(ApiError::Other(format!("File not found: {}", file))); + } + + let metadata = + std::fs::metadata(path).map_err(|e| ApiError::Other(format!("Cannot read file: {}", e)))?; + + let content_length = metadata.len(); + let filename = path + .file_name() + .and_then(|n| n.to_str()) + .unwrap_or(file) + .to_string(); + + // Build import options + let search_replace = match (search_replace_from, search_replace_to) { + (Some(from), Some(to)) => Some(SearchReplace { from, to }), + _ => None, + }; + + let options = if drop_tables || disable_foreign_keys || search_replace.is_some() { + Some(ImportOptions { + drop_tables, + disable_foreign_keys, + search_replace, + }) + } else { + None + }; + + let body = CreateImportSessionRequest { + filename: Some(filename.clone()), + content_length: Some(content_length), + options, + }; + + // Step 1: Create import session + if format == OutputFormat::Table { + print_message("Creating import session..."); + } + + let response: Value = + client.post(&format!("/api/v1/vector/sites/{}/imports", site_id), &body)?; + + let data = &response["data"]; + let import_id = data["id"] + .as_str() + .ok_or_else(|| ApiError::Other("Missing import ID in response".to_string()))?; + let upload_url = data["upload_url"] + .as_str() + .ok_or_else(|| ApiError::Other("Missing upload URL in response".to_string()))?; + + if format == OutputFormat::Table { + print_message(&format!("Import ID: {}", import_id)); + } + + // Step 2: Upload file + let size_mb = content_length as f64 / 1_048_576.0; + if format == OutputFormat::Table { + print_message(&format!("Uploading {} ({:.1} MB)...", filename, size_mb)); + } + + let file_handle = std::fs::File::open(path) + .map_err(|e| ApiError::Other(format!("Cannot open file: {}", e)))?; + + client.put_file(upload_url, file_handle, content_length)?; + + if format == OutputFormat::Table { + print_message("Upload complete."); + } + + // Step 3: Trigger import + if format == OutputFormat::Table { + print_message("Starting import..."); + } + + let run_response: Value = client.post_empty(&format!( + "/api/v1/vector/sites/{}/imports/{}/run", + site_id, import_id + ))?; + + if format == OutputFormat::Table { + print_message("Import started."); + } + + // Step 4: Poll if --wait + if wait { + if format == OutputFormat::Table { + print_message("\nWaiting for import to complete..."); + } + + loop { + thread::sleep(Duration::from_secs(poll_interval)); + + let status_response: Value = client.get(&format!( + "/api/v1/vector/sites/{}/imports/{}", + site_id, import_id + ))?; + + let status_data = &status_response["data"]; + let status = status_data["status"].as_str().unwrap_or("unknown"); + + match status { + "completed" => { + if format == OutputFormat::Json { + print_json(&status_response); + } else { + let duration = format_option( + &status_data["duration_ms"].as_u64().map(|v| v.to_string()), + ); + print_message(&format!("Status: completed (duration: {}ms)", duration)); + } + return Ok(()); + } + "failed" => { + if format == OutputFormat::Json { + print_json(&status_response); + return Ok(()); + } + let error_msg = + format_option(&status_data["error_message"].as_str().map(String::from)); + return Err(ApiError::Other(format!("Import failed: {}", error_msg))); + } + _ => { + if format == OutputFormat::Table { + print_message(&format!("Status: {}", status)); + } + } + } + } + } + + // Final output + if format == OutputFormat::Json { + print_json(&run_response); + } else { + print_key_value(vec![ + ("Import ID", import_id.to_string()), + ( + "Status", + run_response["data"]["status"] + .as_str() + .unwrap_or("-") + .to_string(), + ), + ]); + print_message("\nCheck status with:"); + print_message(&format!( + " vector db import-session status {} {}", + site_id, import_id + )); + } + + Ok(()) +} diff --git a/src/commands/mod.rs b/src/commands/mod.rs index 1e06179..9838f02 100644 --- a/src/commands/mod.rs +++ b/src/commands/mod.rs @@ -1,4 +1,5 @@ pub mod account; +pub mod archive; pub mod auth; pub mod backup; pub mod db; diff --git a/src/main.rs b/src/main.rs index f09ce30..a549cd8 100644 --- a/src/main.rs +++ b/src/main.rs @@ -11,15 +11,15 @@ use std::process; use api::{ApiClient, ApiError, EXIT_SUCCESS}; use cli::{ AccountApiKeyCommands, AccountCommands, AccountSecretCommands, AccountSshKeyCommands, - AuthCommands, BackupCommands, BackupDownloadCommands, Cli, Commands, DbCommands, - DbExportCommands, DbImportSessionCommands, DeployCommands, EnvCommands, EnvDbCommands, - EnvDomainChangeCommands, EnvSecretCommands, EventCommands, McpCommands, RestoreCommands, - SiteCommands, SiteSshKeyCommands, SslCommands, WafAllowedReferrerCommands, + ArchiveCommands, AuthCommands, BackupCommands, BackupDownloadCommands, Cli, Commands, + DbCommands, DbExportCommands, DbImportSessionCommands, DeployCommands, EnvCommands, + EnvDbCommands, EnvDomainChangeCommands, EnvSecretCommands, EventCommands, McpCommands, + RestoreCommands, SiteCommands, SiteSshKeyCommands, SslCommands, WafAllowedReferrerCommands, WafBlockedIpCommands, WafBlockedReferrerCommands, WafCommands, WafRateLimitCommands, WebhookCommands, }; use commands::{ - account, auth, backup, db, deploy, env, event, mcp, restore, site, ssl, waf, webhook, + account, archive, auth, backup, db, deploy, env, event, mcp, restore, site, ssl, waf, webhook, }; use config::{Config, Credentials}; use output::{OutputFormat, print_error, print_json, print_message, print_table}; @@ -47,6 +47,7 @@ fn run(command: Commands, format: OutputFormat) -> Result<(), ApiError> { Commands::Deploy { command } => run_deploy(command, format), Commands::Ssl { command } => run_ssl(command, format), Commands::Db { command } => run_db(command, format), + Commands::Archive { command } => run_archive(command, format), Commands::Waf { command } => run_waf(command, format), Commands::Account { command } => run_account(command, format), Commands::Backup { command } => run_backup(command, format), @@ -349,6 +350,34 @@ fn run_db_export( } } +fn run_archive(command: ArchiveCommands, format: OutputFormat) -> Result<(), ApiError> { + let client = get_client()?; + + match command { + ArchiveCommands::Import { + site_id, + file, + drop_tables, + disable_foreign_keys, + search_replace_from, + search_replace_to, + wait, + poll_interval, + } => archive::import( + &client, + &site_id, + &file, + drop_tables, + disable_foreign_keys, + search_replace_from, + search_replace_to, + wait, + poll_interval, + format, + ), + } +} + fn run_waf(command: WafCommands, format: OutputFormat) -> Result<(), ApiError> { let client = get_client()?; diff --git a/tests/cli.rs b/tests/cli.rs index ad10b5b..bfdb81f 100644 --- a/tests/cli.rs +++ b/tests/cli.rs @@ -24,6 +24,7 @@ fn test_help() { assert!(stdout.contains("ssl")); assert!(stdout.contains("mcp")); assert!(stdout.contains("restore")); + assert!(stdout.contains("archive")); } #[test] @@ -370,6 +371,45 @@ fn test_restore_create_scope_files_requires_auth() { assert_eq!(output.status.code(), Some(2)); } +#[test] +fn test_archive_help() { + let output = vector_cmd() + .args(["archive", "--help"]) + .output() + .expect("Failed to run"); + assert!(output.status.success()); + let stdout = String::from_utf8_lossy(&output.stdout); + assert!(stdout.contains("import")); +} + +#[test] +fn test_archive_import_help() { + let output = vector_cmd() + .args(["archive", "import", "--help"]) + .output() + .expect("Failed to run"); + assert!(output.status.success()); + let stdout = String::from_utf8_lossy(&output.stdout); + assert!(stdout.contains("--drop-tables")); + assert!(stdout.contains("--disable-foreign-keys")); + assert!(stdout.contains("--search-replace-from")); + assert!(stdout.contains("--search-replace-to")); + assert!(stdout.contains("--wait")); + assert!(stdout.contains("--poll-interval")); +} + +#[test] +fn test_archive_import_requires_auth() { + let output = vector_cmd() + .args(["archive", "import", "test-site", "test-file.tar.gz"]) + .env("VECTOR_CONFIG_DIR", &nonexistent_config_dir()) + .env_remove("VECTOR_API_KEY") + .output() + .expect("Failed to run"); + assert!(!output.status.success()); + assert_eq!(output.status.code(), Some(2)); // EXIT_AUTH_ERROR +} + #[test] fn test_invalid_subcommand() { let output = vector_cmd()