/* * Author: Ruben Fiszel * Copyright: Windmill Labs, Inc 2022 * This file and its contents are licensed under the AGPLv3 License. * Please see the included NOTICE for copyright information and * LICENSE-AGPL for a copy of the license. */ use quick_cache::sync::Cache; use std::{ future::Future, hash::{Hash, Hasher}, net::SocketAddr, str::FromStr, sync::{ atomic::{AtomicBool, Ordering}, Arc, }, }; use tokio::{spawn, sync::broadcast}; use ee_oss::CriticalErrorChannel; use error::Error; use scripts::ScriptLang; use sqlx::{Acquire, Postgres}; pub mod agent_workers; pub mod ai_providers; pub mod apps; pub mod assets; pub mod auth; #[cfg(feature = "benchmark")] pub mod bench; pub mod cache; pub mod client; pub mod db; #[cfg(feature = "private")] pub mod ee; pub mod ee_oss; #[cfg(feature = "private")] pub mod email_ee; pub mod email_oss; pub mod error; pub mod external_ip; pub mod flow_status; pub mod flows; pub mod global_settings; pub mod indexer; pub mod job_metrics; #[cfg(all(feature = "parquet", feature = "private"))] pub mod job_s3_helpers_ee; #[cfg(feature = "parquet")] pub mod job_s3_helpers_oss; pub mod jobs; pub mod jwt; pub mod more_serde; pub mod oauth2; #[cfg(all(feature = "enterprise", feature = "openidconnect", feature = "private"))] pub mod oidc_ee; #[cfg(all(feature = "enterprise", feature = "openidconnect"))] pub mod oidc_oss; #[cfg(feature = "private")] pub mod otel_ee; pub mod otel_oss; pub mod queue; pub mod result_stream; pub mod s3_helpers; pub mod schedule; pub mod schema; pub mod scripts; pub mod server; #[cfg(feature = "private")] pub mod stats_ee; pub mod stats_oss; pub mod stream; #[cfg(feature = "private")] pub mod teams_ee; pub mod teams_oss; pub mod tracing_init; pub mod triggers; pub mod users; pub mod utils; pub mod variables; pub mod worker; pub mod worker_group_job_stats; pub mod workspaces; pub const DEFAULT_MAX_CONNECTIONS_SERVER: u32 = 50; pub const DEFAULT_MAX_CONNECTIONS_WORKER: u32 = 5; pub const DEFAULT_MAX_CONNECTIONS_INDEXER: u32 = 5; pub const DEFAULT_HUB_BASE_URL: &str = "https://hub.windmill.dev"; pub const SERVICE_LOG_RETENTION_SECS: i64 = 60 * 60 * 24 * 14; // 2 weeks retention period for logs #[macro_export] macro_rules! add_time { ($bench:expr, $name:expr) => { #[cfg(feature = "benchmark")] { $bench.add_timing($name); // println!("{}: {:?}", $z, $y.elapsed()); } }; } lazy_static::lazy_static! { pub static ref METRICS_PORT: u16 = std::env::var("METRICS_PORT") .ok() .and_then(|s| s.parse::().ok()) .unwrap_or(8001); pub static ref METRICS_ADDR: SocketAddr = std::env::var("METRICS_ADDR") .ok() .map(|s| { s.parse::() .map(|b| b.then(|| SocketAddr::from(([0, 0, 0, 0], *METRICS_PORT)))) .or_else(|_| s.parse::().map(Some)) }) .transpose().ok() .flatten() .flatten() .unwrap_or_else(|| SocketAddr::from(([0, 0, 0, 0], *METRICS_PORT))); pub static ref METRICS_ENABLED: AtomicBool = AtomicBool::new(std::env::var("METRICS_PORT").is_ok() || std::env::var("METRICS_ADDR").is_ok()); pub static ref OTEL_METRICS_ENABLED: AtomicBool = AtomicBool::new(std::env::var("OTEL_METRICS").is_ok()); pub static ref OTEL_TRACING_ENABLED: AtomicBool = AtomicBool::new(std::env::var("OTEL_TRACING").is_ok()); pub static ref OTEL_LOGS_ENABLED: AtomicBool = AtomicBool::new(std::env::var("OTEL_LOGS").is_ok()); pub static ref METRICS_DEBUG_ENABLED: AtomicBool = AtomicBool::new(false); pub static ref CRITICAL_ALERT_MUTE_UI_ENABLED: AtomicBool = AtomicBool::new(false); pub static ref BASE_URL: Arc> = Arc::new(RwLock::new("".to_string())); pub static ref IS_READY: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false); pub static ref BASE_INTERNAL_URL: String = std::env::var("BASE_INTERNAL_URL").unwrap_or("http://localhost:8000".to_string()); pub static ref HUB_BASE_URL: Arc> = Arc::new(RwLock::new(DEFAULT_HUB_BASE_URL.to_string())); pub static ref CRITICAL_ERROR_CHANNELS: Arc>> = Arc::new(RwLock::new(vec![])); pub static ref CRITICAL_ALERTS_ON_DB_OVERSIZE: Arc>> = Arc::new(RwLock::new(None)); pub static ref JOB_RETENTION_SECS: Arc> = Arc::new(RwLock::new(0)); pub static ref MONITOR_LOGS_ON_OBJECT_STORE: Arc> = Arc::new(RwLock::new(false)); pub static ref INSTANCE_NAME: String = rd_string(5); pub static ref DEPLOYED_SCRIPT_HASH_CACHE: Cache<(String, String), ExpiringLatestVersionId> = Cache::new(1000); pub static ref FLOW_VERSION_CACHE: Cache<(String, String), ExpiringLatestVersionId> = Cache::new(1000); pub static ref DYNAMIC_INPUT_CACHE: Cache> = Cache::new(1000); pub static ref DEPLOYED_SCRIPT_INFO_CACHE: Cache<(String, i64), ScriptHashInfo> = Cache::new(1000); pub static ref FLOW_INFO_CACHE: Cache<(String, i64), FlowVersionInfo> = Cache::new(1000); pub static ref QUIET_LOGS: bool = std::env::var("QUIET_LOGS").map(|s| s.parse::().unwrap_or(false)).unwrap_or(false); } const LATEST_VERSION_ID_CACHE_TTL: std::time::Duration = std::time::Duration::from_secs(60); pub async fn shutdown_signal( tx: KillpillSender, mut rx: tokio::sync::broadcast::Receiver<()>, ) -> anyhow::Result<()> { use std::io; #[cfg(any(target_os = "linux", target_os = "macos"))] async fn terminate() -> io::Result<()> { use tokio::signal::unix::SignalKind; tokio::signal::unix::signal(SignalKind::terminate())? .recv() .await; Ok(()) } #[cfg(any(target_os = "linux", target_os = "macos"))] tokio::select! { _ = terminate() => { tracing::info!("shutdown monitor received terminate"); }, _ = tokio::signal::ctrl_c() => { tracing::info!("shutdown monitor received ctrl-c"); }, _ = rx.recv() => { tracing::info!("shutdown monitor received killpill"); }, } #[cfg(not(any(target_os = "linux", target_os = "macos")))] tokio::select! { _ = tokio::signal::ctrl_c() => {}, _ = rx.recv() => { tracing::info!("shutdown monitor received killpill"); }, } spawn(async move { #[cfg(any(target_os = "linux", target_os = "macos"))] tokio::select! { _ = terminate() => { tracing::info!("2nd shutdown monitor received terminate"); }, _ = tokio::signal::ctrl_c() => { tracing::info!("2nd shutdown monitor received ctrl-c"); }, } #[cfg(not(any(target_os = "linux", target_os = "macos")))] tokio::select! { _ = tokio::signal::ctrl_c() => {}, _ = rx.recv() => { tracing::info!("2nd shutdown monitor received killpill"); }, } tracing::info!("Second terminate signal received, forcefully exiting"); let handle = tokio::runtime::Handle::current(); let metrics = handle.metrics(); tracing::info!( "Alive tasks: {}, global queue depth: {}", metrics.num_alive_tasks(), metrics.global_queue_depth() ); std::process::exit(1); }); tracing::info!("signal received, starting graceful shutdown"); let _ = tx.send(); spawn(async move { tokio::time::sleep(std::time::Duration::from_secs(24 * 7 * 60 * 60)).await; tracing::info!("Forcefully exiting after 7 days"); std::process::exit(1); }); Ok(()) } use tokio::sync::RwLock; use utils::rd_string; #[cfg(feature = "prometheus")] pub async fn serve_metrics( addr: SocketAddr, mut rx: tokio::sync::broadcast::Receiver<()>, ready_worker_endpoint: bool, metrics_endpoint: bool, ) -> anyhow::Result<()> { if !metrics_endpoint && !ready_worker_endpoint { return Ok(()); } use axum::{ routing::{get, post}, Router, }; use hyper::StatusCode; let router = Router::new(); let router = if metrics_endpoint { router .route("/metrics", get(metrics)) .route("/reset", post(reset)) } else { router }; let router = if ready_worker_endpoint { router.route( "/ready", get(|| async { if IS_READY.load(std::sync::atomic::Ordering::Relaxed) { (StatusCode::OK, "ready") } else { (StatusCode::INTERNAL_SERVER_ERROR, "not ready") } }), ) } else { router }; tokio::spawn(async move { tracing::info!("Serving metrics at: {addr}"); let listener = tokio::net::TcpListener::bind(addr).await; if let Err(e) = listener { tracing::error!("Error binding to metrics address: {}", e); return; } if let Err(e) = axum::serve(listener.unwrap(), router.into_make_service()) .with_graceful_shutdown(async move { rx.recv().await.ok(); tracing::info!("Graceful shutdown of metrics"); }) .await { tracing::error!("Error serving metrics: {}", e); } }) .await?; Ok(()) } #[cfg(feature = "prometheus")] async fn metrics() -> Result { let metric_families = prometheus::gather(); Ok(prometheus::TextEncoder::new() .encode_to_string(&metric_families) .map_err(anyhow::Error::from)?) } #[cfg(feature = "prometheus")] async fn reset() -> () { todo!() } pub struct PostgresUrlComponents { pub scheme: String, pub username: Option, pub password: Option, pub host: String, pub port: Option, pub database: String, pub ssl_mode: Option, } pub fn parse_postgres_url(url: &str) -> Result { let parsed_url = url::Url::parse(url).map_err(|_| Error::BadConfig("Invalid PostgreSQL URL".to_string()))?; let scheme = parsed_url.scheme().to_string(); let username = parsed_url.username().to_string(); let password = parsed_url.password().map(|p| p.to_string()); let host = parsed_url .host_str() .ok_or_else(|| Error::BadConfig("Missing host in PostgreSQL URL".to_string()))? .to_string(); let port = parsed_url.port(); let database = parsed_url.path().trim_start_matches('/').to_string(); let mut ssl_mode = None; for query in parsed_url.query_pairs() { if query.0 == "sslmode" { ssl_mode = Some(query.1.to_string()); } } Ok(PostgresUrlComponents { scheme, username: if username.is_empty() { None } else { Some(username) }, password, host, port, database, ssl_mode, }) } pub async fn get_database_url() -> Result { use std::env::var; use tokio::fs::File; use tokio::io::AsyncReadExt; match var("DATABASE_URL_FILE") { Ok(file_path) => { let mut file = File::open(file_path).await?; let mut contents = String::new(); file.read_to_string(&mut contents).await?; Ok(contents.trim().to_string()) } Err(_) => var("DATABASE_URL").map_err(|_| { Error::BadConfig( "Either DATABASE_URL_FILE or DATABASE_URL env var is missing".to_string(), ) }), } } pub async fn initial_connection() -> Result, error::Error> { let database_url = get_database_url().await?; sqlx::postgres::PgPoolOptions::new() .max_connections(2) .connect_with(sqlx::postgres::PgConnectOptions::from_str(&database_url)?) .await .map_err(|err| Error::ConnectingToDatabase(err.to_string())) } pub async fn connect_db( server_mode: bool, indexer_mode: bool, worker_mode: bool, ) -> anyhow::Result> { use anyhow::Context; let database_url = get_database_url().await?; let max_connections = match std::env::var("DATABASE_CONNECTIONS") { Ok(n) => n.parse::().context("invalid DATABASE_CONNECTIONS")?, Err(_) => { if server_mode { DEFAULT_MAX_CONNECTIONS_SERVER } else if indexer_mode { DEFAULT_MAX_CONNECTIONS_INDEXER } else { DEFAULT_MAX_CONNECTIONS_WORKER + std::env::var("NUM_WORKERS") .ok() .map(|x| x.parse().ok()) .flatten() .unwrap_or(1) - 1 } } }; Ok(connect(&database_url, max_connections, worker_mode).await?) } pub async fn connect( database_url: &str, max_connections: u32, worker_mode: bool, ) -> Result, error::Error> { use sqlx::Executor; use std::time::Duration; sqlx::postgres::PgPoolOptions::new() .min_connections((max_connections / 5).clamp(3, max_connections)) .max_connections(max_connections) .max_lifetime(Duration::from_secs(30 * 60)) // 30 mins .after_connect(move |conn, _| { if worker_mode { Box::pin(async move { if let Err(e) = conn .execute( r#" SET enable_seqscan = OFF; SET statement_timeout = '5min'; SET idle_in_transaction_session_timeout = '10min'; SET tcp_keepalives_idle = 300; SET tcp_keepalives_interval = 60; SET tcp_keepalives_count = 10;"#, ) .await { tracing::error!("Error setting postgres settings: {}", e); } Ok(()) }) } else { Box::pin(async move { if let Err(e) = conn .execute( r#" SET statement_timeout = '5min'; SET idle_in_transaction_session_timeout = '10min'; SET tcp_keepalives_idle = 300; SET tcp_keepalives_interval = 60; SET tcp_keepalives_count = 10;"#, ) .await { tracing::error!("Error setting postgres settings: {}", e); } Ok(()) }) } }) .connect_with( sqlx::postgres::PgConnectOptions::from_str(database_url)?.statement_cache_capacity(400), ) .await .map_err(|err| Error::ConnectingToDatabase(err.to_string())) } type Tag = String; pub use db::DB; use crate::{ auth::{PermsCache, FLOW_PERMS_CACHE, HASH_PERMS_CACHE}, db::{AuthedRef, UserDbWithAuthed}, scripts::ScriptHash, }; #[derive(Clone)] pub struct ExpiringLatestVersionId { id: i64, expires_at: std::time::Instant, } #[derive(Clone)] pub struct ScriptHashInfo { pub path: String, pub hash: i64, pub tag: Option, pub concurrency_key: Option, pub concurrent_limit: Option, pub concurrency_time_window_s: Option, pub cache_ttl: Option, pub language: ScriptLang, pub dedicated_worker: Option, pub priority: Option, pub delete_after_use: Option, pub timeout: Option, pub has_preprocessor: Option, pub on_behalf_of_email: Option, pub created_by: String, } pub fn get_latest_deployed_hash_for_path<'e>( db: Option>>, db2: DB, w_id: &'e str, script_path: &'e str, ) -> impl Future> + Send + 'e { async move { let cache_key = (w_id.to_string(), script_path.to_string()); let mut computed_hash = None; let hash = match DEPLOYED_SCRIPT_HASH_CACHE.get(&cache_key) { Some(cached_hash) if cached_hash.expires_at > std::time::Instant::now() && db.as_ref().is_none_or(|x| { let r = HASH_PERMS_CACHE .check_perms_in_cache(x.authed, ScriptHash(cached_hash.id)); computed_hash = Some(r.1); return r.0; }) => { tracing::debug!( "Using cached script hash {} for {script_path}", cached_hash.id ); cached_hash.id } _ => { tracing::debug!("Fetching script hash for {script_path}"); let hash = if let Some(db) = db { let authed = db.authed; let mut conn = db.acquire().await?; let hash = get_latest_script_hash(&mut *conn, script_path, w_id).await?; if let Some(hash) = hash { HASH_PERMS_CACHE.insert( computed_hash.unwrap_or_else(|| PermsCache::compute_hash(authed)), ScriptHash(hash), ); } else { let mut conn = db2.acquire().await?; let exists = get_latest_script_hash(&mut *conn, script_path, w_id) .await? .is_some(); if exists { return Err(Error::NotAuthorized(format!("You are not authorized to access this script: {script_path} (but it exists). Your permissions are: {:?}", authed))); } } hash } else { let mut conn = db2.acquire().await?; get_latest_script_hash(&mut *conn, script_path, w_id).await? }; let hash = utils::not_found_if_none(hash, "script", script_path)?; DEPLOYED_SCRIPT_HASH_CACHE.insert( cache_key, ExpiringLatestVersionId { id: hash, expires_at: std::time::Instant::now() + LATEST_VERSION_ID_CACHE_TTL, }, ); hash } }; get_script_info_for_hash(None, &db2, w_id, hash).await } } pub async fn get_latest_script_hash<'e, E: sqlx::PgExecutor<'e>>( db: E, script_path: &'e str, w_id: &'e str, ) -> error::Result> { let hash = sqlx::query_scalar!( "select hash from script where path = $1 AND workspace_id = $2 AND deleted = false AND lock IS not NULL AND lock_error_logs IS NULL ORDER BY created_at DESC LIMIT 1", script_path, w_id ) .fetch_optional(db) .await?; return Ok(hash); } pub async fn get_script_info_for_hash<'e, E: sqlx::PgExecutor<'e>>( db_authed: Option>>, db: E, w_id: &str, hash: i64, ) -> error::Result { let key = (w_id.to_string(), hash); let mut computed_hash = None; match DEPLOYED_SCRIPT_INFO_CACHE.get(&key) { Some(info) if db_authed.as_ref().is_none_or(|x| { let r = HASH_PERMS_CACHE.check_perms_in_cache(x.authed, scripts::ScriptHash(hash)); computed_hash = Some(r.1); return r.0; }) => { tracing::debug!("Using cached deployed script info for {hash}"); Ok(info) } _ => { tracing::debug!("Fetching deployed script info for {hash}"); let info = if let Some(db_authed) = db_authed { let mut conn = db_authed.acquire().await?; let hash_info = get_script_info_for_hash_inner(&mut *conn, w_id, hash).await?; if hash_info.is_some() { HASH_PERMS_CACHE.insert( computed_hash.unwrap_or_else(|| PermsCache::compute_hash(db_authed.authed)), ScriptHash(hash), ); } hash_info } else { get_script_info_for_hash_inner(db, w_id, hash).await? }; let info = utils::not_found_if_none(info, "script", &hash.to_string())?; DEPLOYED_SCRIPT_INFO_CACHE.insert(key, info.clone()); Ok(info) } } } async fn get_script_info_for_hash_inner<'e, E: sqlx::PgExecutor<'e>>( db: E, w_id: &str, hash: i64, ) -> error::Result> { let r = sqlx::query_as!( ScriptHashInfo, "select hash, tag, concurrency_key, concurrent_limit, concurrency_time_window_s, cache_ttl, language as \"language: ScriptLang\", dedicated_worker, priority, delete_after_use, timeout, has_preprocessor, on_behalf_of_email, created_by, path from script where hash = $1 AND workspace_id = $2", hash, w_id ) .fetch_optional(db) .await?; Ok(r) } #[derive(Clone)] pub struct FlowVersionInfo { pub version: i64, pub tag: Option, pub early_return: Option, pub has_preprocessor: Option, pub on_behalf_of_email: Option, pub edited_by: String, pub dedicated_worker: Option, } struct CachedFlowPath(String); impl Into for CachedFlowPath { fn into(self) -> u64 { let mut hasher = std::collections::hash_map::DefaultHasher::new(); self.0.hash(&mut hasher); hasher.finish() } } pub fn get_latest_flow_version_id_for_path< 'a, 'e, A: sqlx::Acquire<'e, Database = Postgres> + Send + 'a, >( db_authed: Option>>, db: A, w_id: &'a str, path: &'a str, use_cache: bool, ) -> impl Future> + Send + 'a where 'e: 'a, { // as instructed in the docstring of sqlx::Acquire async move { let cache_key = (w_id.to_string(), path.to_string()); let cached_version = if use_cache { FLOW_VERSION_CACHE.get(&cache_key) } else { None }; let mut computed_hash: Option<_> = None; let version = match cached_version { Some(cached_version) if cached_version.expires_at > std::time::Instant::now() && db_authed.as_ref().is_none_or(|x| { let r = FLOW_PERMS_CACHE .check_perms_in_cache(x.authed, CachedFlowPath(path.to_string())); computed_hash = Some(r.1); return r.0; }) => { tracing::debug!("Using cached flow version {} for {path}", cached_version.id); cached_version.id } _ => { tracing::debug!("Fetching flow version for {path}"); let version = if let Some(db_authed) = db_authed { let mut conn = db_authed.acquire().await?; let r = get_latest_flow_version_for_path(&mut *conn, w_id, path).await?; if r.is_some() { FLOW_PERMS_CACHE.insert( computed_hash .unwrap_or_else(|| PermsCache::compute_hash(db_authed.authed)), CachedFlowPath(path.to_string()), ); } else { let mut conn = db.acquire().await?; let exists = get_latest_flow_version_for_path(&mut *conn, w_id, path) .await? .is_some(); if exists { return Err(Error::NotAuthorized(format!( "You are not authorized to access this flow: {path} (but it exists). Your permissions are: {:?}", db_authed.authed ))); } } r } else { let mut conn = db.acquire().await?; get_latest_flow_version_for_path(&mut *conn, w_id, path).await? }; let version = utils::not_found_if_none(version, "flow", path)?; FLOW_VERSION_CACHE.insert( cache_key, ExpiringLatestVersionId { id: version, expires_at: std::time::Instant::now() + LATEST_VERSION_ID_CACHE_TTL, }, ); version } }; Ok(version) } } pub fn get_latest_flow_version_info_for_path_from_version< 'a, 'e, A: sqlx::Acquire<'e, Database = Postgres> + Send + 'a, >( db: A, version: i64, w_id: &'a str, path: &'a str, ) -> impl Future> + Send + 'a { async move { // as instructed in the docstring of sqlx::Acquire let key = (w_id.to_string(), version); match FLOW_INFO_CACHE.get(&key) { Some(info) => { tracing::debug!("Using cached flow version info for {version} ({path})"); Ok(info) } _ => { tracing::debug!("Fetching flow version info for {version} ({path})"); let mut conn = db.acquire().await?; let info = sqlx::query_as!( FlowVersionInfo, "SELECT tag, dedicated_worker, flow_version.value->>'early_return' as early_return, flow_version.value->>'preprocessor_module' IS NOT NULL as has_preprocessor, on_behalf_of_email, edited_by, flow_version.id AS version FROM flow INNER JOIN flow_version ON flow_version.id = $3 WHERE flow.path = $1 and flow.workspace_id = $2", path, w_id, version ) .fetch_optional(&mut *conn) .await?; let info = utils::not_found_if_none(info, "flow", path)?; FLOW_INFO_CACHE.insert(key, info.clone()); Ok(info) } } } } pub async fn get_latest_flow_version_info_for_path<'e>( db_authed: Option>>, db: &DB, w_id: &'e str, path: &'e str, use_cache: bool, ) -> error::Result { // as instructed in the docstring of sqlx::Acquire let version = get_latest_flow_version_id_for_path(db_authed, &db.clone(), w_id, path, use_cache).await?; get_latest_flow_version_info_for_path_from_version(db, version, w_id, path).await } async fn get_latest_flow_version_for_path<'e, E: sqlx::PgExecutor<'e>>( db: E, w_id: &str, path: &str, ) -> error::Result> { let version = sqlx::query_scalar!( "SELECT flow_version.id from flow INNER JOIN flow_version ON flow_version.id = flow.versions[array_upper(flow.versions, 1)] WHERE flow.path = $1 and flow.workspace_id = $2", path, w_id ) .fetch_optional(db) .await?; Ok(version) } pub async fn get_latest_hash_for_path<'c, E: sqlx::PgExecutor<'c>>( db: E, w_id: &str, script_path: &str, require_locked: bool, ) -> error::Result<( scripts::ScriptHash, Option, Option, Option, Option, Option, ScriptLang, Option, Option, Option, Option, String, )> { let r_o = sqlx::query!( "select hash, tag, concurrency_key, concurrent_limit, concurrency_time_window_s, cache_ttl, language as \"language: ScriptLang\", dedicated_worker, priority, timeout, on_behalf_of_email, created_by FROM script WHERE path = $1 AND workspace_id = $2 AND archived = false AND (lock IS NOT NULL OR $3 = false) ORDER BY created_at DESC LIMIT 1", script_path, w_id, require_locked ) .fetch_optional(db) .await?; let script = utils::not_found_if_none(r_o, "script", script_path)?; Ok(( scripts::ScriptHash(script.hash), script.tag, script.concurrency_key, script.concurrent_limit, script.concurrency_time_window_s, script.cache_ttl, script.language, script.dedicated_worker, script.priority, script.timeout, script.on_behalf_of_email, script.created_by, )) } pub struct KillpillSender { tx: broadcast::Sender<()>, already_sent: Arc, } impl Clone for KillpillSender { fn clone(&self) -> Self { KillpillSender { tx: self.tx.clone(), already_sent: self.already_sent.clone() } } } impl KillpillSender { pub fn new(capacity: usize) -> (Self, broadcast::Receiver<()>) { let (tx, rx) = broadcast::channel(capacity); let sender = KillpillSender { tx, already_sent: Arc::new(AtomicBool::new(false)) }; (sender, rx) } pub fn clone(&self) -> Self { KillpillSender { tx: self.tx.clone(), already_sent: self.already_sent.clone() } } pub fn subscribe(&self) -> broadcast::Receiver<()> { self.tx.subscribe() } // Try to send the killpill if it hasn't been sent already pub fn send(&self) -> bool { // Check if it's already been sent, and if not, set the flag to true if !self.already_sent.swap(true, Ordering::SeqCst) { // We're the first to set it to true, so send the signal if let Err(e) = self.tx.send(()) { tracing::error!("failed to send killpill: {:?}", e); } true } else { // Signal was already sent false } } // // Force send a signal regardless of previous sends // fn force_send(&self) -> Result> { // self.already_sent.store(true, Ordering::SeqCst); // self.tx.send(()) // } // // Check if the killpill has been sent // fn is_sent(&self) -> bool { // self.already_sent.load(Ordering::SeqCst) // } }