use std::str::FromStr;
use std::sync::Arc;
use adk_core::{AdkError, Result};
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use cron::Schedule;
use futures::stream::BoxStream;
use tokio::time::sleep;
use super::event_source::{EventSource, TriggerEvent};
use super::watermark::TickWatermark;
const DEFAULT_MAX_CATCH_UP: usize = 64;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum MissedTickPolicy {
#[default]
Skip,
CoalesceOne,
All,
}
pub struct CronTrigger {
expression: String,
schedule: Schedule,
name: String,
missed_tick_policy: MissedTickPolicy,
watermark: Option<Arc<dyn TickWatermark>>,
max_catch_up: usize,
}
impl CronTrigger {
pub fn new(expression: &str) -> Result<Self> {
let schedule = Schedule::from_str(expression)
.map_err(|e| AdkError::agent(format!("invalid cron expression: {e}")))?;
Ok(Self {
expression: expression.to_string(),
schedule,
name: format!("cron:{expression}"),
missed_tick_policy: MissedTickPolicy::default(),
watermark: None,
max_catch_up: DEFAULT_MAX_CATCH_UP,
})
}
pub fn with_missed_tick_policy(mut self, policy: MissedTickPolicy) -> Self {
self.missed_tick_policy = policy;
self
}
pub fn with_watermark(mut self, watermark: Arc<dyn TickWatermark>) -> Self {
self.watermark = Some(watermark);
self
}
pub fn with_max_catch_up(mut self, max_catch_up: usize) -> Self {
self.max_catch_up = max_catch_up.max(1);
self
}
pub fn missed_tick_policy(&self) -> MissedTickPolicy {
self.missed_tick_policy
}
}
fn elapsed_ticks(
schedule: &Schedule,
cursor: DateTime<Utc>,
now: DateTime<Utc>,
cap: usize,
) -> (Vec<DateTime<Utc>>, bool) {
let mut ticks = Vec::new();
let mut truncated = false;
for tick in schedule.after(&cursor).take_while(|tick| *tick <= now) {
if ticks.len() == cap {
truncated = true;
break;
}
ticks.push(tick);
}
(ticks, truncated)
}
#[async_trait]
impl EventSource for CronTrigger {
fn name(&self) -> &str {
&self.name
}
async fn subscribe(&self) -> Result<BoxStream<'static, TriggerEvent>> {
let schedule = self.schedule.clone();
let source_name = self.name.clone();
let expression = self.expression.clone();
let policy = self.missed_tick_policy;
let watermark = self.watermark.clone();
let cap = self.max_catch_up;
let restored = match (&watermark, policy) {
(Some(store), policy) if policy != MissedTickPolicy::Skip => store.read().await?,
_ => None,
};
let stream = async_stream::stream! {
let mut cursor = restored.unwrap_or_else(Utc::now);
loop {
if policy != MissedTickPolicy::Skip {
let now = Utc::now();
let (missed, truncated) = elapsed_ticks(&schedule, cursor, now, cap);
if !missed.is_empty() {
let missed_count = missed.len();
if truncated {
tracing::warn!(
source = %source_name,
replayed = missed_count,
"catch-up cap reached; discarding the remainder of the missed span"
);
}
match policy {
MissedTickPolicy::Skip => unreachable!("guarded above"),
MissedTickPolicy::CoalesceOne => {
let scheduled_for = missed[missed_count - 1];
tracing::info!(
source = %source_name,
missed = missed_count,
"replaying missed ticks as one coalesced event"
);
let accounted_through =
if truncated { now } else { scheduled_for };
if let Some(ref store) = watermark
&& let Err(error) = store.write(accounted_through).await
{
tracing::error!(%error, "failed to persist tick watermark; stopping cron stream before emission");
return;
}
yield TriggerEvent {
source: source_name.clone(),
payload: serde_json::json!({
"expression": expression,
"tick": Utc::now().to_rfc3339(),
"scheduled_for": scheduled_for.to_rfc3339(),
"catch_up": true,
"missed_count": missed_count,
"missed_count_truncated": truncated,
}),
principal: None,
};
}
MissedTickPolicy::All => {
for (index, scheduled_for) in missed.iter().enumerate() {
let accounted_through = if truncated && index + 1 == missed_count {
now
} else {
*scheduled_for
};
if let Some(ref store) = watermark
&& let Err(error) = store.write(accounted_through).await
{
tracing::error!(%error, "failed to persist tick watermark; stopping cron stream before emission");
return;
}
yield TriggerEvent {
source: source_name.clone(),
payload: serde_json::json!({
"expression": expression,
"tick": Utc::now().to_rfc3339(),
"scheduled_for": scheduled_for.to_rfc3339(),
"catch_up": true,
}),
principal: None,
};
}
}
}
}
cursor = if truncated { now } else { missed.last().copied().unwrap_or(cursor) };
} else {
cursor = Utc::now();
}
let Some(next_tick) = schedule.after(&cursor).next() else {
break;
};
let duration = (next_tick - Utc::now()).to_std().unwrap_or_default();
sleep(duration).await;
cursor = next_tick;
if policy != MissedTickPolicy::Skip
&& let Some(ref store) = watermark
&& let Err(error) = store.write(next_tick).await
{
tracing::error!(%error, "failed to persist tick watermark; stopping cron stream before emission");
return;
}
yield TriggerEvent {
source: source_name.clone(),
payload: serde_json::json!({
"expression": expression,
"tick": Utc::now().to_rfc3339(),
"scheduled_for": next_tick.to_rfc3339(),
"catch_up": false,
}),
principal: None,
};
}
};
Ok(Box::pin(stream))
}
}
impl std::fmt::Debug for CronTrigger {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CronTrigger")
.field("expression", &self.expression)
.field("missed_tick_policy", &self.missed_tick_policy)
.field("watermark", &self.watermark.is_some())
.field("max_catch_up", &self.max_catch_up)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::TimeZone;
fn every_minute() -> Schedule {
Schedule::from_str("0 * * * * *").expect("valid expression")
}
#[test]
fn policy_defaults_to_skip() {
let trigger = CronTrigger::new("0 * * * * *").expect("valid expression");
assert_eq!(trigger.missed_tick_policy(), MissedTickPolicy::Skip);
}
#[test]
fn invalid_expressions_are_rejected() {
assert!(CronTrigger::new("not a cron expression").is_err());
}
#[test]
fn max_catch_up_of_zero_is_treated_as_one() {
let trigger =
CronTrigger::new("0 * * * * *").expect("valid expression").with_max_catch_up(0);
assert_eq!(trigger.max_catch_up, 1);
}
#[test]
fn elapsed_ticks_finds_every_tick_in_the_gap() {
let cursor = Utc.with_ymd_and_hms(2026, 8, 22, 10, 0, 0).unwrap();
let now = Utc.with_ymd_and_hms(2026, 8, 22, 10, 5, 0).unwrap();
let (ticks, truncated) = elapsed_ticks(&every_minute(), cursor, now, 64);
assert_eq!(ticks.len(), 5, "10:01 through 10:05 inclusive");
assert!(!truncated);
assert_eq!(ticks[0], Utc.with_ymd_and_hms(2026, 8, 22, 10, 1, 0).unwrap());
assert_eq!(ticks[4], now);
}
#[test]
fn elapsed_ticks_is_empty_when_no_tick_has_come_due() {
let cursor = Utc.with_ymd_and_hms(2026, 8, 22, 10, 0, 0).unwrap();
let now = Utc.with_ymd_and_hms(2026, 8, 22, 10, 0, 30).unwrap();
let (ticks, truncated) = elapsed_ticks(&every_minute(), cursor, now, 64);
assert!(ticks.is_empty());
assert!(!truncated);
}
#[test]
fn elapsed_ticks_reports_truncation_at_the_cap() {
let cursor = Utc.with_ymd_and_hms(2026, 8, 22, 10, 0, 0).unwrap();
let now = Utc.with_ymd_and_hms(2026, 8, 22, 20, 0, 0).unwrap();
let (ticks, truncated) = elapsed_ticks(&every_minute(), cursor, now, 10);
assert_eq!(ticks.len(), 10, "capped rather than replaying ten hours");
assert!(truncated);
}
#[test]
fn elapsed_ticks_excludes_the_cursor_itself() {
let cursor = Utc.with_ymd_and_hms(2026, 8, 22, 10, 0, 0).unwrap();
let now = Utc.with_ymd_and_hms(2026, 8, 22, 10, 1, 0).unwrap();
let (ticks, _) = elapsed_ticks(&every_minute(), cursor, now, 64);
assert_eq!(ticks, vec![now], "the cursor tick already fired");
}
}