diff --git a/Cargo.lock b/Cargo.lock index 88eb207..59c1f57 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1346,6 +1346,7 @@ dependencies = [ "serde", "serde_json", "tokio", + "tokio-util", "toml", "tracing", "tracing-log", diff --git a/Cargo.toml b/Cargo.toml index 3a3d947..355a50b 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -20,6 +20,7 @@ rlimit = "=0.10.2" serde = "=1.0.219" serde_json = "=1.0.140" tokio = { version = "=1.46.1", features = ["full"] } +tokio-util = "=0.7.15" toml = "=0.9.2" tracing = "=0.1.41" tracing-log = "=0.2.0" diff --git a/src/checker.rs b/src/checker.rs index c592187..513ea3d 100644 --- a/src/checker.rs +++ b/src/checker.rs @@ -1,6 +1,6 @@ use std::{collections::HashSet, sync::Arc}; -use color_eyre::eyre::{OptionExt as _, WrapErr as _}; +use color_eyre::eyre::WrapErr as _; #[cfg(feature = "tui")] use crate::event::{AppEvent, Event}; @@ -8,57 +8,65 @@ use crate::{config::Config, proxy::Proxy, utils::pretty_error}; pub async fn check_all( config: Arc, - proxies: HashSet, + proxies: Arc>>, + token: tokio_util::sync::CancellationToken, #[cfg(feature = "tui")] tx: tokio::sync::mpsc::UnboundedSender, -) -> color_eyre::Result> { +) -> color_eyre::Result<()> { let workers_count = - config.checking.max_concurrent_checks.min(proxies.len()); + config.checking.max_concurrent_checks.min(proxies.lock().await.len()); if workers_count == 0 { - return Ok(Vec::new()); + return Ok(()); } - let checked_proxies = Arc::new(tokio::sync::Mutex::new(Vec::new())); - let queue = Arc::new(tokio::sync::Mutex::new( - proxies.into_iter().collect::>(), + proxies.lock().await.drain().collect::>(), )); let mut join_set = tokio::task::JoinSet::>::new(); for _ in 0..workers_count { let queue = Arc::clone(&queue); let config = Arc::clone(&config); - let new_storage = Arc::clone(&checked_proxies); + let proxies = Arc::clone(&proxies); + let token = token.clone(); #[cfg(feature = "tui")] let tx = tx.clone(); join_set.spawn(async move { - loop { - let Some(mut proxy) = queue.lock().await.pop() else { - break Ok(()); - }; - let check_result = proxy.check(&config).await; - #[cfg(feature = "tui")] - drop(tx.send(Event::App(AppEvent::ProxyChecked( - proxy.protocol.clone(), - )))); - match check_result { - Ok(()) => { + tokio::select! { + biased; + res = async move { + loop { + let Some(mut proxy) = queue.lock().await.pop() else { + break Ok(()); + }; + let check_result = proxy.check(&config).await; #[cfg(feature = "tui")] - drop(tx.send(Event::App(AppEvent::ProxyWorking( + drop(tx.send(Event::App(AppEvent::ProxyChecked( proxy.protocol.clone(), )))); - new_storage.lock().await.push(proxy); + match check_result { + Ok(()) => { + #[cfg(feature = "tui")] + drop(tx.send(Event::App(AppEvent::ProxyWorking( + proxy.protocol.clone(), + )))); + proxies.lock().await.insert(proxy); + } + Err(e) + if tracing::event_enabled!( + tracing::Level::DEBUG + ) => + { + tracing::debug!( + "{} | {}", + proxy.as_str(true), + pretty_error(&e) + ); + } + Err(_) => {} + } } - Err(e) - if tracing::event_enabled!(tracing::Level::DEBUG) => - { - tracing::debug!( - "{} | {}", - proxy.as_str(true), - pretty_error(&e) - ); - } - Err(_) => {} - } + } => res, + () = token.cancelled() => Ok(()), } }); } @@ -78,7 +86,5 @@ pub async fn check_all( } } - Ok(Arc::into_inner(checked_proxies) - .ok_or_eyre("failed to unwrap Arc")? - .into_inner()) + Ok(()) } diff --git a/src/main.rs b/src/main.rs index 69d1629..71d5d3d 100644 --- a/src/main.rs +++ b/src/main.rs @@ -16,6 +16,7 @@ clippy::else_if_without_else, clippy::float_arithmetic, clippy::implicit_return, + clippy::integer_division_remainder_used, clippy::iter_over_hash_type, clippy::min_ident_chars, clippy::missing_docs_in_private_items, @@ -48,8 +49,7 @@ mod scraper; #[cfg(feature = "tui")] mod tui; mod utils; - -use std::sync::Arc; +use std::{collections::HashSet, path::Path, sync::Arc}; use color_eyre::eyre::WrapErr as _; use tracing_subscriber::{ @@ -90,54 +90,50 @@ fn create_logging_filter( } } -fn spawn_ip_database_tasks( - config: &Arc, - http_client: &reqwest::Client, - #[cfg(feature = "tui")] tx: &tokio::sync::mpsc::UnboundedSender< +async fn download_output_dependencies( + config: &config::Config, + http_client: reqwest::Client, + token: tokio_util::sync::CancellationToken, + #[cfg(feature = "tui")] tx: tokio::sync::mpsc::UnboundedSender< event::Event, >, -) -> tokio::task::JoinSet> { +) -> color_eyre::Result<()> { let mut output_dependencies_tasks = tokio::task::JoinSet::new(); if config.asn_enabled() { let http_client = http_client.clone(); + let token = token.clone(); #[cfg(feature = "tui")] let tx = tx.clone(); output_dependencies_tasks.spawn(async move { - ipdb::DbType::Asn - .download( + tokio::select! { + biased; + res = ipdb::DbType::Asn.download( http_client, #[cfg(feature = "tui")] tx, - ) - .await + ) => res, + () = token.cancelled() => Ok(()) + } }); } if config.geolocation_enabled() { - let http_client = http_client.clone(); - #[cfg(feature = "tui")] - let tx = tx.clone(); - output_dependencies_tasks.spawn(async move { - ipdb::DbType::Geo - .download( + tokio::select! { + biased; + res = ipdb::DbType::Geo.download( http_client, #[cfg(feature = "tui")] tx, - ) - .await + ) => res, + () = token.cancelled() => Ok(()) + } }); } - output_dependencies_tasks -} - -async fn wait_for_ip_database_tasks( - mut tasks: tokio::task::JoinSet>, -) -> color_eyre::Result<()> { - while let Some(task) = tasks.join_next().await { + while let Some(task) = output_dependencies_tasks.join_next().await { task.wrap_err("output dependencies task panicked or was cancelled")??; } Ok(()) @@ -145,63 +141,64 @@ async fn wait_for_ip_database_tasks( async fn process_proxies( config: Arc, - proxies: std::collections::HashSet, - #[cfg(feature = "tui")] tx: &tokio::sync::mpsc::UnboundedSender< + proxies: Arc>>, + token: tokio_util::sync::CancellationToken, + #[cfg(feature = "tui")] tx: tokio::sync::mpsc::UnboundedSender< event::Event, >, -) -> color_eyre::Result> { +) -> color_eyre::Result<()> { if config.checking.check_url.is_empty() { - Ok(proxies.into_iter().collect()) - } else { - #[cfg(not(feature = "tui"))] - tracing::info!("Started checking {} proxies", proxies.len()); - - checker::check_all( - config, - proxies, - #[cfg(feature = "tui")] - tx.clone(), - ) - .await - .wrap_err("failed to check proxies") + return Ok(()); } + #[cfg(not(feature = "tui"))] + tracing::info!("Started checking {} proxies", proxies.lock().await.len()); + + checker::check_all( + config, + proxies, + token, + #[cfg(feature = "tui")] + tx.clone(), + ) + .await } async fn main_task( config: Arc, + token: tokio_util::sync::CancellationToken, #[cfg(feature = "tui")] tx: tokio::sync::mpsc::UnboundedSender< event::Event, >, ) -> color_eyre::Result<()> { let http_client = create_reqwest_client() .wrap_err("failed to create reqwest HTTP client")?; + let proxies = Arc::new(tokio::sync::Mutex::new(HashSet::new())); - let ip_db_tasks = spawn_ip_database_tasks( - &config, - &http_client, - #[cfg(feature = "tui")] - &tx, - ); + tokio::try_join!( + download_output_dependencies( + &config, + http_client.clone(), + token.clone(), + #[cfg(feature = "tui")] + tx.clone(), + ), + scraper::scrape_all( + Arc::clone(&config), + http_client, + Arc::clone(&proxies), + token.clone(), + #[cfg(feature = "tui")] + tx.clone(), + ), + )?; - let proxies = scraper::scrape_all( + process_proxies( Arc::clone(&config), - http_client.clone(), + Arc::clone(&proxies), + token, #[cfg(feature = "tui")] tx.clone(), ) - .await - .wrap_err("failed to scrape proxies")?; - - drop(http_client); - - wait_for_ip_database_tasks(ip_db_tasks).await?; - - let proxies = process_proxies( - Arc::clone(&config), - proxies, - #[cfg(feature = "tui")] - &tx, - ) .await?; output::save_proxies(config, proxies) @@ -233,15 +230,16 @@ async fn run_with_tui( let terminal_guard = tui::RatatuiRestoreGuard; let (tx, rx) = tokio::sync::mpsc::unbounded_channel(); - let main_task = tokio::task::spawn(main_task(config, tx.clone())); - let main_task_handle = main_task.abort_handle(); + let token = tokio_util::sync::CancellationToken::new(); - tokio::try_join!(async move { main_task.await? }, async move { - let result = tui::run(terminal, tx, rx).await; - drop(terminal_guard); - main_task_handle.abort(); - result - })?; + tokio::try_join!( + main_task(config, token.clone(), tx.clone()), + async move { + let result = tui::run(terminal, token, tx, rx).await; + drop(terminal_guard); + result + } + )?; Ok(()) } @@ -256,16 +254,29 @@ async fn run_without_tui( .with(tracing_subscriber::fmt::layer()) .init(); - main_task(config).await + let token = tokio_util::sync::CancellationToken::new(); + + tokio::try_join!(main_task(config, token.clone()), async move { + match tokio::signal::ctrl_c().await { + Ok(()) => { + tracing::info!("Received Ctrl+C, exiting..."); + token.cancel(); + } + Err(e) => { + tracing::error!("Failed to listen for Ctrl+C: {}", e); + } + } + Ok(()) + })?; + + Ok(()) } async fn load_config() -> color_eyre::Result> { let raw_config_path = raw_config::get_config_path(); - let raw_config = raw_config::read_config(std::path::Path::new( - &raw_config_path, - )) - .await - .wrap_err_with(move || format!("failed to read {raw_config_path}"))?; + let raw_config = raw_config::read_config(Path::new(&raw_config_path)) + .await + .wrap_err_with(move || format!("failed to read {raw_config_path}"))?; let config = config::Config::from_raw_config(raw_config) .await diff --git a/src/output.rs b/src/output.rs index 6d9b314..b57cd45 100644 --- a/src/output.rs +++ b/src/output.rs @@ -1,11 +1,11 @@ use std::{ - collections::HashMap, + collections::{HashMap, HashSet}, io, iter, net::{IpAddr, Ipv4Addr}, sync::Arc, }; -use color_eyre::eyre::WrapErr as _; +use color_eyre::eyre::{OptionExt as _, WrapErr as _}; use crate::{ config::Config, @@ -60,8 +60,14 @@ fn group_proxies<'a>( #[expect(clippy::too_many_lines)] pub async fn save_proxies( config: Arc, - mut proxies: Vec, + proxies: Arc>>, ) -> color_eyre::Result<()> { + let mut proxies: Vec<_> = Arc::into_inner(proxies) + .ok_or_eyre("failed to unwrap Arc")? + .into_inner() + .into_iter() + .filter(|p| config.checking.check_url.is_empty() || p.is_checked()) + .collect(); if config.output.sort_by_speed { proxies.sort_by_key(sort_by_timeout); } else { diff --git a/src/proxy.rs b/src/proxy.rs index a6c510e..a242104 100644 --- a/src/proxy.rs +++ b/src/proxy.rs @@ -82,6 +82,10 @@ impl TryFrom<&mut Proxy> for reqwest::Proxy { } impl Proxy { + pub const fn is_checked(&self) -> bool { + self.timeout.is_some() + } + pub async fn check(&mut self, config: &Config) -> color_eyre::Result<()> { let client = reqwest::Client::builder() .user_agent(USER_AGENT) diff --git a/src/scraper.rs b/src/scraper.rs index 77fdcdf..b820005 100644 --- a/src/scraper.rs +++ b/src/scraper.rs @@ -119,10 +119,10 @@ async fn scrape_one( pub async fn scrape_all( config: Arc, http_client: reqwest::Client, + proxies: Arc>>, + token: tokio_util::sync::CancellationToken, #[cfg(feature = "tui")] tx: tokio::sync::mpsc::UnboundedSender, -) -> color_eyre::Result> { - let proxies = Arc::new(tokio::sync::Mutex::new(HashSet::new())); - +) -> color_eyre::Result<()> { let mut join_set = tokio::task::JoinSet::new(); for (proto, sources) in config.scraping.sources.clone() { #[cfg(feature = "tui")] @@ -135,19 +135,23 @@ pub async fn scrape_all( let http_client = http_client.clone(); let proto = proto.clone(); let proxies = Arc::clone(&proxies); + let token = token.clone(); #[cfg(feature = "tui")] let tx = tx.clone(); join_set.spawn(async move { - scrape_one( - config, - http_client, - proto, - proxies, - &source, - #[cfg(feature = "tui")] - tx, - ) - .await + tokio::select! { + biased; + res = scrape_one( + config, + http_client, + proto, + proxies, + &source, + #[cfg(feature = "tui")] + tx, + ) => res, + () = token.cancelled() => Ok(()) + } }); } } @@ -157,7 +161,5 @@ pub async fn scrape_all( .wrap_err("proxy scraping task failed")?; } - Ok(Arc::into_inner(proxies) - .ok_or_eyre("failed to unwrap Arc")? - .into_inner()) + Ok(()) } diff --git a/src/tui.rs b/src/tui.rs index 7fcf418..fb70811 100644 --- a/src/tui.rs +++ b/src/tui.rs @@ -38,6 +38,7 @@ impl Drop for RatatuiRestoreGuard { pub async fn run( mut terminal: ratatui::DefaultTerminal, + token: tokio_util::sync::CancellationToken, tx: tokio::sync::mpsc::UnboundedSender, mut rx: tokio::sync::mpsc::UnboundedReceiver, ) -> color_eyre::Result<()> { @@ -49,7 +50,8 @@ pub async fn run( let logger_state = TuiWidgetState::default(); while !matches!(app_state.mode, AppMode::Quit) { if let Some(event) = rx.recv().await { - if handle_event(event, &mut app_state, &logger_state).await { + if handle_event(event, &mut app_state, &token, &logger_state).await + { terminal .draw(|frame| draw(frame, &app_state, &logger_state)) .wrap_err("failed to draw tui")?; @@ -70,6 +72,7 @@ pub async fn run( pub enum AppMode { #[default] Running, + Done, Quit, } @@ -96,7 +99,6 @@ async fn tick_event_listener( ) -> Result<(), tokio::sync::mpsc::error::SendError> { let mut tick = tokio::time::interval(tokio::time::Duration::from_secs_f64(1.0 / FPS)); - #[expect(clippy::integer_division_remainder_used)] loop { tokio::select! { biased; @@ -104,7 +106,7 @@ async fn tick_event_listener( break Ok(()); }, _ = tick.tick() =>{ - tx.send(Event::Tick)?; + drop(tx.send(Event::Tick)); } } } @@ -114,7 +116,6 @@ async fn crossterm_event_listener( tx: tokio::sync::mpsc::UnboundedSender, ) -> Result<(), tokio::sync::mpsc::error::SendError> { let mut reader = crossterm::event::EventStream::new(); - #[expect(clippy::integer_division_remainder_used)] loop { tokio::select! { biased; @@ -124,7 +125,7 @@ async fn crossterm_event_listener( maybe = reader.next() => { match maybe { Some(Ok(event)) => { - tx.send(Event::Crossterm(event))?; + drop(tx.send(Event::Crossterm(event))); }, Some(Err(_)) => {}, None => { @@ -276,8 +277,13 @@ fn draw(f: &mut Frame, state: &AppState, logger_state: &TuiWidgetState) { let lines = vec![ Line::from("Up/PageUp/k - scroll logs up"), Line::from("Down/PageDown/j - scroll logs down"), - Line::from("ESC/q/Ctrl-C - exit") - .style(Style::default().fg(Color::Red)), + if matches!(state.mode, AppMode::Running) { + Line::from("ESC/q/Ctrl-C - stop") + .style(Style::default().fg(Color::Yellow)) + } else { + Line::from("ESC/q/Ctrl-C - quit") + .style(Style::default().fg(Color::Red)) + }, ]; f.render_widget(Text::from(lines).centered(), outer_layout[3]); } @@ -289,6 +295,7 @@ async fn is_interactive() -> bool { async fn handle_event( event: Event, state: &mut AppState, + token: &tokio_util::sync::CancellationToken, logger_state: &TuiWidgetState, ) -> bool { match event { @@ -297,12 +304,22 @@ async fn handle_event( match crossterm_event { CrosstermEvent::Key(key_event) => match key_event.code { KeyCode::Esc | KeyCode::Char('q' | 'Q') => { - state.mode = AppMode::Quit; + state.mode = if matches!(state.mode, AppMode::Running) { + AppMode::Done + } else { + AppMode::Quit + }; + token.cancel(); } KeyCode::Char('c' | 'C') if key_event.modifiers == KeyModifiers::CONTROL => { - state.mode = AppMode::Quit; + state.mode = if matches!(state.mode, AppMode::Running) { + AppMode::Done + } else { + AppMode::Quit + }; + token.cancel(); } KeyCode::Up | KeyCode::PageUp | KeyCode::Char('k') => { logger_state.transition(TuiWidgetEvent::PrevPageKey); @@ -369,9 +386,11 @@ async fn handle_event( .or_insert(1); } AppEvent::Done => { - if !is_interactive().await { - state.mode = AppMode::Quit; - } + state.mode = if is_interactive().await { + AppMode::Done + } else { + AppMode::Quit + }; } } false