Skip to main content

ironflow_engine/
accounts.rs

1//! Provider Account selection and injection for agent steps.
2//!
3//! [`AccountAwareProvider`] wraps the worker's [`AgentProvider`]. For each
4//! agent invocation on a provider that can inject a Provider Account
5//! credential ([`AgentProvider::account_kind`]), it:
6//!
7//! 1. lists the candidate accounts of that kind from the store,
8//! 2. picks one with an [`AccountStrategy`],
9//! 3. reads its credential and runs the invocation under it,
10//! 4. records the usage windows observed during the invocation.
11//!
12//! With no account of the kind in the store, or a provider that cannot
13//! inject one, the invocation runs unchanged with the worker environment.
14//!
15//! # Examples
16//!
17//! ```no_run
18//! use std::sync::Arc;
19//! use ironflow_core::account_strategy::Priority;
20//! use ironflow_core::providers::claude::ClaudeCodeProvider;
21//! use ironflow_engine::accounts::AccountAwareProvider;
22//! use ironflow_store::memory::InMemoryStore;
23//!
24//! let provider = AccountAwareProvider::new(
25//!     Arc::new(ClaudeCodeProvider::new()),
26//!     Arc::new(InMemoryStore::new()),
27//! )
28//! .with_strategy(Arc::new(Priority));
29//! # let _ = provider;
30//! ```
31
32use std::collections::HashMap;
33use std::fmt;
34use std::sync::Arc;
35
36use chrono::Utc;
37use ironflow_core::account::{
38    AccountKind, AccountSession, AccountWindow, ClaudeSubscriptionKind, RateLimitRecorder,
39    WindowStatus,
40};
41use ironflow_core::account_strategy::{
42    AccountCandidate, AccountStrategy, LeastUtilized, select_account,
43};
44use ironflow_core::error::AgentError;
45use ironflow_core::provider::{
46    AgentConfig, AgentOutput, AgentProvider, InvokeFuture, LogSink, ReleaseFuture,
47};
48use ironflow_store::entities::{
49    AccountWindowStatus, NewAccountWindow, NewProviderAccountObservation, ProviderAccount,
50    ProviderAccountCandidate, ProviderAccountWindow,
51};
52use ironflow_store::store::Store;
53use tracing::{debug, info, warn};
54
55/// Convert a stored window into the core representation.
56///
57/// # Examples
58///
59/// ```
60/// use chrono::Utc;
61/// use ironflow_engine::accounts::window_from_store;
62/// use ironflow_store::entities::{AccountWindowStatus, ProviderAccountWindow};
63/// use uuid::Uuid;
64///
65/// let window = window_from_store(&ProviderAccountWindow {
66///     account_id: Uuid::now_v7(),
67///     window: "five_hour".to_string(),
68///     utilization: 0.4,
69///     resets_at: None,
70///     status: AccountWindowStatus::Allowed,
71///     model_scope: None,
72///     observed_at: Utc::now(),
73/// });
74/// assert_eq!(window.window, "five_hour");
75/// ```
76pub fn window_from_store(window: &ProviderAccountWindow) -> AccountWindow {
77    AccountWindow {
78        window: window.window.clone(),
79        utilization: window.utilization,
80        resets_at: window.resets_at,
81        status: match window.status {
82            AccountWindowStatus::Allowed => WindowStatus::Allowed,
83            AccountWindowStatus::AllowedWarning => WindowStatus::AllowedWarning,
84            AccountWindowStatus::Rejected => WindowStatus::Rejected,
85        },
86        model_scope: window.model_scope.clone(),
87        observed_at: window.observed_at,
88    }
89}
90
91/// Convert an observed core window into the store representation.
92///
93/// # Examples
94///
95/// ```
96/// use chrono::Utc;
97/// use ironflow_core::account::{AccountWindow, WindowStatus};
98/// use ironflow_engine::accounts::window_to_store;
99/// use ironflow_store::entities::AccountWindowStatus;
100///
101/// let window = window_to_store(AccountWindow {
102///     window: "seven_day".to_string(),
103///     utilization: 1.0,
104///     resets_at: None,
105///     status: WindowStatus::Rejected,
106///     model_scope: Some("opus".to_string()),
107///     observed_at: Utc::now(),
108/// });
109/// assert_eq!(window.status, AccountWindowStatus::Rejected);
110/// ```
111pub fn window_to_store(window: AccountWindow) -> NewAccountWindow {
112    NewAccountWindow {
113        window: window.window,
114        utilization: window.utilization,
115        resets_at: window.resets_at,
116        status: match window.status {
117            WindowStatus::Allowed => AccountWindowStatus::Allowed,
118            WindowStatus::AllowedWarning => AccountWindowStatus::AllowedWarning,
119            WindowStatus::Rejected => AccountWindowStatus::Rejected,
120        },
121        model_scope: window.model_scope,
122        observed_at: window.observed_at,
123    }
124}
125
126fn to_core_candidate(candidate: &ProviderAccountCandidate) -> AccountCandidate {
127    AccountCandidate {
128        id: candidate.account.id.to_string(),
129        name: candidate.account.name.clone(),
130        priority: candidate.account.priority,
131        max_concurrency: candidate.account.max_concurrency,
132        running_steps: candidate.running_steps,
133        windows: candidate.windows.iter().map(window_from_store).collect(),
134    }
135}
136
137fn resolution_error(message: String) -> AgentError {
138    AgentError::ProcessFailed {
139        exit_code: -1,
140        stderr: message,
141    }
142}
143
144/// An [`AgentProvider`] that runs each invocation under a Provider Account.
145///
146/// See the [module documentation](self).
147pub struct AccountAwareProvider {
148    inner: Arc<dyn AgentProvider>,
149    store: Arc<dyn Store>,
150    strategy: Arc<dyn AccountStrategy>,
151    kinds: HashMap<&'static str, Arc<dyn AccountKind>>,
152}
153
154impl fmt::Debug for AccountAwareProvider {
155    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
156        f.debug_struct("AccountAwareProvider")
157            .field("strategy", &self.strategy.name())
158            .field("kinds", &self.kinds.keys().collect::<Vec<_>>())
159            .finish_non_exhaustive()
160    }
161}
162
163impl AccountAwareProvider {
164    /// Wrap `inner`, reading accounts from `store`, with the
165    /// [`LeastUtilized`] strategy and the [`ClaudeSubscriptionKind`] kind.
166    ///
167    /// # Examples
168    ///
169    /// See the [module documentation](self).
170    pub fn new(inner: Arc<dyn AgentProvider>, store: Arc<dyn Store>) -> Self {
171        let claude: Arc<dyn AccountKind> = Arc::new(ClaudeSubscriptionKind::new());
172        Self {
173            inner,
174            store,
175            strategy: Arc::new(LeastUtilized),
176            kinds: HashMap::from([(claude.id(), claude)]),
177        }
178    }
179
180    /// Replace the account selection strategy.
181    ///
182    /// # Examples
183    ///
184    /// See the [module documentation](self).
185    pub fn with_strategy(mut self, strategy: Arc<dyn AccountStrategy>) -> Self {
186        self.strategy = strategy;
187        self
188    }
189
190    /// Register an additional account kind (or replace one with the same id).
191    ///
192    /// # Examples
193    ///
194    /// ```no_run
195    /// use std::sync::Arc;
196    /// use ironflow_core::account::ClaudeSubscriptionKind;
197    /// use ironflow_core::providers::claude::ClaudeCodeProvider;
198    /// use ironflow_engine::accounts::AccountAwareProvider;
199    /// use ironflow_store::memory::InMemoryStore;
200    ///
201    /// let provider = AccountAwareProvider::new(
202    ///     Arc::new(ClaudeCodeProvider::new()),
203    ///     Arc::new(InMemoryStore::new()),
204    /// )
205    /// .with_kind(Arc::new(ClaudeSubscriptionKind::with_api_base("http://proxy:8080")));
206    /// # let _ = provider;
207    /// ```
208    pub fn with_kind(mut self, kind: Arc<dyn AccountKind>) -> Self {
209        self.kinds.insert(kind.id(), kind);
210        self
211    }
212
213    async fn invoke_inner(
214        &self,
215        config: &AgentConfig,
216        sink: Option<Arc<dyn LogSink>>,
217    ) -> Result<AgentOutput, AgentError> {
218        match sink {
219            Some(sink) => self.inner.invoke_with_logs(config, sink).await,
220            None => self.inner.invoke(config).await,
221        }
222    }
223
224    async fn run(
225        &self,
226        config: &AgentConfig,
227        sink: Option<Arc<dyn LogSink>>,
228    ) -> Result<AgentOutput, AgentError> {
229        let Some(kind_id) = self.inner.account_kind() else {
230            return self.invoke_inner(config, sink).await;
231        };
232        let Some(kind) = self.kinds.get(kind_id).cloned() else {
233            debug!(
234                kind = kind_id,
235                "no account kind registered, using worker environment"
236            );
237            return self.invoke_inner(config, sink).await;
238        };
239
240        let candidates = self
241            .store
242            .list_provider_account_candidates(kind_id.to_string())
243            .await
244            .map_err(|e| resolution_error(format!("provider account resolution failed: {e}")))?;
245        if candidates.is_empty() {
246            debug!(
247                kind = kind_id,
248                "no provider account for kind, using worker environment"
249            );
250            return self.invoke_inner(config, sink).await;
251        }
252
253        let core_candidates: Vec<AccountCandidate> =
254            candidates.iter().map(to_core_candidate).collect();
255        let now = Utc::now();
256        let Some(selected) =
257            select_account(self.strategy.as_ref(), &core_candidates, &config.model, now)
258        else {
259            let next_reset = core_candidates
260                .iter()
261                .flat_map(|c| c.windows.iter())
262                .filter(|w| w.applies_to(&config.model) && w.is_exhausted(now))
263                .filter_map(|w| w.resets_at)
264                .min()
265                .map_or_else(|| "unknown".to_string(), |at| at.to_rfc3339());
266            return Err(resolution_error(format!(
267                "no provider account available for {kind_id}: all limited or at max_concurrency (next reset {next_reset})"
268            )));
269        };
270        let account: &ProviderAccount = &candidates
271            .iter()
272            .find(|c| c.account.id.to_string() == selected.id)
273            .ok_or_else(|| resolution_error("selected provider account vanished".to_string()))?
274            .account;
275
276        let missing = || {
277            resolution_error(format!(
278                "credential of provider account '{}' is missing",
279                account.name
280            ))
281        };
282        let secret = match self.store.get_secret(&account.secret_key).await {
283            Ok(Some(secret)) => secret,
284            Ok(None) => return Err(missing()),
285            Err(e) => {
286                warn!(account = %account.name, error = %e, "failed to read provider account credential");
287                return Err(missing());
288            }
289        };
290
291        info!(
292            account = %account.name,
293            strategy = self.strategy.name(),
294            model = %config.model,
295            "selected provider account"
296        );
297
298        let recorder = RateLimitRecorder::default();
299        // Rate-limit events only appear in stream-json, hence verbose.
300        let account_config = config
301            .clone()
302            .verbose(true)
303            .account_session(AccountSession::new(
304                kind.credential(&secret.value),
305                recorder.clone(),
306            ));
307
308        let result = self.invoke_inner(&account_config, sink).await;
309
310        let windows = recorder.take();
311        let auth_failed = matches!(
312            result,
313            Err(AgentError::Api {
314                status: Some(401 | 403),
315                ..
316            })
317        );
318        if !windows.is_empty() || auth_failed {
319            let observation = NewProviderAccountObservation {
320                windows: windows.into_iter().map(window_to_store).collect(),
321                auth_failed,
322            };
323            if let Err(e) = self
324                .store
325                .record_provider_account_observation(account.id, observation)
326                .await
327            {
328                warn!(account = %account.name, error = %e, "failed to record provider account usage");
329            }
330        }
331
332        result.map(|mut output| {
333            output.account_id = Some(account.id.to_string());
334            output
335        })
336    }
337}
338
339impl AgentProvider for AccountAwareProvider {
340    fn invoke<'a>(&'a self, config: &'a AgentConfig) -> InvokeFuture<'a> {
341        Box::pin(self.run(config, None))
342    }
343
344    fn invoke_with_logs<'a>(
345        &'a self,
346        config: &'a AgentConfig,
347        log_sink: Arc<dyn LogSink>,
348    ) -> InvokeFuture<'a> {
349        Box::pin(self.run(config, Some(log_sink)))
350    }
351
352    fn release_run<'a>(&'a self, run_id: &'a str) -> ReleaseFuture<'a> {
353        self.inner.release_run(run_id)
354    }
355
356    fn account_kind(&self) -> Option<&'static str> {
357        self.inner.account_kind()
358    }
359}
360
361#[cfg(test)]
362mod tests {
363    use std::sync::Mutex;
364
365    use chrono::TimeDelta;
366    use ironflow_store::crypto::KeyRing;
367    use ironflow_store::entities::{NewProviderAccount, provider_account_secret_key};
368    use ironflow_store::memory::InMemoryStore;
369    use ironflow_store::provider_account_store::ProviderAccountStore;
370    use ironflow_store::secret_store::SecretStore;
371    use serde_json::json;
372    use uuid::Uuid;
373
374    use super::*;
375
376    const TOKEN: &str = "sk-ant-oat01-test-token-abcdefghijklmnopqrstuvwxyz";
377
378    /// What the test provider does once invoked.
379    #[derive(Clone, Copy)]
380    enum Outcome {
381        Succeed,
382        FailApi(u16),
383    }
384
385    /// A provider that behaves like the Claude transports: it reads the
386    /// injected credential and reports a rate-limit window to the recorder.
387    struct RecordingProvider {
388        kind: Option<&'static str>,
389        outcome: Outcome,
390        seen: Mutex<Vec<(Option<String>, bool)>>,
391    }
392
393    impl RecordingProvider {
394        fn new(kind: Option<&'static str>, outcome: Outcome) -> Self {
395            Self {
396                kind,
397                outcome,
398                seen: Mutex::new(Vec::new()),
399            }
400        }
401
402        fn seen(&self) -> Vec<(Option<String>, bool)> {
403            self.seen.lock().unwrap().clone()
404        }
405    }
406
407    impl AgentProvider for RecordingProvider {
408        fn invoke<'a>(&'a self, config: &'a AgentConfig) -> InvokeFuture<'a> {
409            Box::pin(async move {
410                let credential = config
411                    .account
412                    .as_ref()
413                    .map(|s| s.credential().expose().to_string());
414                self.seen.lock().unwrap().push((credential, config.verbose));
415                if let Some(session) = &config.account {
416                    session.recorder().record(AccountWindow {
417                        window: "five_hour".to_string(),
418                        utilization: 0.42,
419                        resets_at: Some(Utc::now() + TimeDelta::hours(2)),
420                        status: WindowStatus::Allowed,
421                        model_scope: None,
422                        observed_at: Utc::now(),
423                    });
424                }
425                match self.outcome {
426                    Outcome::Succeed => Ok(AgentOutput::new(json!("done"))),
427                    Outcome::FailApi(status) => Err(AgentError::Api {
428                        status: Some(status),
429                        code: None,
430                        message: "API Error".to_string(),
431                    }),
432                }
433            })
434        }
435
436        fn account_kind(&self) -> Option<&'static str> {
437            self.kind
438        }
439    }
440
441    fn store_with_key() -> InMemoryStore {
442        let mut store = InMemoryStore::new();
443        let spec = format!("1:{}", "aa".repeat(32));
444        store.set_key_ring(KeyRing::from_spec(&spec, Some(1)).unwrap());
445        store
446    }
447
448    async fn add_account(store: &InMemoryStore, name: &str, priority: i32) -> ProviderAccount {
449        let id = Uuid::now_v7();
450        let secret_key = provider_account_secret_key(id);
451        store
452            .set_secret(&secret_key, &format!("{TOKEN}-{name}"))
453            .await
454            .unwrap();
455        store
456            .create_provider_account(NewProviderAccount {
457                id,
458                name: name.to_string(),
459                display_name: name.to_string(),
460                kind: ClaudeSubscriptionKind::ID.to_string(),
461                secret_key,
462                enabled: true,
463                priority,
464                tags: Vec::new(),
465                max_concurrency: None,
466                alert_threshold: 0.8,
467                expires_at: Utc::now() + TimeDelta::days(30),
468                plan: None,
469                created_by: None,
470            })
471            .await
472            .unwrap()
473    }
474
475    fn wrap(inner: Arc<RecordingProvider>, store: &Arc<InMemoryStore>) -> AccountAwareProvider {
476        let store: Arc<dyn Store> = store.clone();
477        AccountAwareProvider::new(inner, store)
478    }
479
480    #[tokio::test]
481    async fn account_aware_provider_injects_selected_account() {
482        let store = Arc::new(store_with_key());
483        add_account(&store, "busy", 10).await;
484        let busy = store
485            .find_provider_account_by_name("busy")
486            .await
487            .unwrap()
488            .unwrap();
489        store
490            .record_provider_account_observation(
491                busy.id,
492                NewProviderAccountObservation {
493                    windows: vec![NewAccountWindow {
494                        window: "five_hour".to_string(),
495                        utilization: 0.9,
496                        resets_at: Some(Utc::now() + TimeDelta::hours(1)),
497                        status: AccountWindowStatus::Allowed,
498                        model_scope: None,
499                        observed_at: Utc::now(),
500                    }],
501                    auth_failed: false,
502                },
503            )
504            .await
505            .unwrap();
506        add_account(&store, "idle", 20).await;
507
508        let inner = Arc::new(RecordingProvider::new(
509            Some(ClaudeSubscriptionKind::ID),
510            Outcome::Succeed,
511        ));
512        let provider = wrap(inner.clone(), &store);
513        provider.invoke(&AgentConfig::new("hello")).await.unwrap();
514
515        let seen = inner.seen();
516        assert_eq!(seen.len(), 1);
517        assert_eq!(seen[0].0.as_deref(), Some(format!("{TOKEN}-idle").as_str()));
518        assert!(seen[0].1, "verbose must be forced for rate_limit_event");
519    }
520
521    #[tokio::test]
522    async fn account_aware_provider_passthrough_without_accounts() {
523        let store = Arc::new(store_with_key());
524        let inner = Arc::new(RecordingProvider::new(
525            Some(ClaudeSubscriptionKind::ID),
526            Outcome::Succeed,
527        ));
528        let provider = wrap(inner.clone(), &store);
529        let output = provider.invoke(&AgentConfig::new("hello")).await.unwrap();
530        assert_eq!(output.account_id, None);
531        assert_eq!(inner.seen(), vec![(None, false)]);
532    }
533
534    #[tokio::test]
535    async fn account_aware_provider_passthrough_for_kindless_provider() {
536        let store = Arc::new(store_with_key());
537        add_account(&store, "perso", 10).await;
538        let inner = Arc::new(RecordingProvider::new(None, Outcome::Succeed));
539        let provider = wrap(inner.clone(), &store);
540        let output = provider.invoke(&AgentConfig::new("hello")).await.unwrap();
541        assert_eq!(output.account_id, None);
542        assert_eq!(inner.seen(), vec![(None, false)]);
543        assert_eq!(provider.account_kind(), None);
544    }
545
546    #[tokio::test]
547    async fn account_aware_provider_records_windows_on_error() {
548        let store = Arc::new(store_with_key());
549        let account = add_account(&store, "perso", 10).await;
550        let inner = Arc::new(RecordingProvider::new(
551            Some(ClaudeSubscriptionKind::ID),
552            Outcome::FailApi(500),
553        ));
554        let provider = wrap(inner, &store);
555        let err = provider
556            .invoke(&AgentConfig::new("hello"))
557            .await
558            .unwrap_err();
559        assert!(matches!(
560            err,
561            AgentError::Api {
562                status: Some(500),
563                ..
564            }
565        ));
566
567        let windows = store
568            .list_provider_account_windows(vec![account.id])
569            .await
570            .unwrap();
571        assert_eq!(windows.len(), 1);
572        assert_eq!(windows[0].window, "five_hour");
573        let stored = store
574            .get_provider_account(account.id)
575            .await
576            .unwrap()
577            .unwrap();
578        assert!(stored.auth_failed_at.is_none());
579    }
580
581    #[tokio::test]
582    async fn account_aware_provider_marks_auth_failed_on_401() {
583        let store = Arc::new(store_with_key());
584        let account = add_account(&store, "perso", 10).await;
585        let inner = Arc::new(RecordingProvider::new(
586            Some(ClaudeSubscriptionKind::ID),
587            Outcome::FailApi(401),
588        ));
589        let provider = wrap(inner, &store);
590        provider
591            .invoke(&AgentConfig::new("hello"))
592            .await
593            .unwrap_err();
594
595        let stored = store
596            .get_provider_account(account.id)
597            .await
598            .unwrap()
599            .unwrap();
600        assert!(stored.auth_failed_at.is_some());
601        let candidates = store
602            .list_provider_account_candidates(ClaudeSubscriptionKind::ID.to_string())
603            .await
604            .unwrap();
605        assert!(candidates.is_empty(), "a rejected token is not a candidate");
606    }
607
608    #[tokio::test]
609    async fn account_aware_provider_fails_when_all_exhausted() {
610        let store = Arc::new(store_with_key());
611        let account = add_account(&store, "perso", 10).await;
612        let reset = Utc::now() + TimeDelta::hours(1);
613        store
614            .record_provider_account_observation(
615                account.id,
616                NewProviderAccountObservation {
617                    windows: vec![NewAccountWindow {
618                        window: "five_hour".to_string(),
619                        utilization: 1.0,
620                        resets_at: Some(reset),
621                        status: AccountWindowStatus::Rejected,
622                        model_scope: None,
623                        observed_at: Utc::now(),
624                    }],
625                    auth_failed: false,
626                },
627            )
628            .await
629            .unwrap();
630        let inner = Arc::new(RecordingProvider::new(
631            Some(ClaudeSubscriptionKind::ID),
632            Outcome::Succeed,
633        ));
634        let provider = wrap(inner.clone(), &store);
635        let err = provider
636            .invoke(&AgentConfig::new("hello"))
637            .await
638            .unwrap_err();
639        let AgentError::ProcessFailed { stderr, .. } = err else {
640            panic!("expected ProcessFailed");
641        };
642        assert!(stderr.contains("no provider account available"));
643        assert!(stderr.contains(&reset.to_rfc3339()));
644        assert!(inner.seen().is_empty(), "the agent must not run");
645    }
646
647    #[tokio::test]
648    async fn account_aware_provider_fails_when_credential_missing() {
649        let store = Arc::new(store_with_key());
650        let account = add_account(&store, "perso", 10).await;
651        store.delete_secret(&account.secret_key).await.unwrap();
652        let inner = Arc::new(RecordingProvider::new(
653            Some(ClaudeSubscriptionKind::ID),
654            Outcome::Succeed,
655        ));
656        let provider = wrap(inner, &store);
657        let err = provider
658            .invoke(&AgentConfig::new("hello"))
659            .await
660            .unwrap_err();
661        let message = err.to_string();
662        assert!(message.contains("credential of provider account 'perso' is missing"));
663        assert!(!message.contains(TOKEN));
664    }
665
666    #[tokio::test]
667    async fn account_aware_provider_sets_output_account_id() {
668        let store = Arc::new(store_with_key());
669        let account = add_account(&store, "perso", 10).await;
670        let inner = Arc::new(RecordingProvider::new(
671            Some(ClaudeSubscriptionKind::ID),
672            Outcome::Succeed,
673        ));
674        let provider = wrap(inner, &store);
675        let output = provider.invoke(&AgentConfig::new("hello")).await.unwrap();
676        assert_eq!(output.account_id, Some(account.id.to_string()));
677
678        let windows = store
679            .list_provider_account_windows(vec![account.id])
680            .await
681            .unwrap();
682        assert_eq!(windows.len(), 1);
683        assert!((windows[0].utilization - 0.42).abs() < 1e-9);
684    }
685}