use super::tracker::SpendResult;
use super::types::{AlertSeverity, Budget, BudgetAlert, BudgetAlertType};
use crate::core::net::ProviderEndpointPolicy;
use crate::utils::net::http::ProviderHttpClient;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::RwLock;
use tracing::{debug, error, info, warn};
#[derive(Clone)]
pub struct BudgetAlertManager {
alerts: Arc<RwLock<AlertStorage>>,
webhooks: Arc<RwLock<Vec<RegisteredWebhook>>>,
config: Arc<RwLock<AlertConfig>>,
}
#[derive(Clone)]
struct RegisteredWebhook {
config: WebhookConfig,
client: ProviderHttpClient,
}
#[derive(Debug, Default)]
struct AlertStorage {
alerts: HashMap<String, BudgetAlert>,
alerts_by_budget: HashMap<String, Vec<String>>,
history: Vec<BudgetAlert>,
max_history_size: usize,
}
impl AlertStorage {
fn new(max_history_size: usize) -> Self {
Self {
alerts: HashMap::new(),
alerts_by_budget: HashMap::new(),
history: Vec::new(),
max_history_size,
}
}
fn add_alert(&mut self, alert: BudgetAlert) {
let alert_id = alert.id.clone();
let budget_id = alert.budget_id.clone();
self.alerts.insert(alert_id.clone(), alert.clone());
self.alerts_by_budget
.entry(budget_id)
.or_default()
.push(alert_id);
self.history.push(alert);
if self.history.len() > self.max_history_size {
let excess = self.history.len() - self.max_history_size;
self.history.drain(0..excess);
}
}
fn get_alerts_for_budget(&self, budget_id: &str) -> Vec<&BudgetAlert> {
self.alerts_by_budget
.get(budget_id)
.map(|ids| ids.iter().filter_map(|id| self.alerts.get(id)).collect())
.unwrap_or_default()
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct WebhookConfig {
pub url: String,
pub headers: HashMap<String, String>,
pub severities: Vec<AlertSeverity>,
pub enabled: bool,
pub timeout_secs: u64,
pub max_retries: u32,
}
impl Default for WebhookConfig {
fn default() -> Self {
Self {
url: String::new(),
headers: HashMap::new(),
severities: vec![AlertSeverity::Warning, AlertSeverity::Critical],
enabled: true,
timeout_secs: 30,
max_retries: 3,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
pub enum BudgetWebhookError {
#[error("Budget webhook URL is invalid or disallowed by outbound policy")]
InvalidUrl,
#[error("Failed to create policy-bound budget webhook client")]
ClientConstruction,
}
#[derive(Debug, Clone)]
pub struct AlertConfig {
pub enabled: bool,
pub soft_limit_percentage: f64,
pub warning_thresholds: Vec<f64>,
pub max_history_size: usize,
pub duplicate_suppression_secs: u64,
}
impl Default for AlertConfig {
fn default() -> Self {
Self {
enabled: true,
soft_limit_percentage: 0.8,
warning_thresholds: vec![0.9, 0.95],
max_history_size: 1000,
duplicate_suppression_secs: 3600, }
}
}
impl Default for BudgetAlertManager {
fn default() -> Self {
Self::new()
}
}
impl BudgetAlertManager {
pub fn new() -> Self {
let config = AlertConfig::default();
Self {
alerts: Arc::new(RwLock::new(AlertStorage::new(config.max_history_size))),
webhooks: Arc::new(RwLock::new(Vec::new())),
config: Arc::new(RwLock::new(config)),
}
}
pub fn with_config(config: AlertConfig) -> Self {
Self {
alerts: Arc::new(RwLock::new(AlertStorage::new(config.max_history_size))),
webhooks: Arc::new(RwLock::new(Vec::new())),
config: Arc::new(RwLock::new(config)),
}
}
pub async fn add_webhook(&self, config: WebhookConfig) -> Result<(), BudgetWebhookError> {
let timeout = Duration::from_secs(config.timeout_secs);
let client =
ProviderHttpClient::no_redirect(ProviderEndpointPolicy::public_only(), timeout)
.map_err(|_| BudgetWebhookError::ClientConstruction)?;
self.add_webhook_with_client(config, client).await
}
async fn add_webhook_with_client(
&self,
mut config: WebhookConfig,
client: ProviderHttpClient,
) -> Result<(), BudgetWebhookError> {
let url = reqwest::Url::parse(&config.url).map_err(|_| BudgetWebhookError::InvalidUrl)?;
if !matches!(url.scheme(), "http" | "https")
|| ProviderEndpointPolicy::public_only()
.validate_url_without_resolution(&url)
.is_err()
{
return Err(BudgetWebhookError::InvalidUrl);
}
config.url = url.to_string();
let mut webhooks = self.webhooks.write().await;
webhooks.push(RegisteredWebhook { config, client });
Ok(())
}
#[cfg(test)]
async fn add_webhook_with_client_for_test(
&self,
config: WebhookConfig,
client: ProviderHttpClient,
) -> Result<(), BudgetWebhookError> {
self.add_webhook_with_client(config, client).await
}
pub async fn clear_webhooks(&self) {
let mut webhooks = self.webhooks.write().await;
webhooks.clear();
}
pub async fn process_spend_result(&self, result: &SpendResult, budget: &Budget) {
let config = self.config.read().await;
if !config.enabled {
return;
}
drop(config);
if result.should_alert_soft_limit {
self.create_alert(budget, BudgetAlertType::SoftLimitReached, budget.soft_limit)
.await;
}
if result.should_alert_exceeded {
self.create_alert(budget, BudgetAlertType::BudgetExceeded, budget.max_budget)
.await;
return;
}
let config = self.config.read().await;
for &threshold_pct in &config.warning_thresholds {
let threshold = budget.max_budget * threshold_pct;
if result.current_spend >= threshold
&& result.current_spend - (result.max_budget - result.remaining) < threshold
{
drop(config);
self.create_alert(budget, BudgetAlertType::ApproachingLimit, threshold)
.await;
break;
}
}
}
async fn create_alert(&self, budget: &Budget, alert_type: BudgetAlertType, threshold: f64) {
let alert = BudgetAlert::new(budget, alert_type, threshold);
info!(
"Budget alert created: {} - {} (severity: {:?})",
budget.name, alert.message, alert.severity
);
{
let mut storage = self.alerts.write().await;
storage.add_alert(alert.clone());
}
self.send_webhook_notifications(&alert).await;
}
pub async fn create_reset_alert(&self, budget: &Budget) {
let config = self.config.read().await;
if !config.enabled {
return;
}
drop(config);
let alert = BudgetAlert::new(budget, BudgetAlertType::BudgetReset, 0.0);
info!("Budget reset alert: {}", alert.message);
let mut storage = self.alerts.write().await;
storage.add_alert(alert.clone());
drop(storage);
self.send_webhook_notifications(&alert).await;
}
async fn send_webhook_notifications(&self, alert: &BudgetAlert) {
let webhooks = self.webhooks.read().await;
for webhook in webhooks.iter() {
if !webhook.config.enabled {
continue;
}
if !webhook.config.severities.contains(&alert.severity) {
continue;
}
self.send_single_webhook(webhook, alert).await;
}
}
async fn send_single_webhook(&self, webhook: &RegisteredWebhook, alert: &BudgetAlert) {
let payload = serde_json::json!({
"type": "budget_alert",
"alert": {
"id": alert.id,
"budget_id": alert.budget_id,
"scope": alert.scope.to_string(),
"alert_type": format!("{:?}", alert.alert_type),
"severity": format!("{:?}", alert.severity),
"message": alert.message,
"current_spend": alert.current_spend,
"threshold": alert.threshold,
"max_budget": alert.max_budget,
"created_at": alert.created_at.to_rfc3339()
}
});
let mut retries = 0;
let max_retries = webhook.config.max_retries;
loop {
let mut request = match webhook.client.post(&webhook.config.url) {
Ok(request) => request.json(&payload),
Err(_) => {
error!("Budget alert webhook rejected by outbound endpoint policy");
return;
}
};
for (key, value) in &webhook.config.headers {
request = request.header(key, value);
}
match request.send().await {
Ok(response) => {
if response.status().is_success() {
debug!("Successfully sent budget alert webhook");
return;
} else {
warn!(
status = response.status().as_u16(),
"Budget alert webhook returned error status"
);
}
}
Err(error) => {
if ProviderHttpClient::request_error_is_endpoint_policy(&error) {
error!("Budget alert webhook rejected by outbound endpoint policy");
} else {
error!("Budget alert webhook transport request failed");
}
}
}
retries += 1;
if retries >= max_retries {
error!("Exhausted retries for budget alert webhook");
return;
}
let delay = Duration::from_millis(100 * 2_u64.pow(retries));
tokio::time::sleep(delay).await;
}
}
pub async fn get_alerts_for_budget(&self, budget_id: &str) -> Vec<BudgetAlert> {
let storage = self.alerts.read().await;
storage
.get_alerts_for_budget(budget_id)
.into_iter()
.cloned()
.collect()
}
pub async fn get_all_alerts(&self) -> Vec<BudgetAlert> {
let storage = self.alerts.read().await;
storage.alerts.values().cloned().collect()
}
pub async fn get_unacknowledged_alerts(&self) -> Vec<BudgetAlert> {
let storage = self.alerts.read().await;
storage
.alerts
.values()
.filter(|a| !a.acknowledged)
.cloned()
.collect()
}
pub async fn get_alerts_by_severity(&self, severity: AlertSeverity) -> Vec<BudgetAlert> {
let storage = self.alerts.read().await;
storage
.alerts
.values()
.filter(|a| a.severity == severity)
.cloned()
.collect()
}
pub async fn acknowledge_alert(&self, alert_id: &str) -> bool {
let mut storage = self.alerts.write().await;
if let Some(alert) = storage.alerts.get_mut(alert_id) {
alert.acknowledge();
true
} else {
false
}
}
pub async fn acknowledge_alerts_for_budget(&self, budget_id: &str) -> usize {
let mut storage = self.alerts.write().await;
let mut count = 0;
if let Some(alert_ids) = storage.alerts_by_budget.get(budget_id).cloned() {
for alert_id in alert_ids {
if let Some(alert) = storage.alerts.get_mut(&alert_id)
&& !alert.acknowledged
{
alert.acknowledge();
count += 1;
}
}
}
count
}
pub async fn get_alert_history(&self, limit: Option<usize>) -> Vec<BudgetAlert> {
let storage = self.alerts.read().await;
let limit = limit.unwrap_or(storage.history.len());
storage.history.iter().rev().take(limit).cloned().collect()
}
pub async fn get_alert_stats(&self) -> AlertStats {
let storage = self.alerts.read().await;
let mut stats = AlertStats::default();
for alert in storage.alerts.values() {
stats.total_alerts += 1;
if !alert.acknowledged {
stats.unacknowledged += 1;
}
match alert.severity {
AlertSeverity::Info => stats.info_count += 1,
AlertSeverity::Warning => stats.warning_count += 1,
AlertSeverity::Critical => stats.critical_count += 1,
}
match alert.alert_type {
BudgetAlertType::SoftLimitReached => stats.soft_limit_alerts += 1,
BudgetAlertType::BudgetExceeded => stats.exceeded_alerts += 1,
BudgetAlertType::BudgetReset => stats.reset_alerts += 1,
BudgetAlertType::ApproachingLimit => stats.approaching_limit_alerts += 1,
}
}
stats
}
pub async fn clear_alerts(&self) {
let mut storage = self.alerts.write().await;
storage.alerts.clear();
storage.alerts_by_budget.clear();
}
pub async fn clear_acknowledged_alerts(&self) -> usize {
let mut storage = self.alerts.write().await;
let acknowledged_ids: Vec<String> = storage
.alerts
.iter()
.filter(|(_, alert)| alert.acknowledged)
.map(|(id, _)| id.clone())
.collect();
let count = acknowledged_ids.len();
for id in acknowledged_ids {
storage.alerts.remove(&id);
}
let remaining_ids: std::collections::HashSet<String> =
storage.alerts.keys().cloned().collect();
for alerts in storage.alerts_by_budget.values_mut() {
alerts.retain(|id| remaining_ids.contains(id));
}
count
}
pub async fn update_config(&self, new_config: AlertConfig) {
let mut config = self.config.write().await;
*config = new_config;
}
pub async fn get_config(&self) -> AlertConfig {
self.config.read().await.clone()
}
pub async fn is_enabled(&self) -> bool {
self.config.read().await.enabled
}
pub async fn set_enabled(&self, enabled: bool) {
let mut config = self.config.write().await;
config.enabled = enabled;
}
}
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
pub struct AlertStats {
pub total_alerts: usize,
pub unacknowledged: usize,
pub info_count: usize,
pub warning_count: usize,
pub critical_count: usize,
pub soft_limit_alerts: usize,
pub exceeded_alerts: usize,
pub reset_alerts: usize,
pub approaching_limit_alerts: usize,
}
#[cfg(test)]
#[path = "alerts_tests.rs"]
mod tests;