diff --git a/Cargo.lock b/Cargo.lock index 24000d5..19a68b8 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -75,9 +75,9 @@ dependencies = [ [[package]] name = "anyhow" -version = "1.0.102" +version = "1.0.103" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" +checksum = "2a4385e2e34eb35d6b3efe798b9eb88096925d87726c0798709bf56d9ed84af3" [[package]] name = "arboard" @@ -320,7 +320,7 @@ dependencies = [ [[package]] name = "crowbar" -version = "0.4.1" +version = "0.5.0" dependencies = [ "anyhow", "arboard", diff --git a/Cargo.toml b/Cargo.toml index a3715aa..6f425c8 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "crowbar" -version = "0.4.1" +version = "0.5.0" edition = "2024" description = "A TUI web security testing proxy" license = "MIT" diff --git a/README.md b/README.md index acb868f..341a1fb 100644 --- a/README.md +++ b/README.md @@ -38,6 +38,8 @@ A terminal-based web security testing proxy built in Rust. Intercept, inspect, a - **Editor Modes** — Choose between a standard editor and Vim-style keybindings (with normal/insert modes, motions, and operators); toggle with `F2` or set via config/CLI - **Multi-Instance Support** — Run multiple Crowbar instances simultaneously; automatic port selection finds the next available port if the default is occupied - **Runtime Reconfiguration** — Change the proxy bind address and scope patterns without restarting +- **Resource Limits**: Configurable caps on buffered HTTP body size, WebSocket frame size, retained history entries, and concurrent connections, with a non-loopback bind refused unless explicitly opted into via `--allow-remote` +- **Hardened Local Storage**: CA key material, sessions, rules, and config are written to `~/.crowbar` with private directory/file permissions ## Supported Platforms @@ -79,8 +81,12 @@ Binaries are placed in `dist/` as `crowbar---`. # Start the proxy (default: 127.0.0.1:8080) crowbar -# Custom bind address -crowbar --bind 0.0.0.0:9090 +# Custom loopback bind address +crowbar --bind 127.0.0.1:9090 + +# Remote access is an explicit opt-in because the proxy has no authentication. +# Restrict access with a host firewall or trusted private network. +crowbar --bind 0.0.0.0:9090 --allow-remote # Start with intercept enabled crowbar --intercept @@ -102,8 +108,25 @@ crowbar --load ~/.crowbar/sessions/my-session.json # Use a custom config file crowbar --config /path/to/config.toml + +# Override resource limits (values are bytes or entry/connection counts) +crowbar --max-body-bytes 10485760 \ + --max-ws-frame-bytes 16777216 \ + --max-history-entries 10000 \ + --max-connections 128 ``` +Crowbar refuses a non-loopback bind unless `--allow-remote` is passed or +`allow_remote = true` is set in the configuration file. `--allow-remote` does +not enable authentication: only expose the listener on a network whose clients +and outbound access you trust. Scope patterns limit which traffic is captured; +they are not an access-control or destination policy. + +The default resource limits are 10 MiB per buffered HTTP body, 16 MiB per +WebSocket frame, 10,000 retained history entries, and 128 concurrent client +connections. Lower these values for an untrusted or memory-constrained +environment. The resource-limit settings are currently CLI-only. + ### CA Certificate Crowbar generates a CA certificate on first run and stores it at `~/.crowbar/ca.pem`. Install it in your browser or system trust store to avoid TLS warnings. @@ -146,6 +169,7 @@ Optional config at `~/.crowbar/config.toml`: ```toml bind = "127.0.0.1:8080" +allow_remote = false # required for a non-loopback bind; does not add authentication intercept = false scope = ["*.example.com"] editor_mode = "default" # or "vim" @@ -155,6 +179,11 @@ proto_include = ["./third_party"] # extra import/include paths CLI flags override config file values. +Crowbar creates `~/.crowbar` with private directory permissions and stores CA +key material, sessions, rules, configuration, and logs as private files. Keep +the generated CA key secret: anyone who obtains it can impersonate certificates +to systems that trust the Crowbar CA. + ## TUI Tabs The interface is organized into five tabs, switchable with `Tab`/`Shift+Tab` or number keys `1`–`5`. If the default port is in use, Crowbar automatically tries the next available port (up to 25 consecutive ports from the base). diff --git a/src/app/actions.rs b/src/app/actions.rs index 318f474..ad45008 100644 --- a/src/app/actions.rs +++ b/src/app/actions.rs @@ -109,7 +109,10 @@ impl App { fn macro_send_next(&mut self) { if self.macros.current_step >= self.macros.steps.len() { self.macros.running = false; - self.set_status(format!("Macro complete ({} steps)", self.macros.steps.len())); + self.set_status(format!( + "Macro complete ({} steps)", + self.macros.steps.len() + )); return; } @@ -119,13 +122,12 @@ impl App { let ui_tx = self.ui_tx.clone(); tokio::spawn(async move { - let ui_tx_inner = ui_tx.clone(); match repeater::send_raw_request(request).await { Ok(resp) => { - let _ = ui_tx_inner.send(ProxyToUi::MacroResponse(step_idx, resp)); + let _ = ui_tx.try_send(ProxyToUi::MacroResponse(step_idx, resp)); } Err(e) => { - let _ = ui_tx_inner.send(ProxyToUi::MacroError(step_idx, e)); + let _ = ui_tx.try_send(ProxyToUi::MacroError(step_idx, e)); } } }); diff --git a/src/app/dialogs.rs b/src/app/dialogs.rs index 4a7cc88..688b3f9 100644 --- a/src/app/dialogs.rs +++ b/src/app/dialogs.rs @@ -65,29 +65,23 @@ impl App { } fn save_session_to(&mut self, path: &std::path::Path) { - if let Some(parent) = path.parent() - && let Err(e) = std::fs::create_dir_all(parent) - { - self.set_status(format!("Save failed: {}", e)); - return; - } - let macro_requests: Vec<_> = self.macros.steps.iter().map(|s| s.request.clone()).collect(); - let session = crate::http::session::Session::new(self.store.entries().to_vec(), macro_requests); - match std::fs::File::create(path) { - Ok(file) => { - let writer = std::io::BufWriter::new(file); - match serde_json::to_writer_pretty(writer, &session) { - Ok(()) => { - self.set_status(format!("Saved to {}", path.display())); - } - Err(e) => { - self.set_status(format!("Save failed: {}", e)); - } - } - } - Err(e) => { - self.set_status(format!("Save failed: {}", e)); - } + let macro_requests: Vec<_> = self + .macros + .steps + .iter() + .map(|s| s.request.clone()) + .collect(); + let session = + crate::http::session::Session::new(self.store.entries().to_vec(), macro_requests); + let result = crate::fs_security::write_private_with(path, |file| { + use std::io::Write; + let mut writer = std::io::BufWriter::new(file); + serde_json::to_writer_pretty(&mut writer, &session).map_err(std::io::Error::other)?; + writer.flush() + }); + match result { + Ok(()) => self.set_status(format!("Saved to {}", path.display())), + Err(e) => self.set_status(format!("Save failed: {}", e)), } } @@ -100,7 +94,11 @@ impl App { let name = crate::rules::persist::auto_save_name(); match crate::rules::persist::save(&rules, &name) { Ok(path) => { - self.set_status(format!("Exported {} rules to {}", rules.len(), path.display())); + self.set_status(format!( + "Exported {} rules to {}", + rules.len(), + path.display() + )); } Err(e) => { self.set_status(format!("Export failed: {}", e)); @@ -113,11 +111,7 @@ impl App { Ok(session) => { self.store.load_entries(session.entries); if let Some(saved) = session.macros { - self.macros.steps = saved - .steps - .into_iter() - .map(SequenceStep::new) - .collect(); + self.macros.steps = saved.steps.into_iter().map(SequenceStep::new).collect(); self.macros.selected = 0; self.macros.running = false; self.macros.current_step = 0; diff --git a/src/app/event.rs b/src/app/event.rs index 12eb092..5add515 100644 --- a/src/app/event.rs +++ b/src/app/event.rs @@ -4,6 +4,23 @@ use crate::tui::tabs::Tab; use super::{App, EditorTarget}; +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum InputContext { + Help, + QuitConfirmation, + SaveDialog, + CertificateInfo, + InterceptEditor, + RepeaterEditor, + HistoryFilter, + ToolsEditor, + BindAddressEditor, + ScopeEditor, + RuleEditor, + RuleImport, + Tab, +} + impl App { pub fn handle_event(&mut self, event: Event) { if let Event::Key(key) = event { @@ -11,76 +28,66 @@ impl App { return; } - if self.show_help { - self.show_help = false; - return; - } - - if self.show_quit_confirm { - self.handle_quit_confirm_key(key); - return; - } - - if self.show_save_dialog { - self.handle_save_dialog_key(key); - return; - } - - if self.show_cert_info { - self.handle_cert_overlay_key(key); - return; - } - - if self.intercept_ui.editing { - self.handle_editor_key(key, EditorTarget::Intercept); - return; - } - - if self.repeater.editing { - self.handle_editor_key(key, EditorTarget::Repeater); - return; - } - - if self.history.filtering { - self.handle_filter_key(key); - return; - } - - if self.tools.editing { - self.handle_tools_editor_key(key); - return; - } - - if self.editing_bind_addr { - self.handle_bind_addr_editor_key(key); - return; - } - - if self.editing_scope { - self.handle_scope_editor_key(key); - return; - } - - if self.rules_ui.editing_field.is_some() { - self.handle_rules_editor_key(key); - return; - } - - if self.rules_ui.importing { - self.handle_rules_import_editor_key(key); - return; + match self.input_context() { + InputContext::Help => self.show_help = false, + InputContext::QuitConfirmation => self.handle_quit_confirm_key(key), + InputContext::SaveDialog => self.handle_save_dialog_key(key), + InputContext::CertificateInfo => self.handle_cert_overlay_key(key), + InputContext::InterceptEditor => { + self.handle_editor_key(key, EditorTarget::Intercept); + } + InputContext::RepeaterEditor => { + self.handle_editor_key(key, EditorTarget::Repeater); + } + InputContext::HistoryFilter => self.handle_filter_key(key), + InputContext::ToolsEditor => self.handle_tools_editor_key(key), + InputContext::BindAddressEditor => self.handle_bind_addr_editor_key(key), + InputContext::ScopeEditor => self.handle_scope_editor_key(key), + InputContext::RuleEditor => self.handle_rules_editor_key(key), + InputContext::RuleImport => self.handle_rules_import_editor_key(key), + InputContext::Tab => { + if self.handle_global_key(key) { + return; + } + match self.active_tab { + Tab::History => self.handle_history_key(key), + Tab::Proxy => self.handle_proxy_key(key), + Tab::Repeater => self.handle_repeater_key(key), + Tab::Rules => self.handle_rules_key(key), + Tab::Tools => self.handle_tools_key(key), + } + } } + } + } - if self.handle_global_key(key) { - return; - } - match self.active_tab { - Tab::History => self.handle_history_key(key), - Tab::Proxy => self.handle_proxy_key(key), - Tab::Repeater => self.handle_repeater_key(key), - Tab::Rules => self.handle_rules_key(key), - Tab::Tools => self.handle_tools_key(key), - } + fn input_context(&self) -> InputContext { + if self.show_help { + InputContext::Help + } else if self.show_quit_confirm { + InputContext::QuitConfirmation + } else if self.show_save_dialog { + InputContext::SaveDialog + } else if self.show_cert_info { + InputContext::CertificateInfo + } else if self.intercept_ui.editing { + InputContext::InterceptEditor + } else if self.repeater.editing { + InputContext::RepeaterEditor + } else if self.history.filtering { + InputContext::HistoryFilter + } else if self.tools.editing { + InputContext::ToolsEditor + } else if self.editing_bind_addr { + InputContext::BindAddressEditor + } else if self.editing_scope { + InputContext::ScopeEditor + } else if self.rules_ui.editing_field.is_some() { + InputContext::RuleEditor + } else if self.rules_ui.importing { + InputContext::RuleImport + } else { + InputContext::Tab } } diff --git a/src/app/handlers.rs b/src/app/handlers.rs index 77b0254..a75a74c 100644 --- a/src/app/handlers.rs +++ b/src/app/handlers.rs @@ -36,7 +36,8 @@ impl App { } KeyCode::Char('e') => { if let Some(req) = self.intercept_ui.queue.front() { - self.intercept_ui.editor = TextEditor::new(codec::request_to_lines(req), self.editor_mode); + self.intercept_ui.editor = + TextEditor::new(codec::request_to_lines(req), self.editor_mode); self.intercept_ui.editing = true; self.intercept_ui.scroll = 0; } @@ -76,11 +77,19 @@ impl App { .parse::() .or_else(|_| format!("127.0.0.1:{}", input).parse::()); match parsed { - Ok(addr) => { + Ok(addr) if addr.ip().is_loopback() || self.allow_remote => { self.pending_rebind = Some(addr); self.editing_bind_addr = false; self.bind_addr_buffer.clear(); } + Ok(addr) => { + self.set_status(format!( + "Refusing remote bind {} without --allow-remote", + addr + )); + self.editing_bind_addr = false; + self.bind_addr_buffer.clear(); + } Err(_) => { self.set_status(format!("Invalid address: {}", input)); self.editing_bind_addr = false; @@ -118,7 +127,11 @@ impl App { self.set_status(if count == 0 { "Scope cleared — capturing all traffic".to_string() } else { - format!("Scope updated ({} pattern{})", count, if count == 1 { "" } else { "s" }) + format!( + "Scope updated ({} pattern{})", + count, + if count == 1 { "" } else { "s" } + ) }); } KeyCode::Char(c) => { @@ -150,14 +163,12 @@ impl App { self.history.selected += 1; } } - KeyCode::Home | KeyCode::Char('g') - if !self.history.detail_open => { - self.history.selected = 0; - } - KeyCode::End | KeyCode::Char('G') - if !self.history.detail_open && entry_count > 0 => { - self.history.selected = entry_count - 1; - } + KeyCode::Home | KeyCode::Char('g') if !self.history.detail_open => { + self.history.selected = 0; + } + KeyCode::End | KeyCode::Char('G') if !self.history.detail_open && entry_count > 0 => { + self.history.selected = entry_count - 1; + } KeyCode::Enter => { if self.history.detail_open { self.history.detail_open = false; @@ -167,56 +178,56 @@ impl App { self.history.scroll = 0; } } - KeyCode::Esc - if self.history.detail_open => { - self.history.detail_open = false; - self.history.scroll = 0; - } - KeyCode::Char('r') - if entry_count > 0 => { - self.send_to_repeater(); - } - KeyCode::Char('/') - if !self.history.detail_open => { - self.history.filtering = true; - } + KeyCode::Esc if self.history.detail_open => { + self.history.detail_open = false; + self.history.scroll = 0; + } + KeyCode::Char('r') if entry_count > 0 => { + self.send_to_repeater(); + } + KeyCode::Char('/') if !self.history.detail_open => { + self.history.filtering = true; + } KeyCode::Char('c') => { if entry_count > 0 - && let Some(entry) = self.store.filtered_entry(self.history.selected) { - let curl = crate::http::export::to_curl(entry); - self.export_to_file("curl", "sh", &curl); - } + && let Some(entry) = self.store.filtered_entry(self.history.selected) + { + let curl = crate::http::export::to_curl(entry); + self.export_to_file("curl", "sh", &curl); + } } KeyCode::Char('w') => { if entry_count > 0 - && let Some(entry) = self.store.filtered_entry(self.history.selected) { - let raw = crate::http::export::to_raw(entry); - self.export_to_file("raw", "txt", &raw); - } - } - KeyCode::Char('h') - if !self.history.detail_open => { - let entries: Vec<_> = self.store.filtered_entries_iter() - .cloned() - .collect(); - let har = crate::http::export::to_har(&entries); - self.export_to_file("har", "har", &har); + && let Some(entry) = self.store.filtered_entry(self.history.selected) + { + let raw = crate::http::export::to_raw(entry); + self.export_to_file("raw", "txt", &raw); } + } + KeyCode::Char('h') if !self.history.detail_open => { + let entries: Vec<_> = self.store.filtered_entries_iter().cloned().collect(); + let har = crate::http::export::to_har(&entries); + self.export_to_file("har", "har", &har); + } KeyCode::Char('m') => { if entry_count > 0 - && let Some(entry) = self.store.filtered_entry(self.history.selected) { - self.macros.steps.push(SequenceStep::new(entry.request.clone())); - self.set_status(format!("Added to macro ({} steps)", self.macros.steps.len())); - } + && let Some(entry) = self.store.filtered_entry(self.history.selected) + { + self.macros + .steps + .push(SequenceStep::new(entry.request.clone())); + self.set_status(format!( + "Added to macro ({} steps)", + self.macros.steps.len() + )); + } } _ => {} } } pub(super) fn handle_rules_key(&mut self, key: KeyEvent) { - let rules = self.rules.read(); - let count = rules.len(); - drop(rules); + let count = self.rules.read().len(); match key.code { KeyCode::Char('a') => { @@ -225,64 +236,59 @@ impl App { rules.push(crate::rules::Rule::new(name)); self.rules_ui.selected = rules.len() - 1; } - KeyCode::Char('x') - if count > 0 => { - let mut rules = self.rules.write(); - rules.remove(self.rules_ui.selected); - if self.rules_ui.selected >= rules.len() && !rules.is_empty() { - self.rules_ui.selected = rules.len() - 1; - } - } - KeyCode::Enter - if count > 0 => { - let mut rules = self.rules.write(); - rules[self.rules_ui.selected].enabled = !rules[self.rules_ui.selected].enabled; - } - KeyCode::Char('t') - if count > 0 => { - let mut rules = self.rules.write(); - rules[self.rules_ui.selected].target = rules[self.rules_ui.selected].target.next(); - } - KeyCode::Char('s') - if count > 0 => { - let mut rules = self.rules.write(); - rules[self.rules_ui.selected].scope = rules[self.rules_ui.selected].scope.next(); - } - KeyCode::Char('R') - if count > 0 => { - let mut rules = self.rules.write(); - rules[self.rules_ui.selected].is_regex = !rules[self.rules_ui.selected].is_regex; - rules[self.rules_ui.selected].invalidate_regex(); + KeyCode::Char('x') if count > 0 => { + let mut rules = self.rules.write(); + rules.remove(self.rules_ui.selected); + if self.rules_ui.selected >= rules.len() && !rules.is_empty() { + self.rules_ui.selected = rules.len() - 1; } - KeyCode::Char('n') - if count > 0 => { + } + KeyCode::Enter if count > 0 => { + let mut rules = self.rules.write(); + rules[self.rules_ui.selected].enabled = !rules[self.rules_ui.selected].enabled; + } + KeyCode::Char('t') if count > 0 => { + let mut rules = self.rules.write(); + rules[self.rules_ui.selected].target = rules[self.rules_ui.selected].target.next(); + } + KeyCode::Char('s') if count > 0 => { + let mut rules = self.rules.write(); + rules[self.rules_ui.selected].scope = rules[self.rules_ui.selected].scope.next(); + } + KeyCode::Char('R') if count > 0 => { + let mut rules = self.rules.write(); + rules[self.rules_ui.selected].is_regex = !rules[self.rules_ui.selected].is_regex; + rules[self.rules_ui.selected].invalidate_regex(); + } + KeyCode::Char('n') if count > 0 => { + self.rules_ui.edit_buffer = { let rules = self.rules.read(); - self.rules_ui.edit_buffer = rules[self.rules_ui.selected].name.clone(); - drop(rules); - self.rules_ui.editing_field = Some(RuleField::Name); - } - KeyCode::Char('p') - if count > 0 => { + rules[self.rules_ui.selected].name.clone() + }; + self.rules_ui.editing_field = Some(RuleField::Name); + } + KeyCode::Char('p') if count > 0 => { + self.rules_ui.edit_buffer = { let rules = self.rules.read(); - self.rules_ui.edit_buffer = rules[self.rules_ui.selected].match_pattern.clone(); - drop(rules); - self.rules_ui.editing_field = Some(RuleField::Pattern); - } - KeyCode::Char('e') - if count > 0 => { + rules[self.rules_ui.selected].match_pattern.clone() + }; + self.rules_ui.editing_field = Some(RuleField::Pattern); + } + KeyCode::Char('e') if count > 0 => { + self.rules_ui.edit_buffer = { let rules = self.rules.read(); - self.rules_ui.edit_buffer = rules[self.rules_ui.selected].replacement.clone(); - drop(rules); - self.rules_ui.editing_field = Some(RuleField::Replacement); - } - KeyCode::Up | KeyCode::Char('k') - if self.rules_ui.selected > 0 => { - self.rules_ui.selected -= 1; - } + rules[self.rules_ui.selected].replacement.clone() + }; + self.rules_ui.editing_field = Some(RuleField::Replacement); + } + KeyCode::Up | KeyCode::Char('k') if self.rules_ui.selected > 0 => { + self.rules_ui.selected -= 1; + } KeyCode::Down | KeyCode::Char('j') - if count > 0 && self.rules_ui.selected < count - 1 => { - self.rules_ui.selected += 1; - } + if count > 0 && self.rules_ui.selected < count - 1 => + { + self.rules_ui.selected += 1; + } KeyCode::Char('E') => { self.export_rules(); } @@ -310,7 +316,11 @@ impl App { Ok(imported) => { let count = imported.len(); self.rules.write().extend(imported); - self.set_status(format!("Imported {} rules from {}", count, expanded.display())); + self.set_status(format!( + "Imported {} rules from {}", + count, + expanded.display() + )); } Err(e) => { self.set_status(format!("Import failed: {}", e)); @@ -341,14 +351,17 @@ impl App { if self.rules_ui.selected < rules.len() { match field { RuleField::Name => { - rules[self.rules_ui.selected].name = self.rules_ui.edit_buffer.clone(); + rules[self.rules_ui.selected].name = + self.rules_ui.edit_buffer.clone(); } RuleField::Pattern => { - rules[self.rules_ui.selected].match_pattern = self.rules_ui.edit_buffer.clone(); + rules[self.rules_ui.selected].match_pattern = + self.rules_ui.edit_buffer.clone(); rules[self.rules_ui.selected].invalidate_regex(); } RuleField::Replacement => { - rules[self.rules_ui.selected].replacement = self.rules_ui.edit_buffer.clone(); + rules[self.rules_ui.selected].replacement = + self.rules_ui.edit_buffer.clone(); } } } @@ -431,7 +444,9 @@ impl App { match self.tools.mode { ToolsMode::UrlEncode => super::encode::url_encode(&input), ToolsMode::UrlDecode => crate::http::url_decode(&input), - ToolsMode::Base64Encode => base64::engine::general_purpose::STANDARD.encode(input.as_bytes()), + ToolsMode::Base64Encode => { + base64::engine::general_purpose::STANDARD.encode(input.as_bytes()) + } ToolsMode::Base64Decode => { match base64::engine::general_purpose::STANDARD.decode(input.trim()) { Ok(bytes) => String::from_utf8_lossy(&bytes).into_owned(), @@ -439,12 +454,10 @@ impl App { } } ToolsMode::HexEncode => super::encode::hex_encode(input.as_bytes()), - ToolsMode::HexDecode => { - match super::encode::hex_decode(input.trim()) { - Ok(bytes) => String::from_utf8_lossy(&bytes).into_owned(), - Err(e) => format!("Error: {}", e), - } - } + ToolsMode::HexDecode => match super::encode::hex_decode(input.trim()) { + Ok(bytes) => String::from_utf8_lossy(&bytes).into_owned(), + Err(e) => format!("Error: {}", e), + }, } } @@ -480,33 +493,37 @@ impl App { self.repeater_send(); } } - (KeyModifiers::NONE, KeyCode::Char('e')) - if self.repeater.editor.has_content() => { - self.repeater.editing = true; - self.repeater.editor.cursor_line = 0; - self.repeater.editor.cursor_col = 0; - if self.editor_mode == EditorMode::Vim { - self.repeater.editor.vim_mode = crate::editor::VimMode::Insert; - } + (KeyModifiers::NONE, KeyCode::Char('e')) if self.repeater.editor.has_content() => { + self.repeater.editing = true; + self.repeater.editor.cursor_line = 0; + self.repeater.editor.cursor_col = 0; + if self.editor_mode == EditorMode::Vim { + self.repeater.editor.vim_mode = crate::editor::VimMode::Insert; } + } (KeyModifiers::NONE, KeyCode::Char('d')) - if !self.macros.show && self.repeater.original.is_some() => { - self.repeater.show_diff = !self.repeater.show_diff; - } + if !self.macros.show && self.repeater.original.is_some() => + { + self.repeater.show_diff = !self.repeater.show_diff; + } (KeyModifiers::NONE, KeyCode::Enter) - if self.macros.show && !self.macros.steps.is_empty() && !self.macros.running => { - self.load_macro_step(true); - } + if self.macros.show && !self.macros.steps.is_empty() && !self.macros.running => + { + self.load_macro_step(true); + } (KeyModifiers::NONE, KeyCode::Char('e')) - if self.macros.show && !self.macros.steps.is_empty() && !self.macros.running => { - self.load_macro_step(false); - } + if self.macros.show && !self.macros.steps.is_empty() && !self.macros.running => + { + self.load_macro_step(false); + } (KeyModifiers::SHIFT, KeyCode::Char('M')) => { self.macros.show = !self.macros.show; } (KeyModifiers::NONE, KeyCode::Char('j') | KeyCode::Down) => { if self.macros.show { - if !self.macros.steps.is_empty() && self.macros.selected < self.macros.steps.len() - 1 { + if !self.macros.steps.is_empty() + && self.macros.selected < self.macros.steps.len() - 1 + { self.macros.selected += 1; } } else { @@ -529,17 +546,21 @@ impl App { self.repeater.resp_scroll = self.repeater.resp_scroll.saturating_sub(1); } (KeyModifiers::NONE, KeyCode::Char('x')) - if self.macros.show && !self.macros.steps.is_empty() && !self.macros.running => { - self.macros.steps.remove(self.macros.selected); - if self.macros.selected >= self.macros.steps.len() && !self.macros.steps.is_empty() { - self.macros.selected = self.macros.steps.len() - 1; - } - } - (KeyModifiers::NONE, KeyCode::Char('X')) | (KeyModifiers::SHIFT, KeyCode::Char('X')) - if self.macros.show && !self.macros.running => { - self.macros.steps.clear(); - self.macros.selected = 0; + if self.macros.show && !self.macros.steps.is_empty() && !self.macros.running => + { + self.macros.steps.remove(self.macros.selected); + if self.macros.selected >= self.macros.steps.len() && !self.macros.steps.is_empty() + { + self.macros.selected = self.macros.steps.len() - 1; } + } + (KeyModifiers::NONE, KeyCode::Char('X')) + | (KeyModifiers::SHIFT, KeyCode::Char('X')) + if self.macros.show && !self.macros.running => + { + self.macros.steps.clear(); + self.macros.selected = 0; + } _ => {} } } @@ -554,34 +575,28 @@ impl App { match action { EditorAction::Consumed => {} - EditorAction::ExitEditor => { - match target { - EditorTarget::Intercept => { - self.intercept_ui.editing = false; - self.intercept_ui.editor = TextEditor::new(vec![], self.editor_mode); - } - EditorTarget::Repeater => { - self.repeater.editing = false; - } - } - } - EditorAction::Enter => { - match target { - EditorTarget::Intercept => self.forward_edited_intercept(), - EditorTarget::Repeater => { - self.repeater.editor.insert_newline(); - } - } - } - EditorAction::CtrlEnter => { - match target { - EditorTarget::Intercept => self.forward_edited_intercept(), - EditorTarget::Repeater => { - self.repeater.editing = false; - self.repeater_send(); - } + EditorAction::ExitEditor => match target { + EditorTarget::Intercept => { + self.intercept_ui.editing = false; + self.intercept_ui.editor = TextEditor::new(vec![], self.editor_mode); + } + EditorTarget::Repeater => { + self.repeater.editing = false; + } + }, + EditorAction::Enter => match target { + EditorTarget::Intercept => self.forward_edited_intercept(), + EditorTarget::Repeater => { + self.repeater.editor.insert_newline(); + } + }, + EditorAction::CtrlEnter => match target { + EditorTarget::Intercept => self.forward_edited_intercept(), + EditorTarget::Repeater => { + self.repeater.editing = false; + self.repeater_send(); } - } + }, EditorAction::Custom(_) => {} } } diff --git a/src/app/mod.rs b/src/app/mod.rs index 0a77433..2bd4ad2 100644 --- a/src/app/mod.rs +++ b/src/app/mod.rs @@ -78,7 +78,7 @@ pub struct App { pub bind_addr: SocketAddr, pub intercept_state: Arc, pub scope: Arc, - pub ui_tx: mpsc::UnboundedSender, + pub ui_tx: mpsc::Sender, pub history: HistoryState, pub intercept_ui: InterceptUiState, @@ -92,6 +92,8 @@ pub struct App { pub rules: SharedRules, pub editor_mode: EditorMode, pub proxy_running: bool, + pub allow_remote: bool, + pub proxy_limits: crate::proxy::ProxyLimits, // Bind address editing pub editing_bind_addr: bool, @@ -116,6 +118,18 @@ pub struct App { pub show_quit_confirm: bool, } +pub struct AppInit { + pub bind_addr: SocketAddr, + pub intercept_state: Arc, + pub scope: Arc, + pub rules: SharedRules, + pub ui_tx: mpsc::Sender, + pub editor_mode: EditorMode, + pub allow_remote: bool, + pub proxy_limits: crate::proxy::ProxyLimits, + pub max_history_entries: usize, +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum ToolsMode { UrlEncode, @@ -159,18 +173,22 @@ impl ToolsMode { } impl App { - pub fn new( - bind_addr: SocketAddr, - intercept_state: Arc, - scope: Arc, - rules: SharedRules, - ui_tx: mpsc::UnboundedSender, - editor_mode: EditorMode, - ) -> Self { + pub fn new(init: AppInit) -> Self { + let AppInit { + bind_addr, + intercept_state, + scope, + rules, + ui_tx, + editor_mode, + allow_remote, + proxy_limits, + max_history_entries, + } = init; Self { active_tab: Tab::History, should_quit: false, - store: InMemoryStore::new(), + store: InMemoryStore::new(max_history_entries), bind_addr, intercept_state, scope, @@ -224,6 +242,8 @@ impl App { }, editor_mode, proxy_running: true, + allow_remote, + proxy_limits, editing_bind_addr: false, bind_addr_buffer: String::new(), pending_rebind: None, @@ -242,7 +262,6 @@ impl App { pub fn intercept_enabled(&self) -> bool { self.intercept_state.is_enabled() } - } pub(super) enum EditorTarget { @@ -256,4 +275,3 @@ pub enum RuleField { Pattern, Replacement, } - diff --git a/src/app/render.rs b/src/app/render.rs index 1a87b1a..71c7633 100644 --- a/src/app/render.rs +++ b/src/app/render.rs @@ -1,17 +1,16 @@ +use ratatui::Frame; use ratatui::layout::{Constraint, Layout, Rect}; use ratatui::style::{Color, Modifier, Style}; use ratatui::text::{Line, Span, Text}; use ratatui::widgets::{Block, Borders, Clear, Paragraph, Tabs, Wrap}; -use ratatui::Frame; use crate::editor::EditorMode; -use crate::http::models::EntryState; +use crate::tui::tabs::Tab; use crate::tui::tabs::history_tab; use crate::tui::tabs::proxy_tab; use crate::tui::tabs::repeater_tab; use crate::tui::tabs::rules_tab; use crate::tui::tabs::tools_tab; -use crate::tui::tabs::Tab; use super::App; @@ -64,9 +63,7 @@ impl App { if *tab == Tab::Proxy && !self.intercept_ui.queue.is_empty() { spans.push(Span::styled( format!(" ({})", self.intercept_ui.queue.len()), - Style::default() - .fg(Color::Red) - .add_modifier(Modifier::BOLD), + Style::default().fg(Color::Red).add_modifier(Modifier::BOLD), )); } @@ -81,7 +78,11 @@ impl App { let version = env!("CARGO_PKG_VERSION"); let tabs = Tabs::new(titles) - .block(Block::default().borders(Borders::ALL).title(format!(" crowbar v{version} "))) + .block( + Block::default() + .borders(Borders::ALL) + .title(format!(" crowbar v{version} ")), + ) .select(selected) .highlight_style( Style::default() @@ -104,23 +105,18 @@ impl App { fn render_status_bar(&self, frame: &mut Frame, area: Rect) { if let Some((msg, when)) = &self.status_message - && when.elapsed() < std::time::Duration::from_secs(3) { - let line = Line::from(Span::styled( - format!(" {} ", msg), - Style::default().fg(Color::Yellow), - )); - frame.render_widget(Paragraph::new(line), area); - return; - } + && when.elapsed() < std::time::Duration::from_secs(3) + { + let line = Line::from(Span::styled( + format!(" {} ", msg), + Style::default().fg(Color::Yellow), + )); + frame.render_widget(Paragraph::new(line), area); + return; + } let total = self.store.len(); - let (complete, errors) = self.store.entries().iter().fold((0, 0), |(c, e), entry| { - match entry.state { - EntryState::Complete => (c + 1, e), - EntryState::Error => (c, e + 1), - _ => (c, e), - } - }); + let (complete, errors) = self.store.state_counts(); let intercept_span = if !self.proxy_running { Span::styled( @@ -150,7 +146,8 @@ impl App { }; let editor_mode_span = if self.editor_mode == EditorMode::Vim { - let active_editing = self.tools.editing || self.intercept_ui.editing || self.repeater.editing; + let active_editing = + self.tools.editing || self.intercept_ui.editing || self.repeater.editing; if active_editing { let editor = if self.tools.editing { &self.tools.editor @@ -213,9 +210,13 @@ impl App { let y = (area.height.saturating_sub(height)) / 2; let popup = Rect::new(x, y, width, height); - let key = Style::default().fg(Color::Yellow).add_modifier(Modifier::BOLD); + let key = Style::default() + .fg(Color::Yellow) + .add_modifier(Modifier::BOLD); let dim = Style::default().fg(Color::DarkGray); - let section = Style::default().fg(Color::Cyan).add_modifier(Modifier::BOLD); + let section = Style::default() + .fg(Color::Cyan) + .add_modifier(Modifier::BOLD); let path_style = Style::default().fg(Color::Green); let cert_path = dirs::home_dir() @@ -250,9 +251,7 @@ impl App { cert_display )), ]), - Line::from(Span::raw( - " && sudo update-ca-certificates", - )), + Line::from(Span::raw(" && sudo update-ca-certificates")), Line::raw(""), Line::from(vec![ Span::styled(" Firefox:", key), @@ -318,10 +317,7 @@ impl App { Span::styled(&self.save_buffer, Style::default().fg(Color::White)), Span::styled("\u{2588}", Style::default().fg(Color::Yellow)), ]), - Line::from(Span::styled( - " Enter to save, Esc to cancel", - dim, - )), + Line::from(Span::styled(" Enter to save, Esc to cancel", dim)), ]; frame.render_widget(Clear, popup); @@ -346,7 +342,9 @@ impl App { let y = (area.height.saturating_sub(height)) / 2; let popup = Rect::new(x, y, width, height); - let key = Style::default().fg(Color::Yellow).add_modifier(Modifier::BOLD); + let key = Style::default() + .fg(Color::Yellow) + .add_modifier(Modifier::BOLD); let dim = Style::default().fg(Color::DarkGray); let mut lines = vec![Line::raw("")]; @@ -399,70 +397,212 @@ impl App { let y = (area.height.saturating_sub(height)) / 2; let popup = Rect::new(x, y, width, height); - let key = Style::default().fg(Color::Yellow).add_modifier(Modifier::BOLD); + let key = Style::default() + .fg(Color::Yellow) + .add_modifier(Modifier::BOLD); let dim = Style::default().fg(Color::DarkGray); - let section = Style::default().fg(Color::Cyan).add_modifier(Modifier::BOLD); + let section = Style::default() + .fg(Color::Cyan) + .add_modifier(Modifier::BOLD); let lines = vec![ Line::from(Span::styled("Global", section)), - Line::from(vec![Span::styled(" Tab/Shift+Tab ", key), Span::raw("Cycle tabs")]), - Line::from(vec![Span::styled(" 1-5 ", key), Span::raw("Proxy/History/Repeater/Rules/Tools")]), - Line::from(vec![Span::styled(" ? ", key), Span::raw("Toggle this help")]), - Line::from(vec![Span::styled(" Ctrl+S ", key), Span::raw("Save session")]), - Line::from(vec![Span::styled(" q / Ctrl+C ", key), Span::raw("Quit")]), + Line::from(vec![ + Span::styled(" Tab/Shift+Tab ", key), + Span::raw("Cycle tabs"), + ]), + Line::from(vec![ + Span::styled(" 1-5 ", key), + Span::raw("Proxy/History/Repeater/Rules/Tools"), + ]), + Line::from(vec![ + Span::styled(" ? ", key), + Span::raw("Toggle this help"), + ]), + Line::from(vec![ + Span::styled(" Ctrl+S ", key), + Span::raw("Save session"), + ]), + Line::from(vec![ + Span::styled(" q / Ctrl+C ", key), + Span::raw("Quit"), + ]), Line::raw(""), Line::from(Span::styled("Proxy (Intercept)", section)), - Line::from(vec![Span::styled(" i ", key), Span::raw("Toggle intercept on/off")]), - Line::from(vec![Span::styled(" f ", key), Span::raw("Forward intercepted request")]), - Line::from(vec![Span::styled(" d ", key), Span::raw("Drop intercepted request")]), - Line::from(vec![Span::styled(" e ", key), Span::raw("Edit intercepted request")]), - Line::from(vec![Span::styled(" b ", key), Span::raw("Change bind address")]), - Line::from(vec![Span::styled(" s ", key), Span::raw("Edit scope patterns")]), - Line::from(vec![Span::styled(" C ", key), Span::raw("Export CA certificate")]), - Line::from(vec![Span::styled(" j/k ", key), Span::raw("Scroll request")]), + Line::from(vec![ + Span::styled(" i ", key), + Span::raw("Toggle intercept on/off"), + ]), + Line::from(vec![ + Span::styled(" f ", key), + Span::raw("Forward intercepted request"), + ]), + Line::from(vec![ + Span::styled(" d ", key), + Span::raw("Drop intercepted request"), + ]), + Line::from(vec![ + Span::styled(" e ", key), + Span::raw("Edit intercepted request"), + ]), + Line::from(vec![ + Span::styled(" b ", key), + Span::raw("Change bind address"), + ]), + Line::from(vec![ + Span::styled(" s ", key), + Span::raw("Edit scope patterns"), + ]), + Line::from(vec![ + Span::styled(" C ", key), + Span::raw("Export CA certificate"), + ]), + Line::from(vec![ + Span::styled(" j/k ", key), + Span::raw("Scroll request"), + ]), Line::raw(""), Line::from(Span::styled("History", section)), - Line::from(vec![Span::styled(" j/k ", key), Span::raw("Navigate / scroll")]), - Line::from(vec![Span::styled(" g/G ", key), Span::raw("Jump to first / last")]), - Line::from(vec![Span::styled(" / ", key), Span::raw("Filter by host, path, method, status")]), - Line::from(vec![Span::styled(" Enter ", key), Span::raw("Toggle detail view")]), - Line::from(vec![Span::styled(" r ", key), Span::raw("Send to repeater")]), - Line::from(vec![Span::styled(" m ", key), Span::raw("Add to macro sequence")]), - Line::from(vec![Span::styled(" c ", key), Span::raw("Export as curl")]), - Line::from(vec![Span::styled(" w ", key), Span::raw("Export as raw HTTP")]), - Line::from(vec![Span::styled(" h ", key), Span::raw("Export all as HAR")]), + Line::from(vec![ + Span::styled(" j/k ", key), + Span::raw("Navigate / scroll"), + ]), + Line::from(vec![ + Span::styled(" g/G ", key), + Span::raw("Jump to first / last"), + ]), + Line::from(vec![ + Span::styled(" / ", key), + Span::raw("Filter by host, path, method, status"), + ]), + Line::from(vec![ + Span::styled(" Enter ", key), + Span::raw("Toggle detail view"), + ]), + Line::from(vec![ + Span::styled(" r ", key), + Span::raw("Send to repeater"), + ]), + Line::from(vec![ + Span::styled(" m ", key), + Span::raw("Add to macro sequence"), + ]), + Line::from(vec![ + Span::styled(" c ", key), + Span::raw("Export as curl"), + ]), + Line::from(vec![ + Span::styled(" w ", key), + Span::raw("Export as raw HTTP"), + ]), + Line::from(vec![ + Span::styled(" h ", key), + Span::raw("Export all as HAR"), + ]), Line::raw(""), Line::from(Span::styled("Repeater", section)), - Line::from(vec![Span::styled(" Ctrl+Enter ", key), Span::raw("Send request")]), - Line::from(vec![Span::styled(" e ", key), Span::raw("Edit request")]), - Line::from(vec![Span::styled(" d ", key), Span::raw("Toggle diff view")]), - Line::from(vec![Span::styled(" M ", key), Span::raw("Toggle macro view")]), - Line::from(vec![Span::styled(" j/k ", key), Span::raw("Scroll request")]), - Line::from(vec![Span::styled(" J/K ", key), Span::raw("Scroll response")]), + Line::from(vec![ + Span::styled(" Ctrl+Enter ", key), + Span::raw("Send request"), + ]), + Line::from(vec![ + Span::styled(" e ", key), + Span::raw("Edit request"), + ]), + Line::from(vec![ + Span::styled(" d ", key), + Span::raw("Toggle diff view"), + ]), + Line::from(vec![ + Span::styled(" M ", key), + Span::raw("Toggle macro view"), + ]), + Line::from(vec![ + Span::styled(" j/k ", key), + Span::raw("Scroll request"), + ]), + Line::from(vec![ + Span::styled(" J/K ", key), + Span::raw("Scroll response"), + ]), Line::raw(""), Line::from(Span::styled("Rules", section)), - Line::from(vec![Span::styled(" a ", key), Span::raw("Add rule")]), - Line::from(vec![Span::styled(" x ", key), Span::raw("Delete rule")]), - Line::from(vec![Span::styled(" Enter ", key), Span::raw("Toggle enabled")]), - Line::from(vec![Span::styled(" n/p/e ", key), Span::raw("Edit name / pattern / replacement")]), - Line::from(vec![Span::styled(" t/s/R ", key), Span::raw("Cycle target / scope / regex")]), + Line::from(vec![ + Span::styled(" a ", key), + Span::raw("Add rule"), + ]), + Line::from(vec![ + Span::styled(" x ", key), + Span::raw("Delete rule"), + ]), + Line::from(vec![ + Span::styled(" Enter ", key), + Span::raw("Toggle enabled"), + ]), + Line::from(vec![ + Span::styled(" n/p/e ", key), + Span::raw("Edit name / pattern / replacement"), + ]), + Line::from(vec![ + Span::styled(" t/s/R ", key), + Span::raw("Cycle target / scope / regex"), + ]), Line::raw(""), Line::from(Span::styled("Tools", section)), - Line::from(vec![Span::styled(" h/l ", key), Span::raw("Switch tool")]), - Line::from(vec![Span::styled(" e ", key), Span::raw("Edit input")]), - Line::from(vec![Span::styled(" j/k ", key), Span::raw("Scroll output")]), - Line::from(vec![Span::styled(" Ctrl+U ", key), Span::raw("Clear input")]), - Line::from(vec![Span::styled(" Ctrl+Y ", key), Span::raw("Copy output to clipboard")]), + Line::from(vec![ + Span::styled(" h/l ", key), + Span::raw("Switch tool"), + ]), + Line::from(vec![ + Span::styled(" e ", key), + Span::raw("Edit input"), + ]), + Line::from(vec![ + Span::styled(" j/k ", key), + Span::raw("Scroll output"), + ]), + Line::from(vec![ + Span::styled(" Ctrl+U ", key), + Span::raw("Clear input"), + ]), + Line::from(vec![ + Span::styled(" Ctrl+Y ", key), + Span::raw("Copy output to clipboard"), + ]), Line::raw(""), Line::from(Span::styled("Editor", section)), - Line::from(vec![Span::styled(" F2 ", key), Span::raw("Toggle vim/default mode")]), - Line::from(vec![Span::styled(" Ctrl+Home/End ", key), Span::raw("Jump to start/end of input")]), - Line::from(vec![Span::styled(" Vim: Esc ", key), Span::raw("Normal mode / exit edit")]), - Line::from(vec![Span::styled(" Vim: i/a/o ", key), Span::raw("Enter insert mode")]), - Line::from(vec![Span::styled(" Vim: hjkl ", key), Span::raw("Movement (normal mode)")]), - Line::from(vec![Span::styled(" Vim: gg/G ", key), Span::raw("Jump to start/end of input")]), - Line::from(vec![Span::styled(" Vim: w/b ", key), Span::raw("Word forward/backward")]), - Line::from(vec![Span::styled(" Vim: dd/x/u ", key), Span::raw("Delete line/char, undo")]), + Line::from(vec![ + Span::styled(" F2 ", key), + Span::raw("Toggle vim/default mode"), + ]), + Line::from(vec![ + Span::styled(" Ctrl+Home/End ", key), + Span::raw("Jump to start/end of input"), + ]), + Line::from(vec![ + Span::styled(" Vim: Esc ", key), + Span::raw("Normal mode / exit edit"), + ]), + Line::from(vec![ + Span::styled(" Vim: i/a/o ", key), + Span::raw("Enter insert mode"), + ]), + Line::from(vec![ + Span::styled(" Vim: hjkl ", key), + Span::raw("Movement (normal mode)"), + ]), + Line::from(vec![ + Span::styled(" Vim: gg/G ", key), + Span::raw("Jump to start/end of input"), + ]), + Line::from(vec![ + Span::styled(" Vim: w/b ", key), + Span::raw("Word forward/backward"), + ]), + Line::from(vec![ + Span::styled(" Vim: dd/x/u ", key), + Span::raw("Delete line/char, undo"), + ]), Line::raw(""), Line::from(Span::styled("Press any key to close", dim)), ]; diff --git a/src/config.rs b/src/config.rs index 8fd0aab..933dbcf 100644 --- a/src/config.rs +++ b/src/config.rs @@ -13,13 +13,43 @@ struct Cli { #[arg(short, long, help = "Proxy bind address (default: 127.0.0.1:8080)")] pub bind: Option, + #[arg( + long, + help = "Allow binding the unauthenticated proxy to a non-loopback address" + )] + pub allow_remote: bool, + + #[arg(long, default_value_t = 10 * 1024 * 1024, help = "Maximum buffered HTTP body size in bytes")] + pub max_body_bytes: usize, + + #[arg(long, default_value_t = 16 * 1024 * 1024, help = "Maximum WebSocket frame size in bytes")] + pub max_ws_frame_bytes: usize, + + #[arg( + long, + default_value_t = 10_000, + help = "Maximum number of requests retained in history" + )] + pub max_history_entries: usize, + + #[arg( + long, + default_value_t = 128, + help = "Maximum concurrent client connections" + )] + pub max_connections: usize, + #[arg(long, help = "Start with intercept mode enabled")] pub intercept: bool, #[arg(short, long, help = "Path to config file")] pub config: Option, - #[arg(short, long, help = "Scope pattern (e.g. *.example.com). Repeat for multiple.")] + #[arg( + short, + long, + help = "Scope pattern (e.g. *.example.com). Repeat for multiple." + )] pub scope: Vec, #[arg(short, long, help = "Load a saved session file")] @@ -28,7 +58,10 @@ struct Cli { #[arg(long, help = "Editor mode: 'default' or 'vim'")] pub editor_mode: Option, - #[arg(long, help = "Directory of .proto files for gRPC decoding. Repeat for multiple.")] + #[arg( + long, + help = "Directory of .proto files for gRPC decoding. Repeat for multiple." + )] pub proto_dir: Vec, #[arg(long, help = "Extra .proto import/include path. Repeat for multiple.")] @@ -40,7 +73,10 @@ struct Cli { #[derive(Subcommand, Debug)] pub enum Command { - #[command(name = "ca-export", about = "Export the CA certificate for browser/OS trust store installation")] + #[command( + name = "ca-export", + about = "Export the CA certificate for browser/OS trust store installation" + )] CaExport { #[arg(help = "Output file path (omit to print to stdout)")] output: Option, @@ -49,7 +85,11 @@ pub enum Command { Import { #[arg(help = "Path to HAR file")] input: PathBuf, - #[arg(short, long, help = "Output session name (default: derived from input filename)")] + #[arg( + short, + long, + help = "Output session name (default: derived from input filename)" + )] name: Option, }, #[command(name = "rules-export", about = "Export rules to a JSON file")] @@ -67,6 +107,7 @@ pub enum Command { #[derive(Debug, Deserialize, Default)] struct FileConfig { bind: Option, + allow_remote: Option, intercept: Option, scope: Option>, editor_mode: Option, @@ -77,6 +118,9 @@ struct FileConfig { #[derive(Debug)] pub struct Config { pub bind: SocketAddr, + pub allow_remote: bool, + pub limits: crate::proxy::ProxyLimits, + pub max_history_entries: usize, pub intercept: bool, pub scope: Vec, pub load: Option, @@ -105,20 +149,12 @@ impl Config { fc } Err(e) => { - eprintln!( - "Warning: failed to parse {}: {}", - config_path.display(), - e - ); + eprintln!("Warning: failed to parse {}: {}", config_path.display(), e); FileConfig::default() } }, Err(e) => { - eprintln!( - "Warning: failed to read {}: {}", - config_path.display(), - e - ); + eprintln!("Warning: failed to read {}: {}", config_path.display(), e); FileConfig::default() } } @@ -128,9 +164,7 @@ impl Config { let default_bind: SocketAddr = "127.0.0.1:8080".parse().unwrap(); - let file_bind = file_config - .bind - .and_then(|s| s.parse::().ok()); + let file_bind = file_config.bind.and_then(|s| s.parse::().ok()); let scope = if !cli.scope.is_empty() { cli.scope @@ -138,11 +172,13 @@ impl Config { file_config.scope.unwrap_or_default() }; - let cli_editor_mode = cli.editor_mode.and_then(|s| match s.to_lowercase().as_str() { - "vim" => Some(EditorMode::Vim), - "default" => Some(EditorMode::Default), - _ => None, - }); + let cli_editor_mode = cli + .editor_mode + .and_then(|s| match s.to_lowercase().as_str() { + "vim" => Some(EditorMode::Vim), + "default" => Some(EditorMode::Default), + _ => None, + }); let editor_mode = cli_editor_mode .or(file_config.editor_mode) @@ -162,6 +198,13 @@ impl Config { Config { bind: cli.bind.or(file_bind).unwrap_or(default_bind), + allow_remote: cli.allow_remote || file_config.allow_remote.unwrap_or(false), + limits: crate::proxy::ProxyLimits { + max_body_bytes: cli.max_body_bytes.max(1), + max_ws_frame_bytes: cli.max_ws_frame_bytes.max(1), + max_connections: cli.max_connections.max(1), + }, + max_history_entries: cli.max_history_entries.max(1), intercept: cli.intercept || file_config.intercept.unwrap_or(false), scope, load: cli.load, @@ -172,3 +215,59 @@ impl Config { } } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn non_loopback_bind_does_not_implicitly_enable_remote_access() { + let cli = Cli::try_parse_from(["crowbar", "--bind", "0.0.0.0:9090"]).unwrap(); + + assert_eq!(cli.bind.unwrap(), "0.0.0.0:9090".parse().unwrap()); + assert!(!cli.allow_remote); + } + + #[test] + fn remote_access_requires_explicit_cli_opt_in() { + let cli = + Cli::try_parse_from(["crowbar", "--bind", "0.0.0.0:9090", "--allow-remote"]).unwrap(); + + assert!(cli.allow_remote); + } + + #[test] + fn resource_limit_flags_are_parsed_independently() { + let cli = Cli::try_parse_from([ + "crowbar", + "--max-body-bytes", + "1024", + "--max-ws-frame-bytes", + "2048", + "--max-history-entries", + "50", + "--max-connections", + "8", + ]) + .unwrap(); + + assert_eq!(cli.max_body_bytes, 1024); + assert_eq!(cli.max_ws_frame_bytes, 2048); + assert_eq!(cli.max_history_entries, 50); + assert_eq!(cli.max_connections, 8); + } + + #[test] + fn file_config_can_explicitly_allow_remote_access() { + let config: FileConfig = toml::from_str( + r#" + bind = "192.0.2.10:8080" + allow_remote = true + "#, + ) + .unwrap(); + + assert_eq!(config.bind.as_deref(), Some("192.0.2.10:8080")); + assert_eq!(config.allow_remote, Some(true)); + } +} diff --git a/src/editor.rs b/src/editor.rs index 0805411..f16406b 100644 --- a/src/editor.rs +++ b/src/editor.rs @@ -49,7 +49,21 @@ pub struct TextEditor { pub mode: EditorMode, pub vim_mode: VimMode, pending_key: Option, - undo_stack: std::collections::VecDeque<(Vec, usize, usize)>, + undo_stack: std::collections::VecDeque, +} + +enum UndoEntry { + Full { + lines: Vec, + cursor_line: usize, + cursor_col: usize, + }, + Line { + index: usize, + content: String, + cursor_line: usize, + cursor_col: usize, + }, } impl TextEditor { @@ -135,10 +149,7 @@ impl TextEditor { let mut result = Vec::with_capacity(self.lines.len()); for (i, line) in self.lines.iter().enumerate() { - let num_span = Span::styled( - format!("{:>width$} ", i + 1, width = gw), - num_style, - ); + let num_span = Span::styled(format!("{:>width$} ", i + 1, width = gw), num_style); if editing && i == self.cursor_line { let char_count = line.chars().count(); @@ -147,7 +158,11 @@ impl TextEditor { let col_byte = indices.nth(col).map(|(i, _)| i).unwrap_or(line.len()); let before = &line[..col_byte]; let (cursor_char, after) = if col < char_count { - let next_byte = line[col_byte..].chars().next().map(|c| col_byte + c.len_utf8()).unwrap_or(line.len()); + let next_byte = line[col_byte..] + .chars() + .next() + .map(|c| col_byte + c.len_utf8()) + .unwrap_or(line.len()); (&line[col_byte..next_byte], &line[next_byte..]) } else { (" ", "") @@ -162,10 +177,7 @@ impl TextEditor { Span::raw(after), ])); } else { - result.push(Line::from(vec![ - num_span, - Span::raw(line.as_str()), - ])); + result.push(Line::from(vec![num_span, Span::raw(line.as_str())])); } } result @@ -198,7 +210,7 @@ impl TextEditor { match (key.modifiers, key.code) { (_, KeyCode::Enter) => EditorAction::Enter, (_, KeyCode::Char(c)) if !key.modifiers.contains(KeyModifiers::CONTROL) => { - self.save_undo(); + self.save_line_undo(); let col = self.cursor_col.min(self.current_line_len()); let byte_col = self.char_to_byte(col); self.lines[self.cursor_line].insert(byte_col, c); @@ -354,7 +366,7 @@ impl TextEditor { // Deletion (_, KeyCode::Char('x')) => { - self.save_undo(); + self.save_line_undo(); let len = self.current_line_len(); if len > 0 && self.cursor_col < len { let byte_col = self.char_to_byte(self.cursor_col); @@ -369,7 +381,7 @@ impl TextEditor { EditorAction::Consumed } (KeyModifiers::SHIFT, KeyCode::Char('D')) => { - self.save_undo(); + self.save_line_undo(); let col = self.cursor_col.min(self.current_line_len()); let byte_col = self.char_to_byte(col); self.lines[self.cursor_line].truncate(byte_col); @@ -411,13 +423,14 @@ impl TextEditor { match (pending, key.code) { // dd - delete line ('d', KeyCode::Char('d')) => { - self.save_undo(); if self.lines.len() > 1 { + self.save_undo(); self.lines.remove(self.cursor_line); if self.cursor_line >= self.lines.len() { self.cursor_line = self.lines.len().saturating_sub(1); } } else { + self.save_line_undo(); self.lines[0].clear(); } self.clamp_cursor_normal(); @@ -425,14 +438,14 @@ impl TextEditor { } // dw - delete word ('d', KeyCode::Char('w')) => { - self.save_undo(); + self.save_line_undo(); self.delete_word_forward(); self.clamp_cursor_normal(); EditorAction::Consumed } // d$ - delete to end of line ('d', KeyCode::Char('$')) => { - self.save_undo(); + self.save_line_undo(); let col = self.cursor_col.min(self.current_line_len()); let byte_col = self.char_to_byte(col); self.lines[self.cursor_line].truncate(byte_col); @@ -454,7 +467,7 @@ impl TextEditor { fn handle_backspace(&mut self) { if self.cursor_col > 0 && self.cursor_line < self.lines.len() { - self.save_undo(); + self.save_line_undo(); self.cursor_col -= 1; let byte_col = self.char_to_byte(self.cursor_col); let next_byte = self.lines[self.cursor_line][byte_col..] @@ -478,7 +491,7 @@ impl TextEditor { } let line_len = self.current_line_len(); if self.cursor_col < line_len { - self.save_undo(); + self.save_line_undo(); let byte_col = self.char_to_byte(self.cursor_col); let next_byte = self.lines[self.cursor_line][byte_col..] .chars() @@ -551,27 +564,31 @@ impl TextEditor { // --- Word motion --- fn word_forward(&mut self) { - let line: Vec = self.lines[self.cursor_line].chars().collect(); + let line = &self.lines[self.cursor_line]; + let line_len = line.chars().count(); let mut col = self.cursor_col; + let mut chars = line.chars().skip(col).peekable(); - if col < line.len() { - let start_is_word = is_word_char(line[col]); - while col < line.len() - && is_word_char(line[col]) == start_is_word - && !line[col].is_whitespace() + if let Some(&first) = chars.peek() { + let start_is_word = is_word_char(first); + while chars + .next_if(|ch| is_word_char(*ch) == start_is_word && !ch.is_whitespace()) + .is_some() { col += 1; } - while col < line.len() && line[col].is_whitespace() { + while chars.next_if(|ch| ch.is_whitespace()).is_some() { col += 1; } } - if col >= line.len() && self.cursor_line + 1 < self.lines.len() { + if col >= line_len && self.cursor_line + 1 < self.lines.len() { self.cursor_line += 1; - let next_line: Vec = self.lines[self.cursor_line].chars().collect(); col = 0; - while col < next_line.len() && next_line[col].is_whitespace() { + for ch in self.lines[self.cursor_line].chars() { + if !ch.is_whitespace() { + break; + } col += 1; } } @@ -580,56 +597,38 @@ impl TextEditor { } fn word_backward(&mut self) { - let line: Vec = self.lines[self.cursor_line].chars().collect(); let mut col = self.cursor_col; if col == 0 { if self.cursor_line > 0 { self.cursor_line -= 1; - let prev_line: Vec = self.lines[self.cursor_line].chars().collect(); - col = prev_line.len(); - while col > 0 && prev_line[col - 1].is_whitespace() { - col -= 1; - } - if col > 0 { - let is_word = is_word_char(prev_line[col - 1]); - while col > 0 && is_word_char(prev_line[col - 1]) == is_word && !prev_line[col - 1].is_whitespace() { - col -= 1; - } - } + col = self.lines[self.cursor_line].chars().count(); + col = word_start_before(&self.lines[self.cursor_line], col); } self.cursor_col = col; return; } - while col > 0 && line[col - 1].is_whitespace() { - col -= 1; - } - if col > 0 { - let is_word = is_word_char(line[col - 1]); - while col > 0 && is_word_char(line[col - 1]) == is_word && !line[col - 1].is_whitespace() { - col -= 1; - } - } - - self.cursor_col = col; + self.cursor_col = word_start_before(&self.lines[self.cursor_line], col); } fn word_end(&mut self) { - let line: Vec = self.lines[self.cursor_line].chars().collect(); + let line = &self.lines[self.cursor_line]; + let line_len = line.chars().count(); let mut col = self.cursor_col; - if col + 1 < line.len() { + if col + 1 < line_len { col += 1; } - while col < line.len() && line[col].is_whitespace() { + let mut chars = line.chars().skip(col).peekable(); + while chars.next_if(|ch| ch.is_whitespace()).is_some() { col += 1; } - if col < line.len() { - let is_word = is_word_char(line[col]); - while col + 1 < line.len() - && is_word_char(line[col + 1]) == is_word - && !line[col + 1].is_whitespace() + if let Some(first) = chars.next() { + let is_word = is_word_char(first); + while chars + .next_if(|ch| is_word_char(*ch) == is_word && !ch.is_whitespace()) + .is_some() { col += 1; } @@ -639,19 +638,19 @@ impl TextEditor { } fn delete_word_forward(&mut self) { - let line: Vec = self.lines[self.cursor_line].chars().collect(); let start = self.cursor_col; let mut end = start; + let mut chars = self.lines[self.cursor_line].chars().skip(start).peekable(); - if end < line.len() { - let start_is_word = is_word_char(line[end]); - while end < line.len() - && is_word_char(line[end]) == start_is_word - && !line[end].is_whitespace() + if let Some(&first) = chars.peek() { + let start_is_word = is_word_char(first); + while chars + .next_if(|ch| is_word_char(*ch) == start_is_word && !ch.is_whitespace()) + .is_some() { end += 1; } - while end < line.len() && line[end].is_whitespace() { + while chars.next_if(|ch| ch.is_whitespace()).is_some() { end += 1; } } @@ -666,21 +665,51 @@ impl TextEditor { // --- Undo --- fn save_undo(&mut self) { + self.push_undo(UndoEntry::Full { + lines: self.lines.clone(), + cursor_line: self.cursor_line, + cursor_col: self.cursor_col, + }); + } + + fn save_line_undo(&mut self) { + self.push_undo(UndoEntry::Line { + index: self.cursor_line, + content: self.lines[self.cursor_line].clone(), + cursor_line: self.cursor_line, + cursor_col: self.cursor_col, + }); + } + + fn push_undo(&mut self, entry: UndoEntry) { if self.undo_stack.len() >= 100 { self.undo_stack.pop_front(); } - self.undo_stack.push_back(( - self.lines.clone(), - self.cursor_line, - self.cursor_col, - )); + self.undo_stack.push_back(entry); } fn undo(&mut self) { - if let Some((lines, line, col)) = self.undo_stack.pop_back() { - self.lines = lines; - self.cursor_line = line; - self.cursor_col = col; + match self.undo_stack.pop_back() { + Some(UndoEntry::Full { + lines, + cursor_line, + cursor_col, + }) => { + self.lines = lines; + self.cursor_line = cursor_line; + self.cursor_col = cursor_col; + } + Some(UndoEntry::Line { + index, + content, + cursor_line, + cursor_col, + }) => { + self.lines[index] = content; + self.cursor_line = cursor_line; + self.cursor_col = cursor_col; + } + None => {} } } } @@ -689,6 +718,29 @@ fn is_word_char(c: char) -> bool { c.is_alphanumeric() || c == '_' } +fn word_start_before(line: &str, col: usize) -> usize { + let byte_col = line + .char_indices() + .nth(col) + .map_or(line.len(), |(index, _)| index); + let mut chars = line[..byte_col].chars().rev().peekable(); + let mut col = col; + + while chars.next_if(|ch| ch.is_whitespace()).is_some() { + col -= 1; + } + if let Some(&last) = chars.peek() { + let is_word = is_word_char(last); + while chars + .next_if(|ch| is_word_char(*ch) == is_word && !ch.is_whitespace()) + .is_some() + { + col -= 1; + } + } + col +} + #[cfg(test)] mod tests { use super::*; @@ -708,12 +760,26 @@ mod tests { fn editor(text: &str) -> TextEditor { let lines: Vec = text.lines().map(String::from).collect(); - TextEditor::new(if lines.is_empty() { vec![String::new()] } else { lines }, EditorMode::Default) + TextEditor::new( + if lines.is_empty() { + vec![String::new()] + } else { + lines + }, + EditorMode::Default, + ) } fn vim_editor(text: &str) -> TextEditor { let lines: Vec = text.lines().map(String::from).collect(); - TextEditor::new(if lines.is_empty() { vec![String::new()] } else { lines }, EditorMode::Vim) + TextEditor::new( + if lines.is_empty() { + vec![String::new()] + } else { + lines + }, + EditorMode::Vim, + ) } // --- Constructor / basic state --- @@ -922,7 +988,10 @@ mod tests { #[test] fn ctrl_enter() { let mut ed = editor("test"); - assert_eq!(ed.handle_key(ctrl_key(KeyCode::Enter)), EditorAction::CtrlEnter); + assert_eq!( + ed.handle_key(ctrl_key(KeyCode::Enter)), + EditorAction::CtrlEnter + ); } // --- Gutter width --- @@ -959,6 +1028,26 @@ mod tests { assert_eq!(ed2.lines[0], "hello!"); } + #[test] + fn undo_line_edit_preserves_other_lines() { + let mut ed = editor("first\nsecond"); + ed.cursor_line = 1; + ed.cursor_col = 6; + ed.handle_key(key(KeyCode::Char('!'))); + ed.undo(); + assert_eq!(ed.lines, vec!["first", "second"]); + } + + #[test] + fn undo_restores_structural_edit() { + let mut ed = editor("hello world"); + ed.cursor_col = 5; + ed.insert_newline(); + ed.undo(); + assert_eq!(ed.lines, vec!["hello world"]); + assert_eq!((ed.cursor_line, ed.cursor_col), (0, 5)); + } + // --- Vim normal mode --- #[test] @@ -1153,6 +1242,17 @@ mod tests { assert_eq!(ed.cursor_col, 4); } + #[test] + fn vim_word_motions_preserve_character_columns() { + let mut ed = vim_editor("héllo 世界"); + ed.handle_key(key(KeyCode::Char('w'))); + assert_eq!(ed.cursor_col, 6); + ed.handle_key(key(KeyCode::Char('e'))); + assert_eq!(ed.cursor_col, 7); + ed.handle_key(key(KeyCode::Char('b'))); + assert_eq!(ed.cursor_col, 6); + } + #[test] fn vim_u_undo() { let mut ed = vim_editor("hello"); @@ -1165,7 +1265,10 @@ mod tests { #[test] fn vim_q_exits() { let mut ed = vim_editor("test"); - assert_eq!(ed.handle_key(key(KeyCode::Char('q'))), EditorAction::ExitEditor); + assert_eq!( + ed.handle_key(key(KeyCode::Char('q'))), + EditorAction::ExitEditor + ); } #[test] diff --git a/src/event.rs b/src/event.rs index 74f7113..dede21f 100644 --- a/src/event.rs +++ b/src/event.rs @@ -14,18 +14,18 @@ pub enum AppEvent { } pub struct EventLoop { - rx: mpsc::UnboundedReceiver, + rx: mpsc::Receiver, } impl EventLoop { - pub fn new(mut proxy_rx: mpsc::UnboundedReceiver) -> Self { - let (tx, rx) = mpsc::unbounded_channel(); + pub fn new(mut proxy_rx: mpsc::Receiver) -> Self { + let (tx, rx) = mpsc::channel(1_024); let input_tx = tx.clone(); tokio::spawn(async move { let mut stream = EventStream::new(); while let Some(Ok(event)) = stream.next().await { - if input_tx.send(AppEvent::Input(event)).is_err() { + if input_tx.send(AppEvent::Input(event)).await.is_err() { break; } } @@ -34,7 +34,7 @@ impl EventLoop { let proxy_tx = tx.clone(); tokio::spawn(async move { while let Some(msg) = proxy_rx.recv().await { - if proxy_tx.send(AppEvent::Proxy(msg)).is_err() { + if proxy_tx.send(AppEvent::Proxy(msg)).await.is_err() { break; } } @@ -44,7 +44,7 @@ impl EventLoop { let mut interval = tokio::time::interval(Duration::from_millis(250)); loop { interval.tick().await; - if tx.send(AppEvent::Tick).is_err() { + if tx.send(AppEvent::Tick).await.is_err() { break; } } diff --git a/src/fs_security.rs b/src/fs_security.rs new file mode 100644 index 0000000..473d324 --- /dev/null +++ b/src/fs_security.rs @@ -0,0 +1,221 @@ +use std::path::Path; + +/// Create a directory intended to hold captured traffic and private key material. +pub fn ensure_private_dir(path: &Path) -> std::io::Result<()> { + #[cfg(unix)] + { + use std::os::unix::fs::{DirBuilderExt, PermissionsExt}; + + let mut builder = std::fs::DirBuilder::new(); + builder.recursive(true).mode(0o700).create(path)?; + std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o700))?; + } + + #[cfg(not(unix))] + std::fs::create_dir_all(path)?; + + Ok(()) +} + +/// Atomically replace a sensitive file without following a destination symlink. +pub fn write_private(path: &Path, contents: impl AsRef<[u8]>) -> std::io::Result<()> { + write_private_with(path, |file| { + use std::io::Write; + file.write_all(contents.as_ref()) + }) +} + +pub fn write_private_with( + path: &Path, + write: impl FnOnce(&mut std::fs::File) -> std::io::Result<()>, +) -> std::io::Result<()> { + if let Some(parent) = path.parent() { + if parent.exists() { + if !parent.is_dir() { + return Err(std::io::Error::new( + std::io::ErrorKind::NotADirectory, + format!("{} is not a directory", parent.display()), + )); + } + } else { + ensure_private_dir(parent)?; + } + } + + let file_name = path + .file_name() + .and_then(|name| name.to_str()) + .unwrap_or("private"); + let mut attempt = 0u32; + let (temporary, mut file) = loop { + let candidate = path.with_file_name(format!( + ".{file_name}.{}.{}.tmp", + std::process::id(), + attempt + )); + let mut options = std::fs::OpenOptions::new(); + options.write(true).create_new(true); + #[cfg(unix)] + { + use std::os::unix::fs::OpenOptionsExt; + options.mode(0o600); + } + match options.open(&candidate) { + Ok(file) => break (candidate, file), + Err(error) if error.kind() == std::io::ErrorKind::AlreadyExists => { + attempt = attempt.checked_add(1).ok_or(error)?; + } + Err(error) => return Err(error), + } + }; + + let result = (|| { + write(&mut file)?; + file.sync_all()?; + drop(file); + std::fs::rename(&temporary, path)?; + set_private_file_permissions(path) + })(); + + if result.is_err() { + let _ = std::fs::remove_file(&temporary); + } + result +} + +pub fn harden_private_tree(path: &Path) -> std::io::Result<()> { + if !path.exists() { + return ensure_private_dir(path); + } + if std::fs::symlink_metadata(path)?.file_type().is_symlink() { + return Err(std::io::Error::new( + std::io::ErrorKind::PermissionDenied, + format!("refusing symlinked private directory {}", path.display()), + )); + } + harden_entry(path) +} + +fn harden_entry(path: &Path) -> std::io::Result<()> { + let metadata = std::fs::symlink_metadata(path)?; + if metadata.file_type().is_symlink() { + return Ok(()); + } + if metadata.is_dir() { + ensure_private_dir(path)?; + for entry in std::fs::read_dir(path)? { + harden_entry(&entry?.path())?; + } + } else if metadata.is_file() { + set_private_file_permissions(path)?; + } + Ok(()) +} + +fn set_private_file_permissions(path: &Path) -> std::io::Result<()> { + #[cfg(unix)] + { + use std::os::unix::fs::PermissionsExt; + std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?; + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn test_root(name: &str) -> std::path::PathBuf { + use std::sync::atomic::{AtomicU64, Ordering}; + + static NEXT_ID: AtomicU64 = AtomicU64::new(0); + std::env::temp_dir().join(format!( + "crowbar-{name}-{}-{}", + std::process::id(), + NEXT_ID.fetch_add(1, Ordering::Relaxed) + )) + } + + #[cfg(unix)] + #[test] + fn private_files_and_directories_have_restricted_modes() { + use std::os::unix::fs::PermissionsExt; + + let root = test_root("private-files"); + let file = root.join("nested/session.json"); + write_private(&file, b"secret").unwrap(); + assert_eq!( + std::fs::metadata(&root).unwrap().permissions().mode() & 0o777, + 0o700 + ); + assert_eq!( + std::fs::metadata(root.join("nested")) + .unwrap() + .permissions() + .mode() + & 0o777, + 0o700 + ); + assert_eq!( + std::fs::metadata(file).unwrap().permissions().mode() & 0o777, + 0o600 + ); + std::fs::remove_dir_all(root).unwrap(); + } + + #[cfg(unix)] + #[test] + fn harden_private_tree_repairs_existing_permissions() { + use std::os::unix::fs::PermissionsExt; + + let root = test_root("harden-tree"); + let nested = root.join("sessions"); + let file = nested.join("session.json"); + std::fs::create_dir_all(&nested).unwrap(); + std::fs::write(&file, b"captured credentials").unwrap(); + std::fs::set_permissions(&root, std::fs::Permissions::from_mode(0o755)).unwrap(); + std::fs::set_permissions(&nested, std::fs::Permissions::from_mode(0o777)).unwrap(); + std::fs::set_permissions(&file, std::fs::Permissions::from_mode(0o644)).unwrap(); + + harden_private_tree(&root).unwrap(); + + assert_eq!( + std::fs::metadata(&root).unwrap().permissions().mode() & 0o777, + 0o700 + ); + assert_eq!( + std::fs::metadata(&nested).unwrap().permissions().mode() & 0o777, + 0o700 + ); + assert_eq!( + std::fs::metadata(&file).unwrap().permissions().mode() & 0o777, + 0o600 + ); + std::fs::remove_dir_all(root).unwrap(); + } + + #[cfg(unix)] + #[test] + fn private_write_replaces_symlink_without_overwriting_target() { + use std::os::unix::fs::symlink; + + let root = test_root("symlink-write"); + let target = root.join("outside.txt"); + let private = root.join("ca.key"); + std::fs::create_dir_all(&root).unwrap(); + std::fs::write(&target, b"do not replace").unwrap(); + symlink(&target, &private).unwrap(); + + write_private(&private, b"private key").unwrap(); + + assert_eq!(std::fs::read(&target).unwrap(), b"do not replace"); + assert_eq!(std::fs::read(&private).unwrap(), b"private key"); + assert!( + !std::fs::symlink_metadata(&private) + .unwrap() + .file_type() + .is_symlink() + ); + std::fs::remove_dir_all(root).unwrap(); + } +} diff --git a/src/http/codec.rs b/src/http/codec.rs index 1a8a059..648f412 100644 --- a/src/http/codec.rs +++ b/src/http/codec.rs @@ -85,7 +85,8 @@ pub fn lines_to_request(lines: &[String], original: &RequestData) -> RequestData let body = if body_lines.is_empty() { Bytes::new() } else if original.is_grpc { - encode_grpc_body_lines(&body_lines, &uri).unwrap_or_else(|| Bytes::from(body_lines.join("\n"))) + encode_grpc_body_lines(&body_lines, &uri) + .unwrap_or_else(|| Bytes::from(body_lines.join("\n"))) } else { Bytes::from(body_lines.join("\n")) }; @@ -220,9 +221,7 @@ mod tests { let uri = "https://example.com/sample.UserService/GetUser"; let desc = proto_schema::request_type(uri).expect("schema resolves"); let frame = |n: &str| { - protobuf::encode_grpc_frame( - &proto_schema::encode_message_text(&desc, &[n]).unwrap(), - ) + protobuf::encode_grpc_frame(&proto_schema::encode_message_text(&desc, &[n]).unwrap()) }; let mut body = frame("1 user_id int: 1"); body.extend_from_slice(&frame("1 user_id int: 2")); @@ -230,9 +229,18 @@ mod tests { // Both frames decode, separated by the `---` marker. let lines = request_to_lines(&req); - assert!(lines.iter().any(|l| l == "---"), "expected separator: {lines:?}"); - assert!(lines.iter().any(|l| l.contains("user_id int: 1")), "{lines:?}"); - assert!(lines.iter().any(|l| l.contains("user_id int: 2")), "{lines:?}"); + assert!( + lines.iter().any(|l| l == "---"), + "expected separator: {lines:?}" + ); + assert!( + lines.iter().any(|l| l.contains("user_id int: 1")), + "{lines:?}" + ); + assert!( + lines.iter().any(|l| l.contains("user_id int: 2")), + "{lines:?}" + ); // And the two-frame body round-trips byte-for-byte. let rebuilt = lines_to_request(&lines, &req); diff --git a/src/http/export.rs b/src/http/export.rs index 6050280..f50dfbb 100644 --- a/src/http/export.rs +++ b/src/http/export.rs @@ -62,7 +62,11 @@ pub fn to_raw(entry: &HistoryEntry) -> String { if let Some(resp) = &entry.response { output.push_str("\r\n---\r\n\r\n"); - let _ = write!(output, "{} {} {}\r\n", resp.version, resp.status, resp.reason); + let _ = write!( + output, + "{} {} {}\r\n", + resp.version, resp.status, resp.reason + ); for (key, value) in &resp.headers { let _ = write!(output, "{}: {}\r\n", key, value); @@ -240,9 +244,9 @@ fn extract_path(uri: &str) -> &str { #[cfg(test)] mod tests { use super::*; + use crate::http::models::*; use bytes::Bytes; use std::time::{Duration, UNIX_EPOCH}; - use crate::http::models::*; fn make_request(method: &str, uri: &str, host: &str, is_tls: bool) -> RequestData { RequestData { @@ -414,7 +418,9 @@ mod tests { let entry = make_entry(req, Some(resp)); let har = to_har(&[entry]); let parsed: serde_json::Value = serde_json::from_str(&har).unwrap(); - let started = parsed["log"]["entries"][0]["startedDateTime"].as_str().unwrap(); + let started = parsed["log"]["entries"][0]["startedDateTime"] + .as_str() + .unwrap(); assert!(started.ends_with('Z')); assert!(started.contains('T')); assert_eq!(started.len(), 20); diff --git a/src/http/import.rs b/src/http/import.rs index 3adf11e..5e1da48 100644 --- a/src/http/import.rs +++ b/src/http/import.rs @@ -15,13 +15,14 @@ pub fn load_file(path: &Path) -> anyhow::Result { match ext.as_str() { "har" => load_har(path).map(|entries| super::session::Session::new(entries, Vec::new())), - "json" => { - match super::session::load(path) { - Ok(session) => Ok(session), - Err(session_err) => load_har(path).map(|entries| super::session::Session::new(entries, Vec::new())) - .map_err(|har_err| anyhow::anyhow!("Failed as session ({session_err}) and as HAR ({har_err})")), - } - } + "json" => match super::session::load(path) { + Ok(session) => Ok(session), + Err(session_err) => load_har(path) + .map(|entries| super::session::Session::new(entries, Vec::new())) + .map_err(|har_err| { + anyhow::anyhow!("Failed as session ({session_err}) and as HAR ({har_err})") + }), + }, _ => super::session::load(path), } } @@ -30,12 +31,7 @@ fn load_har(path: &Path) -> anyhow::Result> { let content = std::fs::read_to_string(path)?; let har: HarFile = serde_json::from_str(&content)?; - let entries = har - .log - .entries - .into_iter() - .map(convert_har_entry) - .collect(); + let entries = har.log.entries.into_iter().map(convert_har_entry).collect(); Ok(entries) } @@ -60,8 +56,7 @@ fn convert_har_entry(entry: HarEntry) -> HistoryEntry { .map(|pd| Bytes::from(pd.text.clone().unwrap_or_default())) .unwrap_or_default(); - let timestamp = parse_iso_timestamp(&entry.started_date_time) - .unwrap_or(SystemTime::now()); + let timestamp = parse_iso_timestamp(&entry.started_date_time).unwrap_or(SystemTime::now()); let request_data = RequestData { id: RequestId::next(), @@ -118,9 +113,7 @@ fn convert_har_entry(entry: HarEntry) -> HistoryEntry { fn extract_host(url: &str) -> String { if let Some(pos) = url.find("://") { let after_scheme = &url[pos + 3..]; - let end = after_scheme - .find('/') - .unwrap_or(after_scheme.len()); + let end = after_scheme.find('/').unwrap_or(after_scheme.len()); let host_port = &after_scheme[..end]; host_port.split(':').next().unwrap_or(host_port).to_string() } else { @@ -293,17 +286,22 @@ mod tests { method: "POST".into(), url: "https://api.example.com/v1/users".into(), http_version: "HTTP/2".into(), - headers: vec![ - HarHeader { name: "content-type".into(), value: "application/json".into() }, - ], - post_data: Some(HarPostData { text: Some("{\"name\":\"test\"}".into()) }), + headers: vec![HarHeader { + name: "content-type".into(), + value: "application/json".into(), + }], + post_data: Some(HarPostData { + text: Some("{\"name\":\"test\"}".into()), + }), }, response: HarResponse { status: 201, status_text: "Created".into(), http_version: "HTTP/2".into(), headers: vec![], - content: HarContent { text: Some("{\"id\":1}".into()) }, + content: HarContent { + text: Some("{\"id\":1}".into()), + }, }, }; diff --git a/src/http/mod.rs b/src/http/mod.rs index ca702e1..cd8bba5 100644 --- a/src/http/mod.rs +++ b/src/http/mod.rs @@ -50,7 +50,10 @@ pub(crate) fn date_to_days(year: u64, month: u64, day: u64) -> u64 { for y in 1970..year { days += if is_leap(y) { 366 } else { 365 }; } - for m in month_lengths(year).iter().take((month as usize).saturating_sub(1)) { + for m in month_lengths(year) + .iter() + .take((month as usize).saturating_sub(1)) + { days += m; } days + day.saturating_sub(1) @@ -75,15 +78,14 @@ pub(crate) fn url_decode(input: &str) -> String { let bytes = input.as_bytes(); let mut i = 0; while i < bytes.len() { - if bytes[i] == b'%' && i + 2 < bytes.len() - && let Ok(byte) = u8::from_str_radix( - &input[i + 1..i + 3], - 16, - ) { - result.push(byte); - i += 3; - continue; - } + if bytes[i] == b'%' + && i + 2 < bytes.len() + && let Ok(byte) = u8::from_str_radix(&input[i + 1..i + 3], 16) + { + result.push(byte); + i += 3; + continue; + } if bytes[i] == b'+' { result.push(b' '); } else { @@ -164,7 +166,11 @@ mod tests { fn date_days_roundtrip() { for (y, m, d) in [(1970, 1, 1), (2000, 6, 15), (2024, 12, 31), (1999, 2, 28)] { let days = date_to_days(y, m, d); - assert_eq!(days_to_date(days), (y, m, d), "roundtrip failed for {y}-{m}-{d}"); + assert_eq!( + days_to_date(days), + (y, m, d), + "roundtrip failed for {y}-{m}-{d}" + ); } } @@ -190,7 +196,10 @@ mod tests { #[test] fn extract_path_with_query() { - assert_eq!(extract_path("https://example.com/search?q=test"), "/search?q=test"); + assert_eq!( + extract_path("https://example.com/search?q=test"), + "/search?q=test" + ); } #[test] diff --git a/src/http/models.rs b/src/http/models.rs index 087e4ce..9f4b55e 100644 --- a/src/http/models.rs +++ b/src/http/models.rs @@ -230,19 +230,20 @@ pub struct HistoryEntry { impl HistoryEntry { pub fn matches_filter(&self, filter: &str) -> bool { let req = &self.request; - if req.method.to_lowercase().contains(filter) { + if contains_case_insensitive(&req.method, filter) { return true; } - if req.host.to_lowercase().contains(filter) { + if contains_case_insensitive(&req.host, filter) { return true; } - if req.uri.to_lowercase().contains(filter) { + if contains_case_insensitive(&req.uri, filter) { return true; } if let Some(resp) = &self.response - && resp.status.to_string().contains(filter) { - return true; - } + && resp.status.to_string().contains(filter) + { + return true; + } if self.request.is_grpc && "grpc".starts_with(filter) { return true; } @@ -250,6 +251,20 @@ impl HistoryEntry { } } +fn contains_case_insensitive(haystack: &str, needle: &str) -> bool { + if !haystack.is_ascii() || !needle.is_ascii() { + return haystack.to_lowercase().contains(&needle.to_lowercase()); + } + let needle = needle.as_bytes(); + if needle.is_empty() { + return true; + } + haystack + .as_bytes() + .windows(needle.len()) + .any(|candidate| candidate.eq_ignore_ascii_case(needle)) +} + #[cfg(test)] mod tests { use super::*; @@ -276,14 +291,26 @@ mod tests { #[test] fn http_version_from_hyper() { - assert_eq!(HttpVersion::from(hyper::Version::HTTP_10), HttpVersion::Http10); - assert_eq!(HttpVersion::from(hyper::Version::HTTP_11), HttpVersion::Http11); - assert_eq!(HttpVersion::from(hyper::Version::HTTP_2), HttpVersion::Http2); + assert_eq!( + HttpVersion::from(hyper::Version::HTTP_10), + HttpVersion::Http10 + ); + assert_eq!( + HttpVersion::from(hyper::Version::HTTP_11), + HttpVersion::Http11 + ); + assert_eq!( + HttpVersion::from(hyper::Version::HTTP_2), + HttpVersion::Http2 + ); } #[test] fn http_version_unknown_defaults_to_http11() { - assert_eq!(HttpVersion::from(hyper::Version::HTTP_3), HttpVersion::Http11); + assert_eq!( + HttpVersion::from(hyper::Version::HTTP_3), + HttpVersion::Http11 + ); } #[test] @@ -416,7 +443,13 @@ mod tests { assert_eq!(msg.text(), None); } - fn make_entry(method: &str, host: &str, uri: &str, is_grpc: bool, status: Option) -> HistoryEntry { + fn make_entry( + method: &str, + host: &str, + uri: &str, + is_grpc: bool, + status: Option, + ) -> HistoryEntry { let request = RequestData { id: RequestId(1), method: method.into(), @@ -439,7 +472,11 @@ mod tests { duration: Duration::from_millis(10), timing: None, }); - let state = if response.is_some() { EntryState::Complete } else { EntryState::Pending }; + let state = if response.is_some() { + EntryState::Complete + } else { + EntryState::Pending + }; HistoryEntry { request, response, @@ -464,6 +501,12 @@ mod tests { assert!(entry.matches_filter("example")); } + #[test] + fn matches_filter_unicode_case_insensitively() { + let entry = make_entry("GET", "BÜCHER.example", "/", false, Some(200)); + assert!(entry.matches_filter("bücher")); + } + #[test] fn matches_filter_by_uri() { let entry = make_entry("GET", "example.com", "/api/users", false, Some(200)); @@ -496,7 +539,11 @@ mod tests { map.insert("x-custom", "value".parse().unwrap()); let headers = extract_headers(&map); assert_eq!(headers.len(), 2); - assert!(headers.iter().any(|(k, v)| k == "content-type" && v == "text/html")); + assert!( + headers + .iter() + .any(|(k, v)| k == "content-type" && v == "text/html") + ); } #[test] @@ -523,7 +570,9 @@ mod tests { data: Bytes, } - let original = Wrapper { data: Bytes::from("hello world") }; + let original = Wrapper { + data: Bytes::from("hello world"), + }; let json = serde_json::to_string(&original).unwrap(); let restored: Wrapper = serde_json::from_str(&json).unwrap(); assert_eq!(original, restored); diff --git a/src/http/proto_schema.rs b/src/http/proto_schema.rs index 1bd6964..eb892eb 100644 --- a/src/http/proto_schema.rs +++ b/src/http/proto_schema.rs @@ -178,7 +178,13 @@ fn render_singular(fd: &FieldDescriptor, value: &Value, indent: usize, out: &mut .get_value(n) .map(|v| v.name().to_string()) .unwrap_or_else(|| n.to_string()); - out.push(format!("{}{} {} enum: {}", prefix, fd.number(), fd.name(), name)); + out.push(format!( + "{}{} {} enum: {}", + prefix, + fd.number(), + fd.name(), + name + )); } kind => { out.push(format!( @@ -207,10 +213,8 @@ fn render_map_entries( }; let prefix = pad(indent); // Sort entries by key text for deterministic output. - let mut pairs: Vec<(String, &Value)> = entries - .iter() - .map(|(k, v)| (mapkey_repr(k), v)) - .collect(); + let mut pairs: Vec<(String, &Value)> = + entries.iter().map(|(k, v)| (mapkey_repr(k), v)).collect(); pairs.sort_by(|a, b| a.0.cmp(&b.0)); for (key, value) in pairs { match val_fd.kind() { @@ -497,7 +501,9 @@ mod tests { } fn user_msg() -> MessageDescriptor { - pool().get_message_by_name("sample.User").expect("User type") + pool() + .get_message_by_name("sample.User") + .expect("User type") } /// Encode `msg`, decode to text, re-encode the text, and assert the two @@ -635,9 +641,7 @@ mod tests { "double_val" => msg.set_field(&f, Value::F64(2.25)), "uint32_val" => msg.set_field(&f, Value::U32(123)), "uint64_val" => msg.set_field(&f, Value::U64(456)), - "bytes_val" => { - msg.set_field(&f, Value::Bytes(vec![0xde, 0xad, 0xbe, 0xef].into())) - } + "bytes_val" => msg.set_field(&f, Value::Bytes(vec![0xde, 0xad, 0xbe, 0xef].into())), _ => {} } } @@ -648,9 +652,18 @@ mod tests { // Tags must reflect the schema's wire type, and signed/zigzag values // must survive — exactly what the heuristic decoder cannot do. assert!(text.contains("sint32_val sint: -7"), "got:\n{text}"); - assert!(text.contains("sint64_val sint: -9000000000"), "got:\n{text}"); - assert!(text.contains("fixed32_val fixed: 4294967295"), "got:\n{text}"); - assert!(text.contains("sfixed32_val sfixed: -2147483648"), "got:\n{text}"); + assert!( + text.contains("sint64_val sint: -9000000000"), + "got:\n{text}" + ); + assert!( + text.contains("fixed32_val fixed: 4294967295"), + "got:\n{text}" + ); + assert!( + text.contains("sfixed32_val sfixed: -2147483648"), + "got:\n{text}" + ); assert!(text.contains("float_val f32: 1.5"), "got:\n{text}"); assert!(text.contains("double_val f64: 2.25"), "got:\n{text}"); assert!(text.contains("uint32_val uint: 123"), "got:\n{text}"); @@ -670,7 +683,10 @@ mod tests { fn unknown_enum_number_renders_as_number() { let desc = user_msg(); let mut msg = DynamicMessage::new(desc.clone()); - msg.set_field(&desc.get_field_by_name("role").unwrap(), Value::EnumNumber(99)); + msg.set_field( + &desc.get_field_by_name("role").unwrap(), + Value::EnumNumber(99), + ); let text = decode_message_text(&desc, &msg.encode_to_vec(), 0) .unwrap() .join("\n"); @@ -700,9 +716,10 @@ mod tests { let reg = ProtoRegistry { pool: pool() }; assert!(reg.method_for_uri("/sample.UserService/GetUser").is_some()); // Full URL with a query string. - assert!(reg - .method_for_uri("https://h/sample.UserService/GetUser?x=1") - .is_some()); + assert!( + reg.method_for_uri("https://h/sample.UserService/GetUser?x=1") + .is_some() + ); // Unknown method on a known service. assert!(reg.method_for_uri("/sample.UserService/Nope").is_none()); // Missing the method segment entirely. @@ -714,7 +731,10 @@ mod tests { let dir: PathBuf = concat!(env!("CARGO_MANIFEST_DIR"), "/tests/fixtures/proto").into(); let mut files = Vec::new(); collect_protos(&dir, &mut files).unwrap(); - assert!(files.iter().any(|p| p.ends_with("sample.proto")), "files: {files:?}"); + assert!( + files.iter().any(|p| p.ends_with("sample.proto")), + "files: {files:?}" + ); assert!( files.iter().any(|p| p.ends_with("nested/extra.proto")), "recursion missed nested dir; files: {files:?}" @@ -729,7 +749,8 @@ mod tests { concat!(env!("CARGO_MANIFEST_DIR"), "/tests/fixtures/proto_common").into(); // With the extra include root, the cross-dir import resolves. - let pool = build_pool(&[importer.clone()], &[common]).expect("compile with include"); + let pool = + build_pool(std::slice::from_ref(&importer), &[common]).expect("compile with include"); assert!(pool.get_message_by_name("importer.Wrapper").is_some()); assert!(pool.get_message_by_name("common.Common").is_some()); diff --git a/src/http/protobuf.rs b/src/http/protobuf.rs index a652d8d..567befd 100644 --- a/src/http/protobuf.rs +++ b/src/http/protobuf.rs @@ -170,22 +170,28 @@ fn read_varint(data: &[u8]) -> Option<(u64, usize)> { fn interpret_length_delimited(data: &[u8], depth: usize) -> ProtoValue { if let Ok(s) = std::str::from_utf8(data) - && !s.is_empty() && is_likely_text(s) { - return ProtoValue::String(s.to_string()); - } + && !s.is_empty() + && is_likely_text(s) + { + return ProtoValue::String(s.to_string()); + } if data.len() >= 2 && let Some(fields) = decode_raw_inner(data, depth + 1) - && fields.iter().all(|f| f.number < 1000) { - return ProtoValue::Message(fields); - } + && fields.iter().all(|f| f.number < 1000) + { + return ProtoValue::Message(fields); + } ProtoValue::Bytes(data.to_vec()) } fn is_likely_text(s: &str) -> bool { let (total, printable) = s.chars().fold((0usize, 0usize), |(t, p), c| { - (t + 1, p + usize::from(c.is_ascii_graphic() || c.is_ascii_whitespace())) + ( + t + 1, + p + usize::from(c.is_ascii_graphic() || c.is_ascii_whitespace()), + ) }); total > 0 && (printable * 100 / total) >= 90 } @@ -217,17 +223,26 @@ pub fn format_proto_text(fields: &[ProtoField], indent: usize) -> Vec { lines.push(format!("{}{} int: {}", prefix, field.number, v)); } ProtoValue::Fixed64(v) => { - lines.push(format!("{}{} f64: {}", prefix, field.number, format_fixed64(*v))); + lines.push(format!( + "{}{} f64: {}", + prefix, + field.number, + format_fixed64(*v) + )); } ProtoValue::Fixed32(v) => { - lines.push(format!("{}{} f32: {}", prefix, field.number, format_fixed32(*v))); + lines.push(format!( + "{}{} f32: {}", + prefix, + field.number, + format_fixed32(*v) + )); } ProtoValue::String(s) => { lines.push(format!("{}{} str: {}", prefix, field.number, s)); } ProtoValue::Bytes(b) => { - let hex: String = - b.iter().map(|byte| format!("{:02x}", byte)).collect(); + let hex: String = b.iter().map(|byte| format!("{:02x}", byte)).collect(); lines.push(format!("{}{} hex: {}", prefix, field.number, hex)); } ProtoValue::Message(sub_fields) => { @@ -248,11 +263,7 @@ pub fn parse_proto_text(lines: &[&str]) -> Option> { Some(fields) } -fn parse_proto_at_depth( - lines: &[&str], - pos: &mut usize, - depth: usize, -) -> Option> { +fn parse_proto_at_depth(lines: &[&str], pos: &mut usize, depth: usize) -> Option> { let mut fields = Vec::new(); while *pos < lines.len() { @@ -482,8 +493,14 @@ mod tests { #[test] fn encode_decode_roundtrip() { let original = vec![ - ProtoField { number: 1, value: ProtoValue::Varint(150) }, - ProtoField { number: 2, value: ProtoValue::String("testing".into()) }, + ProtoField { + number: 1, + value: ProtoValue::Varint(150), + }, + ProtoField { + number: 2, + value: ProtoValue::String("testing".into()), + }, ]; let encoded = encode_raw(&original); let decoded = decode_raw(&encoded).unwrap(); @@ -495,14 +512,24 @@ mod tests { #[test] fn text_format_roundtrip() { let fields = vec![ - ProtoField { number: 1, value: ProtoValue::Varint(42) }, - ProtoField { number: 2, value: ProtoValue::String("hello".into()) }, - ProtoField { number: 3, value: ProtoValue::Bytes(vec![0xde, 0xad]) }, + ProtoField { + number: 1, + value: ProtoValue::Varint(42), + }, + ProtoField { + number: 2, + value: ProtoValue::String("hello".into()), + }, + ProtoField { + number: 3, + value: ProtoValue::Bytes(vec![0xde, 0xad]), + }, ProtoField { number: 4, - value: ProtoValue::Message(vec![ - ProtoField { number: 1, value: ProtoValue::Varint(99) }, - ]), + value: ProtoValue::Message(vec![ProtoField { + number: 1, + value: ProtoValue::Varint(99), + }]), }, ]; let text = format_proto_text(&fields, 0); @@ -515,9 +542,10 @@ mod tests { #[test] fn grpc_frame_encode_roundtrip() { - let payload = encode_raw(&[ - ProtoField { number: 1, value: ProtoValue::Varint(150) }, - ]); + let payload = encode_raw(&[ProtoField { + number: 1, + value: ProtoValue::Varint(150), + }]); let frame = encode_grpc_frame(&payload); let messages = decode_grpc_body(&frame); assert_eq!(messages.len(), 1); diff --git a/src/http/session.rs b/src/http/session.rs index d0a0a64..8e624ec 100644 --- a/src/http/session.rs +++ b/src/http/session.rs @@ -34,7 +34,9 @@ impl Session { let macros = if macro_requests.is_empty() { None } else { - Some(SavedMacro { steps: macro_requests }) + Some(SavedMacro { + steps: macro_requests, + }) }; Self { version: 2, @@ -49,16 +51,24 @@ pub fn sessions_dir() -> anyhow::Result { .ok_or_else(|| anyhow::anyhow!("Cannot find home directory"))? .join(".crowbar") .join("sessions"); - std::fs::create_dir_all(&dir)?; + crate::fs_security::ensure_private_dir(&dir)?; Ok(dir) } -pub fn save(entries: Vec, macro_requests: Vec, name: &str) -> anyhow::Result { +pub fn save( + entries: Vec, + macro_requests: Vec, + name: &str, +) -> anyhow::Result { let dir = sessions_dir()?; let path = dir.join(format!("{}.json", safe_name(name)?)); let session = Session::new(entries, macro_requests); - let json = serde_json::to_string_pretty(&session)?; - std::fs::write(&path, json)?; + crate::fs_security::write_private_with(&path, |file| { + use std::io::Write; + let mut writer = std::io::BufWriter::new(file); + serde_json::to_writer_pretty(&mut writer, &session).map_err(std::io::Error::other)?; + writer.flush() + })?; Ok(path) } @@ -89,7 +99,15 @@ mod tests { #[test] fn rejects_traversal_and_separators() { - for bad in ["..", ".", "../../etc/cron.d/evil", "a/b", "/etc/passwd", "", "foo/"] { + for bad in [ + "..", + ".", + "../../etc/cron.d/evil", + "a/b", + "/etc/passwd", + "", + "foo/", + ] { assert!(safe_name(bad).is_err(), "should reject {bad:?}"); } } diff --git a/src/http/store.rs b/src/http/store.rs index 9d8224f..e0073bb 100644 --- a/src/http/store.rs +++ b/src/http/store.rs @@ -2,7 +2,11 @@ use std::collections::HashMap; use crate::scanning::Finding; -use super::models::{EntryState, GrpcMessage, HistoryEntry, RequestData, RequestId, ResponseData, WsMessage}; +use super::models::{ + EntryState, GrpcMessage, HistoryEntry, RequestData, RequestId, ResponseData, WsMessage, +}; + +const MAX_STREAM_MESSAGES_PER_ENTRY: usize = 1_000; #[derive(Default)] struct FilterCache { @@ -11,16 +15,25 @@ struct FilterCache { entry_count: usize, } -#[derive(Default)] pub struct InMemoryStore { entries: Vec, index: HashMap, filter_cache: FilterCache, + max_entries: usize, + complete_count: usize, + error_count: usize, } impl InMemoryStore { - pub fn new() -> Self { - Self::default() + pub fn new(max_entries: usize) -> Self { + Self { + entries: Vec::new(), + index: HashMap::new(), + filter_cache: FilterCache::default(), + max_entries: max_entries.max(1), + complete_count: 0, + error_count: 0, + } } pub fn insert(&mut self, request: RequestData) { @@ -36,22 +49,35 @@ impl InMemoryStore { findings: Vec::new(), }); self.index.insert(id, idx); + self.evict_oldest_if_needed(); + } + + fn evict_oldest_if_needed(&mut self) { + if self.entries.len() <= self.max_entries { + return; + } + let drop_count = (self.max_entries / 10).max(1); + self.entries.drain(..drop_count); + self.index.clear(); + for (idx, entry) in self.entries.iter().enumerate() { + self.index.insert(entry.request.id, idx); + } + self.filter_cache.entry_count = usize::MAX; + self.recount_states(); } pub fn update_response(&mut self, id: RequestId, response: ResponseData) { - if let Some(&idx) = self.index.get(&id) - && let Some(entry) = self.entries.get_mut(idx) - { - entry.response = Some(response); - entry.state = EntryState::Complete; + if let Some(&idx) = self.index.get(&id) { + self.set_state(idx, EntryState::Complete); + if let Some(entry) = self.entries.get_mut(idx) { + entry.response = Some(response); + } } } pub fn mark_dropped(&mut self, id: RequestId) { - if let Some(&idx) = self.index.get(&id) - && let Some(entry) = self.entries.get_mut(idx) - { - entry.state = EntryState::Dropped; + if let Some(&idx) = self.index.get(&id) { + self.set_state(idx, EntryState::Dropped); } } @@ -60,6 +86,9 @@ impl InMemoryStore { && let Some(entry) = self.entries.get_mut(idx) { entry.ws_messages.push(msg); + if entry.ws_messages.len() > MAX_STREAM_MESSAGES_PER_ENTRY { + entry.ws_messages.drain(..100); + } } } @@ -68,15 +97,19 @@ impl InMemoryStore { && let Some(entry) = self.entries.get_mut(idx) { entry.grpc_messages.push(msg); + if entry.grpc_messages.len() > MAX_STREAM_MESSAGES_PER_ENTRY { + entry.grpc_messages.drain(..100); + } } } pub fn update_trailers(&mut self, id: RequestId, trailers: Vec<(String, String)>) { if let Some(&idx) = self.index.get(&id) && let Some(entry) = self.entries.get_mut(idx) - && let Some(resp) = &mut entry.response { - resp.trailers = trailers; - } + && let Some(resp) = &mut entry.response + { + resp.trailers = trailers; + } } pub fn set_findings(&mut self, id: RequestId, findings: Vec) { @@ -88,11 +121,11 @@ impl InMemoryStore { } pub fn mark_error(&mut self, id: RequestId, error: String) { - if let Some(&idx) = self.index.get(&id) - && let Some(entry) = self.entries.get_mut(idx) - { - entry.state = EntryState::Error; - entry.error_message = Some(error); + if let Some(&idx) = self.index.get(&id) { + self.set_state(idx, EntryState::Error); + if let Some(entry) = self.entries.get_mut(idx) { + entry.error_message = Some(error); + } } } @@ -104,6 +137,40 @@ impl InMemoryStore { &self.entries } + pub fn state_counts(&self) -> (usize, usize) { + (self.complete_count, self.error_count) + } + + fn set_state(&mut self, idx: usize, state: EntryState) { + let Some(entry) = self.entries.get_mut(idx) else { + return; + }; + match entry.state { + EntryState::Complete => self.complete_count = self.complete_count.saturating_sub(1), + EntryState::Error => self.error_count = self.error_count.saturating_sub(1), + _ => {} + } + entry.state = state; + match state { + EntryState::Complete => self.complete_count += 1, + EntryState::Error => self.error_count += 1, + _ => {} + } + } + + fn recount_states(&mut self) { + self.complete_count = self + .entries + .iter() + .filter(|entry| entry.state == EntryState::Complete) + .count(); + self.error_count = self + .entries + .iter() + .filter(|entry| entry.state == EntryState::Error) + .count(); + } + pub fn refresh_filter_cache(&mut self, filter: &str) { let cache = &self.filter_cache; if cache.filter == filter && cache.entry_count == self.entries.len() { @@ -160,6 +227,7 @@ impl InMemoryStore { self.entries.push(entry); self.index.insert(id, idx); } + self.recount_states(); } pub fn len(&self) -> usize { @@ -170,3 +238,40 @@ impl InMemoryStore { self.entries.is_empty() } } + +#[cfg(test)] +mod tests { + use super::*; + use bytes::Bytes; + use std::time::SystemTime; + + fn request(id: u64) -> RequestData { + RequestData { + id: RequestId(id), + method: "GET".into(), + uri: "/".into(), + host: "example.com".into(), + version: super::super::models::HttpVersion::Http11, + headers: Vec::new(), + body: Bytes::new(), + is_tls: false, + is_grpc: false, + timestamp: SystemTime::now(), + } + } + + #[test] + fn history_is_bounded_and_state_counts_track_transitions() { + let mut store = InMemoryStore::new(3); + for id in 1..=4 { + store.insert(request(id)); + } + assert_eq!(store.len(), 3); + assert!(store.get(RequestId(1)).is_none()); + + store.mark_error(RequestId(2), "failure".into()); + assert_eq!(store.state_counts(), (0, 1)); + store.mark_dropped(RequestId(2)); + assert_eq!(store.state_counts(), (0, 0)); + } +} diff --git a/src/main.rs b/src/main.rs index ac0005f..10b7af2 100644 --- a/src/main.rs +++ b/src/main.rs @@ -3,6 +3,7 @@ mod channel; mod config; mod editor; mod event; +mod fs_security; mod http; mod proxy; mod rules; @@ -18,14 +19,14 @@ use tokio_util::sync::CancellationToken; use tracing::info; use tracing_subscriber::EnvFilter; -use crate::app::App; +use crate::app::{App, AppInit}; use crate::channel::ProxyToUi; use crate::config::Config; use crate::event::{AppEvent, EventLoop}; +use crate::proxy::ProxyContext; use crate::proxy::intercept::InterceptState; use crate::proxy::scope::Scope; use crate::proxy::server::ProxyServer; -use crate::proxy::ProxyContext; use crate::rules::SharedRules; use crate::tls::ca::CertificateAuthority; use crate::tls::cert_cache::CertCache; @@ -35,12 +36,19 @@ use crate::tui::terminal; async fn main() -> anyhow::Result<()> { let config = Config::parse(); + let log_dir = dirs::home_dir().unwrap_or_default().join(".crowbar"); + crate::fs_security::harden_private_tree(&log_dir)?; + if let Some(cmd) = config.command { return handle_command(cmd); } - let log_dir = dirs::home_dir().unwrap_or_default().join(".crowbar"); - std::fs::create_dir_all(&log_dir)?; + if !config.bind.ip().is_loopback() && !config.allow_remote { + anyhow::bail!( + "refusing non-loopback bind {}; pass --allow-remote after securing network access", + config.bind + ); + } let file_appender = tracing_appender::rolling::never(&log_dir, "crowbar.log"); let (non_blocking, _guard) = tracing_appender::non_blocking(file_appender); @@ -54,6 +62,7 @@ async fn main() -> anyhow::Result<()> { .init(); info!("Starting crowbar proxy on {}", config.bind); + crate::fs_security::harden_private_tree(&log_dir)?; if !config.proto_dir.is_empty() { match crate::http::proto_schema::init(&config.proto_dir, &config.proto_include) { @@ -74,7 +83,7 @@ async fn main() -> anyhow::Result<()> { let mut cancel = CancellationToken::new(); - let (ui_tx, ui_rx) = mpsc::unbounded_channel::(); + let (ui_tx, ui_rx) = mpsc::channel::(1_024); let app_tx = ui_tx.clone(); let bound = { @@ -109,12 +118,9 @@ async fn main() -> anyhow::Result<()> { intercept: intercept.clone(), scope: scope.clone(), rules: rules.clone(), + limits: config.limits, }; - let server = ProxyServer::new( - bind_addr, - ctx, - cancel.clone(), - ); + let server = ProxyServer::new(bind_addr, ctx, cancel.clone()); tokio::spawn(async move { if let Err(e) = server.run(listener).await { tracing::error!("Proxy server error: {}", e); @@ -129,7 +135,17 @@ async fn main() -> anyhow::Result<()> { } let mut tui = terminal::init()?; - let mut app = App::new(bind_addr, intercept.clone(), scope.clone(), rules.clone(), app_tx, config.editor_mode); + let mut app = App::new(AppInit { + bind_addr, + intercept_state: intercept.clone(), + scope: scope.clone(), + rules: rules.clone(), + ui_tx: app_tx, + editor_mode: config.editor_mode, + allow_remote: config.allow_remote, + proxy_limits: config.limits, + max_history_entries: config.max_history_entries, + }); app.proxy_running = proxy_running; if !proxy_running { app.status_message = Some(( @@ -197,12 +213,9 @@ async fn run_app( intercept: app.intercept_state.clone(), scope: app.scope.clone(), rules: app.rules.clone(), + limits: app.proxy_limits, }; - let server = ProxyServer::new( - new_addr, - ctx, - cancel.clone(), - ); + let server = ProxyServer::new(new_addr, ctx, cancel.clone()); tokio::spawn(async move { if let Err(e) = server.run(listener).await { tracing::error!("Proxy server error: {}", e); @@ -254,8 +267,14 @@ fn handle_command(cmd: crate::config::Command) -> anyhow::Result<()> { eprintln!("CA certificate written to {}", path.display()); eprintln!(); eprintln!("To trust this certificate:"); - eprintln!(" macOS: security add-trusted-cert -d -r trustRoot -k ~/Library/Keychains/login.keychain-db {}", path.display()); - eprintln!(" Linux: sudo cp {} /usr/local/share/ca-certificates/crowbar.crt && sudo update-ca-certificates", path.display()); + eprintln!( + " macOS: security add-trusted-cert -d -r trustRoot -k ~/Library/Keychains/login.keychain-db {}", + path.display() + ); + eprintln!( + " Linux: sudo cp {} /usr/local/share/ca-certificates/crowbar.crt && sudo update-ca-certificates", + path.display() + ); eprintln!(" Firefox: Settings > Privacy & Security > Certificates > Import"); } None => { @@ -273,10 +292,7 @@ fn handle_command(cmd: crate::config::Command) -> anyhow::Result<()> { .to_string_lossy() .into_owned() }); - let macro_requests = session - .macros - .map(|m| m.steps) - .unwrap_or_default(); + let macro_requests = session.macros.map(|m| m.steps).unwrap_or_default(); let entry_count = session.entries.len(); let path = crate::http::session::save(session.entries, macro_requests, &session_name)?; eprintln!( diff --git a/src/proxy/handler.rs b/src/proxy/handler.rs index 13f0f00..eed812b 100644 --- a/src/proxy/handler.rs +++ b/src/proxy/handler.rs @@ -6,7 +6,6 @@ use http_body_util::{BodyExt, Full}; use hyper::body::Incoming; use hyper::{Method, Request, Response}; use hyper_util::rt::TokioIo; -use tokio::net::TcpStream; use tracing::warn; use crate::channel::ProxyToUi; @@ -25,6 +24,10 @@ impl ProxyHandler { Self { ctx } } + pub fn limits(&self) -> crate::proxy::ProxyLimits { + self.ctx.limits + } + pub async fn handle( &self, req: Request, @@ -78,7 +81,8 @@ impl ProxyHandler { req.headers() .get(hyper::header::HOST) .and_then(|v| v.to_str().ok()) - .map(|s| s.split(':').next().unwrap_or(s).to_string()) + .and_then(|value| value.parse::().ok()) + .map(|authority| authority.host().to_string()) }) .unwrap_or_default(); @@ -89,7 +93,13 @@ impl ProxyHandler { } let (parts, body) = req.into_parts(); - let body_bytes = body.collect().await?.to_bytes(); + let body_bytes = match http_body_util::Limited::new(body, self.ctx.limits.max_body_bytes) + .collect() + .await + { + Ok(body) => body.to_bytes(), + Err(_) => return Ok(crate::proxy::payload_too_large()), + }; let in_scope = self.ctx.scope.is_in_scope(&host); @@ -108,26 +118,31 @@ impl ProxyHandler { if in_scope { let _ = self - .ctx.ui_tx - .send(ProxyToUi::RequestCaptured(request_data.clone())); + .ctx + .ui_tx + .try_send(ProxyToUi::RequestCaptured(request_data.clone())); } if in_scope - && let Some(rx) = self.ctx.intercept.intercept_request(&request_data, &self.ctx.ui_tx) { - match rx.await { - Ok(InterceptDecision::Drop) => { - return Ok(Response::builder() - .status(503) - .body(Full::new(Bytes::from("Request dropped by interceptor"))) - .unwrap()); - } - Ok(InterceptDecision::ForwardEdited(edited)) => { - request_data = edited; - } - Ok(InterceptDecision::Forward) => {} - Err(_) => {} + && let Some(rx) = self + .ctx + .intercept + .intercept_request(&request_data, &self.ctx.ui_tx) + { + match rx.await { + Ok(InterceptDecision::Drop) => { + return Ok(Response::builder() + .status(503) + .body(Full::new(Bytes::from("Request dropped by interceptor"))) + .unwrap()); } + Ok(InterceptDecision::ForwardEdited(edited)) => { + request_data = edited; + } + Ok(InterceptDecision::Forward) => {} + Err(_) => {} } + } rules::apply_request_rules( &self.ctx.rules, @@ -136,18 +151,31 @@ impl ProxyHandler { &mut request_data.body, ); - let upstream_host = uri.host().unwrap_or(&host); - let upstream_port = uri.port_u16().unwrap_or(80); - let addr = format!("{}:{}", upstream_host, upstream_port); + let fallback_authority = request_data + .headers + .iter() + .find(|(name, _)| name.eq_ignore_ascii_case("host")) + .map_or(request_data.host.as_str(), |(_, value)| value.as_str()); + let target = match crate::proxy::resolve_upstream_target( + &request_data.uri, + fallback_authority, + 80, + ) { + Ok(target) => target, + Err(error) => return Ok(crate::proxy::bad_gateway(&error.to_string())), + }; - let upstream = match TcpStream::connect(&addr).await { + let upstream = match crate::proxy::connect_tcp(&target.host, target.port).await { Ok(s) => { timing.tcp_connected = Some(Instant::now()); s } Err(e) => { - warn!("Failed to connect to upstream {}: {}", addr, e); - let _ = self.ctx.ui_tx.send(ProxyToUi::RequestError( + warn!( + "Failed to connect to upstream {}:{}: {}", + target.host, target.port, e + ); + let _ = self.ctx.ui_tx.try_send(ProxyToUi::RequestError( request_id, format!("Connection failed: {}", e), )); @@ -159,16 +187,10 @@ impl ProxyHandler { }; let io = TokioIo::new(upstream); - let path_and_query = parts - .uri - .path_and_query() - .map(|pq| pq.as_str()) - .unwrap_or("/"); - Ok(crate::proxy::forward_h1( io, parts.version, - path_and_query, + &target.path_and_query, &request_data, timing, in_scope, @@ -193,9 +215,7 @@ impl ProxyHandler { let full_uri = format!( "ws://{}{}", host, - uri.path_and_query() - .map(|pq| pq.as_str()) - .unwrap_or("/") + uri.path_and_query().map(|pq| pq.as_str()).unwrap_or("/") ); let request_data = RequestData { @@ -215,16 +235,18 @@ impl ProxyHandler { let _ = self .ctx .ui_tx - .send(ProxyToUi::RequestCaptured(request_data)); + .try_send(ProxyToUi::RequestCaptured(request_data)); } let upstream_host = uri.host().unwrap_or(host); - let addr = format!("{}:{}", upstream_host, port); - let tcp = match TcpStream::connect(&addr).await { + let tcp = match crate::proxy::connect_tcp(upstream_host, port).await { Ok(s) => s, Err(e) => { - warn!("WebSocket: failed to connect to upstream {}: {}", addr, e); - let _ = self.ctx.ui_tx.send(ProxyToUi::RequestError( + warn!( + "WebSocket: failed to connect to upstream {}:{}: {}", + upstream_host, port, e + ); + let _ = self.ctx.ui_tx.try_send(ProxyToUi::RequestError( request_id, format!("Connection failed: {}", e), )); @@ -269,8 +291,8 @@ impl ProxyHandler { .uri(&path_and_query) .version(req.version()); - for (key, value) in req.headers() { - upstream_req = upstream_req.header(key, value); + for (key, value) in crate::proxy::end_to_end_headers(&headers, true) { + upstream_req = upstream_req.header(key.as_str(), value.as_str()); } let client_req_for_upgrade = req; @@ -283,10 +305,7 @@ impl ProxyHandler { Ok(resp) => resp, Err(e) => { warn!("WebSocket: upstream request failed: {}", e); - return Ok(crate::proxy::bad_gateway(&format!( - "Request failed: {}", - e - ))); + return Ok(crate::proxy::bad_gateway(&format!("Request failed: {}", e))); } }; @@ -294,7 +313,8 @@ impl ProxyHandler { if resp_status != 101 { let resp_headers = crate::http::models::extract_headers(upstream_resp.headers()); - let resp_body = upstream_resp + let resp_body = upstream_resp.into_body(); + let resp_body = http_body_util::Limited::new(resp_body, self.ctx.limits.max_body_bytes) .collect() .await .map(|b| b.to_bytes()) @@ -323,11 +343,12 @@ impl ProxyHandler { let _ = self .ctx .ui_tx - .send(ProxyToUi::ResponseReceived(request_id, response_data)); + .try_send(ProxyToUi::ResponseReceived(request_id, response_data)); } let ui_tx_clone = self.ctx.ui_tx.clone(); let host_owned = host.to_string(); + let max_ws_frame_bytes = self.ctx.limits.max_ws_frame_bytes; tokio::spawn(async move { let upstream_upgraded = match hyper::upgrade::on(upstream_resp).await { @@ -344,11 +365,7 @@ impl ProxyHandler { let client_upgraded = match hyper::upgrade::on(client_req_for_upgrade).await { Ok(u) => u, Err(e) => { - tracing::debug!( - "WebSocket client upgrade failed for {}: {}", - host_owned, - e - ); + tracing::debug!("WebSocket client upgrade failed for {}: {}", host_owned, e); return; } }; @@ -362,6 +379,7 @@ impl ProxyHandler { request_id, ui_tx_clone, in_scope, + max_ws_frame_bytes, ) .await; }); diff --git a/src/proxy/intercept.rs b/src/proxy/intercept.rs index e5f3193..6f82c18 100644 --- a/src/proxy/intercept.rs +++ b/src/proxy/intercept.rs @@ -42,7 +42,7 @@ impl InterceptState { pub fn intercept_request( &self, request: &RequestData, - ui_tx: &mpsc::UnboundedSender, + ui_tx: &mpsc::Sender, ) -> Option> { if !self.is_enabled() { return None; @@ -52,7 +52,7 @@ impl InterceptState { let id = request.id; if ui_tx - .send(ProxyToUi::InterceptedRequest(request.clone())) + .try_send(ProxyToUi::InterceptedRequest(request.clone())) .is_err() { return None; @@ -82,5 +82,4 @@ impl InterceptState { } } } - } diff --git a/src/proxy/mod.rs b/src/proxy/mod.rs index 3e63320..ecd7d7a 100644 --- a/src/proxy/mod.rs +++ b/src/proxy/mod.rs @@ -6,6 +6,7 @@ pub mod server; pub mod tunnel; pub mod ws_relay; +use std::collections::HashSet; use std::sync::Arc; use std::time::{Duration, Instant}; @@ -18,6 +19,8 @@ use hyper::client::conn::http1::Builder as ClientBuilder; use hyper_util::rt::TokioIo; use tracing::debug; +const CONNECT_TIMEOUT: Duration = Duration::from_secs(10); + use crate::channel::ProxyToUi; use crate::http::models::{RequestId, ResponseData, TimingData}; use crate::proxy::intercept::InterceptState; @@ -33,6 +36,78 @@ pub(crate) struct TimingContext { pub first_byte: Option, } +#[derive(Clone, Copy, Debug)] +pub struct ProxyLimits { + pub max_body_bytes: usize, + pub max_ws_frame_bytes: usize, + pub max_connections: usize, +} + +#[derive(Debug, PartialEq, Eq)] +pub(crate) struct UpstreamTarget { + pub host: String, + pub port: u16, + pub path_and_query: String, +} + +pub(crate) fn resolve_upstream_target( + uri: &str, + fallback_authority: &str, + default_port: u16, +) -> anyhow::Result { + let parsed: http::Uri = uri + .parse() + .map_err(|error| anyhow::anyhow!("invalid request URI {uri:?}: {error}"))?; + let fallback: http::uri::Authority = fallback_authority + .parse() + .map_err(|error| anyhow::anyhow!("invalid authority {fallback_authority:?}: {error}"))?; + let authority = parsed.authority().unwrap_or(&fallback); + Ok(UpstreamTarget { + host: authority + .host() + .trim_start_matches('[') + .trim_end_matches(']') + .to_string(), + port: authority.port_u16().unwrap_or(default_port), + path_and_query: parsed + .path_and_query() + .map_or_else(|| "/".to_string(), ToString::to_string), + }) +} + +pub(crate) async fn connect_tcp(host: &str, port: u16) -> anyhow::Result { + tokio::time::timeout( + CONNECT_TIMEOUT, + tokio::net::TcpStream::connect((host, port)), + ) + .await + .map_err(|_| anyhow::anyhow!("connection to {host}:{port} timed out"))? + .map_err(Into::into) +} + +pub(crate) fn format_authority(host: &str, port: u16, default_port: u16) -> String { + let host = if host.contains(':') && !host.starts_with('[') { + format!("[{host}]") + } else { + host.to_string() + }; + if port == default_port { + host + } else { + format!("{host}:{port}") + } +} + +impl Default for ProxyLimits { + fn default() -> Self { + Self { + max_body_bytes: 10 * 1024 * 1024, + max_ws_frame_bytes: 16 * 1024 * 1024, + max_connections: 128, + } + } +} + impl TimingContext { pub fn new() -> Self { Self { @@ -72,11 +147,12 @@ impl TimingContext { #[derive(Clone)] pub struct ProxyContext { - pub ui_tx: mpsc::UnboundedSender, + pub ui_tx: mpsc::Sender, pub cert_cache: Arc, pub intercept: Arc, pub scope: Arc, pub rules: SharedRules, + pub limits: ProxyLimits, } pub(crate) fn build_forwarding_request( @@ -84,14 +160,41 @@ pub(crate) fn build_forwarding_request( uri: &str, headers: &[(String, String)], body: Bytes, -) -> Request> { +) -> Result>, http::Error> { let mut builder = Request::builder().method(method).uri(uri); - for (key, value) in headers { + for (key, value) in end_to_end_headers(headers, false) { builder = builder.header(key.as_str(), value.as_str()); } - builder - .body(Full::new(body)) - .expect("building forwarding request") + builder.body(Full::new(body)) +} + +pub(crate) fn end_to_end_headers( + headers: &[(String, String)], + preserve_upgrade: bool, +) -> impl Iterator { + let connection_tokens: HashSet = headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case("connection")) + .flat_map(|(_, value)| value.split(',')) + .map(|token| token.trim().to_ascii_lowercase()) + .collect(); + + headers.iter().filter(move |(name, _)| { + let lower = name.to_ascii_lowercase(); + if matches!( + lower.as_str(), + "proxy-authorization" | "proxy-authenticate" | "proxy-connection" + ) { + return false; + } + if preserve_upgrade && matches!(lower.as_str(), "connection" | "upgrade") { + return true; + } + !matches!( + lower.as_str(), + "connection" | "keep-alive" | "te" | "trailer" | "transfer-encoding" | "upgrade" + ) && !connection_tokens.contains(&lower) + }) } pub(crate) fn build_client_response( @@ -100,10 +203,17 @@ pub(crate) fn build_client_response( body: Bytes, ) -> Response> { let mut builder = Response::builder().status(status); - for (key, value) in headers { + for (key, value) in end_to_end_headers(headers, false) { builder = builder.header(key.as_str(), value.as_str()); } - builder.body(Full::new(body)).unwrap() + builder.body(Full::new(body)).unwrap_or_else(|error| { + Response::builder() + .status(http::StatusCode::BAD_GATEWAY) + .body(Full::new(Bytes::from(format!( + "Invalid rewritten response: {error}" + )))) + .expect("static 502 response is valid") + }) } pub(crate) async fn process_h1_response( @@ -111,23 +221,31 @@ pub(crate) async fn process_h1_response( request_id: RequestId, timing: TimingContext, in_scope: bool, - shared_rules: &SharedRules, - ui_tx: &mpsc::UnboundedSender, + ctx: &ProxyContext, ) -> Response> { let resp_status = upstream_resp.status().as_u16(); let resp_version = upstream_resp.version().into(); let mut resp_headers = crate::http::models::extract_headers(upstream_resp.headers()); let body_start = Instant::now(); - let mut resp_body = upstream_resp - .collect() - .await - .map(|b| b.to_bytes()) - .unwrap_or_default(); + let mut resp_body = + match http_body_util::Limited::new(upstream_resp.into_body(), ctx.limits.max_body_bytes) + .collect() + .await + { + Ok(body) => body.to_bytes(), + Err(error) => { + let _ = ctx.ui_tx.try_send(ProxyToUi::RequestError( + request_id, + format!("Response body rejected: {error}"), + )); + return bad_gateway(&format!("Response body rejected: {error}")); + } + }; let content_transfer = body_start.elapsed(); let duration = timing.start.elapsed(); - rules::apply_response_rules(shared_rules, &mut resp_headers, &mut resp_body); + rules::apply_response_rules(&ctx.rules, &mut resp_headers, &mut resp_body); let timing_data = timing.finish(Some(content_transfer)); @@ -143,7 +261,9 @@ pub(crate) async fn process_h1_response( }; if in_scope { - let _ = ui_tx.send(ProxyToUi::ResponseReceived(request_id, response_data)); + let _ = ctx + .ui_tx + .try_send(ProxyToUi::ResponseReceived(request_id, response_data)); } build_client_response(resp_status, &resp_headers, resp_body) @@ -156,6 +276,86 @@ pub(crate) fn bad_gateway(msg: &str) -> Response> { .unwrap() } +pub(crate) fn payload_too_large() -> Response> { + Response::builder() + .status(http::StatusCode::PAYLOAD_TOO_LARGE) + .body(Full::new(Bytes::from_static( + b"Body exceeds configured limit", + ))) + .expect("static 413 response is valid") +} + +pub(crate) async fn forward_h1( + io: TokioIo, + version: hyper::Version, + path_and_query: &str, + request_data: &crate::http::models::RequestData, + mut timing: TimingContext, + in_scope: bool, + ctx: &ProxyContext, +) -> Response> +where + IO: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static, +{ + let request_id = request_data.id; + + let (mut sender, conn) = match ClientBuilder::new() + .preserve_header_case(true) + .title_case_headers(true) + .handshake(io) + .await + { + Ok(pair) => pair, + Err(e) => { + let _ = ctx.ui_tx.try_send(ProxyToUi::RequestError( + request_id, + format!("HTTP handshake failed: {}", e), + )); + return bad_gateway(&format!("HTTP handshake failed: {}", e)); + } + }; + + timing.http_handshake_done = Some(Instant::now()); + + tokio::spawn(async move { + if let Err(e) = conn.await { + debug!("Upstream connection ended: {}", e); + } + }); + + let mut upstream_req = match build_forwarding_request( + &request_data.method, + path_and_query, + &request_data.headers, + request_data.body.clone(), + ) { + Ok(request) => request, + Err(error) => { + let _ = ctx.ui_tx.try_send(ProxyToUi::RequestError( + request_id, + format!("Invalid edited request: {error}"), + )); + return bad_gateway(&format!("Invalid edited request: {error}")); + } + }; + *upstream_req.version_mut() = version; + + let upstream_resp = match sender.send_request(upstream_req).await { + Ok(resp) => resp, + Err(e) => { + let _ = ctx.ui_tx.try_send(ProxyToUi::RequestError( + request_id, + format!("Request failed: {}", e), + )); + return bad_gateway(&format!("Upstream request failed: {}", e)); + } + }; + + timing.first_byte = Some(Instant::now()); + + process_h1_response(upstream_resp, request_id, timing, in_scope, ctx).await +} + #[cfg(test)] mod tests { use super::*; @@ -239,11 +439,15 @@ mod tests { ("x-custom".into(), "value".into()), ], Bytes::from("{\"a\":1}"), - ); + ) + .unwrap(); assert_eq!(req.method(), "POST"); assert_eq!(req.uri(), "/api/v1"); assert_eq!(req.headers().len(), 2); - assert_eq!(req.headers().get("content-type").unwrap(), "application/json"); + assert_eq!( + req.headers().get("content-type").unwrap(), + "application/json" + ); } #[test] @@ -262,74 +466,57 @@ mod tests { let resp = bad_gateway("upstream failed"); assert_eq!(resp.status(), 502); } -} -pub(crate) async fn forward_h1( - io: TokioIo, - version: hyper::Version, - path_and_query: &str, - request_data: &crate::http::models::RequestData, - mut timing: TimingContext, - in_scope: bool, - ctx: &ProxyContext, -) -> Response> -where - IO: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static, -{ - let request_id = request_data.id; - - let (mut sender, conn) = match ClientBuilder::new() - .preserve_header_case(true) - .title_case_headers(true) - .handshake(io) - .await - { - Ok(pair) => pair, - Err(e) => { - let _ = ctx.ui_tx.send(ProxyToUi::RequestError( - request_id, - format!("HTTP handshake failed: {}", e), - )); - return bad_gateway(&format!("HTTP handshake failed: {}", e)); - } - }; - - timing.http_handshake_done = Some(Instant::now()); - - tokio::spawn(async move { - if let Err(e) = conn.await { - debug!("Upstream connection ended: {}", e); - } - }); - - let mut upstream_req = build_forwarding_request( - &request_data.method, - path_and_query, - &request_data.headers, - request_data.body.clone(), - ); - *upstream_req.version_mut() = version; - - let upstream_resp = match sender.send_request(upstream_req).await { - Ok(resp) => resp, - Err(e) => { - let _ = ctx.ui_tx.send(ProxyToUi::RequestError( - request_id, - format!("Request failed: {}", e), - )); - return bad_gateway(&format!("Upstream request failed: {}", e)); - } - }; + #[test] + fn forwarding_strips_proxy_and_connection_headers() { + let req = build_forwarding_request( + "GET", + "/", + &[ + ("proxy-authorization".into(), "secret".into()), + ("connection".into(), "x-private, keep-alive".into()), + ("x-private".into(), "remove-me".into()), + ("x-public".into(), "keep-me".into()), + ], + Bytes::new(), + ) + .unwrap(); + assert!(req.headers().get("proxy-authorization").is_none()); + assert!(req.headers().get("connection").is_none()); + assert!(req.headers().get("x-private").is_none()); + assert_eq!(req.headers().get("x-public").unwrap(), "keep-me"); + } - timing.first_byte = Some(Instant::now()); + #[test] + fn resolves_absolute_and_ipv6_targets() { + assert_eq!( + resolve_upstream_target("http://example.com:8080/a?q=1", "fallback", 80).unwrap(), + UpstreamTarget { + host: "example.com".into(), + port: 8080, + path_and_query: "/a?q=1".into(), + } + ); + assert_eq!( + resolve_upstream_target("/v1", "[::1]:9090", 80).unwrap(), + UpstreamTarget { + host: "[::1]".trim_matches(['[', ']']).to_string(), + port: 9090, + path_and_query: "/v1".into(), + } + ); + } - process_h1_response( - upstream_resp, - request_id, - timing, - in_scope, - &ctx.rules, - &ctx.ui_tx, - ) - .await + #[test] + fn invalid_edited_request_is_an_error() { + assert!( + build_forwarding_request( + "NOT A METHOD", + "/", + &[("bad header".into(), "value".into())], + Bytes::new(), + ) + .is_err() + ); + } } diff --git a/src/proxy/repeater.rs b/src/proxy/repeater.rs index 0b8e66c..960ee35 100644 --- a/src/proxy/repeater.rs +++ b/src/proxy/repeater.rs @@ -4,7 +4,6 @@ use bytes::Bytes; use http_body_util::BodyExt; use hyper::client::conn::http1::Builder as ClientBuilder; use hyper_util::rt::{TokioExecutor, TokioIo}; -use tokio::net::TcpStream; use tokio::sync::mpsc; use tokio_rustls::TlsConnector; use tracing::debug; @@ -13,16 +12,13 @@ use crate::channel::ProxyToUi; use crate::http::models::{HttpVersion, RequestData, ResponseData}; use crate::proxy::TimingContext; -pub async fn send_request( - request: RequestData, - ui_tx: mpsc::UnboundedSender, -) { +pub async fn send_request(request: RequestData, ui_tx: mpsc::Sender) { match send_raw_request(request).await { Ok(resp) => { - let _ = ui_tx.send(ProxyToUi::RepeaterResponse(resp)); + let _ = ui_tx.try_send(ProxyToUi::RepeaterResponse(resp)); } Err(e) => { - let _ = ui_tx.send(ProxyToUi::RepeaterError(e)); + let _ = ui_tx.try_send(ProxyToUi::RepeaterError(e)); } } } @@ -39,30 +35,39 @@ pub async fn send_raw_request(request: RequestData) -> Result anyhow::Result { - let host = strip_port(&request.host); - let port = extract_port(&request.uri).unwrap_or(80); + let target = crate::proxy::resolve_upstream_target(&request.uri, &request.host, 80)?; let mut timing = TimingContext::new(); - let tcp = TcpStream::connect(format!("{}:{}", host, port)).await?; + let tcp = crate::proxy::connect_tcp(&target.host, target.port).await?; timing.tcp_connected = Some(Instant::now()); - send_h1_via(TokioIo::new(tcp), request, timing).await + send_h1_via(TokioIo::new(tcp), request, &target.path_and_query, timing).await } async fn send_https(request: &RequestData) -> anyhow::Result { - let host = strip_port(&request.host); - let port = extract_port(&request.uri).unwrap_or(443); + let target = crate::proxy::resolve_upstream_target(&request.uri, &request.host, 443)?; let mut timing = TimingContext::new(); - let tcp = TcpStream::connect(format!("{}:{}", host, port)).await?; + let tcp = crate::proxy::connect_tcp(&target.host, target.port).await?; timing.tcp_connected = Some(Instant::now()); - let server_name = crate::tls::server_name_or_localhost(host); + let server_name = crate::tls::server_name_or_localhost(&target.host); let connector = TlsConnector::from(crate::tls::build_tls_client_config()); let tls_stream = connector.connect(server_name, tcp).await?; timing.tls_done = Some(Instant::now()); - send_h1_via(TokioIo::new(tls_stream), request, timing).await + send_h1_via( + TokioIo::new(tls_stream), + request, + &target.path_and_query, + timing, + ) + .await } -async fn send_h1_via(io: TokioIo, request: &RequestData, mut timing: TimingContext) -> anyhow::Result +async fn send_h1_via( + io: TokioIo, + request: &RequestData, + path: &str, + mut timing: TimingContext, +) -> anyhow::Result where IO: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static, { @@ -80,35 +85,34 @@ where } }); - let path = crate::http::extract_path(&request.uri); let req = crate::proxy::build_forwarding_request( - &request.method, path, &request.headers, request.body.clone(), - ); + &request.method, + path, + &request.headers, + request.body.clone(), + )?; let resp = sender.send_request(req).await?; timing.first_byte = Some(Instant::now()); parse_response(resp, timing).await } async fn send_h2(request: &RequestData) -> anyhow::Result { - let host = strip_port(&request.host); - let port = extract_port(&request.uri).unwrap_or(443); - let addr = format!("{}:{}", host, port); + let target = crate::proxy::resolve_upstream_target(&request.uri, &request.host, 443)?; let mut timing = TimingContext::new(); let client_config = crate::tls::build_tls_h2_client_config(); - let tcp = TcpStream::connect(&addr).await?; + let tcp = crate::proxy::connect_tcp(&target.host, target.port).await?; timing.tcp_connected = Some(Instant::now()); - let server_name = crate::tls::server_name_or_localhost(host); + let server_name = crate::tls::server_name_or_localhost(&target.host); let connector = TlsConnector::from(client_config); let tls_stream = connector.connect(server_name, tcp).await?; timing.tls_done = Some(Instant::now()); let io = TokioIo::new(tls_stream); - let (mut sender, conn) = - hyper::client::conn::http2::Builder::new(TokioExecutor::new()) - .handshake(io) - .await?; + let (mut sender, conn) = hyper::client::conn::http2::Builder::new(TokioExecutor::new()) + .handshake(io) + .await?; timing.http_handshake_done = Some(Instant::now()); @@ -118,16 +122,15 @@ async fn send_h2(request: &RequestData) -> anyhow::Result { } }); - let path = crate::http::extract_path(&request.uri); - let upstream_uri = if port == 443 { - format!("https://{}{}", host, path) - } else { - format!("https://{}:{}{}", host, port, path) - }; + let authority = crate::proxy::format_authority(&target.host, target.port, 443); + let upstream_uri = format!("https://{}{}", authority, target.path_and_query); let req = crate::proxy::build_forwarding_request( - &request.method, &upstream_uri, &request.headers, request.body.clone(), - ); + &request.method, + &upstream_uri, + &request.headers, + request.body.clone(), + )?; let resp = sender.send_request(req).await?; timing.first_byte = Some(Instant::now()); @@ -143,7 +146,14 @@ async fn parse_response( let headers = crate::http::models::extract_headers(resp.headers()); let body_start = Instant::now(); - let body = resp.collect().await?.to_bytes(); + let body = http_body_util::Limited::new( + resp.into_body(), + crate::proxy::ProxyLimits::default().max_body_bytes, + ) + .collect() + .await + .map_err(|error| anyhow::anyhow!(error.to_string()))? + .to_bytes(); let content_transfer = body_start.elapsed(); let duration = timing.start.elapsed(); let timing_data = timing.finish(Some(content_transfer)); @@ -168,7 +178,13 @@ async fn parse_h2_response( let headers = crate::http::models::extract_headers(resp.headers()); let body_start = Instant::now(); - let collected = resp.into_body().collect().await?; + let collected = http_body_util::Limited::new( + resp.into_body(), + crate::proxy::ProxyLimits::default().max_body_bytes, + ) + .collect() + .await + .map_err(|error| anyhow::anyhow!(error.to_string()))?; let content_transfer = body_start.elapsed(); let trailers_hm = collected.trailers().cloned(); let body = collected.to_bytes(); @@ -195,22 +211,3 @@ async fn parse_h2_response( timing: Some(timing_data), }) } - -fn strip_port(host: &str) -> &str { - if let Some(bracket) = host.find(']') { - // IPv6: [::1]:port - return &host[..bracket + 1]; - } - host.split(':').next().unwrap_or(host) -} - -fn extract_port(uri: &str) -> Option { - if let Some(pos) = uri.find("://") { - let after_scheme = &uri[pos + 3..]; - let authority = after_scheme.split('/').next().unwrap_or(after_scheme); - if let Some(colon) = authority.rfind(':') { - return authority[colon + 1..].parse().ok(); - } - } - None -} diff --git a/src/proxy/server.rs b/src/proxy/server.rs index 62b24ff..1bbd548 100644 --- a/src/proxy/server.rs +++ b/src/proxy/server.rs @@ -1,15 +1,17 @@ use std::net::SocketAddr; use std::sync::Arc; +use std::time::Duration; use hyper::server::conn::http1; use hyper::service::service_fn; -use hyper_util::rt::TokioIo; +use hyper_util::rt::{TokioIo, TokioTimer}; use tokio::net::TcpListener; +use tokio::sync::Semaphore; use tokio_util::sync::CancellationToken; use tracing::{error, info}; -use crate::proxy::handler::ProxyHandler; use crate::proxy::ProxyContext; +use crate::proxy::handler::ProxyHandler; pub struct ProxyServer { bind_addr: SocketAddr, @@ -18,11 +20,7 @@ pub struct ProxyServer { } impl ProxyServer { - pub fn new( - bind_addr: SocketAddr, - ctx: ProxyContext, - cancel: CancellationToken, - ) -> Self { + pub fn new(bind_addr: SocketAddr, ctx: ProxyContext, cancel: CancellationToken) -> Self { Self { bind_addr, ctx, @@ -34,14 +32,21 @@ impl ProxyServer { info!("Proxy listening on {}", self.bind_addr); let handler = Arc::new(ProxyHandler::new(self.ctx)); + let connections = Arc::new(Semaphore::new(handler.limits().max_connections)); loop { tokio::select! { result = listener.accept() => { let (stream, client_addr) = result?; let handler = handler.clone(); + let permit = connections + .clone() + .acquire_owned() + .await + .map_err(|_| anyhow::anyhow!("connection limiter closed"))?; tokio::spawn(async move { + let _permit = permit; let io = TokioIo::new(stream); let svc = service_fn(move |req| { let handler = handler.clone(); @@ -51,6 +56,8 @@ impl ProxyServer { if let Err(e) = http1::Builder::new() .preserve_header_case(true) .title_case_headers(true) + .timer(TokioTimer::new()) + .header_read_timeout(Duration::from_secs(15)) .serve_connection(io, svc) .with_upgrades() .await diff --git a/src/proxy/tunnel.rs b/src/proxy/tunnel.rs index 6b640ae..1194df0 100644 --- a/src/proxy/tunnel.rs +++ b/src/proxy/tunnel.rs @@ -6,19 +6,20 @@ use std::time::Instant; use bytes::{Bytes, BytesMut}; use http_body::Frame; use http_body_util::{BodyExt, Full}; +use hyper::Request; use hyper::body::Incoming; use hyper::client::conn::http1::Builder as ClientBuilder; use hyper::server::conn::http1::Builder as ServerBuilder; use hyper::service::service_fn; -use hyper::Request; use hyper_util::rt::{TokioExecutor, TokioIo}; -use tokio::net::TcpStream; use tokio::sync::mpsc; use tokio_rustls::{TlsAcceptor, TlsConnector}; use tracing::{debug, error, warn}; use crate::channel::ProxyToUi; -use crate::http::models::{GrpcDirection, GrpcMessage, HttpVersion, RequestData, RequestId, ResponseData}; +use crate::http::models::{ + GrpcDirection, GrpcMessage, HttpVersion, RequestData, RequestId, ResponseData, +}; use crate::proxy::intercept::InterceptDecision; use crate::proxy::{ProxyContext, TimingContext}; use crate::rules; @@ -140,21 +141,26 @@ async fn handle_websocket_upgrade( }; if in_scope { - let _ = ctx.ui_tx.send(ProxyToUi::RequestCaptured(request_data)); + let _ = ctx.ui_tx.try_send(ProxyToUi::RequestCaptured(request_data)); } let client_config = crate::tls::build_tls_client_config(); - let addr = format!("{}:{}", host, port); - let tcp = match TcpStream::connect(&addr).await { + let tcp = match crate::proxy::connect_tcp(host, port).await { Ok(s) => s, Err(e) => { - warn!("WebSocket: failed to connect to upstream {}: {}", addr, e); - let _ = ctx.ui_tx.send(ProxyToUi::RequestError( + warn!( + "WebSocket: failed to connect to upstream {}:{}: {}", + host, port, e + ); + let _ = ctx.ui_tx.try_send(ProxyToUi::RequestError( request_id, format!("Connection failed: {}", e), )); - return Ok(crate::proxy::bad_gateway(&format!("Connection failed: {}", e))); + return Ok(crate::proxy::bad_gateway(&format!( + "Connection failed: {}", + e + ))); } }; @@ -163,12 +169,18 @@ async fn handle_websocket_upgrade( let tls_stream = match connector.connect(server_name, tcp).await { Ok(s) => s, Err(e) => { - warn!("WebSocket: TLS handshake failed with {}: {}", addr, e); - let _ = ctx.ui_tx.send(ProxyToUi::RequestError( + warn!( + "WebSocket: TLS handshake failed with {}:{}: {}", + host, port, e + ); + let _ = ctx.ui_tx.try_send(ProxyToUi::RequestError( request_id, format!("TLS handshake failed: {}", e), )); - return Ok(crate::proxy::bad_gateway(&format!("TLS handshake failed: {}", e))); + return Ok(crate::proxy::bad_gateway(&format!( + "TLS handshake failed: {}", + e + ))); } }; @@ -182,7 +194,10 @@ async fn handle_websocket_upgrade( Ok(pair) => pair, Err(e) => { warn!("WebSocket: upstream handshake failed: {}", e); - return Ok(crate::proxy::bad_gateway(&format!("HTTP handshake failed: {}", e))); + return Ok(crate::proxy::bad_gateway(&format!( + "HTTP handshake failed: {}", + e + ))); } }; @@ -203,8 +218,8 @@ async fn handle_websocket_upgrade( .uri(&path_and_query) .version(req.version()); - for (key, value) in req.headers() { - upstream_req = upstream_req.header(key, value); + for (key, value) in crate::proxy::end_to_end_headers(&headers, true) { + upstream_req = upstream_req.header(key.as_str(), value.as_str()); } let client_req_for_upgrade = req; @@ -225,7 +240,8 @@ async fn handle_websocket_upgrade( if resp_status != 101 { let resp_headers = crate::http::models::extract_headers(upstream_resp.headers()); - let resp_body = upstream_resp + let resp_body = upstream_resp.into_body(); + let resp_body = http_body_util::Limited::new(resp_body, ctx.limits.max_body_bytes) .collect() .await .map(|b| b.to_bytes()) @@ -251,17 +267,23 @@ async fn handle_websocket_upgrade( timing: None, }; if in_scope { - let _ = ctx.ui_tx.send(ProxyToUi::ResponseReceived(request_id, response_data)); + let _ = ctx + .ui_tx + .try_send(ProxyToUi::ResponseReceived(request_id, response_data)); } let ui_tx_clone = ctx.ui_tx.clone(); let host_owned = host.to_string(); + let max_ws_frame_bytes = ctx.limits.max_ws_frame_bytes; tokio::spawn(async move { let upstream_upgraded = match hyper::upgrade::on(upstream_resp).await { Ok(u) => u, Err(e) => { - debug!("WebSocket upstream upgrade failed for {}: {}", host_owned, e); + debug!( + "WebSocket upstream upgrade failed for {}: {}", + host_owned, e + ); return; } }; @@ -282,6 +304,7 @@ async fn handle_websocket_upgrade( request_id, ui_tx_clone, in_scope, + max_ws_frame_bytes, ) .await; }); @@ -310,11 +333,17 @@ async fn handle_http_request( let headers = crate::http::models::extract_headers(req.headers()); let (parts, body) = req.into_parts(); - let body_bytes = match body.collect().await { + let body_bytes = match http_body_util::Limited::new(body, ctx.limits.max_body_bytes) + .collect() + .await + { Ok(b) => b.to_bytes(), Err(e) => { warn!("Failed to read request body: {}", e); - return Ok(crate::proxy::bad_gateway(&format!("Failed to read request body: {}", e))); + return Ok(crate::proxy::bad_gateway(&format!( + "Failed to read request body: {}", + e + ))); } }; @@ -342,7 +371,9 @@ async fn handle_http_request( }; if in_scope { - let _ = ctx.ui_tx.send(ProxyToUi::RequestCaptured(request_data.clone())); + let _ = ctx + .ui_tx + .try_send(ProxyToUi::RequestCaptured(request_data.clone())); if let Some(rx) = ctx.intercept.intercept_request(&request_data, &ctx.ui_tx) { match rx.await { @@ -368,25 +399,41 @@ async fn handle_http_request( &mut request_data.body, ); + let fallback_authority = request_data + .headers + .iter() + .find(|(name, _)| name.eq_ignore_ascii_case("host")) + .map_or(request_data.host.as_str(), |(_, value)| value.as_str()); + let target = + match crate::proxy::resolve_upstream_target(&request_data.uri, fallback_authority, port) { + Ok(target) => target, + Err(error) => return Ok(crate::proxy::bad_gateway(&error.to_string())), + }; + let client_config = crate::tls::build_tls_client_config(); - let addr = format!("{}:{}", host, port); - let tcp = match TcpStream::connect(&addr).await { + let tcp = match crate::proxy::connect_tcp(&target.host, target.port).await { Ok(s) => { timing.tcp_connected = Some(Instant::now()); s } Err(e) => { - warn!("Failed to connect to upstream {}: {}", addr, e); - let _ = ctx.ui_tx.send(ProxyToUi::RequestError( + warn!( + "Failed to connect to upstream {}:{}: {}", + target.host, target.port, e + ); + let _ = ctx.ui_tx.try_send(ProxyToUi::RequestError( request_id, format!("Connection failed: {}", e), )); - return Ok(crate::proxy::bad_gateway(&format!("Connection failed: {}", e))); + return Ok(crate::proxy::bad_gateway(&format!( + "Connection failed: {}", + e + ))); } }; - let server_name = crate::tls::server_name_or_localhost(host); + let server_name = crate::tls::server_name_or_localhost(&target.host); let connector = TlsConnector::from(client_config); let tls_stream = match connector.connect(server_name, tcp).await { Ok(s) => { @@ -394,26 +441,26 @@ async fn handle_http_request( s } Err(e) => { - warn!("TLS handshake failed with {}: {}", addr, e); - let _ = ctx.ui_tx.send(ProxyToUi::RequestError( + warn!( + "TLS handshake failed with {}:{}: {}", + target.host, target.port, e + ); + let _ = ctx.ui_tx.try_send(ProxyToUi::RequestError( request_id, format!("TLS handshake failed: {}", e), )); - return Ok(crate::proxy::bad_gateway(&format!("TLS handshake failed: {}", e))); + return Ok(crate::proxy::bad_gateway(&format!( + "TLS handshake failed: {}", + e + ))); } }; let io = TokioIo::new(tls_stream); - let path_and_query = parts - .uri - .path_and_query() - .map(|pq| pq.as_str()) - .unwrap_or("/"); - Ok(crate::proxy::forward_h1( io, parts.version, - path_and_query, + &target.path_and_query, &request_data, timing, in_scope, @@ -472,9 +519,10 @@ struct GrpcTeeBody { inner: Incoming, request_id: RequestId, direction: GrpcDirection, - ui_tx: mpsc::UnboundedSender, + ui_tx: mpsc::Sender, buffer: BytesMut, in_scope: bool, + max_message_bytes: usize, } impl http_body::Body for GrpcTeeBody { @@ -496,11 +544,14 @@ impl http_body::Body for GrpcTeeBody { this.request_id, this.direction, &this.ui_tx, + this.max_message_bytes, ); } if let Some(trailers) = frame.trailers_ref() { let pairs = crate::http::models::extract_trailers(Some(trailers)); - let _ = this.ui_tx.send(ProxyToUi::GrpcTrailers(this.request_id, pairs)); + let _ = this + .ui_tx + .try_send(ProxyToUi::GrpcTrailers(this.request_id, pairs)); } } Poll::Ready(Some(Ok(frame))) @@ -518,16 +569,24 @@ fn extract_grpc_messages( buffer: &mut BytesMut, request_id: RequestId, direction: GrpcDirection, - ui_tx: &mpsc::UnboundedSender, + ui_tx: &mpsc::Sender, + max_message_bytes: usize, ) { while buffer.len() >= 5 { let compressed = buffer[0] != 0; - let len = - u32::from_be_bytes([buffer[1], buffer[2], buffer[3], buffer[4]]) as usize; - if buffer.len() < 5 + len { + let len = u32::from_be_bytes([buffer[1], buffer[2], buffer[3], buffer[4]]) as usize; + if len > max_message_bytes { + buffer.clear(); + return; + } + let Some(frame_len) = 5usize.checked_add(len) else { + buffer.clear(); + return; + }; + if buffer.len() < frame_len { break; } - let frame_data = buffer.split_to(5 + len); + let frame_data = buffer.split_to(frame_len); let payload = Bytes::copy_from_slice(&frame_data[5..]); let msg = GrpcMessage { @@ -536,7 +595,7 @@ fn extract_grpc_messages( payload, timestamp: std::time::SystemTime::now(), }; - let _ = ui_tx.send(ProxyToUi::GrpcFrame(request_id, msg)); + let _ = ui_tx.try_send(ProxyToUi::GrpcFrame(request_id, msg)); } } @@ -583,7 +642,10 @@ async fn run_h2_tunnel( let upstream_sender = match establish_h2_upstream(host, port).await { Ok(s) => s, Err(e) => { - warn!("Failed to establish h2 upstream to {}:{}: {}", host, port, e); + warn!( + "Failed to establish h2 upstream to {}:{}: {}", + host, port, e + ); return Err(e); } }; @@ -593,9 +655,7 @@ async fn run_h2_tunnel( let host = host.clone(); let ctx = ctx.clone(); let sender = upstream_sender.clone(); - async move { - handle_h2_request(req, &host, port, &ctx, sender).await - } + async move { handle_h2_request(req, &host, port, &ctx, sender).await } }); hyper::server::conn::http2::Builder::new(TokioExecutor::new()) @@ -611,8 +671,7 @@ async fn establish_h2_upstream( ) -> anyhow::Result>> { let client_config = crate::tls::build_tls_h2_client_config(); - let addr = format!("{}:{}", host, port); - let tcp = TcpStream::connect(&addr).await?; + let tcp = crate::proxy::connect_tcp(host, port).await?; let server_name = crate::tls::server_name_or_localhost(host); let connector = TlsConnector::from(client_config); @@ -649,11 +708,17 @@ async fn handle_h2_request( let headers = crate::http::models::extract_headers(req.headers()); let (parts, body) = req.into_parts(); - let body_bytes = match body.collect().await { + let body_bytes = match http_body_util::Limited::new(body, ctx.limits.max_body_bytes) + .collect() + .await + { Ok(b) => b.to_bytes(), Err(e) => { warn!("Failed to read h2 request body: {}", e); - return Ok(h2_bad_gateway(&format!("Failed to read request body: {}", e))); + return Ok(h2_bad_gateway(&format!( + "Failed to read request body: {}", + e + ))); } }; @@ -681,30 +746,37 @@ async fn handle_h2_request( if in_scope { if is_grpc && !body_bytes.is_empty() { let mut buf = BytesMut::from(body_bytes.as_ref()); - extract_grpc_messages(&mut buf, request_id, GrpcDirection::ClientToServer, &ctx.ui_tx); + extract_grpc_messages( + &mut buf, + request_id, + GrpcDirection::ClientToServer, + &ctx.ui_tx, + ctx.limits.max_body_bytes, + ); } - let _ = ctx.ui_tx.send(ProxyToUi::RequestCaptured(request_data.clone())); - - if !is_grpc - && let Some(rx) = ctx.intercept.intercept_request(&request_data, &ctx.ui_tx) { - match rx.await { - Ok(InterceptDecision::Drop) => { - return Ok(hyper::Response::builder() - .status(503) - .body(H2RespBody::Buffered(H2Body { - data: Some(Bytes::from("Request dropped by interceptor")), - trailers: None, - })) - .unwrap()); - } - Ok(InterceptDecision::ForwardEdited(edited)) => { - request_data = edited; - } - Ok(InterceptDecision::Forward) => {} - Err(_) => {} + let _ = ctx + .ui_tx + .try_send(ProxyToUi::RequestCaptured(request_data.clone())); + + if !is_grpc && let Some(rx) = ctx.intercept.intercept_request(&request_data, &ctx.ui_tx) { + match rx.await { + Ok(InterceptDecision::Drop) => { + return Ok(hyper::Response::builder() + .status(503) + .body(H2RespBody::Buffered(H2Body { + data: Some(Bytes::from("Request dropped by interceptor")), + trailers: None, + })) + .unwrap()); } + Ok(InterceptDecision::ForwardEdited(edited)) => { + request_data = edited; + } + Ok(InterceptDecision::Forward) => {} + Err(_) => {} } + } } if !is_grpc { @@ -716,37 +788,39 @@ async fn handle_h2_request( ); } - let fwd_path = { - let uri: http::Uri = request_data.uri.parse().unwrap_or_else(|_| { - http::Uri::builder() - .path_and_query("/") - .build() - .unwrap() - }); - uri.path_and_query() - .map(|pq| pq.as_str()) - .unwrap_or("/") - .to_string() - }; - - let upstream_uri = if port == 443 { - format!("https://{}{}", host, fwd_path) - } else { - format!("https://{}:{}{}", host, port, fwd_path) - }; + let fallback_authority = request_data + .headers + .iter() + .find(|(name, _)| name.eq_ignore_ascii_case("host")) + .map_or(host, |(_, value)| value.as_str()); + let target = + match crate::proxy::resolve_upstream_target(&request_data.uri, fallback_authority, port) { + Ok(target) => target, + Err(error) => return Ok(h2_bad_gateway(&error.to_string())), + }; + if !target.host.eq_ignore_ascii_case(host) || target.port != port { + return Ok(h2_bad_gateway( + "Changing authority is not supported on an existing HTTP/2 tunnel", + )); + } + let authority = crate::proxy::format_authority(host, port, 443); + let upstream_uri = format!("https://{}{}", authority, target.path_and_query); - let upstream_req = crate::proxy::build_forwarding_request( - parts.method.as_str(), + let upstream_req = match crate::proxy::build_forwarding_request( + &request_data.method, &upstream_uri, &request_data.headers, request_data.body.clone(), - ); + ) { + Ok(request) => request, + Err(error) => return Ok(h2_bad_gateway(&format!("Invalid edited request: {error}"))), + }; let upstream_resp = match upstream.send_request(upstream_req).await { Ok(resp) => resp, Err(e) => { warn!("H2 upstream request to {}:{} failed: {}", host, port, e); - let _ = ctx.ui_tx.send(ProxyToUi::RequestError( + let _ = ctx.ui_tx.try_send(ProxyToUi::RequestError( request_id, format!("Request failed: {}", e), )); @@ -759,13 +833,26 @@ async fn handle_h2_request( if is_grpc { return Ok(build_grpc_response( - upstream_resp, resp_status, resp_headers, timing, request_id, in_scope, ctx, + upstream_resp, + resp_status, + resp_headers, + timing, + request_id, + in_scope, + ctx, )); } build_buffered_response( - upstream_resp, resp_status, resp_headers, timing, request_id, in_scope, ctx, - ).await + upstream_resp, + resp_status, + resp_headers, + timing, + request_id, + in_scope, + ctx, + ) + .await } fn build_grpc_response( @@ -792,7 +879,9 @@ fn build_grpc_response( duration, timing: Some(timing_data), }; - let _ = ctx.ui_tx.send(ProxyToUi::ResponseReceived(request_id, response_data)); + let _ = ctx + .ui_tx + .try_send(ProxyToUi::ResponseReceived(request_id, response_data)); } let (_, resp_body_incoming) = upstream_resp.into_parts(); @@ -803,6 +892,7 @@ fn build_grpc_response( ui_tx: ctx.ui_tx.clone(), buffer: BytesMut::new(), in_scope, + max_message_bytes: ctx.limits.max_body_bytes, }; let mut response = hyper::Response::builder().status(resp_status); @@ -825,17 +915,21 @@ async fn build_buffered_response( let (_, resp_body_incoming) = upstream_resp.into_parts(); let body_start = Instant::now(); - let collected = match resp_body_incoming.collect().await { - Ok(c) => c, - Err(e) => { - error!("Failed to read h2 upstream response body: {}", e); - let _ = ctx.ui_tx.send(ProxyToUi::RequestError( - request_id, - format!("Response body error: {}", e), - )); - return Ok(h2_bad_gateway(&format!("Response body error: {}", e))); - } - }; + let collected = + match http_body_util::Limited::new(resp_body_incoming, ctx.limits.max_body_bytes) + .collect() + .await + { + Ok(c) => c, + Err(e) => { + error!("Failed to read h2 upstream response body: {}", e); + let _ = ctx.ui_tx.try_send(ProxyToUi::RequestError( + request_id, + format!("Response body error: {}", e), + )); + return Ok(h2_bad_gateway(&format!("Response body error: {}", e))); + } + }; let content_transfer = body_start.elapsed(); let resp_trailers = collected.trailers().cloned(); @@ -861,7 +955,9 @@ async fn build_buffered_response( }; if in_scope { - let _ = ctx.ui_tx.send(ProxyToUi::ResponseReceived(request_id, response_data)); + let _ = ctx + .ui_tx + .try_send(ProxyToUi::ResponseReceived(request_id, response_data)); } let mut response = hyper::Response::builder().status(resp_status); diff --git a/src/proxy/ws_relay.rs b/src/proxy/ws_relay.rs index 12c8882..5702e8a 100644 --- a/src/proxy/ws_relay.rs +++ b/src/proxy/ws_relay.rs @@ -1,3 +1,4 @@ +use std::time::Duration; use std::time::SystemTime; use bytes::{Bytes, BytesMut}; @@ -8,12 +9,15 @@ use tracing::debug; use crate::channel::ProxyToUi; use crate::http::models::{RequestId, WsDirection, WsMessage}; +type ParsedFrame = (Bytes, u8, Vec); + pub async fn relay( mut client: C, mut server: S, request_id: RequestId, - ui_tx: mpsc::UnboundedSender, + ui_tx: mpsc::Sender, in_scope: bool, + max_frame_bytes: usize, ) where C: AsyncRead + AsyncWrite + Unpin, S: AsyncRead + AsyncWrite + Unpin, @@ -22,15 +26,24 @@ pub async fn relay( let mut server_buf = BytesMut::with_capacity(8192); let mut client_tmp = [0u8; 8192]; let mut server_tmp = [0u8; 8192]; + let idle = tokio::time::sleep(Duration::from_secs(120)); + tokio::pin!(idle); loop { tokio::select! { + _ = &mut idle => break, result = client.read(&mut client_tmp) => { match result { Ok(0) | Err(_) => break, Ok(n) => { + idle.as_mut().reset(tokio::time::Instant::now() + Duration::from_secs(120)); client_buf.extend_from_slice(&client_tmp[..n]); - while let Some((raw, opcode, payload)) = try_parse_frame(&mut client_buf) { + loop { + let (raw, opcode, payload) = match parse_next_frame(&mut client_buf, max_frame_bytes) { + ParseNext::Frame(frame) => frame, + ParseNext::Incomplete => break, + ParseNext::Invalid => return, + }; if in_scope && is_data_frame(opcode) { let msg = WsMessage { direction: WsDirection::ClientToServer, @@ -38,7 +51,7 @@ pub async fn relay( payload: Bytes::from(payload), timestamp: SystemTime::now(), }; - let _ = ui_tx.send(ProxyToUi::WebSocketFrame(request_id, msg)); + let _ = ui_tx.try_send(ProxyToUi::WebSocketFrame(request_id, msg)); } if server.write_all(&raw).await.is_err() { return; @@ -51,8 +64,14 @@ pub async fn relay( match result { Ok(0) | Err(_) => break, Ok(n) => { + idle.as_mut().reset(tokio::time::Instant::now() + Duration::from_secs(120)); server_buf.extend_from_slice(&server_tmp[..n]); - while let Some((raw, opcode, payload)) = try_parse_frame(&mut server_buf) { + loop { + let (raw, opcode, payload) = match parse_next_frame(&mut server_buf, max_frame_bytes) { + ParseNext::Frame(frame) => frame, + ParseNext::Incomplete => break, + ParseNext::Invalid => return, + }; if in_scope && is_data_frame(opcode) { let msg = WsMessage { direction: WsDirection::ServerToClient, @@ -60,7 +79,7 @@ pub async fn relay( payload: Bytes::from(payload), timestamp: SystemTime::now(), }; - let _ = ui_tx.send(ProxyToUi::WebSocketFrame(request_id, msg)); + let _ = ui_tx.try_send(ProxyToUi::WebSocketFrame(request_id, msg)); } if client.write_all(&raw).await.is_err() { return; @@ -75,13 +94,34 @@ pub async fn relay( debug!("WebSocket relay for request {} finished", request_id); } +enum ParseNext { + Incomplete, + Frame(ParsedFrame), + Invalid, +} + +fn parse_next_frame(buf: &mut BytesMut, max_frame_bytes: usize) -> ParseNext { + match try_parse_frame(buf, max_frame_bytes) { + Ok(Some(frame)) => ParseNext::Frame(frame), + Ok(None) => ParseNext::Incomplete, + Err(error) => { + debug!("Closing WebSocket relay: {error}"); + buf.clear(); + ParseNext::Invalid + } + } +} + fn is_data_frame(opcode: u8) -> bool { opcode == 1 || opcode == 2 } -fn try_parse_frame(buf: &mut BytesMut) -> Option<(Bytes, u8, Vec)> { +fn try_parse_frame( + buf: &mut BytesMut, + max_frame_bytes: usize, +) -> Result, &'static str> { if buf.len() < 2 { - return None; + return Ok(None); } let b0 = buf[0]; @@ -94,13 +134,13 @@ fn try_parse_frame(buf: &mut BytesMut) -> Option<(Bytes, u8, Vec)> { if payload_len == 126 { if buf.len() < 4 { - return None; + return Ok(None); } payload_len = u16::from_be_bytes([buf[2], buf[3]]) as u64; offset = 4; } else if payload_len == 127 { if buf.len() < 10 { - return None; + return Ok(None); } payload_len = u64::from_be_bytes([ buf[2], buf[3], buf[4], buf[5], buf[6], buf[7], buf[8], buf[9], @@ -108,17 +148,30 @@ fn try_parse_frame(buf: &mut BytesMut) -> Option<(Bytes, u8, Vec)> { offset = 10; } - let mask_size = if masked { 4 } else { 0 }; - let total = offset + mask_size + payload_len as usize; + let payload_len = + usize::try_from(payload_len).map_err(|_| "frame length does not fit usize")?; + if payload_len > max_frame_bytes { + return Err("frame exceeds configured limit"); + } + let mask_size = if masked { 4usize } else { 0 }; + let total = offset + .checked_add(mask_size) + .and_then(|value| value.checked_add(payload_len)) + .ok_or("frame length overflow")?; if buf.len() < total { - return None; + return Ok(None); } let raw = buf.split_to(total); let mask_key = if masked { - Some([raw[offset], raw[offset + 1], raw[offset + 2], raw[offset + 3]]) + Some([ + raw[offset], + raw[offset + 1], + raw[offset + 2], + raw[offset + 3], + ]) } else { None }; @@ -132,5 +185,33 @@ fn try_parse_frame(buf: &mut BytesMut) -> Option<(Bytes, u8, Vec)> { } } - Some((raw.freeze(), opcode, payload)) + Ok(Some((raw.freeze(), opcode, payload))) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn rejects_frames_above_limit() { + let mut frame = BytesMut::from(&b"\x82\x7e\x00\x20"[..]); + assert_eq!( + try_parse_frame(&mut frame, 16), + Err("frame exceeds configured limit") + ); + } + + #[test] + fn rejects_overflowing_64_bit_length() { + let mut frame = BytesMut::from(&b"\x82\x7f\xff\xff\xff\xff\xff\xff\xff\xff"[..]); + assert!(try_parse_frame(&mut frame, usize::MAX).is_err()); + } + + #[test] + fn parses_small_frame() { + let mut frame = BytesMut::from(&b"\x81\x02ok"[..]); + let (_, opcode, payload) = try_parse_frame(&mut frame, 16).unwrap().unwrap(); + assert_eq!(opcode, 1); + assert_eq!(payload, b"ok"); + } } diff --git a/src/rules/mod.rs b/src/rules/mod.rs index 66a1ddf..74788ca 100644 --- a/src/rules/mod.rs +++ b/src/rules/mod.rs @@ -160,13 +160,16 @@ fn apply_rule( if let Some(uri) = uri && (scope == RuleScope::Url || scope == RuleScope::All) - && let Cow::Owned(s) = replace_in_str(uri, &rule.match_pattern, &rule.replacement, compiled) { - *uri = s; - } + && let Cow::Owned(s) = replace_in_str(uri, &rule.match_pattern, &rule.replacement, compiled) + { + *uri = s; + } if scope == RuleScope::Headers || scope == RuleScope::All { for (_key, value) in headers.iter_mut() { - if let Cow::Owned(s) = replace_in_str(value, &rule.match_pattern, &rule.replacement, compiled) { + if let Cow::Owned(s) = + replace_in_str(value, &rule.match_pattern, &rule.replacement, compiled) + { *value = s; } } @@ -174,12 +177,19 @@ fn apply_rule( if (scope == RuleScope::Body || scope == RuleScope::All) && let Ok(text) = std::str::from_utf8(body) - && let Cow::Owned(s) = replace_in_str(text, &rule.match_pattern, &rule.replacement, compiled) { - *body = Bytes::from(s); - } + && let Cow::Owned(s) = + replace_in_str(text, &rule.match_pattern, &rule.replacement, compiled) + { + *body = Bytes::from(s); + } } -fn replace_in_str<'a>(input: &'a str, pattern: &str, replacement: &str, compiled: Option<&Regex>) -> Cow<'a, str> { +fn replace_in_str<'a>( + input: &'a str, + pattern: &str, + replacement: &str, + compiled: Option<&Regex>, +) -> Cow<'a, str> { match compiled { Some(re) => re.replace_all(input, replacement), None => { diff --git a/src/rules/persist.rs b/src/rules/persist.rs index 9e143e6..a7c93e5 100644 --- a/src/rules/persist.rs +++ b/src/rules/persist.rs @@ -36,7 +36,7 @@ pub fn rules_dir() -> anyhow::Result { .ok_or_else(|| anyhow::anyhow!("Cannot find home directory"))? .join(".crowbar") .join("rules"); - std::fs::create_dir_all(&dir)?; + crate::fs_security::ensure_private_dir(&dir)?; Ok(dir) } @@ -44,15 +44,23 @@ pub fn save(rules: &[Rule], name: &str) -> anyhow::Result { let dir = rules_dir()?; let path = dir.join(format!("{}.json", safe_name(name)?)); let file = RulesFile::from_rules(rules); - let json = serde_json::to_string_pretty(&file)?; - std::fs::write(&path, json)?; + write_rules_file(&path, &file)?; Ok(path) } pub fn save_to(rules: &[Rule], path: &Path) -> anyhow::Result<()> { let file = RulesFile::from_rules(rules); - let json = serde_json::to_string_pretty(&file)?; - std::fs::write(path, json)?; + write_rules_file(path, &file)?; + Ok(()) +} + +fn write_rules_file(path: &Path, rules: &RulesFile) -> anyhow::Result<()> { + crate::fs_security::write_private_with(path, |file| { + use std::io::Write; + let mut writer = std::io::BufWriter::new(file); + serde_json::to_writer_pretty(&mut writer, rules).map_err(std::io::Error::other)?; + writer.flush() + })?; Ok(()) } diff --git a/src/scanning.rs b/src/scanning.rs index 46ab81d..b5033f7 100644 --- a/src/scanning.rs +++ b/src/scanning.rs @@ -42,7 +42,10 @@ fn has_header(headers: &[(String, String)], name: &str) -> bool { headers.iter().any(|(k, _)| k.eq_ignore_ascii_case(name)) } -fn get_headers<'a>(headers: &'a [(String, String)], name: &'a str) -> impl Iterator { +fn get_headers<'a>( + headers: &'a [(String, String)], + name: &'a str, +) -> impl Iterator { headers .iter() .filter(move |(k, _)| k.eq_ignore_ascii_case(name)) @@ -105,11 +108,7 @@ fn check_info_disclosure(response: &ResponseData, findings: &mut Vec) { } } -fn check_cookie_flags( - request: &RequestData, - response: &ResponseData, - findings: &mut Vec, -) { +fn check_cookie_flags(request: &RequestData, response: &ResponseData, findings: &mut Vec) { for cookie in get_headers(&response.headers, "set-cookie") { let lower = cookie.to_ascii_lowercase(); @@ -189,9 +188,9 @@ fn truncate_cookie(cookie: &str) -> String { #[cfg(test)] mod tests { use super::*; + use crate::http::models::{HttpVersion, RequestId}; use bytes::Bytes; use std::time::{Duration, SystemTime}; - use crate::http::models::{HttpVersion, RequestId}; fn make_request(is_tls: bool) -> RequestData { RequestData { @@ -213,7 +212,10 @@ mod tests { status: 200, reason: "OK".into(), version: HttpVersion::Http11, - headers: headers.into_iter().map(|(k, v)| (k.into(), v.into())).collect(), + headers: headers + .into_iter() + .map(|(k, v)| (k.into(), v.into())) + .collect(), body: Bytes::from(body.to_string()), trailers: Vec::new(), duration: Duration::from_millis(50), @@ -226,7 +228,11 @@ mod tests { let req = make_request(true); let resp = make_response(vec![], ""); let findings = scan_response(&req, &resp); - assert!(findings.iter().any(|f| f.title.contains("Strict-Transport-Security"))); + assert!( + findings + .iter() + .any(|f| f.title.contains("Strict-Transport-Security")) + ); } #[test] @@ -234,7 +240,11 @@ mod tests { let req = make_request(false); let resp = make_response(vec![], ""); let findings = scan_response(&req, &resp); - assert!(!findings.iter().any(|f| f.title.contains("Strict-Transport-Security"))); + assert!( + !findings + .iter() + .any(|f| f.title.contains("Strict-Transport-Security")) + ); } #[test] @@ -258,7 +268,10 @@ mod tests { #[test] fn stack_trace_detection() { let req = make_request(false); - let resp = make_response(vec![], "Error: Traceback (most recent call last):\n File app.py"); + let resp = make_response( + vec![], + "Error: Traceback (most recent call last):\n File app.py", + ); let findings = scan_response(&req, &resp); assert!(findings.iter().any(|f| f.severity == Severity::High)); } @@ -376,7 +389,11 @@ mod tests { let req = make_request(false); let resp = make_response(vec![], ""); let findings = scan_response(&req, &resp); - assert!(findings.iter().any(|f| f.title.contains("Content-Security-Policy"))); + assert!( + findings + .iter() + .any(|f| f.title.contains("Content-Security-Policy")) + ); } #[test] @@ -392,7 +409,11 @@ mod tests { let req = make_request(false); let resp = make_response(vec![], ""); let findings = scan_response(&req, &resp); - assert!(findings.iter().any(|f| f.title.contains("X-Content-Type-Options"))); + assert!( + findings + .iter() + .any(|f| f.title.contains("X-Content-Type-Options")) + ); } #[test] @@ -415,7 +436,10 @@ mod tests { fn cookie_with_all_flags_no_findings() { let req = make_request(true); let resp = make_response( - vec![("Set-Cookie", "id=abc; Secure; HttpOnly; SameSite=Strict; Path=/")], + vec![( + "Set-Cookie", + "id=abc; Secure; HttpOnly; SameSite=Strict; Path=/", + )], "", ); let findings = scan_response(&req, &resp); diff --git a/src/tls/ca.rs b/src/tls/ca.rs index a522c1b..9e1000b 100644 --- a/src/tls/ca.rs +++ b/src/tls/ca.rs @@ -18,7 +18,7 @@ pub struct CertificateAuthority { impl CertificateAuthority { pub fn load_or_generate() -> anyhow::Result { let dir = Self::config_dir()?; - std::fs::create_dir_all(&dir)?; + crate::fs_security::harden_private_tree(&dir)?; let cert_path = dir.join("ca.pem"); let key_path = dir.join("ca.key"); @@ -78,8 +78,8 @@ impl CertificateAuthority { } fn save_to_disk(&self, cert_path: &Path, key_path: &Path) -> anyhow::Result<()> { - std::fs::write(cert_path, &self.ca_cert_pem)?; - std::fs::write(key_path, self.ca_key.serialize_pem())?; + crate::fs_security::write_private(cert_path, &self.ca_cert_pem)?; + crate::fs_security::write_private(key_path, self.ca_key.serialize_pem())?; Ok(()) } diff --git a/src/tls/cert_gen.rs b/src/tls/cert_gen.rs index b1d3ee5..5127d67 100644 --- a/src/tls/cert_gen.rs +++ b/src/tls/cert_gen.rs @@ -30,7 +30,9 @@ pub fn generate_leaf_cert( Ok(CertifiedKey::new(cert_chain, signing_key)) } -pub fn build_server_config(certified_key: Arc) -> anyhow::Result> { +pub fn build_server_config( + certified_key: Arc, +) -> anyhow::Result> { let mut config = rustls::ServerConfig::builder() .with_no_client_auth() .with_cert_resolver(Arc::new(SingleCertResolver(certified_key))); @@ -44,10 +46,7 @@ pub fn build_server_config(certified_key: Arc) -> anyhow::Result); impl rustls::server::ResolvesServerCert for SingleCertResolver { - fn resolve( - &self, - _client_hello: rustls::server::ClientHello<'_>, - ) -> Option> { + fn resolve(&self, _client_hello: rustls::server::ClientHello<'_>) -> Option> { Some(Arc::clone(&self.0)) } } diff --git a/src/tls/mod.rs b/src/tls/mod.rs index 4742b3f..c8eb568 100644 --- a/src/tls/mod.rs +++ b/src/tls/mod.rs @@ -23,6 +23,7 @@ pub fn build_tls_h2_client_config() -> Arc { } pub fn server_name_or_localhost(host: &str) -> rustls::pki_types::ServerName<'static> { - rustls::pki_types::ServerName::try_from(host.to_owned()) - .unwrap_or_else(|_| rustls::pki_types::ServerName::try_from("localhost".to_owned()).unwrap()) + rustls::pki_types::ServerName::try_from(host.to_owned()).unwrap_or_else(|_| { + rustls::pki_types::ServerName::try_from("localhost".to_owned()).unwrap() + }) } diff --git a/src/tui/tabs/history_tab.rs b/src/tui/tabs/history_tab.rs index d8b6215..1197c64 100644 --- a/src/tui/tabs/history_tab.rs +++ b/src/tui/tabs/history_tab.rs @@ -1,8 +1,11 @@ +use ratatui::Frame; use ratatui::layout::{Constraint, Layout, Rect}; use ratatui::style::{Color, Modifier, Style}; use ratatui::text::{Line, Span, Text}; -use ratatui::widgets::{Block, Borders, Cell, Paragraph, Row, Scrollbar, ScrollbarOrientation, ScrollbarState, Table, TableState, Wrap}; -use ratatui::Frame; +use ratatui::widgets::{ + Block, Borders, Cell, Paragraph, Row, Scrollbar, ScrollbarOrientation, ScrollbarState, Table, + TableState, Wrap, +}; use crate::app::App; use crate::http::models::EntryState; @@ -12,9 +15,7 @@ pub fn render(app: &App, frame: &mut Frame, area: Rect) { let filtered = app.store.filtered_entries_all(); if app.store.is_empty() { - let block = Block::default() - .borders(Borders::ALL) - .title(" History "); + let block = Block::default().borders(Borders::ALL).title(" History "); let inner = block.inner(area); frame.render_widget(block, area); logo::render(frame, inner); @@ -24,11 +25,7 @@ pub fn render(app: &App, frame: &mut Frame, area: Rect) { let has_filter = !app.history.filter.is_empty() || app.history.filtering; let (filter_area, content_area) = if has_filter { - let chunks = Layout::vertical([ - Constraint::Length(1), - Constraint::Min(0), - ]) - .split(area); + let chunks = Layout::vertical([Constraint::Length(1), Constraint::Min(0)]).split(area); (Some(chunks[0]), chunks[1]) } else { (None, area) @@ -45,10 +42,7 @@ pub fn render(app: &App, frame: &mut Frame, area: Rect) { let has_findings = selected.is_some_and(|e| !e.findings.is_empty()); if has_ws || has_grpc || has_findings { - let mut constraints = vec![ - Constraint::Percentage(25), - Constraint::Percentage(35), - ]; + let mut constraints = vec![Constraint::Percentage(25), Constraint::Percentage(35)]; let extra_panes = has_ws as usize + has_grpc as usize + has_findings as usize; let remaining = 40u16 / extra_panes as u16; for _ in 0..extra_panes { @@ -72,11 +66,8 @@ pub fn render(app: &App, frame: &mut Frame, area: Rect) { render_grpc_messages(app, &filtered, frame, chunks[pane_idx]); } } else { - let chunks = Layout::vertical([ - Constraint::Percentage(40), - Constraint::Percentage(60), - ]) - .split(content_area); + let chunks = Layout::vertical([Constraint::Percentage(40), Constraint::Percentage(60)]) + .split(content_area); render_table_filtered(app, &filtered, frame, chunks[0]); render_detail_filtered(app, &filtered, frame, chunks[1]); @@ -103,7 +94,12 @@ fn render_filter_bar(app: &App, frame: &mut Frame, area: Rect, match_count: usiz frame.render_widget(Paragraph::new(line), area); } -fn render_table_filtered(app: &App, filtered: &[&crate::http::models::HistoryEntry], frame: &mut Frame, area: Rect) { +fn render_table_filtered( + app: &App, + filtered: &[&crate::http::models::HistoryEntry], + frame: &mut Frame, + area: Rect, +) { let header = Row::new(vec![ Cell::from("#"), Cell::from("Method"), @@ -113,7 +109,11 @@ fn render_table_filtered(app: &App, filtered: &[&crate::http::models::HistoryEnt Cell::from("Size"), Cell::from("Time"), ]) - .style(Style::default().fg(Color::Yellow).add_modifier(Modifier::BOLD)) + .style( + Style::default() + .fg(Color::Yellow) + .add_modifier(Modifier::BOLD), + ) .height(1); let visible_height = area.height.saturating_sub(3) as usize; @@ -137,7 +137,11 @@ fn render_table_filtered(app: &App, filtered: &[&crate::http::models::HistoryEnt &req.uri }; let path = if path.chars().count() > 50 { - let end = path.char_indices().nth(47).map(|(i, _)| i).unwrap_or(path.len()); + let end = path + .char_indices() + .nth(47) + .map(|(i, _)| i) + .unwrap_or(path.len()); format!("{}...", &path[..end]) } else { path.to_string() @@ -243,17 +247,19 @@ fn render_table_filtered(app: &App, filtered: &[&crate::http::models::HistoryEnt frame.render_stateful_widget(table, area, &mut state); } -fn render_detail_filtered(app: &App, filtered: &[&crate::http::models::HistoryEntry], frame: &mut Frame, area: Rect) { +fn render_detail_filtered( + app: &App, + filtered: &[&crate::http::models::HistoryEntry], + frame: &mut Frame, + area: Rect, +) { let entry = match filtered.get(app.history.selected) { Some(e) => e, None => return, }; - let chunks = Layout::horizontal([ - Constraint::Percentage(50), - Constraint::Percentage(50), - ]) - .split(area); + let chunks = + Layout::horizontal([Constraint::Percentage(50), Constraint::Percentage(50)]).split(area); // Request pane let mut req_lines: Vec = Vec::new(); @@ -275,7 +281,9 @@ fn render_detail_filtered(app: &App, filtered: &[&crate::http::models::HistoryEn if !req.body.is_empty() { req_lines.push(Line::raw("")); - let content_type = req.headers.iter() + let content_type = req + .headers + .iter() .find(|(k, _)| k.eq_ignore_ascii_case("content-type")) .map(|(_, v)| v.as_str()); let proto_type = if req.is_grpc { @@ -284,16 +292,15 @@ fn render_detail_filtered(app: &App, filtered: &[&crate::http::models::HistoryEn None }; req_lines.extend(body_view::body_lines_with_schema( - &req.body, content_type, 100, proto_type.as_ref(), + &req.body, + content_type, + 100, + proto_type.as_ref(), )); } let req_paragraph = Paragraph::new(Text::from(req_lines)) - .block( - Block::default() - .borders(Borders::ALL) - .title(" Request "), - ) + .block(Block::default().borders(Borders::ALL).title(" Request ")) .wrap(Wrap { trim: false }) .scroll((app.history.scroll, 0)); @@ -329,25 +336,27 @@ fn render_detail_filtered(app: &App, filtered: &[&crate::http::models::HistoryEn ])); if entry.request.is_grpc - && let Some((code, name)) = resp.grpc_status() { - let grpc_style = if code == 0 { - Style::default().fg(Color::Green).bold() - } else { - Style::default().fg(Color::Red).bold() - }; - let mut grpc_spans = vec![ - Span::styled("gRPC ", Style::default().fg(Color::DarkGray)), - Span::styled(format!("{} {}", code, name), grpc_style), - ]; - if let Some(msg) = resp.grpc_message() - && !msg.is_empty() { - grpc_spans.push(Span::styled( - format!(" {}", msg), - Style::default().fg(Color::Yellow), - )); - } - resp_lines.push(Line::from(grpc_spans)); + && let Some((code, name)) = resp.grpc_status() + { + let grpc_style = if code == 0 { + Style::default().fg(Color::Green).bold() + } else { + Style::default().fg(Color::Red).bold() + }; + let mut grpc_spans = vec![ + Span::styled("gRPC ", Style::default().fg(Color::DarkGray)), + Span::styled(format!("{} {}", code, name), grpc_style), + ]; + if let Some(msg) = resp.grpc_message() + && !msg.is_empty() + { + grpc_spans.push(Span::styled( + format!(" {}", msg), + Style::default().fg(Color::Yellow), + )); } + resp_lines.push(Line::from(grpc_spans)); + } if let Some(timing) = &resp.timing { resp_lines.push(Line::raw("")); @@ -360,7 +369,9 @@ fn render_detail_filtered(app: &App, filtered: &[&crate::http::models::HistoryEn if !resp.body.is_empty() { resp_lines.push(Line::raw("")); - let content_type = resp.headers.iter() + let content_type = resp + .headers + .iter() .find(|(k, _)| k.eq_ignore_ascii_case("content-type")) .map(|(_, v)| v.as_str()); let proto_type = if entry.request.is_grpc { @@ -369,7 +380,10 @@ fn render_detail_filtered(app: &App, filtered: &[&crate::http::models::HistoryEn None }; resp_lines.extend(body_view::body_lines_with_schema( - &resp.body, content_type, 200, proto_type.as_ref(), + &resp.body, + content_type, + 200, + proto_type.as_ref(), )); } @@ -381,16 +395,10 @@ fn render_detail_filtered(app: &App, filtered: &[&crate::http::models::HistoryEn let msg = match entry.state { EntryState::Pending => "Awaiting response...", EntryState::Dropped => "Request was dropped", - EntryState::Error => entry - .error_message - .as_deref() - .unwrap_or("Unknown error"), + EntryState::Error => entry.error_message.as_deref().unwrap_or("Unknown error"), EntryState::Complete => "No response data", }; - resp_lines.push(Line::styled( - msg, - Style::default().fg(Color::DarkGray), - )); + resp_lines.push(Line::styled(msg, Style::default().fg(Color::DarkGray))); } } @@ -398,29 +406,26 @@ fn render_detail_filtered(app: &App, filtered: &[&crate::http::models::HistoryEn let visible_height = chunks[1].height.saturating_sub(2); let resp_paragraph = Paragraph::new(Text::from(resp_lines)) - .block( - Block::default() - .borders(Borders::ALL) - .title(" Response "), - ) + .block(Block::default().borders(Borders::ALL).title(" Response ")) .wrap(Wrap { trim: false }) .scroll((app.history.scroll, 0)); frame.render_widget(resp_paragraph, chunks[1]); if content_height > visible_height { - let mut scrollbar_state = ScrollbarState::new(content_height as usize) - .position(app.history.scroll as usize); + let mut scrollbar_state = + ScrollbarState::new(content_height as usize).position(app.history.scroll as usize); let scrollbar = Scrollbar::new(ScrollbarOrientation::VerticalRight); - frame.render_stateful_widget( - scrollbar, - chunks[1], - &mut scrollbar_state, - ); + frame.render_stateful_widget(scrollbar, chunks[1], &mut scrollbar_state); } } -fn render_findings(app: &App, filtered: &[&crate::http::models::HistoryEntry], frame: &mut Frame, area: Rect) { +fn render_findings( + app: &App, + filtered: &[&crate::http::models::HistoryEntry], + frame: &mut Frame, + area: Rect, +) { let entry = match filtered.get(app.history.selected) { Some(e) => e, None => return, @@ -442,7 +447,10 @@ fn render_findings(app: &App, filtered: &[&crate::http::models::HistoryEntry], f format!(" [{:>4}] ", finding.severity.label()), severity_style, ), - Span::styled(&finding.title, Style::default().add_modifier(Modifier::BOLD)), + Span::styled( + &finding.title, + Style::default().add_modifier(Modifier::BOLD), + ), ])); lines.push(Line::from(vec![ Span::raw(" "), @@ -452,18 +460,19 @@ fn render_findings(app: &App, filtered: &[&crate::http::models::HistoryEntry], f let title = format!(" Findings ({}) ", entry.findings.len()); let widget = Paragraph::new(Text::from(lines)) - .block( - Block::default() - .borders(Borders::ALL) - .title(title), - ) + .block(Block::default().borders(Borders::ALL).title(title)) .wrap(Wrap { trim: false }) .scroll((app.history.scroll, 0)); frame.render_widget(widget, area); } -fn render_ws_messages(app: &App, filtered: &[&crate::http::models::HistoryEntry], frame: &mut Frame, area: Rect) { +fn render_ws_messages( + app: &App, + filtered: &[&crate::http::models::HistoryEntry], + frame: &mut Frame, + area: Rect, +) { let entry = match filtered.get(app.history.selected) { Some(e) => e, None => return, @@ -474,14 +483,8 @@ fn render_ws_messages(app: &App, filtered: &[&crate::http::models::HistoryEntry] let mut lines: Vec = Vec::new(); for (i, msg) in entry.ws_messages.iter().enumerate() { let dir_span = match msg.direction { - WsDirection::ClientToServer => Span::styled( - ">>> ", - Style::default().fg(Color::Green), - ), - WsDirection::ServerToClient => Span::styled( - "<<< ", - Style::default().fg(Color::Cyan), - ), + WsDirection::ClientToServer => Span::styled(">>> ", Style::default().fg(Color::Green)), + WsDirection::ServerToClient => Span::styled("<<< ", Style::default().fg(Color::Cyan)), }; let type_label = if msg.is_text() { @@ -525,18 +528,19 @@ fn render_ws_messages(app: &App, filtered: &[&crate::http::models::HistoryEntry] let title = format!(" WebSocket ({} messages) ", entry.ws_messages.len()); let widget = Paragraph::new(Text::from(lines)) - .block( - Block::default() - .borders(Borders::ALL) - .title(title), - ) + .block(Block::default().borders(Borders::ALL).title(title)) .wrap(Wrap { trim: false }) .scroll((app.history.scroll, 0)); frame.render_widget(widget, area); } -fn render_grpc_messages(app: &App, filtered: &[&crate::http::models::HistoryEntry], frame: &mut Frame, area: Rect) { +fn render_grpc_messages( + app: &App, + filtered: &[&crate::http::models::HistoryEntry], + frame: &mut Frame, + area: Rect, +) { let entry = match filtered.get(app.history.selected) { Some(e) => e, None => return, @@ -547,14 +551,10 @@ fn render_grpc_messages(app: &App, filtered: &[&crate::http::models::HistoryEntr let mut lines: Vec = Vec::new(); for (i, msg) in entry.grpc_messages.iter().enumerate() { let dir_span = match msg.direction { - GrpcDirection::ClientToServer => Span::styled( - ">>> ", - Style::default().fg(Color::Green), - ), - GrpcDirection::ServerToClient => Span::styled( - "<<< ", - Style::default().fg(Color::Cyan), - ), + GrpcDirection::ClientToServer => { + Span::styled(">>> ", Style::default().fg(Color::Green)) + } + GrpcDirection::ServerToClient => Span::styled("<<< ", Style::default().fg(Color::Cyan)), }; let size = format_size(msg.payload.len()); @@ -639,11 +639,7 @@ fn render_grpc_messages(app: &App, filtered: &[&crate::http::models::HistoryEntr let title = format!(" gRPC Messages ({}) ", entry.grpc_messages.len()); let widget = Paragraph::new(Text::from(lines)) - .block( - Block::default() - .borders(Borders::ALL) - .title(title), - ) + .block(Block::default().borders(Borders::ALL).title(title)) .wrap(Wrap { trim: false }) .scroll((app.history.scroll, 0)); diff --git a/src/tui/tabs/mod.rs b/src/tui/tabs/mod.rs index 21d9d53..8bb6378 100644 --- a/src/tui/tabs/mod.rs +++ b/src/tui/tabs/mod.rs @@ -14,7 +14,13 @@ pub enum Tab { } impl Tab { - pub const ALL: [Tab; 5] = [Tab::Proxy, Tab::History, Tab::Repeater, Tab::Rules, Tab::Tools]; + pub const ALL: [Tab; 5] = [ + Tab::Proxy, + Tab::History, + Tab::Repeater, + Tab::Rules, + Tab::Tools, + ]; pub fn title(self) -> &'static str { match self { diff --git a/src/tui/tabs/proxy_tab.rs b/src/tui/tabs/proxy_tab.rs index 7147468..d137cab 100644 --- a/src/tui/tabs/proxy_tab.rs +++ b/src/tui/tabs/proxy_tab.rs @@ -1,8 +1,8 @@ +use ratatui::Frame; use ratatui::layout::{Constraint, Layout, Rect}; use ratatui::style::{Color, Modifier, Style}; use ratatui::text::{Line, Span, Text}; use ratatui::widgets::{Block, Borders, Paragraph, Wrap}; -use ratatui::Frame; use crate::app::App; use crate::tui::widgets::{dim_style as dim_key_style, key_style, logo}; @@ -48,7 +48,10 @@ fn render_status_line(app: &App, frame: &mut Frame, area: Rect) { Span::raw(" Scope: "), Span::styled(&app.scope_buffer, Style::default().fg(Color::White)), Span::styled("\u{2588}", Style::default().fg(Color::Yellow)), - Span::styled(" (comma-separated, e.g. *.example.com, api.test.com)", Style::default().fg(Color::DarkGray)), + Span::styled( + " (comma-separated, e.g. *.example.com, api.test.com)", + Style::default().fg(Color::DarkGray), + ), ]); let widget = Paragraph::new(line).block( Block::default() @@ -71,15 +74,17 @@ fn render_status_line(app: &App, frame: &mut Frame, area: Rect) { } else { Span::styled( " INTERCEPT OFF ", - Style::default() - .bg(Color::DarkGray) - .fg(Color::White), + Style::default().bg(Color::DarkGray).fg(Color::White), ) }; let queue_count = app.intercept_ui.queue.len(); let queue_text = if queue_count > 0 { - format!(" {} request{} queued", queue_count, if queue_count == 1 { "" } else { "s" }) + format!( + " {} request{} queued", + queue_count, + if queue_count == 1 { "" } else { "s" } + ) } else { String::new() }; @@ -90,11 +95,8 @@ fn render_status_line(app: &App, frame: &mut Frame, area: Rect) { Span::styled(queue_text, Style::default().fg(Color::Yellow)), ]); - let widget = Paragraph::new(line).block( - Block::default() - .borders(Borders::ALL) - .title(" Proxy "), - ); + let widget = + Paragraph::new(line).block(Block::default().borders(Borders::ALL).title(" Proxy ")); frame.render_widget(widget, area); } @@ -107,7 +109,10 @@ fn render_current_request(app: &App, frame: &mut Frame, area: Rect) { Span::raw(" "), Span::raw(&req.uri), Span::raw(" "), - Span::styled(req.version.to_string(), Style::default().fg(Color::DarkGray)), + Span::styled( + req.version.to_string(), + Style::default().fg(Color::DarkGray), + ), ])); lines.push(Line::raw("")); @@ -226,10 +231,7 @@ fn render_actions(app: &App, frame: &mut Frame, area: Rect) { ]) }; - let widget = Paragraph::new(actions).block( - Block::default() - .borders(Borders::ALL) - .title(" Actions "), - ); + let widget = + Paragraph::new(actions).block(Block::default().borders(Borders::ALL).title(" Actions ")); frame.render_widget(widget, area); } diff --git a/src/tui/tabs/repeater_tab.rs b/src/tui/tabs/repeater_tab.rs index 4e881c0..b6eea32 100644 --- a/src/tui/tabs/repeater_tab.rs +++ b/src/tui/tabs/repeater_tab.rs @@ -1,37 +1,29 @@ +use ratatui::Frame; use ratatui::layout::{Constraint, Layout, Rect}; use ratatui::style::{Color, Modifier, Style}; use ratatui::text::{Line, Span, Text}; use ratatui::widgets::{Block, Borders, Paragraph, Wrap}; -use ratatui::Frame; use crate::app::App; use crate::http::codec; use crate::http::sequence::StepState; -use crate::tui::widgets::{body_view, diff_view, dim_style, format_size, key_style, logo, timing_view}; +use crate::tui::widgets::{ + body_view, diff_view, dim_style, format_size, key_style, logo, timing_view, +}; pub fn render(app: &App, frame: &mut Frame, area: Rect) { - let chunks = Layout::vertical([ - Constraint::Min(0), - Constraint::Length(3), - ]) - .split(area); + let chunks = Layout::vertical([Constraint::Min(0), Constraint::Length(3)]).split(area); if app.macros.show { - let panes = Layout::horizontal([ - Constraint::Percentage(50), - Constraint::Percentage(50), - ]) - .split(chunks[0]); + let panes = Layout::horizontal([Constraint::Percentage(50), Constraint::Percentage(50)]) + .split(chunks[0]); render_macro_list(app, frame, panes[0]); render_macro_detail(app, frame, panes[1]); render_macro_actions(app, frame, chunks[1]); } else { - let panes = Layout::horizontal([ - Constraint::Percentage(50), - Constraint::Percentage(50), - ]) - .split(chunks[0]); + let panes = Layout::horizontal([Constraint::Percentage(50), Constraint::Percentage(50)]) + .split(chunks[0]); render_request_editor(app, frame, panes[0]); render_response(app, frame, panes[1]); @@ -43,9 +35,7 @@ fn render_request_editor(app: &App, frame: &mut Frame, area: Rect) { let has_content = app.repeater.editor.has_content(); if !has_content { - let block = Block::default() - .borders(Borders::ALL) - .title(" Request "); + let block = Block::default().borders(Borders::ALL).title(" Request "); let inner = block.inner(area); frame.render_widget(block, area); logo::render(frame, inner); @@ -68,16 +58,19 @@ fn render_request_editor(app: &App, frame: &mut Frame, area: Rect) { let num_style = Style::default().fg(Color::DarkGray); for (i, line) in app.repeater.editor.lines.iter().enumerate() { - let num_span = Span::styled( - format!("{:>width$} ", i + 1, width = gw), - num_style, - ); + let num_span = Span::styled(format!("{:>width$} ", i + 1, width = gw), num_style); if app.repeater.editing && i == app.repeater.editor.cursor_line { - let char_indices: Vec<(usize, char)> = line.char_indices().collect(); - let col = app.repeater.editor.cursor_col.min(char_indices.len()); - let byte_start = char_indices.get(col).map_or(line.len(), |&(i, _)| i); - let byte_end = char_indices.get(col + 1).map_or(line.len(), |&(i, _)| i); + let char_count = line.chars().count(); + let col = app.repeater.editor.cursor_col.min(char_count); + let byte_start = line + .char_indices() + .nth(col) + .map_or(line.len(), |(index, _)| index); + let byte_end = line[byte_start..] + .chars() + .next() + .map_or(line.len(), |ch| byte_start + ch.len_utf8()); let before = &line[..byte_start]; let cursor_char = if byte_start < line.len() { &line[byte_start..byte_end] @@ -96,27 +89,26 @@ fn render_request_editor(app: &App, frame: &mut Frame, area: Rect) { Span::raw(after), ])); } else if i == 0 { - let parts: Vec<&str> = line.splitn(3, ' ').collect(); - if parts.len() >= 2 { - lines.push(Line::from(vec![ - num_span, - Span::styled(parts[0], Style::default().fg(Color::Green).bold()), - Span::raw(" "), - Span::raw(parts[1..].join(" ")), - ])); - } else { - lines.push(Line::from(vec![num_span, Span::raw(line.as_str())])); - } - } else if let Some((key, value)) = line.split_once(':') { + if let Some((method, rest)) = line.split_once(' ') { lines.push(Line::from(vec![ num_span, - Span::styled(key, Style::default().fg(Color::Cyan)), - Span::raw(":"), - Span::raw(value), + Span::styled(method, Style::default().fg(Color::Green).bold()), + Span::raw(" "), + Span::raw(rest), ])); } else { lines.push(Line::from(vec![num_span, Span::raw(line.as_str())])); } + } else if let Some((key, value)) = line.split_once(':') { + lines.push(Line::from(vec![ + num_span, + Span::styled(key, Style::default().fg(Color::Cyan)), + Span::raw(":"), + Span::raw(value), + ])); + } else { + lines.push(Line::from(vec![num_span, Span::raw(line.as_str())])); + } } let title = if app.repeater.editing { @@ -167,11 +159,7 @@ fn render_diff(app: &App, frame: &mut Frame, area: Rect) { fn render_response(app: &App, frame: &mut Frame, area: Rect) { if app.repeater.pending { let msg = Paragraph::new("Sending request...") - .block( - Block::default() - .borders(Borders::ALL) - .title(" Response "), - ) + .block(Block::default().borders(Borders::ALL).title(" Response ")) .style(Style::default().fg(Color::Yellow)); frame.render_widget(msg, area); return; @@ -179,11 +167,7 @@ fn render_response(app: &App, frame: &mut Frame, area: Rect) { if let Some(ref error) = app.repeater.error { let msg = Paragraph::new(format!("Error: {}", error)) - .block( - Block::default() - .borders(Borders::ALL) - .title(" Response "), - ) + .block(Block::default().borders(Borders::ALL).title(" Response ")) .style(Style::default().fg(Color::Red)) .wrap(Wrap { trim: false }); frame.render_widget(msg, area); @@ -202,7 +186,10 @@ fn render_response(app: &App, frame: &mut Frame, area: Rect) { }; lines.push(Line::from(vec![ - Span::styled(resp.version.to_string(), Style::default().fg(Color::DarkGray)), + Span::styled( + resp.version.to_string(), + Style::default().fg(Color::DarkGray), + ), Span::raw(" "), Span::styled(resp.status.to_string(), status_style), Span::raw(" "), @@ -230,14 +217,22 @@ fn render_response(app: &App, frame: &mut Frame, area: Rect) { if !resp.body.is_empty() { lines.push(Line::raw("")); - let content_type = resp.headers.iter() + let content_type = resp + .headers + .iter() .find(|(k, _)| k.eq_ignore_ascii_case("content-type")) .map(|(_, v)| v.as_str()); - let proto_type = app.repeater.original.as_ref() + let proto_type = app + .repeater + .original + .as_ref() .filter(|r| r.is_grpc) .and_then(|r| crate::http::proto_schema::response_type(&r.uri)); lines.extend(body_view::body_lines_with_schema( - &resp.body, content_type, 500, proto_type.as_ref(), + &resp.body, + content_type, + 500, + proto_type.as_ref(), )); } @@ -254,11 +249,7 @@ fn render_response(app: &App, frame: &mut Frame, area: Rect) { }; let widget = Paragraph::new(content) - .block( - Block::default() - .borders(Borders::ALL) - .title(" Response "), - ) + .block(Block::default().borders(Borders::ALL).title(" Response ")) .wrap(Wrap { trim: false }) .scroll((app.repeater.resp_scroll, 0)); @@ -278,7 +269,11 @@ fn render_actions(app: &App, frame: &mut Frame, area: Rect) { Span::raw("navigate"), ]) } else { - let diff_label = if app.repeater.show_diff { "d:raw" } else { "d:diff" }; + let diff_label = if app.repeater.show_diff { + "d:raw" + } else { + "d:diff" + }; Line::from(vec![ if has_request { Span::styled(" Ctrl+Enter ", key_style()) @@ -307,11 +302,8 @@ fn render_actions(app: &App, frame: &mut Frame, area: Rect) { ]) }; - let widget = Paragraph::new(actions).block( - Block::default() - .borders(Borders::ALL) - .title(" Actions "), - ); + let widget = + Paragraph::new(actions).block(Block::default().borders(Borders::ALL).title(" Actions ")); frame.render_widget(widget, area); } @@ -368,13 +360,19 @@ fn render_macro_list(app: &App, frame: &mut Frame, area: Rect) { step.request.uri.clone() }; - lines.push(Line::from(vec![ - Span::styled(format!(" {:>2}. ", i + 1), Style::default().fg(Color::DarkGray)), - Span::styled(format!("[{}] ", state_icon), state_style), - Span::styled(format!("{:<7}", step.request.method), method_style), - Span::raw(format!("{:<36} ", path)), - Span::raw(status_str), - ]).style(row_style)); + lines.push( + Line::from(vec![ + Span::styled( + format!(" {:>2}. ", i + 1), + Style::default().fg(Color::DarkGray), + ), + Span::styled(format!("[{}] ", state_icon), state_style), + Span::styled(format!("{:<7}", step.request.method), method_style), + Span::raw(format!("{:<36} ", path)), + Span::raw(status_str), + ]) + .style(row_style), + ); } let title = if app.macros.running { @@ -409,7 +407,11 @@ fn render_macro_detail(app: &App, frame: &mut Frame, area: Rect) { "Select a step to view details", Style::default().fg(Color::DarkGray), )) - .block(Block::default().borders(Borders::ALL).title(" Step Detail ")); + .block( + Block::default() + .borders(Borders::ALL) + .title(" Step Detail "), + ); frame.render_widget(widget, area); return; } @@ -418,7 +420,10 @@ fn render_macro_detail(app: &App, frame: &mut Frame, area: Rect) { let mut lines: Vec = Vec::new(); lines.push(Line::from(vec![ - Span::styled(&step.request.method, Style::default().fg(Color::Green).bold()), + Span::styled( + &step.request.method, + Style::default().fg(Color::Green).bold(), + ), Span::raw(" "), Span::raw(&step.request.uri), ])); @@ -435,9 +440,15 @@ fn render_macro_detail(app: &App, frame: &mut Frame, area: Rect) { lines.push(Line::from(vec![ Span::styled(format!("{} {}", resp.status, resp.reason), status_style), Span::raw(" "), - Span::styled(format!("{:.0?}", resp.duration), Style::default().fg(Color::DarkGray)), + Span::styled( + format!("{:.0?}", resp.duration), + Style::default().fg(Color::DarkGray), + ), Span::raw(" "), - Span::styled(format_size(resp.body.len()), Style::default().fg(Color::DarkGray)), + Span::styled( + format_size(resp.body.len()), + Style::default().fg(Color::DarkGray), + ), ])); if let Some(timing) = &resp.timing { @@ -449,7 +460,9 @@ fn render_macro_detail(app: &App, frame: &mut Frame, area: Rect) { if !resp.body.is_empty() { lines.push(Line::raw("")); - let ct = resp.headers.iter() + let ct = resp + .headers + .iter() .find(|(k, _)| k.eq_ignore_ascii_case("content-type")) .map(|(_, v)| v.as_str()); let proto_type = if step.request.is_grpc { @@ -458,7 +471,10 @@ fn render_macro_detail(app: &App, frame: &mut Frame, area: Rect) { None }; lines.extend(body_view::body_lines_with_schema( - &resp.body, ct, 100, proto_type.as_ref(), + &resp.body, + ct, + 100, + proto_type.as_ref(), )); } @@ -474,7 +490,11 @@ fn render_macro_detail(app: &App, frame: &mut Frame, area: Rect) { } let widget = Paragraph::new(Text::from(lines)) - .block(Block::default().borders(Borders::ALL).title(" Step Detail ")) + .block( + Block::default() + .borders(Borders::ALL) + .title(" Step Detail "), + ) .wrap(Wrap { trim: false }) .scroll((app.repeater.resp_scroll, 0)); diff --git a/src/tui/tabs/rules_tab.rs b/src/tui/tabs/rules_tab.rs index fce9ff9..c564669 100644 --- a/src/tui/tabs/rules_tab.rs +++ b/src/tui/tabs/rules_tab.rs @@ -1,17 +1,13 @@ +use ratatui::Frame; use ratatui::layout::{Constraint, Layout, Rect}; use ratatui::style::{Color, Modifier, Style}; use ratatui::text::{Line, Span}; use ratatui::widgets::{Block, Borders, Paragraph, Row, Table, Wrap}; -use ratatui::Frame; use crate::app::{App, RuleField}; pub fn render(app: &App, frame: &mut Frame, area: Rect) { - let chunks = Layout::vertical([ - Constraint::Min(0), - Constraint::Length(3), - ]) - .split(area); + let chunks = Layout::vertical([Constraint::Min(0), Constraint::Length(3)]).split(area); let rules = app.rules.read(); @@ -20,14 +16,15 @@ pub fn render(app: &App, frame: &mut Frame, area: Rect) { "No rules configured. Press 'a' to add a rule.", Style::default().fg(Color::DarkGray), )) - .block(Block::default().borders(Borders::ALL).title(" Match & Replace Rules ")); + .block( + Block::default() + .borders(Borders::ALL) + .title(" Match & Replace Rules "), + ); frame.render_widget(msg, chunks[0]); } else { - let detail_split = Layout::vertical([ - Constraint::Min(0), - Constraint::Length(8), - ]) - .split(chunks[0]); + let detail_split = + Layout::vertical([Constraint::Min(0), Constraint::Length(8)]).split(chunks[0]); render_table(app, &rules, frame, detail_split[0]); render_detail(app, &rules, frame, detail_split[1]); @@ -37,12 +34,11 @@ pub fn render(app: &App, frame: &mut Frame, area: Rect) { } fn render_table(app: &App, rules: &[crate::rules::Rule], frame: &mut Frame, area: Rect) { - let header = Row::new(["", "Name", "Target", "Scope", "Match", "Replace", "Regex"]) - .style( - Style::default() - .fg(Color::Yellow) - .add_modifier(Modifier::BOLD), - ); + let header = Row::new(["", "Name", "Target", "Scope", "Match", "Replace", "Regex"]).style( + Style::default() + .fg(Color::Yellow) + .add_modifier(Modifier::BOLD), + ); let rows: Vec = rules .iter() @@ -86,9 +82,11 @@ fn render_table(app: &App, rules: &[crate::rules::Rule], frame: &mut Frame, area Constraint::Length(5), ]; - let table = Table::new(rows, widths) - .header(header) - .block(Block::default().borders(Borders::ALL).title(" Match & Replace Rules ")); + let table = Table::new(rows, widths).header(header).block( + Block::default() + .borders(Borders::ALL) + .title(" Match & Replace Rules "), + ); frame.render_widget(table, area); } @@ -203,11 +201,8 @@ fn render_actions(app: &App, frame: &mut Frame, area: Rect) { Span::raw(":export"), ]); - let widget = Paragraph::new(line).block( - Block::default() - .borders(Borders::ALL) - .title(" Actions "), - ); + let widget = + Paragraph::new(line).block(Block::default().borders(Borders::ALL).title(" Actions ")); frame.render_widget(widget, area); } @@ -233,6 +228,10 @@ fn truncate(s: &str, max: usize) -> String { if s.chars().count() <= max { return s.to_string(); } - let end = s.char_indices().nth(max - 3).map(|(i, _)| i).unwrap_or(s.len()); + let end = s + .char_indices() + .nth(max - 3) + .map(|(i, _)| i) + .unwrap_or(s.len()); format!("{}...", &s[..end]) } diff --git a/src/tui/tabs/tools_tab.rs b/src/tui/tabs/tools_tab.rs index b3b7941..964a25d 100644 --- a/src/tui/tabs/tools_tab.rs +++ b/src/tui/tabs/tools_tab.rs @@ -1,25 +1,18 @@ +use ratatui::Frame; use ratatui::layout::{Constraint, Layout, Rect}; use ratatui::style::{Color, Modifier, Style}; use ratatui::text::{Line, Span, Text}; use ratatui::widgets::{Block, Borders, Paragraph, Wrap}; -use ratatui::Frame; use crate::app::{App, ToolsMode}; pub fn render(app: &App, frame: &mut Frame, area: Rect) { - let chunks = Layout::vertical([ - Constraint::Length(3), - Constraint::Min(0), - ]) - .split(area); + let chunks = Layout::vertical([Constraint::Length(3), Constraint::Min(0)]).split(area); render_mode_selector(app, frame, chunks[0]); - let panes = Layout::horizontal([ - Constraint::Percentage(50), - Constraint::Percentage(50), - ]) - .split(chunks[1]); + let panes = Layout::horizontal([Constraint::Percentage(50), Constraint::Percentage(50)]) + .split(chunks[1]); render_input(app, frame, panes[0]); render_output(app, frame, panes[1]); @@ -73,7 +66,9 @@ fn render_input(app: &App, frame: &mut Frame, area: Rect) { let mut lines: Vec = app.tools.editor.render_lines(app.tools.editing); - if lines.is_empty() || (lines.len() == 1 && app.tools.editor.lines[0].is_empty() && !app.tools.editing) { + if lines.is_empty() + || (lines.len() == 1 && app.tools.editor.lines[0].is_empty() && !app.tools.editing) + { lines = vec![Line::styled( "Press 'e' to edit input", Style::default().fg(Color::DarkGray), @@ -105,11 +100,7 @@ fn render_output(app: &App, frame: &mut Frame, area: Rect) { }; let widget = Paragraph::new(Text::from(lines)) - .block( - Block::default() - .borders(Borders::ALL) - .title(" Output "), - ) + .block(Block::default().borders(Borders::ALL).title(" Output ")) .wrap(Wrap { trim: false }) .scroll((app.tools.scroll, 0)); diff --git a/src/tui/terminal.rs b/src/tui/terminal.rs index 418e4f9..df989af 100644 --- a/src/tui/terminal.rs +++ b/src/tui/terminal.rs @@ -4,8 +4,8 @@ use crossterm::{ execute, terminal::{self, EnterAlternateScreen, LeaveAlternateScreen}, }; -use ratatui::backend::CrosstermBackend; use ratatui::Terminal; +use ratatui::backend::CrosstermBackend; pub type Tui = Terminal>; diff --git a/src/tui/widgets/body_view.rs b/src/tui/widgets/body_view.rs index 96ecf07..bef9961 100644 --- a/src/tui/widgets/body_view.rs +++ b/src/tui/widgets/body_view.rs @@ -5,8 +5,8 @@ use ratatui::text::{Line, Span}; use prost_reflect::MessageDescriptor; -use crate::http::protobuf::{self, ProtoField, ProtoValue}; use crate::http::proto_schema; +use crate::http::protobuf::{self, ProtoField, ProtoValue}; use crate::tui::widgets::hex_view; thread_local! { @@ -80,9 +80,11 @@ fn render_json<'a>(text: &str, max_lines: usize) -> Vec> { let pretty = JSON_CACHE.with(|cache| { let cached = cache.borrow(); if let Some((ptr, len, ref s)) = *cached - && ptr == key.0 && len == key.1 { - return Some(s.clone()); - } + && ptr == key.0 + && len == key.1 + { + return Some(s.clone()); + } None }); @@ -118,17 +120,21 @@ fn colorize_json_line<'a>(line: &str) -> Line<'a> { let trimmed = line.trim_start(); if trimmed.starts_with('"') - && let Some(colon_pos) = trimmed.find("\": ") { - let indent = &line[..line.len() - trimmed.len()]; - let key = &trimmed[..colon_pos + 1]; - let rest = &trimmed[colon_pos + 1..]; - - return Line::from(vec![ - Span::raw(indent.to_string()), - Span::styled(key.to_string(), Style::default().fg(Color::Cyan)), - Span::styled(rest.to_string(), value_style(rest.trim_start().trim_start_matches(": "))), - ]); - } + && let Some(colon_pos) = trimmed.find("\": ") + { + let indent = &line[..line.len() - trimmed.len()]; + let key = &trimmed[..colon_pos + 1]; + let rest = &trimmed[colon_pos + 1..]; + + return Line::from(vec![ + Span::raw(indent.to_string()), + Span::styled(key.to_string(), Style::default().fg(Color::Cyan)), + Span::styled( + rest.to_string(), + value_style(rest.trim_start().trim_start_matches(": ")), + ), + ]); + } if trimmed.starts_with('"') { return Line::styled(line.to_string(), Style::default().fg(Color::Green)); @@ -472,7 +478,15 @@ fn render_proto_fields<'a>( max_lines: usize, lines: &mut Vec>, ) { - let indent = " ".repeat(depth); + // Borrow common indentation depths from static storage to avoid allocating + // and cloning a new String for every rendered protobuf field. + const SPACES: &str = " "; + let width = depth.saturating_mul(2); + let indent = if width <= SPACES.len() { + std::borrow::Cow::Borrowed(&SPACES[..width]) + } else { + std::borrow::Cow::Owned(" ".repeat(width)) + }; for field in fields { if lines.len() >= max_lines { @@ -487,10 +501,7 @@ fn render_proto_fields<'a>( ProtoValue::Varint(v) => { lines.push(Line::from(vec![ Span::raw(indent.clone()), - Span::styled( - field.number.to_string(), - Style::default().fg(Color::Cyan), - ), + Span::styled(field.number.to_string(), Style::default().fg(Color::Cyan)), Span::styled(": ", Style::default().fg(Color::DarkGray)), Span::styled(v.to_string(), Style::default().fg(Color::Magenta)), ])); @@ -500,10 +511,7 @@ fn render_proto_fields<'a>( let is_double = display.contains('.'); let mut spans = vec![ Span::raw(indent.clone()), - Span::styled( - field.number.to_string(), - Style::default().fg(Color::Cyan), - ), + Span::styled(field.number.to_string(), Style::default().fg(Color::Cyan)), Span::styled(": ", Style::default().fg(Color::DarkGray)), Span::styled(display, Style::default().fg(Color::Magenta)), ]; @@ -520,10 +528,7 @@ fn render_proto_fields<'a>( let is_float = display.contains('.'); let mut spans = vec![ Span::raw(indent.clone()), - Span::styled( - field.number.to_string(), - Style::default().fg(Color::Cyan), - ), + Span::styled(field.number.to_string(), Style::default().fg(Color::Cyan)), Span::styled(": ", Style::default().fg(Color::DarkGray)), Span::styled(display, Style::default().fg(Color::Magenta)), ]; @@ -538,10 +543,7 @@ fn render_proto_fields<'a>( ProtoValue::String(s) => { lines.push(Line::from(vec![ Span::raw(indent.clone()), - Span::styled( - field.number.to_string(), - Style::default().fg(Color::Cyan), - ), + Span::styled(field.number.to_string(), Style::default().fg(Color::Cyan)), Span::styled(": ", Style::default().fg(Color::DarkGray)), Span::styled( format!("\"{}\"", truncate_string(s, 200)), @@ -552,10 +554,7 @@ fn render_proto_fields<'a>( ProtoValue::Message(sub_fields) => { lines.push(Line::from(vec![ Span::raw(indent.clone()), - Span::styled( - field.number.to_string(), - Style::default().fg(Color::Cyan), - ), + Span::styled(field.number.to_string(), Style::default().fg(Color::Cyan)), Span::styled(" {", Style::default().fg(Color::DarkGray)), ])); render_proto_fields(sub_fields, depth + 1, max_lines, lines); @@ -569,10 +568,7 @@ fn render_proto_fields<'a>( ProtoValue::Bytes(data) => { lines.push(Line::from(vec![ Span::raw(indent.clone()), - Span::styled( - field.number.to_string(), - Style::default().fg(Color::Cyan), - ), + Span::styled(field.number.to_string(), Style::default().fg(Color::Cyan)), Span::styled(": ", Style::default().fg(Color::DarkGray)), Span::styled( format!("<{} bytes>", data.len()), @@ -588,7 +584,11 @@ fn truncate_string(s: &str, max_len: usize) -> String { if s.chars().count() <= max_len { s.to_string() } else { - let end = s.char_indices().nth(max_len).map(|(i, _)| i).unwrap_or(s.len()); + let end = s + .char_indices() + .nth(max_len) + .map(|(i, _)| i) + .unwrap_or(s.len()); format!("{}...", &s[..end]) } } diff --git a/src/tui/widgets/diff_view.rs b/src/tui/widgets/diff_view.rs index c2e7897..9c61dae 100644 --- a/src/tui/widgets/diff_view.rs +++ b/src/tui/widgets/diff_view.rs @@ -3,10 +3,11 @@ use ratatui::text::{Line, Span}; use similar::{ChangeTag, TextDiff}; pub fn diff_lines<'a>(original: &[String], modified: &[String]) -> Vec> { - let old_text = original.join("\n"); - let new_text = modified.join("\n"); - - let diff = TextDiff::from_lines(&old_text, &new_text); + // Diff the existing line slices directly. Joining both inputs copied every + // request byte on each frame while the diff view was visible. + let original: Vec<&str> = original.iter().map(String::as_str).collect(); + let modified: Vec<&str> = modified.iter().map(String::as_str).collect(); + let diff = TextDiff::from_slices(&original, &modified); let mut lines = Vec::new(); for change in diff.iter_all_changes() { @@ -16,12 +17,9 @@ pub fn diff_lines<'a>(original: &[String], modified: &[String]) -> Vec> ChangeTag::Equal => (" ", Style::default().fg(Color::DarkGray)), }; - let text = change.as_str().unwrap_or("").trim_end_matches('\n'); + let text = change.as_str().unwrap_or(""); lines.push(Line::from(vec![ - Span::styled( - prefix, - style.add_modifier(Modifier::BOLD), - ), + Span::styled(prefix, style.add_modifier(Modifier::BOLD)), Span::styled(format!(" {}", text), style), ])); } @@ -55,9 +53,9 @@ mod tests { let old = vec!["line1".into()]; let new = vec!["line1".into(), "line2".into()]; let result = diff_lines(&old, &new); - let has_addition = result.iter().any(|l| { - l.spans.first().map(|s| s.content.as_ref()) == Some("+") - }); + let has_addition = result + .iter() + .any(|l| l.spans.first().map(|s| s.content.as_ref()) == Some("+")); assert!(has_addition); } @@ -66,9 +64,9 @@ mod tests { let old = vec!["line1".into(), "line2".into()]; let new = vec!["line1".into()]; let result = diff_lines(&old, &new); - let has_deletion = result.iter().any(|l| { - l.spans.first().map(|s| s.content.as_ref()) == Some("-") - }); + let has_deletion = result + .iter() + .any(|l| l.spans.first().map(|s| s.content.as_ref()) == Some("-")); assert!(has_deletion); } } diff --git a/src/tui/widgets/hex_view.rs b/src/tui/widgets/hex_view.rs index 3230781..d11f0d0 100644 --- a/src/tui/widgets/hex_view.rs +++ b/src/tui/widgets/hex_view.rs @@ -36,7 +36,13 @@ pub fn hex_lines(data: &[u8], max_lines: usize) -> Vec> { spans.push(Span::raw(" |")); let ascii: String = chunk .iter() - .map(|&b| if b.is_ascii_graphic() || b == b' ' { b as char } else { '.' }) + .map(|&b| { + if b.is_ascii_graphic() || b == b' ' { + b as char + } else { + '.' + } + }) .collect(); spans.push(Span::raw(ascii)); spans.push(Span::raw("|")); diff --git a/src/tui/widgets/logo.rs b/src/tui/widgets/logo.rs index 4974182..7c1c456 100644 --- a/src/tui/widgets/logo.rs +++ b/src/tui/widgets/logo.rs @@ -1,25 +1,43 @@ +use ratatui::Frame; use ratatui::layout::Rect; use ratatui::style::{Color, Style}; use ratatui::text::{Line, Span}; -use ratatui::Frame; +use std::sync::LazyLock; const LOGO: &str = include_str!("../../../assets/logo.txt"); +const VERSION_TEXT: &str = concat!("v", env!("CARGO_PKG_VERSION")); + +struct LogoGeometry { + lines: Vec<(&'static str, u16)>, + max_width: u16, + total_height: u16, +} + +static LOGO_GEOMETRY: LazyLock = LazyLock::new(|| { + let lines: Vec<_> = LOGO + .lines() + .map(|line| (line, line.chars().count() as u16)) + .collect(); + let max_width = lines.iter().map(|(_, width)| *width).max().unwrap_or(0); + let total_height = lines.len() as u16 + 2; + LogoGeometry { + lines, + max_width, + total_height, + } +}); pub fn render(frame: &mut Frame, area: Rect) { - let logo_lines: Vec<&str> = LOGO.lines().collect(); - let version_text = format!("v{}", env!("CARGO_PKG_VERSION")); - let total_height = logo_lines.len() as u16 + 2; - let max_width = logo_lines.iter().map(|l| l.chars().count()).max().unwrap_or(0) as u16; + let geometry = &*LOGO_GEOMETRY; - if area.height < total_height + 2 || area.width < max_width { + if area.height < geometry.total_height + 2 || area.width < geometry.max_width { return; } - let y_offset = (area.height.saturating_sub(total_height)) / 2; + let y_offset = (area.height.saturating_sub(geometry.total_height)) / 2; - for (i, line) in logo_lines.iter().enumerate() { - let line_width = line.chars().count() as u16; - let x_offset = (area.width.saturating_sub(line_width)) / 2; + for (i, (line, line_width)) in geometry.lines.iter().enumerate() { + let x_offset = (area.width.saturating_sub(*line_width)) / 2; let y = area.y + y_offset + i as u16; if y >= area.y + area.height { @@ -37,15 +55,25 @@ pub fn render(frame: &mut Frame, area: Rect) { let span = Span::styled(*line, style); let buf = frame.buffer_mut(); - buf.set_line(area.x + x_offset, y, &Line::from(span), area.width.saturating_sub(x_offset)); + buf.set_line( + area.x + x_offset, + y, + &Line::from(span), + area.width.saturating_sub(x_offset), + ); } - let version_y = area.y + y_offset + logo_lines.len() as u16 + 1; + let version_y = area.y + y_offset + geometry.lines.len() as u16 + 1; if version_y < area.y + area.height { - let version_width = version_text.chars().count() as u16; + let version_width = VERSION_TEXT.len() as u16; let version_x = (area.width.saturating_sub(version_width)) / 2; - let version_span = Span::styled(version_text, Style::default().fg(Color::DarkGray)); + let version_span = Span::styled(VERSION_TEXT, Style::default().fg(Color::DarkGray)); let buf = frame.buffer_mut(); - buf.set_line(area.x + version_x, version_y, &Line::from(version_span), area.width.saturating_sub(version_x)); + buf.set_line( + area.x + version_x, + version_y, + &Line::from(version_span), + area.width.saturating_sub(version_x), + ); } } diff --git a/src/tui/widgets/mod.rs b/src/tui/widgets/mod.rs index 47a6b77..042bc40 100644 --- a/src/tui/widgets/mod.rs +++ b/src/tui/widgets/mod.rs @@ -44,10 +44,7 @@ pub fn header_lines<'a>(headers: &'a [(String, String)]) -> Vec> { pub fn trailer_lines<'a>(trailers: &'a [(String, String)]) -> Vec> { let mut lines = vec![ Line::raw(""), - Line::styled( - "──── Trailers ────", - Style::default().fg(Color::DarkGray), - ), + Line::styled("──── Trailers ────", Style::default().fg(Color::DarkGray)), ]; for (key, value) in trailers { let value_style = if key == "grpc-status" { diff --git a/src/tui/widgets/timing_view.rs b/src/tui/widgets/timing_view.rs index 4b01d2e..462699a 100644 --- a/src/tui/widgets/timing_view.rs +++ b/src/tui/widgets/timing_view.rs @@ -31,7 +31,10 @@ pub fn timing_lines(timing: &TimingData, total: Duration) -> Vec> if let Some(dur) = duration { let ms = dur.as_secs_f64() * 1000.0; let fraction = ms / total_ms; - let filled = (fraction * bar_width as f64).round().max(1.0).min(bar_width as f64) as usize; + let filled = (fraction * bar_width as f64) + .round() + .max(1.0) + .min(bar_width as f64) as usize; let bar: String = "\u{2588}".repeat(filled); let pad: String = "\u{2591}".repeat(bar_width.saturating_sub(filled));