1use std::collections::HashMap;
45use std::fmt;
46use std::sync::Arc;
47use std::time::Duration;
48
49use chrono::{DateTime, TimeDelta, Utc};
50use ironflow_core::account::{
51 AccountKind, AccountSession, AccountWindow, ClaudeSubscriptionKind, RateLimitRecorder,
52 WindowStatus,
53};
54use ironflow_core::account_strategy::{
55 AccountCandidate, AccountStrategy, LeastUtilized, select_account,
56};
57use ironflow_core::error::AgentError;
58use ironflow_core::provider::{
59 AgentConfig, AgentOutput, AgentProvider, InvokeFuture, LogSink, ReleaseFuture,
60};
61use ironflow_store::entities::{
62 AccountWindowStatus, NewAccountWindow, NewProviderAccountObservation, ProviderAccount,
63 ProviderAccountCandidate, ProviderAccountWindow,
64};
65use ironflow_store::store::Store;
66use tracing::{debug, info, warn};
67use uuid::Uuid;
68
69pub fn window_from_store(window: &ProviderAccountWindow) -> AccountWindow {
91 AccountWindow {
92 window: window.window.clone(),
93 utilization: window.utilization,
94 resets_at: window.resets_at,
95 status: match window.status {
96 AccountWindowStatus::Allowed => WindowStatus::Allowed,
97 AccountWindowStatus::AllowedWarning => WindowStatus::AllowedWarning,
98 AccountWindowStatus::Rejected => WindowStatus::Rejected,
99 },
100 model_scope: window.model_scope.clone(),
101 observed_at: window.observed_at,
102 }
103}
104
105pub fn window_to_store(window: AccountWindow) -> NewAccountWindow {
126 NewAccountWindow {
127 window: window.window,
128 utilization: window.utilization,
129 resets_at: window.resets_at,
130 status: match window.status {
131 WindowStatus::Allowed => AccountWindowStatus::Allowed,
132 WindowStatus::AllowedWarning => AccountWindowStatus::AllowedWarning,
133 WindowStatus::Rejected => AccountWindowStatus::Rejected,
134 },
135 model_scope: window.model_scope,
136 observed_at: window.observed_at,
137 }
138}
139
140fn to_core_candidate(candidate: &ProviderAccountCandidate) -> AccountCandidate {
141 AccountCandidate {
142 id: candidate.account.id.to_string(),
143 name: candidate.account.name.clone(),
144 priority: candidate.account.priority,
145 max_concurrency: candidate.account.max_concurrency,
146 running_steps: candidate.running_steps,
147 windows: candidate.windows.iter().map(window_from_store).collect(),
148 }
149}
150
151fn resolution_error(message: String) -> AgentError {
152 AgentError::ProcessFailed {
153 exit_code: -1,
154 stderr: message,
155 }
156}
157
158pub const DEFAULT_MAX_CAPACITY_WAIT: Duration = Duration::from_secs(6 * 3600);
171
172const CONCURRENCY_RETRY: Duration = Duration::from_secs(60);
175
176#[derive(Debug)]
178struct AccountAttempt {
179 account_id: Uuid,
180 outcome: &'static str,
181}
182
183fn to_time_delta(duration: Duration) -> TimeDelta {
184 TimeDelta::from_std(duration).unwrap_or(TimeDelta::MAX)
185}
186
187fn capacity_decision(
193 kind: &str,
194 next_reset: Option<DateTime<Utc>>,
195 concurrency_blocked: bool,
196 wait: Duration,
197 since: Option<DateTime<Utc>>,
198 now: DateTime<Utc>,
199) -> AgentError {
200 let kind = kind.to_string();
201 if wait.is_zero() {
202 return AgentError::NoCapacity { kind, next_reset };
203 }
204 let budget = to_time_delta(wait);
205 let since = since.unwrap_or(now);
206 match next_reset {
207 Some(reset) if reset - since <= budget => AgentError::CapacityWait {
208 kind,
209 wake_at: reset,
210 },
211 Some(_) => AgentError::NoCapacity { kind, next_reset },
212 None if concurrency_blocked => {
213 let wake_at = now + to_time_delta(CONCURRENCY_RETRY);
214 if wake_at - since <= budget {
215 AgentError::CapacityWait { kind, wake_at }
216 } else {
217 AgentError::NoCapacity {
218 kind,
219 next_reset: None,
220 }
221 }
222 }
223 None => AgentError::NoCapacity {
224 kind,
225 next_reset: None,
226 },
227 }
228}
229
230fn earliest_reset<'a>(
232 windows: impl Iterator<Item = &'a AccountWindow>,
233 model: &str,
234 now: DateTime<Utc>,
235) -> Option<DateTime<Utc>> {
236 windows
237 .filter(|w| w.applies_to(model) && w.is_exhausted(now))
238 .filter_map(|w| w.resets_at)
239 .min()
240}
241
242fn is_targeted(config: &AgentConfig, account: &ProviderAccount) -> bool {
244 if let Some(name) = &config.account_name {
245 return &account.name == name;
246 }
247 if let Some(pool) = &config.account_pool {
248 return account.tags.iter().any(|tag| tag == pool);
249 }
250 true
251}
252
253pub struct AccountAwareProvider {
257 inner: Arc<dyn AgentProvider>,
258 store: Arc<dyn Store>,
259 strategy: Arc<dyn AccountStrategy>,
260 kinds: HashMap<&'static str, Arc<dyn AccountKind>>,
261 max_capacity_wait: Duration,
262}
263
264impl fmt::Debug for AccountAwareProvider {
265 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
266 f.debug_struct("AccountAwareProvider")
267 .field("strategy", &self.strategy.name())
268 .field("kinds", &self.kinds.keys().collect::<Vec<_>>())
269 .field("max_capacity_wait", &self.max_capacity_wait)
270 .finish_non_exhaustive()
271 }
272}
273
274impl AccountAwareProvider {
275 pub fn new(inner: Arc<dyn AgentProvider>, store: Arc<dyn Store>) -> Self {
283 let claude: Arc<dyn AccountKind> = Arc::new(ClaudeSubscriptionKind::new());
284 Self {
285 inner,
286 store,
287 strategy: Arc::new(LeastUtilized),
288 kinds: HashMap::from([(claude.id(), claude)]),
289 max_capacity_wait: DEFAULT_MAX_CAPACITY_WAIT,
290 }
291 }
292
293 pub fn with_strategy(mut self, strategy: Arc<dyn AccountStrategy>) -> Self {
299 self.strategy = strategy;
300 self
301 }
302
303 pub fn with_max_capacity_wait(mut self, wait: Duration) -> Self {
326 self.max_capacity_wait = wait;
327 self
328 }
329
330 pub fn with_kind(mut self, kind: Arc<dyn AccountKind>) -> Self {
349 self.kinds.insert(kind.id(), kind);
350 self
351 }
352
353 async fn invoke_inner(
354 &self,
355 config: &AgentConfig,
356 sink: Option<Arc<dyn LogSink>>,
357 ) -> Result<AgentOutput, AgentError> {
358 match sink {
359 Some(sink) => self.inner.invoke_with_logs(config, sink).await,
360 None => self.inner.invoke(config).await,
361 }
362 }
363
364 async fn run(
365 &self,
366 config: &AgentConfig,
367 sink: Option<Arc<dyn LogSink>>,
368 ) -> Result<AgentOutput, AgentError> {
369 let Some(kind_id) = self.inner.account_kind_for(config) else {
370 return self.invoke_inner(config, sink).await;
371 };
372 let Some(kind) = self.kinds.get(kind_id).cloned() else {
373 debug!(
374 kind = kind_id,
375 "no account kind registered, using worker environment"
376 );
377 return self.invoke_inner(config, sink).await;
378 };
379 let wait = config.max_capacity_wait.unwrap_or(self.max_capacity_wait);
380
381 let mut attempts: Vec<AccountAttempt> = Vec::new();
384 loop {
385 let candidates = self
386 .store
387 .list_provider_account_candidates(kind_id.to_string())
388 .await
389 .map_err(|e| {
390 resolution_error(format!("provider account resolution failed: {e}"))
391 })?;
392 if candidates.is_empty()
393 && config.account_name.is_none()
394 && config.account_pool.is_none()
395 {
396 debug!(
397 kind = kind_id,
398 "no provider account for kind, using worker environment"
399 );
400 return self.run_with_worker_env(kind_id, config, sink, wait).await;
401 }
402
403 let targeted: Vec<&ProviderAccountCandidate> = candidates
404 .iter()
405 .filter(|c| is_targeted(config, &c.account))
406 .collect();
407 if targeted.is_empty() {
408 if let Some(name) = &config.account_name {
409 return Err(AgentError::AccountNotFound { name: name.clone() });
410 }
411 return Err(AgentError::NoCapacity {
412 kind: kind_id.to_string(),
413 next_reset: None,
414 });
415 }
416
417 let all: Vec<_> = targeted.iter().map(|c| to_core_candidate(c)).collect();
418 let untried: Vec<AccountCandidate> = all
419 .iter()
420 .filter(|c| !attempts.iter().any(|a| a.account_id.to_string() == c.id))
421 .cloned()
422 .collect();
423 let now = Utc::now();
424 let Some(selected) =
425 select_account(self.strategy.as_ref(), &untried, &config.model, now)
426 else {
427 let next_reset = earliest_reset(
428 all.iter().flat_map(|c| c.windows.iter()),
429 &config.model,
430 now,
431 );
432 let concurrency_blocked = all.iter().any(|c| {
433 c.max_concurrency.is_some_and(|max| c.running_steps >= max)
434 && !is_rejected(&c.windows, &config.model, now)
435 });
436 let decision = capacity_decision(
437 kind_id,
438 next_reset,
439 concurrency_blocked,
440 wait,
441 config.capacity_wait_since,
442 now,
443 );
444 info!(
445 kind = kind_id,
446 attempts = ?attempts,
447 decision = %decision,
448 "no targeted provider account has capacity"
449 );
450 return Err(decision);
451 };
452 let account: &ProviderAccount = &targeted
453 .iter()
454 .find(|c| c.account.id.to_string() == selected.id)
455 .ok_or_else(|| resolution_error("selected provider account vanished".to_string()))?
456 .account;
457
458 let result = self
459 .run_with_account(kind.as_ref(), account, config, sink.clone())
460 .await?;
461 match result {
462 AccountRun::Done(result) => return result,
463 AccountRun::Rejected => {
464 let attempt = AccountAttempt {
465 account_id: account.id,
466 outcome: "rate limited",
467 };
468 warn!(
469 account = %account.name,
470 account_id = %attempt.account_id,
471 outcome = attempt.outcome,
472 "provider account rate limited mid-step, trying the next account"
473 );
474 attempts.push(attempt);
475 }
476 }
477 }
478 }
479
480 async fn run_with_account(
486 &self,
487 kind: &dyn AccountKind,
488 account: &ProviderAccount,
489 config: &AgentConfig,
490 sink: Option<Arc<dyn LogSink>>,
491 ) -> Result<AccountRun, AgentError> {
492 let missing = || {
493 resolution_error(format!(
494 "credential of provider account '{}' is missing",
495 account.name
496 ))
497 };
498 let secret = match self.store.get_secret(&account.secret_key).await {
499 Ok(Some(secret)) => secret,
500 Ok(None) => return Err(missing()),
501 Err(e) => {
502 warn!(account = %account.name, error = %e, "failed to read provider account credential");
503 return Err(missing());
504 }
505 };
506
507 info!(
508 account = %account.name,
509 strategy = self.strategy.name(),
510 model = %config.model,
511 "selected provider account"
512 );
513
514 let recorder = RateLimitRecorder::default();
515 let account_config = config
517 .clone()
518 .verbose(true)
519 .account_session(AccountSession::new(
520 kind.credential(&secret.value),
521 recorder.clone(),
522 ));
523
524 let result = self.invoke_inner(&account_config, sink).await;
525
526 let windows = recorder.take();
527 let now = Utc::now();
528 let rejected = result.is_err() && is_rejected(&windows, &config.model, now);
529 let auth_failed = matches!(
530 result,
531 Err(AgentError::Api {
532 status: Some(401 | 403),
533 ..
534 })
535 );
536 if !windows.is_empty() || auth_failed {
537 let observation = NewProviderAccountObservation {
538 windows: windows.into_iter().map(window_to_store).collect(),
539 auth_failed,
540 };
541 if let Err(e) = self
542 .store
543 .record_provider_account_observation(account.id, observation)
544 .await
545 {
546 warn!(account = %account.name, error = %e, "failed to record provider account usage");
547 }
548 }
549
550 if rejected {
551 return Ok(AccountRun::Rejected);
552 }
553 Ok(AccountRun::Done(result.map(|mut output| {
554 output.account_id = Some(account.id.to_string());
555 output
556 })))
557 }
558
559 async fn run_with_worker_env(
565 &self,
566 kind_id: &str,
567 config: &AgentConfig,
568 sink: Option<Arc<dyn LogSink>>,
569 wait: Duration,
570 ) -> Result<AgentOutput, AgentError> {
571 let recorder = RateLimitRecorder::default();
572 let env_config = config
574 .clone()
575 .verbose(true)
576 .rate_limit_recorder(recorder.clone());
577 let result = self.invoke_inner(&env_config, sink).await;
578 if result.is_ok() {
579 return result;
580 }
581
582 let windows = recorder.take();
583 let now = Utc::now();
584 if !is_rejected(&windows, &config.model, now) {
585 return result;
586 }
587 let next_reset = earliest_reset(windows.iter(), &config.model, now);
588 let decision = capacity_decision(
589 kind_id,
590 next_reset,
591 false,
592 wait,
593 config.capacity_wait_since,
594 now,
595 );
596 info!(
597 kind = kind_id,
598 decision = %decision,
599 "worker environment credential rate limited"
600 );
601 Err(decision)
602 }
603}
604
605#[allow(clippy::large_enum_variant)]
607enum AccountRun {
608 Done(Result<AgentOutput, AgentError>),
611 Rejected,
613}
614
615fn is_rejected(windows: &[AccountWindow], model: &str, now: DateTime<Utc>) -> bool {
617 windows
618 .iter()
619 .any(|w| w.applies_to(model) && w.is_exhausted(now))
620}
621
622impl AgentProvider for AccountAwareProvider {
623 fn invoke<'a>(&'a self, config: &'a AgentConfig) -> InvokeFuture<'a> {
624 Box::pin(self.run(config, None))
625 }
626
627 fn invoke_with_logs<'a>(
628 &'a self,
629 config: &'a AgentConfig,
630 log_sink: Arc<dyn LogSink>,
631 ) -> InvokeFuture<'a> {
632 Box::pin(self.run(config, Some(log_sink)))
633 }
634
635 fn release_run<'a>(&'a self, run_id: &'a str) -> ReleaseFuture<'a> {
636 self.inner.release_run(run_id)
637 }
638
639 fn account_kind(&self) -> Option<&'static str> {
640 self.inner.account_kind()
641 }
642
643 fn account_kind_for(&self, config: &AgentConfig) -> Option<&'static str> {
644 self.inner.account_kind_for(config)
645 }
646}
647
648#[cfg(test)]
649mod tests {
650 use std::sync::Mutex;
651
652 use ironflow_core::account_strategy::Priority;
653 use ironflow_core::providers::router::{ProviderMatcher, ProviderRouter};
654 use ironflow_store::crypto::KeyRing;
655 use ironflow_store::entities::{
656 NewProviderAccount, NewRun, ProviderKind, RunStatus, RunUpdate, TriggerKind,
657 provider_account_secret_key,
658 };
659 use ironflow_store::memory::InMemoryStore;
660 use ironflow_store::provider_account_store::ProviderAccountStore;
661 use ironflow_store::secret_store::SecretStore;
662 use ironflow_store::store::RunStore;
663 use serde_json::json;
664
665 use super::*;
666
667 const TOKEN: &str = "sk-ant-oat01-test-token-abcdefghijklmnopqrstuvwxyz";
668
669 #[derive(Clone, Copy)]
671 enum Outcome {
672 Succeed,
673 FailApi(u16),
674 Reject {
677 resets_at: DateTime<Utc>,
678 },
679 RejectOnce {
681 resets_at: DateTime<Utc>,
682 },
683 }
684
685 struct RecordingProvider {
688 kind: Option<&'static str>,
689 outcome: Outcome,
690 seen: Mutex<Vec<(Option<String>, bool)>>,
691 }
692
693 impl RecordingProvider {
694 fn new(kind: Option<&'static str>, outcome: Outcome) -> Self {
695 Self {
696 kind,
697 outcome,
698 seen: Mutex::new(Vec::new()),
699 }
700 }
701
702 fn seen(&self) -> Vec<(Option<String>, bool)> {
703 self.seen.lock().unwrap().clone()
704 }
705 }
706
707 impl AgentProvider for RecordingProvider {
708 fn invoke<'a>(&'a self, config: &'a AgentConfig) -> InvokeFuture<'a> {
709 Box::pin(async move {
710 let credential = config
711 .account
712 .as_ref()
713 .map(|s| s.credential().expose().to_string());
714 let previous = {
715 let mut seen = self.seen.lock().unwrap();
716 seen.push((credential, config.verbose));
717 seen.len() - 1
718 };
719 let rejection = match self.outcome {
720 Outcome::Reject { resets_at } => Some(resets_at),
721 Outcome::RejectOnce { resets_at } if previous == 0 => Some(resets_at),
722 _ => None,
723 };
724 if let Some(resets_at) = rejection {
725 let recorder = config
726 .account
727 .as_ref()
728 .map(AccountSession::recorder)
729 .or(config.rate_limits.as_ref())
730 .expect("AccountAwareProvider always sets a recorder");
731 recorder.record(AccountWindow {
732 window: "five_hour".to_string(),
733 utilization: 1.0,
734 resets_at: Some(resets_at),
735 status: WindowStatus::Rejected,
736 model_scope: None,
737 observed_at: Utc::now(),
738 });
739 return Err(AgentError::Api {
740 status: Some(429),
741 code: None,
742 message: "rate limited".to_string(),
743 });
744 }
745 if let Some(session) = &config.account {
746 session.recorder().record(AccountWindow {
747 window: "five_hour".to_string(),
748 utilization: 0.42,
749 resets_at: Some(Utc::now() + TimeDelta::hours(2)),
750 status: WindowStatus::Allowed,
751 model_scope: None,
752 observed_at: Utc::now(),
753 });
754 }
755 match self.outcome {
756 Outcome::Succeed | Outcome::Reject { .. } | Outcome::RejectOnce { .. } => {
757 Ok(AgentOutput::new(json!("done")))
758 }
759 Outcome::FailApi(status) => Err(AgentError::Api {
760 status: Some(status),
761 code: None,
762 message: "API Error".to_string(),
763 }),
764 }
765 })
766 }
767
768 fn account_kind(&self) -> Option<&'static str> {
769 self.kind
770 }
771 }
772
773 fn store_with_key() -> InMemoryStore {
774 let mut store = InMemoryStore::new();
775 let spec = format!("1:{}", "aa".repeat(32));
776 store.set_key_ring(KeyRing::from_spec(&spec, Some(1)).unwrap());
777 store
778 }
779
780 async fn add_account(store: &InMemoryStore, name: &str, priority: i32) -> ProviderAccount {
781 add_account_with(store, name, priority, &[], None).await
782 }
783
784 async fn add_account_with(
785 store: &InMemoryStore,
786 name: &str,
787 priority: i32,
788 tags: &[&str],
789 max_concurrency: Option<u32>,
790 ) -> ProviderAccount {
791 let id = Uuid::now_v7();
792 let secret_key = provider_account_secret_key(id);
793 store
794 .set_secret(&secret_key, &format!("{TOKEN}-{name}"))
795 .await
796 .unwrap();
797 store
798 .create_provider_account(NewProviderAccount {
799 id,
800 name: name.to_string(),
801 display_name: name.to_string(),
802 kind: ClaudeSubscriptionKind::ID.to_string(),
803 secret_key,
804 enabled: true,
805 priority,
806 tags: tags.iter().map(|t| t.to_string()).collect(),
807 max_concurrency,
808 alert_threshold: 0.8,
809 expires_at: Utc::now() + TimeDelta::days(30),
810 plan: None,
811 created_by: None,
812 })
813 .await
814 .unwrap()
815 }
816
817 fn wrap(inner: Arc<RecordingProvider>, store: &Arc<InMemoryStore>) -> AccountAwareProvider {
818 let store: Arc<dyn Store> = store.clone();
819 AccountAwareProvider::new(inner, store)
820 }
821
822 #[tokio::test]
823 async fn account_aware_provider_injects_selected_account() {
824 let store = Arc::new(store_with_key());
825 add_account(&store, "busy", 10).await;
826 let busy = store
827 .find_provider_account_by_name("busy")
828 .await
829 .unwrap()
830 .unwrap();
831 store
832 .record_provider_account_observation(
833 busy.id,
834 NewProviderAccountObservation {
835 windows: vec![NewAccountWindow {
836 window: "five_hour".to_string(),
837 utilization: 0.9,
838 resets_at: Some(Utc::now() + TimeDelta::hours(1)),
839 status: AccountWindowStatus::Allowed,
840 model_scope: None,
841 observed_at: Utc::now(),
842 }],
843 auth_failed: false,
844 },
845 )
846 .await
847 .unwrap();
848 add_account(&store, "idle", 20).await;
849
850 let inner = Arc::new(RecordingProvider::new(
851 Some(ClaudeSubscriptionKind::ID),
852 Outcome::Succeed,
853 ));
854 let provider = wrap(inner.clone(), &store);
855 provider.invoke(&AgentConfig::new("hello")).await.unwrap();
856
857 let seen = inner.seen();
858 assert_eq!(seen.len(), 1);
859 assert_eq!(seen[0].0.as_deref(), Some(format!("{TOKEN}-idle").as_str()));
860 assert!(seen[0].1, "verbose must be forced for rate_limit_event");
861 }
862
863 #[tokio::test]
864 async fn account_is_injected_behind_a_router() {
865 let store = Arc::new(store_with_key());
866 let account = add_account(&store, "perso", 10).await;
867 let claude = Arc::new(RecordingProvider::new(
868 Some(ClaudeSubscriptionKind::ID),
869 Outcome::Succeed,
870 ));
871 let router = ProviderRouter::new(claude.clone());
872 let dyn_store: Arc<dyn Store> = store.clone();
873 let provider = AccountAwareProvider::new(Arc::new(router), dyn_store);
874
875 let output = provider
876 .invoke(&AgentConfig::new("p").model("sonnet"))
877 .await
878 .unwrap();
879
880 assert_eq!(output.account_id, Some(account.id.to_string()));
881 let seen = claude.seen();
882 assert_eq!(seen.len(), 1);
883 assert_eq!(
884 seen[0].0.as_deref(),
885 Some(format!("{TOKEN}-perso").as_str())
886 );
887 }
888
889 #[tokio::test]
890 async fn mixed_router_injects_an_account_only_on_claude_routes() {
891 let store = Arc::new(store_with_key());
892 let account = add_account(&store, "perso", 10).await;
893 let claude = Arc::new(RecordingProvider::new(
894 Some(ClaudeSubscriptionKind::ID),
895 Outcome::Succeed,
896 ));
897 let http = Arc::new(RecordingProvider::new(None, Outcome::Succeed));
898 let router = ProviderRouter::new(claude.clone())
899 .route(ProviderMatcher::ModelPrefix("gpt-".into()), http.clone());
900 let dyn_store: Arc<dyn Store> = store.clone();
901 let provider = AccountAwareProvider::new(Arc::new(router), dyn_store);
902
903 let claude_output = provider
904 .invoke(&AgentConfig::new("p").model("sonnet"))
905 .await
906 .unwrap();
907 assert_eq!(claude_output.account_id, Some(account.id.to_string()));
908 assert!(claude.seen()[0].0.is_some());
909
910 let http_output = provider
911 .invoke(&AgentConfig::new("p").model("gpt-5"))
912 .await
913 .unwrap();
914 assert_eq!(http_output.account_id, None);
915 assert_eq!(http.seen(), vec![(None, false)]);
916 }
917
918 #[tokio::test]
919 async fn account_aware_provider_passthrough_without_accounts() {
920 let store = Arc::new(store_with_key());
921 let inner = Arc::new(RecordingProvider::new(
922 Some(ClaudeSubscriptionKind::ID),
923 Outcome::Succeed,
924 ));
925 let provider = wrap(inner.clone(), &store);
926 let output = provider.invoke(&AgentConfig::new("hello")).await.unwrap();
927 assert_eq!(output.account_id, None);
928 assert_eq!(inner.seen(), vec![(None, true)]);
930 }
931
932 #[tokio::test]
933 async fn account_aware_provider_passthrough_for_kindless_provider() {
934 let store = Arc::new(store_with_key());
935 add_account(&store, "perso", 10).await;
936 let inner = Arc::new(RecordingProvider::new(None, Outcome::Succeed));
937 let provider = wrap(inner.clone(), &store);
938 let output = provider.invoke(&AgentConfig::new("hello")).await.unwrap();
939 assert_eq!(output.account_id, None);
940 assert_eq!(inner.seen(), vec![(None, false)]);
941 assert_eq!(provider.account_kind(), None);
942 }
943
944 #[tokio::test]
945 async fn account_aware_provider_records_windows_on_error() {
946 let store = Arc::new(store_with_key());
947 let account = add_account(&store, "perso", 10).await;
948 let inner = Arc::new(RecordingProvider::new(
949 Some(ClaudeSubscriptionKind::ID),
950 Outcome::FailApi(500),
951 ));
952 let provider = wrap(inner, &store);
953 let err = provider
954 .invoke(&AgentConfig::new("hello"))
955 .await
956 .unwrap_err();
957 assert!(matches!(
958 err,
959 AgentError::Api {
960 status: Some(500),
961 ..
962 }
963 ));
964
965 let windows = store
966 .list_provider_account_windows(vec![account.id])
967 .await
968 .unwrap();
969 assert_eq!(windows.len(), 1);
970 assert_eq!(windows[0].window, "five_hour");
971 let stored = store
972 .get_provider_account(account.id)
973 .await
974 .unwrap()
975 .unwrap();
976 assert!(stored.auth_failed_at.is_none());
977 }
978
979 #[tokio::test]
980 async fn account_aware_provider_marks_auth_failed_on_401() {
981 let store = Arc::new(store_with_key());
982 let account = add_account(&store, "perso", 10).await;
983 let inner = Arc::new(RecordingProvider::new(
984 Some(ClaudeSubscriptionKind::ID),
985 Outcome::FailApi(401),
986 ));
987 let provider = wrap(inner, &store);
988 provider
989 .invoke(&AgentConfig::new("hello"))
990 .await
991 .unwrap_err();
992
993 let stored = store
994 .get_provider_account(account.id)
995 .await
996 .unwrap()
997 .unwrap();
998 assert!(stored.auth_failed_at.is_some());
999 let candidates = store
1000 .list_provider_account_candidates(ClaudeSubscriptionKind::ID.to_string())
1001 .await
1002 .unwrap();
1003 assert!(candidates.is_empty(), "a rejected token is not a candidate");
1004 }
1005
1006 async fn exhaust(store: &InMemoryStore, account: &ProviderAccount, resets_at: DateTime<Utc>) {
1007 store
1008 .record_provider_account_observation(
1009 account.id,
1010 NewProviderAccountObservation {
1011 windows: vec![NewAccountWindow {
1012 window: "five_hour".to_string(),
1013 utilization: 1.0,
1014 resets_at: Some(resets_at),
1015 status: AccountWindowStatus::Rejected,
1016 model_scope: None,
1017 observed_at: Utc::now(),
1018 }],
1019 auth_failed: false,
1020 },
1021 )
1022 .await
1023 .unwrap();
1024 }
1025
1026 fn succeeding() -> Arc<RecordingProvider> {
1027 Arc::new(RecordingProvider::new(
1028 Some(ClaudeSubscriptionKind::ID),
1029 Outcome::Succeed,
1030 ))
1031 }
1032
1033 #[tokio::test]
1034 async fn pool_exhausted_sleeps_until_next_reset() {
1035 let store = Arc::new(store_with_key());
1036 let soon = Utc::now() + TimeDelta::hours(1);
1037 let later = Utc::now() + TimeDelta::hours(2);
1038 let first = add_account(&store, "first", 10).await;
1039 exhaust(&store, &first, later).await;
1040 let second = add_account(&store, "second", 20).await;
1041 exhaust(&store, &second, soon).await;
1042 let inner = succeeding();
1043 let provider = wrap(inner.clone(), &store);
1044
1045 let err = provider
1046 .invoke(&AgentConfig::new("hello"))
1047 .await
1048 .unwrap_err();
1049
1050 let AgentError::CapacityWait { kind, wake_at } = err else {
1051 panic!("expected CapacityWait, got {err}");
1052 };
1053 assert_eq!(kind, ClaudeSubscriptionKind::ID);
1054 assert_eq!(wake_at, soon, "wakes at the earliest reset of the pool");
1055 assert!(inner.seen().is_empty(), "the agent must not run");
1056 }
1057
1058 #[tokio::test]
1059 async fn zero_capacity_wait_fails_fast_with_no_capacity() {
1060 let store = Arc::new(store_with_key());
1061 let reset = Utc::now() + TimeDelta::hours(1);
1062 let account = add_account(&store, "perso", 10).await;
1063 exhaust(&store, &account, reset).await;
1064 let inner = succeeding();
1065 let provider = wrap(inner.clone(), &store);
1066
1067 let err = provider
1068 .invoke(&AgentConfig::new("hello").max_capacity_wait(Duration::ZERO))
1069 .await
1070 .unwrap_err();
1071
1072 let AgentError::NoCapacity { kind, next_reset } = err else {
1073 panic!("expected NoCapacity, got {err}");
1074 };
1075 assert_eq!(kind, ClaudeSubscriptionKind::ID);
1076 assert_eq!(next_reset, Some(reset));
1077 assert!(inner.seen().is_empty(), "the agent must not run");
1078 }
1079
1080 #[tokio::test]
1081 async fn worker_max_capacity_wait_applies_when_the_step_sets_none() {
1082 let store = Arc::new(store_with_key());
1083 let reset = Utc::now() + TimeDelta::hours(1);
1084 let account = add_account(&store, "perso", 10).await;
1085 exhaust(&store, &account, reset).await;
1086 let provider = wrap(succeeding(), &store).with_max_capacity_wait(Duration::ZERO);
1087
1088 let err = provider
1089 .invoke(&AgentConfig::new("hello"))
1090 .await
1091 .unwrap_err();
1092 assert!(matches!(err, AgentError::NoCapacity { .. }), "got {err}");
1093
1094 let err = provider
1095 .invoke(&AgentConfig::new("hello").max_capacity_wait(Duration::from_secs(7200)))
1096 .await
1097 .unwrap_err();
1098 assert!(
1099 matches!(err, AgentError::CapacityWait { .. }),
1100 "the step wait overrides the worker default, got {err}"
1101 );
1102 }
1103
1104 #[tokio::test]
1105 async fn reset_beyond_capacity_wait_fails_with_no_capacity() {
1106 let store = Arc::new(store_with_key());
1107 let reset = Utc::now() + TimeDelta::hours(8);
1108 let account = add_account(&store, "perso", 10).await;
1109 exhaust(&store, &account, reset).await;
1110 let inner = succeeding();
1111 let provider = wrap(inner.clone(), &store);
1112
1113 let err = provider
1114 .invoke(&AgentConfig::new("hello"))
1115 .await
1116 .unwrap_err();
1117
1118 let AgentError::NoCapacity { next_reset, .. } = err else {
1119 panic!("expected NoCapacity, got {err}");
1120 };
1121 assert_eq!(next_reset, Some(reset));
1122 assert!(inner.seen().is_empty(), "the agent must not run");
1123 }
1124
1125 #[tokio::test]
1126 async fn rejection_mid_step_fails_over_to_next_account() {
1127 let store = Arc::new(store_with_key());
1128 let reset = Utc::now() + TimeDelta::hours(1);
1129 let first = add_account(&store, "first", 10).await;
1130 let second = add_account(&store, "second", 20).await;
1131 let inner = Arc::new(RecordingProvider::new(
1132 Some(ClaudeSubscriptionKind::ID),
1133 Outcome::RejectOnce { resets_at: reset },
1134 ));
1135 let provider = wrap(inner.clone(), &store).with_strategy(Arc::new(Priority));
1136
1137 let output = provider.invoke(&AgentConfig::new("hello")).await.unwrap();
1138
1139 assert_eq!(output.account_id, Some(second.id.to_string()));
1140 let credentials: Vec<Option<String>> = inner.seen().into_iter().map(|s| s.0).collect();
1141 assert_eq!(
1142 credentials,
1143 vec![
1144 Some(format!("{TOKEN}-first")),
1145 Some(format!("{TOKEN}-second")),
1146 ]
1147 );
1148 let windows = store
1149 .list_provider_account_windows(vec![first.id])
1150 .await
1151 .unwrap();
1152 assert_eq!(windows.len(), 1);
1153 assert_eq!(windows[0].status, AccountWindowStatus::Rejected);
1154 assert_eq!(windows[0].resets_at, Some(reset));
1155 }
1156
1157 #[tokio::test]
1158 async fn rejection_on_every_account_sleeps_until_next_reset() {
1159 let store = Arc::new(store_with_key());
1160 let reset = Utc::now() + TimeDelta::hours(1);
1161 add_account(&store, "first", 10).await;
1162 add_account(&store, "second", 20).await;
1163 let inner = Arc::new(RecordingProvider::new(
1164 Some(ClaudeSubscriptionKind::ID),
1165 Outcome::Reject { resets_at: reset },
1166 ));
1167 let provider = wrap(inner.clone(), &store);
1168
1169 let err = provider
1170 .invoke(&AgentConfig::new("hello"))
1171 .await
1172 .unwrap_err();
1173
1174 let AgentError::CapacityWait { wake_at, .. } = err else {
1175 panic!("expected CapacityWait, got {err}");
1176 };
1177 assert_eq!(wake_at, reset);
1178 assert_eq!(inner.seen().len(), 2, "each account is tried once");
1179 }
1180
1181 #[tokio::test]
1182 async fn named_account_waits_without_failover() {
1183 let store = Arc::new(store_with_key());
1184 let reset = Utc::now() + TimeDelta::hours(1);
1185 add_account(&store, "first", 10).await;
1186 add_account(&store, "second", 20).await;
1187 let inner = Arc::new(RecordingProvider::new(
1188 Some(ClaudeSubscriptionKind::ID),
1189 Outcome::Reject { resets_at: reset },
1190 ));
1191 let provider = wrap(inner.clone(), &store);
1192
1193 let err = provider
1194 .invoke(&AgentConfig::new("hello").account("second"))
1195 .await
1196 .unwrap_err();
1197
1198 let AgentError::CapacityWait { wake_at, .. } = err else {
1199 panic!("expected CapacityWait, got {err}");
1200 };
1201 assert_eq!(wake_at, reset);
1202 let credentials: Vec<Option<String>> = inner.seen().into_iter().map(|s| s.0).collect();
1203 assert_eq!(credentials, vec![Some(format!("{TOKEN}-second"))]);
1204 }
1205
1206 #[tokio::test]
1207 async fn named_account_runs_under_that_account_only() {
1208 let store = Arc::new(store_with_key());
1209 add_account(&store, "first", 10).await;
1210 let second = add_account(&store, "second", 20).await;
1211 let inner = succeeding();
1212 let provider = wrap(inner.clone(), &store);
1213
1214 let output = provider
1215 .invoke(&AgentConfig::new("hello").account("second"))
1216 .await
1217 .unwrap();
1218
1219 assert_eq!(output.account_id, Some(second.id.to_string()));
1220 }
1221
1222 #[tokio::test]
1223 async fn unknown_account_name_fails_with_account_not_found() {
1224 let store = Arc::new(store_with_key());
1225 add_account(&store, "perso", 10).await;
1226 let inner = succeeding();
1227 let provider = wrap(inner.clone(), &store);
1228
1229 let err = provider
1230 .invoke(&AgentConfig::new("hello").account("missing"))
1231 .await
1232 .unwrap_err();
1233
1234 let AgentError::AccountNotFound { name } = err else {
1235 panic!("expected AccountNotFound, got {err}");
1236 };
1237 assert_eq!(name, "missing");
1238 assert!(inner.seen().is_empty(), "the agent must not run");
1239 }
1240
1241 #[tokio::test]
1242 async fn account_name_without_any_account_never_falls_back_to_worker_env() {
1243 let store = Arc::new(store_with_key());
1244 let inner = succeeding();
1245 let provider = wrap(inner.clone(), &store);
1246
1247 let err = provider
1248 .invoke(&AgentConfig::new("hello").account("perso"))
1249 .await
1250 .unwrap_err();
1251
1252 assert!(
1253 matches!(err, AgentError::AccountNotFound { .. }),
1254 "got {err}"
1255 );
1256 assert!(inner.seen().is_empty(), "the worker token must not be used");
1257 }
1258
1259 #[tokio::test]
1260 async fn account_pool_only_selects_tagged_accounts() {
1261 let store = Arc::new(store_with_key());
1262 add_account_with(&store, "untagged", 1, &[], None).await;
1263 let tagged = add_account_with(&store, "tagged", 50, &["batch"], None).await;
1264 let inner = succeeding();
1265 let provider = wrap(inner.clone(), &store).with_strategy(Arc::new(Priority));
1266
1267 let output = provider
1268 .invoke(&AgentConfig::new("hello").account_pool("batch"))
1269 .await
1270 .unwrap();
1271
1272 assert_eq!(output.account_id, Some(tagged.id.to_string()));
1273 }
1274
1275 #[tokio::test]
1276 async fn unknown_account_pool_fails_with_no_capacity() {
1277 let store = Arc::new(store_with_key());
1278 add_account_with(&store, "tagged", 10, &["batch"], None).await;
1279 let inner = succeeding();
1280 let provider = wrap(inner.clone(), &store);
1281
1282 let err = provider
1283 .invoke(&AgentConfig::new("hello").account_pool("nightly"))
1284 .await
1285 .unwrap_err();
1286
1287 let AgentError::NoCapacity { kind, next_reset } = err else {
1288 panic!("expected NoCapacity, got {err}");
1289 };
1290 assert_eq!(kind, ClaudeSubscriptionKind::ID);
1291 assert_eq!(next_reset, None);
1292 assert!(inner.seen().is_empty(), "the agent must not run");
1293 }
1294
1295 #[tokio::test]
1296 async fn saturated_account_retries_after_a_minute() {
1297 let store = Arc::new(store_with_key());
1298 add_account_with(&store, "busy", 10, &[], Some(0)).await;
1300 let inner = succeeding();
1301 let provider = wrap(inner.clone(), &store);
1302 let before = Utc::now();
1303
1304 let err = provider
1305 .invoke(&AgentConfig::new("hello"))
1306 .await
1307 .unwrap_err();
1308
1309 let AgentError::CapacityWait { wake_at, .. } = err else {
1310 panic!("expected CapacityWait, got {err}");
1311 };
1312 assert!(wake_at >= before + TimeDelta::seconds(60));
1313 assert!(wake_at <= Utc::now() + TimeDelta::seconds(60));
1314 assert!(inner.seen().is_empty(), "the agent must not run");
1315 }
1316
1317 #[tokio::test]
1318 async fn saturated_account_stops_waiting_past_the_cumulative_bound() {
1319 let store = Arc::new(store_with_key());
1320 add_account_with(&store, "busy", 10, &[], Some(0)).await;
1321 let provider = wrap(succeeding(), &store);
1322 let config = AgentConfig::new("hello")
1323 .max_capacity_wait(Duration::from_secs(3600))
1324 .capacity_wait_since(Utc::now() - TimeDelta::minutes(59) - TimeDelta::seconds(30));
1325
1326 let err = provider.invoke(&config).await.unwrap_err();
1327
1328 let AgentError::NoCapacity { next_reset, .. } = err else {
1329 panic!("expected NoCapacity, got {err}");
1330 };
1331 assert_eq!(next_reset, None);
1332 }
1333
1334 #[tokio::test]
1335 async fn worker_env_rejection_sleeps_until_resets_at() {
1336 let store = Arc::new(store_with_key());
1337 let reset = Utc::now() + TimeDelta::hours(1);
1338 let inner = Arc::new(RecordingProvider::new(
1339 Some(ClaudeSubscriptionKind::ID),
1340 Outcome::Reject { resets_at: reset },
1341 ));
1342 let provider = wrap(inner.clone(), &store);
1343
1344 let err = provider
1345 .invoke(&AgentConfig::new("hello"))
1346 .await
1347 .unwrap_err();
1348
1349 let AgentError::CapacityWait { kind, wake_at } = err else {
1350 panic!("expected CapacityWait, got {err}");
1351 };
1352 assert_eq!(kind, ClaudeSubscriptionKind::ID);
1353 assert_eq!(wake_at, reset);
1354 assert_eq!(inner.seen(), vec![(None, true)]);
1355 }
1356
1357 #[tokio::test]
1358 async fn worker_env_failure_without_rejection_is_returned_unchanged() {
1359 let store = Arc::new(store_with_key());
1360 let inner = Arc::new(RecordingProvider::new(
1361 Some(ClaudeSubscriptionKind::ID),
1362 Outcome::FailApi(500),
1363 ));
1364 let provider = wrap(inner, &store);
1365
1366 let err = provider
1367 .invoke(&AgentConfig::new("hello"))
1368 .await
1369 .unwrap_err();
1370
1371 assert!(matches!(
1372 err,
1373 AgentError::Api {
1374 status: Some(500),
1375 ..
1376 }
1377 ));
1378 }
1379
1380 #[tokio::test]
1381 async fn adding_account_wakes_capacity_sleepers() {
1382 let store = store_with_key();
1383 let far = Utc::now() + TimeDelta::hours(3);
1384 let mut runs = Vec::new();
1385 for kind in [ClaudeSubscriptionKind::ID, "other_kind"] {
1386 let run = store
1387 .create_run(NewRun {
1388 created_by: None,
1389 workflow_name: "capacity".to_string(),
1390 trigger: TriggerKind::Manual,
1391 payload: json!({}),
1392 max_retries: 0,
1393 handler_version: None,
1394 labels: HashMap::new(),
1395 scheduled_at: None,
1396 idempotency_key: None,
1397 concurrency_key: None,
1398 priority: 0,
1399 concurrency_limits: Vec::new(),
1400 max_cost_usd: None,
1401 worker_tags: Vec::new(),
1402 })
1403 .await
1404 .unwrap()
1405 .into_run();
1406 store
1407 .update_run_status(run.id, RunStatus::Running)
1408 .await
1409 .unwrap();
1410 store
1411 .update_run(
1412 run.id,
1413 RunUpdate {
1414 status: Some(RunStatus::Sleeping),
1415 scheduled_at: Some(far),
1416 capacity_wait_kind: Some(ProviderKind::from(kind)),
1417 ..RunUpdate::default()
1418 },
1419 )
1420 .await
1421 .unwrap();
1422 runs.push(run.id);
1423 }
1424
1425 add_account(&store, "fresh", 10).await;
1426 let now = Utc::now();
1427
1428 let woken = store.get_run(runs[0]).await.unwrap().unwrap();
1429 assert_eq!(woken.status.state, RunStatus::Sleeping);
1430 assert!(woken.scheduled_at.is_some_and(|at| at <= now));
1431 let untouched = store.get_run(runs[1]).await.unwrap().unwrap();
1432 assert_eq!(untouched.scheduled_at, Some(far));
1433 assert_eq!(
1434 untouched.capacity_wait_kind,
1435 Some(ProviderKind::from("other_kind"))
1436 );
1437 }
1438
1439 #[test]
1440 fn capacity_decision_zero_wait_fails_fast() {
1441 let now = Utc::now();
1442 let reset = now + TimeDelta::minutes(5);
1443 let err = capacity_decision("claude", Some(reset), true, Duration::ZERO, None, now);
1444 assert!(matches!(
1445 err,
1446 AgentError::NoCapacity { next_reset: Some(at), .. } if at == reset
1447 ));
1448 }
1449
1450 #[test]
1451 fn capacity_decision_waits_for_a_reset_within_the_bound() {
1452 let now = Utc::now();
1453 let reset = now + TimeDelta::hours(1);
1454 let err = capacity_decision(
1455 "claude",
1456 Some(reset),
1457 false,
1458 DEFAULT_MAX_CAPACITY_WAIT,
1459 None,
1460 now,
1461 );
1462 assert!(matches!(
1463 err,
1464 AgentError::CapacityWait { wake_at, .. } if wake_at == reset
1465 ));
1466 }
1467
1468 #[test]
1469 fn capacity_decision_counts_the_time_already_waited() {
1470 let now = Utc::now();
1471 let reset = now + TimeDelta::hours(1);
1472 let since = now - TimeDelta::hours(5) - TimeDelta::minutes(30);
1473 let err = capacity_decision(
1474 "claude",
1475 Some(reset),
1476 false,
1477 DEFAULT_MAX_CAPACITY_WAIT,
1478 Some(since),
1479 now,
1480 );
1481 assert!(matches!(
1482 err,
1483 AgentError::NoCapacity { next_reset: Some(at), .. } if at == reset
1484 ));
1485 }
1486
1487 #[test]
1488 fn capacity_decision_fails_for_a_reset_beyond_the_bound() {
1489 let now = Utc::now();
1490 let reset = now + TimeDelta::hours(7);
1491 let err = capacity_decision(
1492 "claude",
1493 Some(reset),
1494 false,
1495 DEFAULT_MAX_CAPACITY_WAIT,
1496 None,
1497 now,
1498 );
1499 assert!(matches!(
1500 err,
1501 AgentError::NoCapacity { next_reset: Some(at), .. } if at == reset
1502 ));
1503 }
1504
1505 #[test]
1506 fn capacity_decision_retries_concurrency_after_a_minute() {
1507 let now = Utc::now();
1508 let err = capacity_decision("claude", None, true, DEFAULT_MAX_CAPACITY_WAIT, None, now);
1509 assert!(matches!(
1510 err,
1511 AgentError::CapacityWait { wake_at, .. } if wake_at == now + TimeDelta::seconds(60)
1512 ));
1513 }
1514
1515 #[test]
1516 fn capacity_decision_bounds_the_concurrency_wait() {
1517 let now = Utc::now();
1518 let wait = Duration::from_secs(600);
1519 let within = capacity_decision(
1520 "claude",
1521 None,
1522 true,
1523 wait,
1524 Some(now - TimeDelta::minutes(9)),
1525 now,
1526 );
1527 assert!(matches!(within, AgentError::CapacityWait { .. }));
1528
1529 let beyond = capacity_decision(
1530 "claude",
1531 None,
1532 true,
1533 wait,
1534 Some(now - TimeDelta::minutes(9) - TimeDelta::seconds(1)),
1535 now,
1536 );
1537 assert!(matches!(
1538 beyond,
1539 AgentError::NoCapacity {
1540 next_reset: None,
1541 ..
1542 }
1543 ));
1544 }
1545
1546 #[test]
1547 fn capacity_decision_without_reset_or_concurrency_fails() {
1548 let now = Utc::now();
1549 let err = capacity_decision("claude", None, false, DEFAULT_MAX_CAPACITY_WAIT, None, now);
1550 assert!(matches!(
1551 err,
1552 AgentError::NoCapacity {
1553 next_reset: None,
1554 ..
1555 }
1556 ));
1557 }
1558
1559 #[tokio::test]
1560 async fn account_aware_provider_fails_when_credential_missing() {
1561 let store = Arc::new(store_with_key());
1562 let account = add_account(&store, "perso", 10).await;
1563 store.delete_secret(&account.secret_key).await.unwrap();
1564 let inner = Arc::new(RecordingProvider::new(
1565 Some(ClaudeSubscriptionKind::ID),
1566 Outcome::Succeed,
1567 ));
1568 let provider = wrap(inner, &store);
1569 let err = provider
1570 .invoke(&AgentConfig::new("hello"))
1571 .await
1572 .unwrap_err();
1573 let message = err.to_string();
1574 assert!(message.contains("credential of provider account 'perso' is missing"));
1575 assert!(!message.contains(TOKEN));
1576 }
1577
1578 #[tokio::test]
1579 async fn account_aware_provider_sets_output_account_id() {
1580 let store = Arc::new(store_with_key());
1581 let account = add_account(&store, "perso", 10).await;
1582 let inner = Arc::new(RecordingProvider::new(
1583 Some(ClaudeSubscriptionKind::ID),
1584 Outcome::Succeed,
1585 ));
1586 let provider = wrap(inner, &store);
1587 let output = provider.invoke(&AgentConfig::new("hello")).await.unwrap();
1588 assert_eq!(output.account_id, Some(account.id.to_string()));
1589
1590 let windows = store
1591 .list_provider_account_windows(vec![account.id])
1592 .await
1593 .unwrap();
1594 assert_eq!(windows.len(), 1);
1595 assert!((windows[0].utilization - 0.42).abs() < 1e-9);
1596 }
1597}