1use std::collections::HashMap;
4use std::sync::{Arc, Mutex};
5use std::time::{Duration, Instant};
6
7use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc};
8use tokio::time::{sleep, timeout};
9use tokio::{select, spawn};
10use tokio_util::sync::CancellationToken;
11use tracing::{error, info, warn};
12use uuid::Uuid;
13
14use ironflow_core::account_strategy::{AccountStrategy, LeastUtilized};
15use ironflow_core::decision::DecisionProvider;
16#[cfg(feature = "prometheus")]
17use ironflow_core::metric_names::{WORKER_ACTIVE, WORKER_LEASES_LOST_TOTAL, WORKER_POLLS_TOTAL};
18use ironflow_core::provider::AgentProvider;
19use ironflow_engine::accounts::{AccountAwareProvider, DEFAULT_MAX_CAPACITY_WAIT};
20use ironflow_engine::engine::{Engine, chain_root};
21use ironflow_engine::error::EngineError;
22use ironflow_engine::handler::WorkflowHandler;
23use ironflow_engine::log_sender::LogReceiver;
24use ironflow_store::entities::{
25 LeaseRequest, RunStatus, WorkerCapabilities, normalize_worker_tags, validate_worker_tags,
26};
27use ironflow_store::error::StoreError;
28use ironflow_store::store::Store;
29#[cfg(feature = "prometheus")]
30use metrics::{counter, gauge};
31#[cfg(feature = "heartbeat")]
32use reqwest::Client;
33
34use crate::api_store::ApiRunStore;
35use crate::artifact_sink::ApiArtifactSink;
36use crate::error::WorkerError;
37use crate::log_pusher::LogPusher;
38#[cfg(feature = "prometheus")]
39use crate::queue_depth::QueueDepthGauge;
40
41const DEFAULT_CONCURRENCY: usize = 2;
42const DEFAULT_POLL_INTERVAL: Duration = Duration::from_secs(2);
43const DEFAULT_RUN_TIMEOUT: Duration = Duration::from_secs(30 * 60);
44const DEFAULT_MAX_CONSECUTIVE_PANICS: u32 = 3;
45const DEFAULT_PANIC_COOLDOWN: Duration = Duration::from_secs(5 * 60);
46const DEFAULT_LEASE_TTL: Duration = Duration::from_secs(90);
49const DEFAULT_LEASE_REFRESH_INTERVAL: Duration = Duration::from_secs(30);
51#[cfg(feature = "heartbeat")]
52const DEFAULT_HEARTBEAT_INTERVAL: Duration = Duration::from_secs(30);
53
54pub struct WorkerBuilder {
78 api_url: String,
79 worker_token: String,
80 worker_id: String,
81 provider: Option<Arc<dyn AgentProvider>>,
82 decision_provider: Option<Arc<dyn DecisionProvider>>,
83 account_strategy: Option<Arc<dyn AccountStrategy>>,
84 max_capacity_wait: Duration,
85 handlers: Vec<Box<dyn WorkflowHandler>>,
86 concurrency: usize,
87 poll_interval: Duration,
88 run_timeout: Duration,
89 max_consecutive_panics: u32,
90 panic_cooldown: Duration,
91 lease_ttl: Duration,
92 tags: Vec<String>,
93 lease_refresh_interval: Duration,
94 #[cfg(feature = "heartbeat")]
95 heartbeat_url: Option<String>,
96 #[cfg(feature = "heartbeat")]
97 heartbeat_interval: Duration,
98}
99
100impl WorkerBuilder {
101 pub fn new(api_url: &str, worker_token: &str) -> Self {
103 Self {
104 api_url: api_url.to_string(),
105 worker_token: worker_token.to_string(),
106 worker_id: format!("worker-{}", Uuid::now_v7()),
107 provider: None,
108 decision_provider: None,
109 account_strategy: None,
110 max_capacity_wait: DEFAULT_MAX_CAPACITY_WAIT,
111 handlers: Vec::new(),
112 concurrency: DEFAULT_CONCURRENCY,
113 poll_interval: DEFAULT_POLL_INTERVAL,
114 run_timeout: DEFAULT_RUN_TIMEOUT,
115 max_consecutive_panics: DEFAULT_MAX_CONSECUTIVE_PANICS,
116 panic_cooldown: DEFAULT_PANIC_COOLDOWN,
117 lease_ttl: DEFAULT_LEASE_TTL,
118 lease_refresh_interval: DEFAULT_LEASE_REFRESH_INTERVAL,
119 tags: Vec::new(),
120 #[cfg(feature = "heartbeat")]
121 heartbeat_url: None,
122 #[cfg(feature = "heartbeat")]
123 heartbeat_interval: DEFAULT_HEARTBEAT_INTERVAL,
124 }
125 }
126
127 pub fn provider(mut self, provider: Arc<dyn AgentProvider>) -> Self {
129 self.provider = Some(provider);
130 self
131 }
132
133 pub fn decision_provider(mut self, provider: Arc<dyn DecisionProvider>) -> Self {
159 self.decision_provider = Some(provider);
160 self
161 }
162
163 pub fn account_strategy(mut self, strategy: Arc<dyn AccountStrategy>) -> Self {
186 self.account_strategy = Some(strategy);
187 self
188 }
189
190 pub fn max_capacity_wait(mut self, wait: Duration) -> Self {
215 self.max_capacity_wait = wait;
216 self
217 }
218
219 pub fn register(mut self, handler: impl WorkflowHandler + 'static) -> Self {
221 self.handlers.push(Box::new(handler));
222 self
223 }
224
225 pub fn concurrency(mut self, n: usize) -> Self {
227 self.concurrency = n;
228 self
229 }
230
231 pub fn poll_interval(mut self, interval: Duration) -> Self {
233 self.poll_interval = interval;
234 self
235 }
236
237 pub fn run_timeout(mut self, timeout: Duration) -> Self {
254 self.run_timeout = timeout;
255 self
256 }
257
258 pub fn max_consecutive_panics(mut self, n: u32) -> Self {
276 self.max_consecutive_panics = n;
277 self
278 }
279
280 pub fn panic_cooldown(mut self, cooldown: Duration) -> Self {
297 self.panic_cooldown = cooldown;
298 self
299 }
300
301 pub fn worker_id(mut self, worker_id: &str) -> Self {
318 self.worker_id = worker_id.to_string();
319 self
320 }
321
322 pub fn tags<I, S>(mut self, tags: I) -> Self
345 where
346 I: IntoIterator<Item = S>,
347 S: Into<String>,
348 {
349 self.tags.extend(tags.into_iter().map(Into::into));
350 self
351 }
352
353 pub fn lease_ttl(mut self, ttl: Duration) -> Self {
371 self.lease_ttl = ttl;
372 self
373 }
374
375 pub fn lease_refresh_interval(mut self, interval: Duration) -> Self {
392 self.lease_refresh_interval = interval;
393 self
394 }
395
396 #[cfg(feature = "heartbeat")]
415 pub fn heartbeat_url(mut self, url: &str) -> Self {
416 self.heartbeat_url = Some(url.to_string());
417 self
418 }
419
420 #[cfg(feature = "heartbeat")]
438 pub fn heartbeat_interval(mut self, interval: Duration) -> Self {
439 self.heartbeat_interval = interval;
440 self
441 }
442
443 pub fn build(self) -> Result<Worker, WorkerError> {
452 let provider = self
453 .provider
454 .ok_or_else(|| WorkerError::Internal("WorkerBuilder: provider is required".into()))?;
455 validate_worker_tags(&self.tags).map_err(EngineError::InvalidWorkerTag)?;
456 let tags = normalize_worker_tags(self.tags);
457
458 let store: Arc<dyn Store> = Arc::new(ApiRunStore::new(&self.api_url, &self.worker_token));
459 let strategy = self
460 .account_strategy
461 .unwrap_or_else(|| Arc::new(LeastUtilized));
462 let provider: Arc<dyn AgentProvider> = Arc::new(
463 AccountAwareProvider::new(provider, store.clone())
464 .with_strategy(strategy)
465 .with_max_capacity_wait(self.max_capacity_wait),
466 );
467
468 let mut engine = Engine::new(store, provider);
469 if let Some(decision_provider) = self.decision_provider {
470 engine = engine.with_decision_provider(decision_provider);
471 }
472 for handler in self.handlers {
473 engine
474 .register_boxed(handler)
475 .map_err(WorkerError::Engine)?;
476 }
477 let capabilities = WorkerCapabilities::new(registered_workflows(&engine), tags.clone());
478 engine.set_worker_tags(tags);
479
480 let (log_sender, log_receiver) = ironflow_engine::log_sender::channel();
481 engine.set_log_sender(log_sender);
482
483 engine.set_artifact_sink(Arc::new(ApiArtifactSink::new(
487 &self.api_url,
488 &self.worker_token,
489 )));
490
491 #[cfg(feature = "heartbeat")]
492 let heartbeat_client = Client::builder()
493 .timeout(Duration::from_secs(5))
494 .build()
495 .expect("failed to build heartbeat HTTP client");
496
497 Ok(Worker {
498 engine: Arc::new(engine),
499 api_url: self.api_url,
500 worker_token: self.worker_token,
501 worker_id: self.worker_id,
502 capabilities,
503 log_receiver: Mutex::new(Some(log_receiver)),
504 concurrency: self.concurrency,
505 poll_interval: self.poll_interval,
506 run_timeout: self.run_timeout,
507 max_consecutive_panics: self.max_consecutive_panics,
508 panic_cooldown: self.panic_cooldown,
509 lease_ttl: self.lease_ttl,
510 lease_refresh_interval: self.lease_refresh_interval,
511 #[cfg(feature = "heartbeat")]
512 heartbeat_url: self.heartbeat_url,
513 #[cfg(feature = "heartbeat")]
514 heartbeat_interval: self.heartbeat_interval,
515 #[cfg(feature = "heartbeat")]
516 heartbeat_client,
517 })
518 }
519}
520
521pub struct Worker {
523 engine: Arc<Engine>,
524 api_url: String,
525 worker_token: String,
526 worker_id: String,
527 capabilities: WorkerCapabilities,
528 log_receiver: Mutex<Option<LogReceiver>>,
529 concurrency: usize,
530 poll_interval: Duration,
531 run_timeout: Duration,
532 max_consecutive_panics: u32,
533 panic_cooldown: Duration,
534 lease_ttl: Duration,
535 lease_refresh_interval: Duration,
536 #[cfg(feature = "heartbeat")]
537 heartbeat_url: Option<String>,
538 #[cfg(feature = "heartbeat")]
539 heartbeat_interval: Duration,
540 #[cfg(feature = "heartbeat")]
541 heartbeat_client: Client,
542}
543
544struct PoisonPillTracker {
546 max_consecutive: u32,
547 cooldown: Duration,
548 state: HashMap<String, (u32, Instant)>,
550}
551
552impl PoisonPillTracker {
553 fn new(max_consecutive: u32, cooldown: Duration) -> Self {
554 Self {
555 max_consecutive,
556 cooldown,
557 state: HashMap::new(),
558 }
559 }
560
561 fn record_panic(&mut self, workflow: &str) -> bool {
564 let entry = self
565 .state
566 .entry(workflow.to_string())
567 .or_insert((0, Instant::now()));
568 entry.0 += 1;
569 entry.1 = Instant::now();
570 entry.0 >= self.max_consecutive
571 }
572
573 fn record_success(&mut self, workflow: &str) {
575 self.state.remove(workflow);
576 }
577
578 fn is_blocked(&self, workflow: &str) -> bool {
580 self.state.get(workflow).is_some_and(|(count, last_panic)| {
581 *count >= self.max_consecutive && last_panic.elapsed() < self.cooldown
582 })
583 }
584}
585
586impl Worker {
587 pub fn capabilities(&self) -> &WorkerCapabilities {
607 &self.capabilities
608 }
609
610 pub async fn run(&self) -> Result<(), WorkerError> {
619 let semaphore = Arc::new(Semaphore::new(self.concurrency));
620 let shutdown = CancellationToken::new();
621 let mut idle_streak = 0u32;
622 let poison_tracker = Arc::new(Mutex::new(PoisonPillTracker::new(
623 self.max_consecutive_panics,
624 self.panic_cooldown,
625 )));
626 let (outcome_tx, mut outcome_rx) = mpsc::unbounded_channel::<RunOutcome>();
627
628 info!(
629 concurrency = self.concurrency,
630 poll_interval_ms = self.poll_interval.as_millis() as u64,
631 run_timeout_secs = self.run_timeout.as_secs(),
632 "worker started"
633 );
634
635 if let Some(receiver) = self.log_receiver.lock().expect("log_receiver lock").take() {
636 let pusher = LogPusher::new(&self.api_url, &self.worker_token);
637 spawn(pusher.run(receiver));
638 info!("log pusher started");
639 }
640
641 let shutdown_clone = shutdown.clone();
643 spawn(async move {
644 shutdown_signal().await;
645 info!("shutdown signal received, draining in-flight runs...");
646 shutdown_clone.cancel();
647 });
648
649 #[cfg(feature = "heartbeat")]
650 if let Some(ref url) = self.heartbeat_url {
651 let interval = self.heartbeat_interval;
652 let url = url.clone();
653 let client = self.heartbeat_client.clone();
654
655 spawn(async move {
656 let mut ticker = tokio::time::interval(interval);
657 ticker.tick().await;
659 loop {
660 ticker.tick().await;
661 match client.head(&url).send().await {
662 Ok(resp) if resp.status().is_success() => {
663 info!(url = %url, "heartbeat sent");
664 }
665 Ok(resp) => {
666 warn!(
667 url = %url,
668 status = %resp.status(),
669 "heartbeat ping returned non-success status"
670 );
671 }
672 Err(err) => {
673 warn!(
674 url = %url,
675 error = %err,
676 "heartbeat ping failed"
677 );
678 }
679 }
680 }
681 });
682 }
683
684 #[cfg(feature = "prometheus")]
685 let mut queue_depth =
686 QueueDepthGauge::new(&self.api_url, &self.worker_token, &self.capabilities);
687
688 while !shutdown.is_cancelled() {
689 while let Ok(outcome) = outcome_rx.try_recv() {
691 let mut tracker = poison_tracker.lock().expect("poison tracker lock poisoned");
692 match outcome {
693 RunOutcome::Success(ref wf) => tracker.record_success(wf),
694 RunOutcome::LeaseLost(ref wf) => {
697 warn!(workflow = %wf, "run abandoned after losing its lease")
698 }
699 RunOutcome::Failed(ref wf) | RunOutcome::Timeout(ref wf) => {
700 if tracker.record_panic(wf) {
701 warn!(workflow = %wf, "workflow flagged as poison pill after consecutive failures");
702 }
703 }
704 RunOutcome::Panicked(ref wf) => {
705 if tracker.record_panic(wf) {
706 error!(workflow = %wf, "workflow flagged as poison pill after consecutive panics");
707 }
708 }
709 }
710 }
711
712 let Some(permit) = acquire_slot(&semaphore, &shutdown).await? else {
715 break;
716 };
717
718 let run = self
719 .engine
720 .store()
721 .pick_next_pending_for(
722 Some(LeaseRequest {
723 worker_id: self.worker_id.clone(),
724 ttl: self.lease_ttl,
725 }),
726 Some(self.capabilities.clone()),
727 )
728 .await;
729
730 match run {
731 Ok(Some(run)) => {
732 #[cfg(feature = "prometheus")]
733 counter!(WORKER_POLLS_TOTAL, "result" => "hit").increment(1);
734
735 let is_blocked = {
737 let tracker = poison_tracker.lock().expect("poison tracker lock poisoned");
738 tracker.is_blocked(&run.workflow_name)
739 };
740 if is_blocked {
741 warn!(
742 workflow = %run.workflow_name,
743 run_id = %run.id,
744 "skipping run: workflow flagged as poison pill, marking as failed"
745 );
746 if let Err(e) = self
747 .engine
748 .store()
749 .update_run_status(run.id, RunStatus::Failed)
750 .await
751 {
752 error!(run_id = %run.id, error = %e, "failed to mark poisoned run as failed");
753 }
754 drop(permit);
755 continue;
756 }
757
758 idle_streak = 0;
759 let engine = self.engine.clone();
760 let run_id = run.id;
761 let workflow = run.workflow_name.clone();
762 let workflow_for_watcher = workflow.clone();
763 let run_timeout = self.run_timeout;
764
765 info!(run_id = %run_id, workflow = %workflow, "executing run");
766
767 #[cfg(feature = "prometheus")]
768 gauge!(WORKER_ACTIVE).increment(1.0);
769
770 let lease_token = CancellationToken::new();
773 let refresher = spawn(refresh_lease(
774 self.engine.store().clone(),
775 run_id,
776 LeaseRequest {
777 worker_id: self.worker_id.clone(),
778 ttl: self.lease_ttl,
779 },
780 self.lease_refresh_interval,
781 lease_token.clone(),
782 ));
783
784 let handle = spawn(async move {
785 let _permit = permit;
786 let result = select! {
787 biased;
788 _ = lease_token.cancelled() => {
789 refresher.abort();
790 warn!(
793 run_id = %run_id,
794 workflow = %workflow,
795 "abandoning run: worker lease lost"
796 );
797 #[cfg(feature = "prometheus")]
798 counter!(WORKER_LEASES_LOST_TOTAL).increment(1);
799 return RunOutcome::LeaseLost(workflow);
801 }
802 result = timeout(run_timeout, engine.execute_handler_run(run_id)) => result,
803 };
804 refresher.abort();
805
806 match result {
807 Ok(Ok(_)) => {
808 info!(run_id = %run_id, workflow = %workflow, "run completed");
809 RunOutcome::Success(workflow)
810 }
811 Ok(Err(e)) => {
819 error!(run_id = %run_id, workflow = %workflow, error = %e, "run failed");
820 RunOutcome::Failed(workflow)
821 }
822 Err(_) => {
823 error!(
824 run_id = %run_id,
825 workflow = %workflow,
826 timeout_secs = run_timeout.as_secs(),
827 "run timed out"
828 );
829 let timeout_msg =
830 format!("run timed out after {}s", run_timeout.as_secs());
831 if let Err(e) = engine
833 .fail_or_schedule_retry(run_id, &timeout_msg, true, None, None)
834 .await
835 {
836 error!(run_id = %run_id, error = %e, "failed to record timed-out run");
837 }
838 RunOutcome::Timeout(workflow)
839 }
840 }
841 });
842
843 let watcher_engine = self.engine.clone();
845 let tx = outcome_tx.clone();
846 spawn(async move {
847 match handle.await {
848 Ok(outcome) => {
849 let _ = tx.send(outcome);
850 }
851 Err(e) => {
852 error!(run_id = %run_id, "spawned task panicked: {e}");
853 if let Err(store_err) = watcher_engine
857 .fail_or_schedule_retry(
858 run_id,
859 "parent run panicked",
860 true,
861 None,
862 None,
863 )
864 .await
865 {
866 error!(run_id = %run_id, error = %store_err, "failed to record panicked run");
867 }
868 let _ = tx.send(RunOutcome::Panicked(workflow_for_watcher));
869 }
870 }
871 #[cfg(feature = "prometheus")]
872 gauge!(WORKER_ACTIVE).decrement(1.0);
873 });
874 }
875 Ok(None) => {
876 drop(permit);
877 #[cfg(feature = "prometheus")]
878 counter!(WORKER_POLLS_TOTAL, "result" => "miss").increment(1);
879
880 idle_streak += 1;
881 let backoff = if idle_streak > 10 {
882 self.poll_interval * 3
883 } else if idle_streak > 5 {
884 self.poll_interval * 2
885 } else {
886 self.poll_interval
887 };
888 sleep(backoff).await;
889 }
890 Err(e) => {
891 drop(permit);
892 warn!(error = %e, "poll error");
893 sleep(self.poll_interval).await;
894 }
895 }
896
897 #[cfg(feature = "prometheus")]
898 queue_depth.refresh_if_due().await;
899 }
900
901 info!(
903 in_flight = self.concurrency - semaphore.available_permits(),
904 "waiting for in-flight runs to complete..."
905 );
906 let _ = semaphore
907 .acquire_many(self.concurrency as u32)
908 .await
909 .map_err(|_| WorkerError::Shutdown("semaphore closed during drain".to_string()))?;
910
911 info!("all in-flight runs completed, worker shut down");
912 Ok(())
913 }
914}
915
916fn registered_workflows(engine: &Engine) -> Option<Vec<String>> {
923 let mut names: Vec<String> = engine
924 .handler_names()
925 .into_iter()
926 .map(str::to_string)
927 .collect();
928 if let Some(name) = names.iter().find(|name| name.contains(',')) {
929 warn!(
930 workflow = %name,
931 "workflow name contains a comma: the worker will not filter runs by workflow"
932 );
933 return None;
934 }
935 names.sort();
936 Some(names)
937}
938
939enum RunOutcome {
941 Success(String),
943 Failed(String),
945 Timeout(String),
947 Panicked(String),
949 LeaseLost(String),
951}
952
953async fn acquire_slot(
962 semaphore: &Arc<Semaphore>,
963 shutdown: &CancellationToken,
964) -> Result<Option<OwnedSemaphorePermit>, WorkerError> {
965 select! {
966 biased;
967 _ = shutdown.cancelled() => Ok(None),
968 permit = semaphore.clone().acquire_owned() => permit
969 .map(Some)
970 .map_err(|e| WorkerError::Internal(format!("semaphore closed: {e}"))),
971 }
972}
973
974async fn refresh_lease(
986 store: Arc<dyn Store>,
987 run_id: Uuid,
988 lease: LeaseRequest,
989 refresh_interval: Duration,
990 lease_token: CancellationToken,
991) {
992 let ttl = lease.ttl;
993 let mut deadline = Instant::now() + ttl;
994 let mut target = run_id;
995
996 loop {
997 sleep(refresh_interval).await;
998
999 match store.renew_lease(target, lease.clone()).await {
1000 Ok(_) => {
1001 deadline = Instant::now() + ttl;
1002 }
1003 Err(StoreError::LeaseLost { held_by, .. }) => {
1004 let followed = if target == run_id {
1005 follow_lease_to_root(store.as_ref(), run_id, &lease).await
1006 } else {
1007 None
1008 };
1009 if let Some(root) = followed {
1010 info!(
1011 run_id = %run_id,
1012 root_run_id = %root,
1013 "lease followed the run to its root"
1014 );
1015 target = root;
1016 deadline = Instant::now() + ttl;
1017 continue;
1018 }
1019 warn!(
1020 run_id = %target,
1021 held_by = held_by.as_deref().unwrap_or("unknown"),
1022 "lease taken over by another worker"
1023 );
1024 lease_token.cancel();
1025 return;
1026 }
1027 Err(err) if Instant::now() >= deadline => {
1028 warn!(
1031 run_id = %target,
1032 error = %err,
1033 ttl_secs = ttl.as_secs(),
1034 "lease could not be refreshed before it expired"
1035 );
1036 lease_token.cancel();
1037 return;
1038 }
1039 Err(err) => {
1040 warn!(run_id = %target, error = %err, "lease refresh failed, retrying");
1041 }
1042 }
1043 }
1044}
1045
1046async fn follow_lease_to_root(
1053 store: &dyn Store,
1054 run_id: Uuid,
1055 lease: &LeaseRequest,
1056) -> Option<Uuid> {
1057 let run = match store.get_run(run_id).await {
1058 Ok(run) => run?,
1059 Err(err) => {
1060 warn!(run_id = %run_id, error = %err, "could not load run to follow its lease");
1061 return None;
1062 }
1063 };
1064 let root = chain_root(&run)?;
1065 match store.renew_lease(root, lease.clone()).await {
1066 Ok(_) => Some(root),
1067 Err(err) => {
1068 warn!(
1069 run_id = %run_id,
1070 root_run_id = %root,
1071 error = %err,
1072 "lease could not follow the run to its root"
1073 );
1074 None
1075 }
1076 }
1077}
1078
1079async fn shutdown_signal() {
1081 use tokio::signal;
1082
1083 let ctrl_c = async {
1084 signal::ctrl_c()
1085 .await
1086 .expect("failed to install Ctrl+C handler");
1087 };
1088
1089 #[cfg(unix)]
1090 let terminate = async {
1091 use tokio::signal::unix::{SignalKind, signal};
1092
1093 signal(SignalKind::terminate())
1094 .expect("failed to install SIGTERM handler")
1095 .recv()
1096 .await;
1097 };
1098
1099 #[cfg(not(unix))]
1100 let terminate = {
1101 use std::future::pending;
1102 pending::<()>()
1103 };
1104
1105 select! {
1106 () = ctrl_c => {},
1107 () = terminate => {},
1108 }
1109}
1110
1111#[cfg(test)]
1112mod tests {
1113 use super::*;
1114
1115 use ironflow_core::account_strategy::Priority;
1116 use ironflow_core::providers::claude::ClaudeCodeProvider;
1117 use ironflow_core::providers::record_replay_decision::RecordReplayDecisionProvider;
1118 use ironflow_engine::context::WorkflowContext;
1119 use ironflow_engine::handler::HandlerFuture;
1120 use ironflow_store::entities::{MAX_WORKER_TAGS, WorkerTagError};
1121
1122 #[test]
1123 fn builder_new_creates_default_config() {
1124 let builder = WorkerBuilder::new("http://localhost:3000", "my-token");
1125 assert_eq!(builder.api_url, "http://localhost:3000");
1126 assert_eq!(builder.worker_token, "my-token");
1127 assert_eq!(builder.concurrency, DEFAULT_CONCURRENCY);
1128 assert_eq!(builder.poll_interval, DEFAULT_POLL_INTERVAL);
1129 assert_eq!(builder.run_timeout, DEFAULT_RUN_TIMEOUT);
1130 assert_eq!(
1131 builder.max_consecutive_panics,
1132 DEFAULT_MAX_CONSECUTIVE_PANICS
1133 );
1134 assert_eq!(builder.panic_cooldown, DEFAULT_PANIC_COOLDOWN);
1135 assert!(builder.provider.is_none());
1136 }
1137
1138 #[test]
1139 fn builder_with_trailing_slash_normalized() {
1140 let builder = WorkerBuilder::new("http://localhost:3000/", "token");
1141 assert_eq!(builder.api_url, "http://localhost:3000/");
1142 }
1143
1144 #[test]
1145 fn builder_provider_sets_provider() {
1146 let provider = Arc::new(ClaudeCodeProvider::new());
1147 let builder =
1148 WorkerBuilder::new("http://localhost:3000", "token").provider(provider.clone());
1149 assert!(builder.provider.is_some());
1150 }
1151
1152 #[test]
1153 fn builder_account_strategy_sets_strategy() {
1154 let builder = WorkerBuilder::new("http://localhost:3000", "token");
1155 assert!(builder.account_strategy.is_none());
1156 let builder = builder.account_strategy(Arc::new(Priority));
1157 assert_eq!(
1158 builder.account_strategy.as_ref().map(|s| s.name()),
1159 Some("priority")
1160 );
1161 }
1162
1163 #[test]
1164 fn builder_max_capacity_wait_defaults_to_six_hours_and_can_be_set() {
1165 let builder = WorkerBuilder::new("http://localhost:3000", "token");
1166 assert_eq!(builder.max_capacity_wait, DEFAULT_MAX_CAPACITY_WAIT);
1167 assert_eq!(builder.max_capacity_wait, Duration::from_secs(6 * 3600));
1168 let builder = builder.max_capacity_wait(Duration::ZERO);
1169 assert_eq!(builder.max_capacity_wait, Duration::ZERO);
1170 }
1171
1172 #[test]
1173 fn builder_decision_provider_defaults_to_none() {
1174 let builder = WorkerBuilder::new("http://localhost:3000", "token");
1175 assert!(builder.decision_provider.is_none());
1176 }
1177
1178 #[test]
1179 fn builder_decision_provider_sets_provider() {
1180 let builder = WorkerBuilder::new("http://localhost:3000", "token").decision_provider(
1181 Arc::new(RecordReplayDecisionProvider::replay("tests/fixtures")),
1182 );
1183 assert!(builder.decision_provider.is_some());
1184 }
1185
1186 #[test]
1187 fn builder_concurrency_sets_concurrency() {
1188 let builder = WorkerBuilder::new("http://localhost:3000", "token").concurrency(8);
1189 assert_eq!(builder.concurrency, 8);
1190 }
1191
1192 #[test]
1193 fn builder_concurrency_zero_accepted() {
1194 let provider = Arc::new(ClaudeCodeProvider::new());
1195 let builder = WorkerBuilder::new("http://localhost:3000", "token")
1196 .provider(provider)
1197 .concurrency(0);
1198 assert_eq!(builder.concurrency, 0);
1199 }
1200
1201 #[test]
1202 fn builder_poll_interval_sets_interval() {
1203 let interval = Duration::from_secs(5);
1204 let builder = WorkerBuilder::new("http://localhost:3000", "token").poll_interval(interval);
1205 assert_eq!(builder.poll_interval, interval);
1206 }
1207
1208 #[test]
1209 fn builder_run_timeout_sets_timeout() {
1210 let dur = Duration::from_secs(120);
1211 let builder = WorkerBuilder::new("http://localhost:3000", "token").run_timeout(dur);
1212 assert_eq!(builder.run_timeout, dur);
1213 }
1214
1215 #[test]
1216 fn builder_max_consecutive_panics_sets_value() {
1217 let builder =
1218 WorkerBuilder::new("http://localhost:3000", "token").max_consecutive_panics(10);
1219 assert_eq!(builder.max_consecutive_panics, 10);
1220 }
1221
1222 #[test]
1223 fn builder_panic_cooldown_sets_value() {
1224 let dur = Duration::from_secs(600);
1225 let builder = WorkerBuilder::new("http://localhost:3000", "token").panic_cooldown(dur);
1226 assert_eq!(builder.panic_cooldown, dur);
1227 }
1228
1229 #[test]
1230 fn builder_defaults_lease_settings() {
1231 let builder = WorkerBuilder::new("http://localhost:3000", "token");
1232 assert_eq!(builder.lease_ttl, DEFAULT_LEASE_TTL);
1233 assert_eq!(
1234 builder.lease_refresh_interval,
1235 DEFAULT_LEASE_REFRESH_INTERVAL
1236 );
1237 assert!(builder.worker_id.starts_with("worker-"));
1238 }
1239
1240 #[test]
1241 fn builder_generates_a_distinct_worker_id_per_instance() {
1242 let a = WorkerBuilder::new("http://localhost:3000", "token");
1243 let b = WorkerBuilder::new("http://localhost:3000", "token");
1244 assert_ne!(a.worker_id, b.worker_id);
1245 }
1246
1247 #[test]
1248 fn builder_worker_id_overrides_default() {
1249 let builder = WorkerBuilder::new("http://localhost:3000", "token").worker_id("worker-eu-1");
1250 assert_eq!(builder.worker_id, "worker-eu-1");
1251 }
1252
1253 #[test]
1254 fn builder_lease_ttl_sets_value() {
1255 let dur = Duration::from_secs(120);
1256 let builder = WorkerBuilder::new("http://localhost:3000", "token").lease_ttl(dur);
1257 assert_eq!(builder.lease_ttl, dur);
1258 }
1259
1260 #[test]
1261 fn builder_lease_refresh_interval_sets_value() {
1262 let dur = Duration::from_secs(5);
1263 let builder =
1264 WorkerBuilder::new("http://localhost:3000", "token").lease_refresh_interval(dur);
1265 assert_eq!(builder.lease_refresh_interval, dur);
1266 }
1267
1268 #[test]
1269 fn builder_build_without_provider_fails() {
1270 let builder = WorkerBuilder::new("http://localhost:3000", "token");
1271 let result = builder.build();
1272 assert!(result.is_err());
1273 match result {
1274 Err(WorkerError::Internal(msg)) => {
1275 assert!(msg.contains("provider is required"));
1276 }
1277 _ => panic!("expected Internal error about missing provider"),
1278 }
1279 }
1280
1281 #[test]
1282 fn builder_build_with_provider_succeeds() {
1283 let provider = Arc::new(ClaudeCodeProvider::new());
1284 let builder = WorkerBuilder::new("http://localhost:3000", "token").provider(provider);
1285 let result = builder.build();
1286 assert!(result.is_ok());
1287 }
1288
1289 struct Named(&'static str);
1290
1291 impl WorkflowHandler for Named {
1292 fn name(&self) -> &str {
1293 self.0
1294 }
1295
1296 fn execute<'a>(&'a self, _ctx: &'a mut WorkflowContext) -> HandlerFuture<'a> {
1297 Box::pin(async { Ok(()) })
1298 }
1299 }
1300
1301 fn tagged_builder() -> WorkerBuilder {
1302 WorkerBuilder::new("http://localhost:3000", "token")
1303 .provider(Arc::new(ClaudeCodeProvider::new()))
1304 }
1305
1306 #[test]
1307 fn builder_tags_default_to_empty() {
1308 let worker = tagged_builder().build().unwrap();
1309 assert!(worker.capabilities().tags.is_empty());
1310 let engine_tags = worker.engine.worker_tags();
1311 assert!(
1312 engine_tags.is_some_and(<[String]>::is_empty),
1313 "{engine_tags:?}"
1314 );
1315 }
1316
1317 #[test]
1318 fn builder_tags_are_normalized_and_given_to_the_engine() {
1319 let worker = tagged_builder()
1320 .tags(["region:eu", " gpu ", "gpu"])
1321 .build()
1322 .unwrap();
1323 let expected = vec!["gpu".to_string(), "region:eu".to_string()];
1324 assert_eq!(worker.capabilities().tags, expected);
1325 assert_eq!(worker.engine.worker_tags(), Some(expected.as_slice()));
1326 }
1327
1328 #[test]
1329 fn builder_tags_extend_previous_ones() {
1330 let worker = tagged_builder()
1331 .tags(["gpu"])
1332 .tags(vec!["arm".to_string()])
1333 .build()
1334 .unwrap();
1335 assert_eq!(
1336 worker.capabilities().tags,
1337 vec!["arm".to_string(), "gpu".to_string()]
1338 );
1339 }
1340
1341 #[test]
1342 fn builder_tags_capabilities_list_registered_workflows_sorted() {
1343 let worker = tagged_builder()
1344 .register(Named("deploy"))
1345 .register(Named("build"))
1346 .build()
1347 .unwrap();
1348 assert_eq!(
1349 worker.capabilities().workflows,
1350 Some(vec!["build".to_string(), "deploy".to_string()])
1351 );
1352 }
1353
1354 #[test]
1355 fn builder_tags_capabilities_without_handlers_take_no_workflow() {
1356 let worker = tagged_builder().build().unwrap();
1357 assert_eq!(worker.capabilities().workflows, Some(Vec::new()));
1358 }
1359
1360 #[test]
1361 fn builder_tags_workflow_name_with_comma_disables_workflow_filter() {
1362 let worker = tagged_builder()
1363 .register(Named("deploy"))
1364 .register(Named("build,test"))
1365 .tags(["gpu"])
1366 .build()
1367 .unwrap();
1368 assert_eq!(worker.capabilities().workflows, None);
1369 assert_eq!(worker.capabilities().tags, vec!["gpu".to_string()]);
1370 }
1371
1372 #[test]
1373 fn builder_rejects_invalid_tags() {
1374 let invalid: [&[&str]; 3] = [&["bad tag"], &[""], &["gpu", " "]];
1375 for tags in invalid {
1376 let result = tagged_builder().tags(tags.iter().copied()).build();
1377 let Err(WorkerError::Engine(EngineError::InvalidWorkerTag(_))) = result else {
1378 panic!("tags {tags:?} must be rejected");
1379 };
1380 }
1381 }
1382
1383 #[test]
1384 fn builder_rejects_too_many_tags() {
1385 let tags = (0..=MAX_WORKER_TAGS).map(|i| format!("tag-{i}"));
1386 let result = tagged_builder().tags(tags).build();
1387 let Err(WorkerError::Engine(EngineError::InvalidWorkerTag(err))) = result else {
1388 panic!("too many tags must be rejected");
1389 };
1390 assert!(matches!(err, WorkerTagError::TooMany { .. }), "{err:?}");
1391 }
1392
1393 #[test]
1394 fn builder_build_creates_worker_with_correct_concurrency() {
1395 let provider = Arc::new(ClaudeCodeProvider::new());
1396 let builder = WorkerBuilder::new("http://localhost:3000", "token")
1397 .provider(provider)
1398 .concurrency(16);
1399 let worker = builder.build().unwrap();
1400 assert_eq!(worker.concurrency, 16);
1401 }
1402
1403 #[test]
1404 fn builder_build_creates_worker_with_correct_interval() {
1405 let provider = Arc::new(ClaudeCodeProvider::new());
1406 let interval = Duration::from_secs(10);
1407 let builder = WorkerBuilder::new("http://localhost:3000", "token")
1408 .provider(provider)
1409 .poll_interval(interval);
1410 let worker = builder.build().unwrap();
1411 assert_eq!(worker.poll_interval, interval);
1412 }
1413
1414 #[test]
1415 fn builder_build_preserves_timeout() {
1416 let provider = Arc::new(ClaudeCodeProvider::new());
1417 let dur = Duration::from_secs(300);
1418 let worker = WorkerBuilder::new("http://localhost:3000", "token")
1419 .provider(provider)
1420 .run_timeout(dur)
1421 .build()
1422 .unwrap();
1423 assert_eq!(worker.run_timeout, dur);
1424 }
1425
1426 #[test]
1427 fn builder_build_preserves_poison_pill_config() {
1428 let provider = Arc::new(ClaudeCodeProvider::new());
1429 let cooldown = Duration::from_secs(120);
1430 let worker = WorkerBuilder::new("http://localhost:3000", "token")
1431 .provider(provider)
1432 .max_consecutive_panics(7)
1433 .panic_cooldown(cooldown)
1434 .build()
1435 .unwrap();
1436 assert_eq!(worker.max_consecutive_panics, 7);
1437 assert_eq!(worker.panic_cooldown, cooldown);
1438 }
1439
1440 #[test]
1441 fn builder_chaining_works() {
1442 let provider = Arc::new(ClaudeCodeProvider::new());
1443 let result = WorkerBuilder::new("http://localhost:3000", "token")
1444 .provider(provider)
1445 .concurrency(4)
1446 .poll_interval(Duration::from_secs(3))
1447 .run_timeout(Duration::from_secs(600))
1448 .max_consecutive_panics(5)
1449 .panic_cooldown(Duration::from_secs(120))
1450 .build();
1451 assert!(result.is_ok());
1452 let worker = result.unwrap();
1453 assert_eq!(worker.concurrency, 4);
1454 assert_eq!(worker.poll_interval, Duration::from_secs(3));
1455 assert_eq!(worker.run_timeout, Duration::from_secs(600));
1456 assert_eq!(worker.max_consecutive_panics, 5);
1457 assert_eq!(worker.panic_cooldown, Duration::from_secs(120));
1458 }
1459
1460 #[test]
1461 fn builder_empty_api_url_accepted() {
1462 let provider = Arc::new(ClaudeCodeProvider::new());
1463 let builder = WorkerBuilder::new("", "token").provider(provider);
1464 let result = builder.build();
1465 assert!(result.is_ok());
1466 }
1467
1468 #[test]
1469 fn builder_empty_token_accepted() {
1470 let provider = Arc::new(ClaudeCodeProvider::new());
1471 let builder = WorkerBuilder::new("http://localhost:3000", "").provider(provider);
1472 let result = builder.build();
1473 assert!(result.is_ok());
1474 }
1475
1476 #[cfg(feature = "heartbeat")]
1477 #[test]
1478 fn builder_heartbeat_defaults() {
1479 let builder = WorkerBuilder::new("http://localhost:3000", "token");
1480 assert!(builder.heartbeat_url.is_none());
1481 assert_eq!(builder.heartbeat_interval, DEFAULT_HEARTBEAT_INTERVAL);
1482 }
1483
1484 #[cfg(feature = "heartbeat")]
1485 #[test]
1486 fn builder_heartbeat_url_sets_url() {
1487 let builder = WorkerBuilder::new("http://localhost:3000", "token")
1488 .heartbeat_url("https://uptime.betterstack.com/api/v1/heartbeat/abc");
1489 assert_eq!(
1490 builder.heartbeat_url.as_deref(),
1491 Some("https://uptime.betterstack.com/api/v1/heartbeat/abc")
1492 );
1493 }
1494
1495 #[cfg(feature = "heartbeat")]
1496 #[test]
1497 fn builder_heartbeat_custom_interval() {
1498 let interval = Duration::from_secs(10);
1499 let builder =
1500 WorkerBuilder::new("http://localhost:3000", "token").heartbeat_interval(interval);
1501 assert_eq!(builder.heartbeat_interval, interval);
1502 }
1503
1504 #[cfg(feature = "heartbeat")]
1505 #[test]
1506 fn builder_build_preserves_heartbeat_config() {
1507 let provider = Arc::new(ClaudeCodeProvider::new());
1508 let interval = Duration::from_secs(15);
1509 let worker = WorkerBuilder::new("http://localhost:3000", "token")
1510 .provider(provider)
1511 .heartbeat_url("https://example.com/heartbeat")
1512 .heartbeat_interval(interval)
1513 .build()
1514 .unwrap();
1515 assert_eq!(
1516 worker.heartbeat_url.as_deref(),
1517 Some("https://example.com/heartbeat")
1518 );
1519 assert_eq!(worker.heartbeat_interval, interval);
1520 }
1521
1522 #[cfg(feature = "heartbeat")]
1523 #[test]
1524 fn builder_build_without_heartbeat_url_has_none() {
1525 let provider = Arc::new(ClaudeCodeProvider::new());
1526 let worker = WorkerBuilder::new("http://localhost:3000", "token")
1527 .provider(provider)
1528 .build()
1529 .unwrap();
1530 assert!(worker.heartbeat_url.is_none());
1531 }
1532
1533 #[test]
1536 fn poison_tracker_not_blocked_initially() {
1537 let tracker = PoisonPillTracker::new(3, Duration::from_secs(300));
1538 assert!(!tracker.is_blocked("my-workflow"));
1539 }
1540
1541 #[test]
1542 fn poison_tracker_blocked_after_max_panics() {
1543 let mut tracker = PoisonPillTracker::new(3, Duration::from_secs(300));
1544 assert!(!tracker.record_panic("wf"));
1545 assert!(!tracker.record_panic("wf"));
1546 assert!(tracker.record_panic("wf"));
1547 assert!(tracker.is_blocked("wf"));
1548 }
1549
1550 #[test]
1551 fn poison_tracker_success_resets_count() {
1552 let mut tracker = PoisonPillTracker::new(3, Duration::from_secs(300));
1553 tracker.record_panic("wf");
1554 tracker.record_panic("wf");
1555 tracker.record_success("wf");
1556 assert!(!tracker.is_blocked("wf"));
1557 assert!(!tracker.record_panic("wf"));
1559 }
1560
1561 #[test]
1562 fn poison_tracker_independent_per_workflow() {
1563 let mut tracker = PoisonPillTracker::new(2, Duration::from_secs(300));
1564 tracker.record_panic("wf-a");
1565 tracker.record_panic("wf-a");
1566 assert!(tracker.is_blocked("wf-a"));
1567 assert!(!tracker.is_blocked("wf-b"));
1568 }
1569
1570 #[test]
1571 fn poison_tracker_unblocks_after_cooldown() {
1572 let mut tracker = PoisonPillTracker::new(2, Duration::from_millis(0));
1573 tracker.record_panic("wf");
1574 tracker.record_panic("wf");
1575 assert!(!tracker.is_blocked("wf"));
1577 }
1578
1579 #[tokio::test]
1580 async fn acquire_slot_returns_a_permit_when_a_slot_is_free() {
1581 let semaphore = Arc::new(Semaphore::new(1));
1582 let shutdown = CancellationToken::new();
1583
1584 let permit = acquire_slot(&semaphore, &shutdown)
1585 .await
1586 .expect("semaphore open");
1587
1588 assert!(permit.is_some());
1589 assert_eq!(semaphore.available_permits(), 0);
1590 drop(permit);
1591 assert_eq!(semaphore.available_permits(), 1);
1592 }
1593
1594 #[tokio::test]
1595 async fn acquire_slot_gives_up_on_shutdown_while_all_slots_are_busy() {
1596 let semaphore = Arc::new(Semaphore::new(1));
1597 let _busy = semaphore.clone().acquire_owned().await.expect("permit");
1598 let shutdown = CancellationToken::new();
1599
1600 let cancel = shutdown.clone();
1601 spawn(async move {
1602 sleep(Duration::from_millis(20)).await;
1603 cancel.cancel();
1604 });
1605
1606 let permit = timeout(Duration::from_secs(5), acquire_slot(&semaphore, &shutdown))
1607 .await
1608 .expect("waiting for a slot must stop on shutdown")
1609 .expect("semaphore open");
1610
1611 assert!(permit.is_none());
1612 }
1613
1614 #[tokio::test]
1615 async fn acquire_slot_takes_no_permit_once_shutdown_is_requested() {
1616 let semaphore = Arc::new(Semaphore::new(1));
1617 let shutdown = CancellationToken::new();
1618 shutdown.cancel();
1619
1620 let permit = acquire_slot(&semaphore, &shutdown)
1621 .await
1622 .expect("semaphore open");
1623
1624 assert!(permit.is_none());
1625 assert_eq!(semaphore.available_permits(), 1);
1626 }
1627
1628 #[tokio::test]
1629 async fn acquire_slot_fails_when_the_semaphore_is_closed() {
1630 let semaphore = Arc::new(Semaphore::new(1));
1631 semaphore.close();
1632 let shutdown = CancellationToken::new();
1633
1634 let result = acquire_slot(&semaphore, &shutdown).await;
1635
1636 assert!(matches!(result, Err(WorkerError::Internal(_))));
1637 }
1638}