feat: implement graceful shutdown & saving proxies on exit

This commit is contained in:
Almaz
2025-07-16 07:23:10 +00:00
committed by GitHub
parent 46ee500b57
commit 0f477d49d7
8 changed files with 194 additions and 144 deletions
Generated
+1
View File
@@ -1346,6 +1346,7 @@ dependencies = [
"serde",
"serde_json",
"tokio",
"tokio-util",
"toml",
"tracing",
"tracing-log",
+1
View File
@@ -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"
+42 -36
View File
@@ -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<Config>,
proxies: HashSet<Proxy>,
proxies: Arc<tokio::sync::Mutex<HashSet<Proxy>>>,
token: tokio_util::sync::CancellationToken,
#[cfg(feature = "tui")] tx: tokio::sync::mpsc::UnboundedSender<Event>,
) -> color_eyre::Result<Vec<Proxy>> {
) -> 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::<Vec<_>>(),
proxies.lock().await.drain().collect::<Vec<_>>(),
));
let mut join_set = tokio::task::JoinSet::<color_eyre::Result<()>>::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(())
}
+88 -77
View File
@@ -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<config::Config>,
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<()>> {
) -> 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<()>>,
) -> 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<config::Config>,
proxies: std::collections::HashSet<crate::proxy::Proxy>,
#[cfg(feature = "tui")] tx: &tokio::sync::mpsc::UnboundedSender<
proxies: Arc<tokio::sync::Mutex<HashSet<crate::proxy::Proxy>>>,
token: tokio_util::sync::CancellationToken,
#[cfg(feature = "tui")] tx: tokio::sync::mpsc::UnboundedSender<
event::Event,
>,
) -> color_eyre::Result<Vec<crate::proxy::Proxy>> {
) -> 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<config::Config>,
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<Arc<config::Config>> {
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
+9 -3
View File
@@ -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<Config>,
mut proxies: Vec<Proxy>,
proxies: Arc<tokio::sync::Mutex<HashSet<Proxy>>>,
) -> 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 {
+4
View File
@@ -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)
+18 -16
View File
@@ -119,10 +119,10 @@ async fn scrape_one(
pub async fn scrape_all(
config: Arc<Config>,
http_client: reqwest::Client,
proxies: Arc<tokio::sync::Mutex<HashSet<Proxy>>>,
token: tokio_util::sync::CancellationToken,
#[cfg(feature = "tui")] tx: tokio::sync::mpsc::UnboundedSender<Event>,
) -> color_eyre::Result<HashSet<Proxy>> {
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(())
}
+31 -12
View File
@@ -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<Event>,
mut rx: tokio::sync::mpsc::UnboundedReceiver<Event>,
) -> 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<Event>> {
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<Event>,
) -> Result<(), tokio::sync::mpsc::error::SendError<Event>> {
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