#![allow(deprecated)]
use std::collections::HashMap;
use std::time::Duration;
use reinhardt_core::macros::settings;
use serde::{Deserialize, Serialize};
use crate::webhook::{HttpWebhookSender, RetryConfig, WebhookConfig};
use crate::worker::{Worker, WorkerConfig};
fn default_queue_name() -> String {
"default".to_string()
}
fn default_max_retries() -> u32 {
3
}
fn default_worker_name() -> String {
"worker".to_string()
}
fn default_concurrency() -> usize {
4
}
fn default_poll_interval_ms() -> u64 {
1000
}
fn default_webhook_method() -> String {
"POST".to_string()
}
fn default_webhook_timeout_secs() -> u64 {
5
}
fn default_retry_max_retries() -> u32 {
3
}
fn default_retry_initial_backoff_ms() -> u64 {
100
}
fn default_retry_max_backoff_ms() -> u64 {
30_000
}
fn default_retry_backoff_multiplier() -> f64 {
2.0
}
#[settings(fragment = true, section = "tasks_queue")]
#[non_exhaustive]
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct QueueSettings {
#[serde(default = "default_queue_name")]
pub name: String,
#[serde(default = "default_max_retries")]
pub max_retries: u32,
}
impl Default for QueueSettings {
fn default() -> Self {
Self {
name: default_queue_name(),
max_retries: default_max_retries(),
}
}
}
#[settings(fragment = true)]
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct WebhookRetrySettings {
#[serde(default = "default_retry_max_retries")]
pub max_retries: u32,
#[serde(default = "default_retry_initial_backoff_ms")]
pub initial_backoff_ms: u64,
#[serde(default = "default_retry_max_backoff_ms")]
pub max_backoff_ms: u64,
#[serde(default = "default_retry_backoff_multiplier")]
pub backoff_multiplier: f64,
}
impl Default for WebhookRetrySettings {
fn default() -> Self {
Self {
max_retries: default_retry_max_retries(),
initial_backoff_ms: default_retry_initial_backoff_ms(),
max_backoff_ms: default_retry_max_backoff_ms(),
backoff_multiplier: default_retry_backoff_multiplier(),
}
}
}
impl From<&WebhookRetrySettings> for RetryConfig {
fn from(settings: &WebhookRetrySettings) -> Self {
Self {
max_retries: settings.max_retries,
initial_backoff: Duration::from_millis(settings.initial_backoff_ms),
max_backoff: Duration::from_millis(settings.max_backoff_ms),
backoff_multiplier: settings.backoff_multiplier,
}
}
}
#[settings(fragment = true, section = "tasks_webhook")]
#[non_exhaustive]
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct WebhookSettings {
#[serde(default)]
pub url: String,
#[serde(default = "default_webhook_method")]
pub method: String,
#[serde(default)]
pub headers: HashMap<String, String>,
#[serde(default = "default_webhook_timeout_secs")]
pub timeout_secs: u64,
#[setting(node)]
#[serde(default)]
pub retry: WebhookRetrySettings,
}
impl Default for WebhookSettings {
fn default() -> Self {
Self {
url: String::new(),
method: default_webhook_method(),
headers: HashMap::new(),
timeout_secs: default_webhook_timeout_secs(),
retry: WebhookRetrySettings::default(),
}
}
}
impl From<&WebhookSettings> for WebhookConfig {
fn from(settings: &WebhookSettings) -> Self {
Self {
url: settings.url.clone(),
method: settings.method.clone(),
headers: settings.headers.clone(),
timeout: Duration::from_secs(settings.timeout_secs),
retry_config: RetryConfig::from(&settings.retry),
}
}
}
pub fn create_webhook_sender_from_settings(settings: &WebhookSettings) -> HttpWebhookSender {
HttpWebhookSender::new(WebhookConfig::from(settings))
}
#[settings(fragment = true, section = "tasks_worker")]
#[non_exhaustive]
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct WorkerSettings {
#[serde(default = "default_worker_name")]
pub name: String,
#[serde(default = "default_concurrency")]
pub concurrency: usize,
#[serde(default = "default_poll_interval_ms")]
pub poll_interval_ms: u64,
#[setting(node)]
#[serde(default)]
pub webhooks: Vec<WebhookSettings>,
}
impl Default for WorkerSettings {
fn default() -> Self {
Self {
name: default_worker_name(),
concurrency: default_concurrency(),
poll_interval_ms: default_poll_interval_ms(),
webhooks: Vec::new(),
}
}
}
impl From<&WorkerSettings> for WorkerConfig {
fn from(settings: &WorkerSettings) -> Self {
Self {
name: settings.name.clone(),
concurrency: settings.concurrency,
poll_interval: Duration::from_millis(settings.poll_interval_ms),
webhook_configs: settings.webhooks.iter().map(WebhookConfig::from).collect(),
}
}
}
pub fn create_worker_from_settings(settings: &WorkerSettings) -> Worker {
Worker::new(WorkerConfig::from(settings))
}
#[cfg(feature = "sqs-backend")]
fn default_sqs_visibility_timeout() -> i32 {
30
}
#[cfg(feature = "sqs-backend")]
fn default_sqs_max_messages() -> i32 {
1
}
#[cfg(feature = "sqs-backend")]
fn default_sqs_wait_time_seconds() -> i32 {
0
}
#[cfg(feature = "sqs-backend")]
#[settings(fragment = true, section = "tasks_sqs")]
#[non_exhaustive]
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct SqsSettings {
#[serde(default)]
pub queue_url: String,
#[serde(default = "default_sqs_visibility_timeout")]
pub visibility_timeout: i32,
#[serde(default = "default_sqs_max_messages")]
pub max_messages: i32,
#[serde(default = "default_sqs_wait_time_seconds")]
pub wait_time_seconds: i32,
}
#[cfg(feature = "sqs-backend")]
impl Default for SqsSettings {
fn default() -> Self {
Self {
queue_url: String::new(),
visibility_timeout: default_sqs_visibility_timeout(),
max_messages: default_sqs_max_messages(),
wait_time_seconds: default_sqs_wait_time_seconds(),
}
}
}
#[cfg(feature = "sqs-backend")]
impl From<&SqsSettings> for crate::backends::sqs::SqsConfig {
fn from(settings: &SqsSettings) -> Self {
crate::backends::sqs::SqsConfig::new(settings.queue_url.clone())
.with_visibility_timeout(settings.visibility_timeout)
.with_max_messages(settings.max_messages)
.with_wait_time_seconds(settings.wait_time_seconds)
}
}
#[cfg(feature = "sqs-backend")]
pub async fn create_sqs_backend_from_settings(
settings: &SqsSettings,
) -> Result<crate::backends::sqs::SqsBackend, crate::TaskExecutionError> {
crate::backends::sqs::SqsBackend::new(crate::backends::sqs::SqsConfig::from(settings)).await
}
#[cfg(feature = "rabbitmq-backend")]
fn default_rabbitmq_url() -> String {
"amqp://localhost:5672/%2f".to_string()
}
#[cfg(feature = "rabbitmq-backend")]
fn default_rabbitmq_queue_name() -> String {
"reinhardt_tasks".to_string()
}
#[cfg(feature = "rabbitmq-backend")]
fn default_rabbitmq_routing_key() -> String {
"reinhardt_tasks".to_string()
}
#[cfg(feature = "rabbitmq-backend")]
#[settings(fragment = true, section = "tasks_rabbitmq")]
#[non_exhaustive]
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct RabbitMQSettings {
#[serde(default = "default_rabbitmq_url")]
pub url: String,
#[serde(default = "default_rabbitmq_queue_name")]
pub queue_name: String,
#[serde(default)]
pub exchange_name: String,
#[serde(default = "default_rabbitmq_routing_key")]
pub routing_key: String,
}
#[cfg(feature = "rabbitmq-backend")]
impl Default for RabbitMQSettings {
fn default() -> Self {
Self {
url: default_rabbitmq_url(),
queue_name: default_rabbitmq_queue_name(),
exchange_name: String::new(),
routing_key: default_rabbitmq_routing_key(),
}
}
}
#[cfg(feature = "rabbitmq-backend")]
impl From<&RabbitMQSettings> for crate::backends::rabbitmq::RabbitMQConfig {
fn from(settings: &RabbitMQSettings) -> Self {
Self {
url: settings.url.clone(),
queue_name: settings.queue_name.clone(),
exchange_name: settings.exchange_name.clone(),
routing_key: settings.routing_key.clone(),
}
}
}
#[cfg(feature = "rabbitmq-backend")]
pub async fn create_rabbitmq_backend_from_settings(
settings: &RabbitMQSettings,
) -> Result<crate::backends::rabbitmq::RabbitMQBackend, lapin::Error> {
crate::backends::rabbitmq::RabbitMQBackend::new(
crate::backends::rabbitmq::RabbitMQConfig::from(settings),
)
.await
}
#[cfg(test)]
mod tests {
use super::*;
use reinhardt_conf::settings::fragment::SettingsFragment;
#[rstest::rstest]
fn section_names_are_crate_prefixed() {
assert_eq!(QueueSettings::section(), "tasks_queue");
assert_eq!(WorkerSettings::section(), "tasks_worker");
assert_eq!(WebhookSettings::section(), "tasks_webhook");
}
#[rstest::rstest]
fn queue_settings_default_has_expected_values() {
let settings = QueueSettings::default();
assert_eq!(settings.name, "default");
assert_eq!(settings.max_retries, 3);
}
#[rstest::rstest]
fn worker_settings_convert_milliseconds_to_duration() {
let settings = WorkerSettings {
name: "ingest".to_string(),
concurrency: 8,
poll_interval_ms: 2500,
webhooks: Vec::new(),
};
let config = WorkerConfig::from(&settings);
assert_eq!(config.name, "ingest");
assert_eq!(config.concurrency, 8);
assert_eq!(config.poll_interval, Duration::from_millis(2500));
assert!(config.webhook_configs.is_empty());
}
#[rstest::rstest]
fn webhook_settings_convert_seconds_and_nested_retry() {
let settings = WebhookSettings::default();
let config = WebhookConfig::from(&settings);
assert_eq!(config.method, "POST");
assert_eq!(config.timeout, Duration::from_secs(5));
assert_eq!(config.retry_config.max_retries, 3);
assert_eq!(
config.retry_config.initial_backoff,
Duration::from_millis(100)
);
assert_eq!(config.retry_config.max_backoff, Duration::from_secs(30));
assert_eq!(config.retry_config.backoff_multiplier, 2.0);
}
#[rstest::rstest]
fn worker_settings_map_nested_webhooks() {
let settings = WorkerSettings {
name: "w".to_string(),
concurrency: 1,
poll_interval_ms: 100,
webhooks: vec![WebhookSettings {
url: "https://example.com/hook".to_string(),
..WebhookSettings::default()
}],
};
let config = WorkerConfig::from(&settings);
assert_eq!(config.webhook_configs.len(), 1);
assert_eq!(config.webhook_configs[0].url, "https://example.com/hook");
}
#[rstest::rstest]
fn webhook_settings_deserialize_with_defaults() {
let json = r#"{ "url": "https://example.com/hook", "timeout_secs": 10 }"#;
let settings: WebhookSettings = serde_json::from_str(json).unwrap();
let config = WebhookConfig::from(&settings);
assert_eq!(config.url, "https://example.com/hook");
assert_eq!(config.method, "POST");
assert_eq!(config.timeout, Duration::from_secs(10));
assert_eq!(config.retry_config.max_retries, 3);
}
#[cfg(feature = "sqs-backend")]
#[rstest::rstest]
fn sqs_settings_default_converts_to_config() {
let settings = SqsSettings {
queue_url: "https://sqs.example.com/q".to_string(),
..SqsSettings::default()
};
let config = crate::backends::sqs::SqsConfig::from(&settings);
assert!(format!("{config:?}").contains("https://sqs.example.com/q"));
}
#[cfg(feature = "rabbitmq-backend")]
#[rstest::rstest]
fn rabbitmq_settings_default_converts_to_config() {
let settings = RabbitMQSettings::default();
let config = crate::backends::rabbitmq::RabbitMQConfig::from(&settings);
assert_eq!(config.queue_name, "reinhardt_tasks");
assert_eq!(config.routing_key, "reinhardt_tasks");
}
}