1use 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
55pub 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
91pub 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
144pub 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 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 pub fn with_strategy(mut self, strategy: Arc<dyn AccountStrategy>) -> Self {
186 self.strategy = strategy;
187 self
188 }
189
190 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 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 #[derive(Clone, Copy)]
380 enum Outcome {
381 Succeed,
382 FailApi(u16),
383 }
384
385 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}