Skip to main content

ironflow_store/memory/
provider_account_store.rs

1//! In-memory [`ProviderAccountStore`] implementation.
2
3use std::cmp::Ordering;
4use std::collections::{HashMap, HashSet};
5
6use chrono::{DateTime, Utc};
7use uuid::Uuid;
8
9use crate::entities::{
10    NewProviderAccount, NewProviderAccountObservation, Page, ProviderAccount,
11    ProviderAccountCandidate, ProviderAccountUpdate, ProviderAccountUsagePoint,
12    ProviderAccountWindow, StepStatus,
13};
14use crate::error::StoreError;
15use crate::memory::{InMemoryStore, State};
16use crate::provider_account_store::ProviderAccountStore;
17use crate::store::StoreFuture;
18
19fn by_priority_then_name(a: &ProviderAccount, b: &ProviderAccount) -> Ordering {
20    a.priority
21        .cmp(&b.priority)
22        .then_with(|| a.name.cmp(&b.name))
23}
24
25fn windows_of(state: &State, ids: &[Uuid]) -> Vec<ProviderAccountWindow> {
26    let mut windows: Vec<ProviderAccountWindow> = state
27        .provider_account_windows
28        .values()
29        .filter(|w| ids.contains(&w.account_id))
30        .cloned()
31        .collect();
32    windows.sort_by(|a, b| {
33        a.account_id
34            .cmp(&b.account_id)
35            .then_with(|| a.window.cmp(&b.window))
36            .then_with(|| a.model_scope.cmp(&b.model_scope))
37    });
38    windows
39}
40
41/// Running steps per account, counting only runs whose worker lease is live.
42fn running_steps(state: &State, now: DateTime<Utc>) -> HashMap<Uuid, u32> {
43    let mut counts = HashMap::new();
44    for step in state.steps.values() {
45        let Some(account_id) = step.account_id else {
46            continue;
47        };
48        if step.status.state != StepStatus::Running {
49            continue;
50        }
51        let leased = state
52            .runs
53            .get(&step.run_id)
54            .and_then(|run| run.lease_expires_at)
55            .is_some_and(|expires| expires > now);
56        if leased {
57            *counts.entry(account_id).or_insert(0) += 1;
58        }
59    }
60    counts
61}
62
63impl ProviderAccountStore for InMemoryStore {
64    fn create_provider_account(&self, req: NewProviderAccount) -> StoreFuture<'_, ProviderAccount> {
65        Box::pin(async move {
66            let mut state = self.state.write().await;
67            if state.provider_accounts.values().any(|a| a.name == req.name) {
68                return Err(StoreError::DuplicateProviderAccount(req.name));
69            }
70            let now = Utc::now();
71            let account = ProviderAccount {
72                id: req.id,
73                name: req.name,
74                display_name: req.display_name,
75                kind: req.kind,
76                secret_key: req.secret_key,
77                enabled: req.enabled,
78                priority: req.priority,
79                tags: req.tags,
80                max_concurrency: req.max_concurrency,
81                alert_threshold: req.alert_threshold,
82                expires_at: req.expires_at,
83                plan: req.plan,
84                auth_failed_at: None,
85                created_by: req.created_by,
86                created_at: now,
87                updated_at: now,
88            };
89            state.provider_accounts.insert(account.id, account.clone());
90            Ok(account)
91        })
92    }
93
94    fn get_provider_account(&self, id: Uuid) -> StoreFuture<'_, Option<ProviderAccount>> {
95        Box::pin(async move {
96            let state = self.state.read().await;
97            Ok(state.provider_accounts.get(&id).cloned())
98        })
99    }
100
101    fn list_provider_accounts_by_ids(
102        &self,
103        ids: Vec<Uuid>,
104    ) -> StoreFuture<'_, Vec<ProviderAccount>> {
105        Box::pin(async move {
106            let wanted: HashSet<Uuid> = ids.into_iter().collect();
107            let state = self.state.read().await;
108            Ok(state
109                .provider_accounts
110                .values()
111                .filter(|a| wanted.contains(&a.id))
112                .cloned()
113                .collect())
114        })
115    }
116
117    fn find_provider_account_by_name(
118        &self,
119        name: &str,
120    ) -> StoreFuture<'_, Option<ProviderAccount>> {
121        let name = name.to_string();
122        Box::pin(async move {
123            let state = self.state.read().await;
124            Ok(state
125                .provider_accounts
126                .values()
127                .find(|a| a.name == name)
128                .cloned())
129        })
130    }
131
132    fn list_provider_accounts(
133        &self,
134        kind: Option<String>,
135        page: u32,
136        per_page: u32,
137    ) -> StoreFuture<'_, Page<ProviderAccount>> {
138        Box::pin(async move {
139            let state = self.state.read().await;
140            let mut all: Vec<ProviderAccount> = state
141                .provider_accounts
142                .values()
143                .filter(|a| kind.as_ref().is_none_or(|k| &a.kind == k))
144                .cloned()
145                .collect();
146            all.sort_by(by_priority_then_name);
147            let total = all.len() as u64;
148            let start = (page.saturating_sub(1) as usize) * (per_page as usize);
149            let items = all
150                .into_iter()
151                .skip(start)
152                .take(per_page as usize)
153                .collect();
154            Ok(Page {
155                items,
156                total,
157                page,
158                per_page,
159            })
160        })
161    }
162
163    fn update_provider_account(
164        &self,
165        id: Uuid,
166        update: ProviderAccountUpdate,
167    ) -> StoreFuture<'_, ProviderAccount> {
168        Box::pin(async move {
169            let mut state = self.state.write().await;
170            let account = state
171                .provider_accounts
172                .get_mut(&id)
173                .ok_or(StoreError::ProviderAccountNotFound(id))?;
174            if let Some(display_name) = update.display_name {
175                account.display_name = display_name;
176            }
177            if let Some(enabled) = update.enabled {
178                account.enabled = enabled;
179            }
180            if let Some(priority) = update.priority {
181                account.priority = priority;
182            }
183            if let Some(tags) = update.tags {
184                account.tags = tags;
185            }
186            if let Some(max_concurrency) = update.max_concurrency {
187                account.max_concurrency = max_concurrency;
188            }
189            if let Some(alert_threshold) = update.alert_threshold {
190                account.alert_threshold = alert_threshold;
191            }
192            if let Some(expires_at) = update.expires_at {
193                account.expires_at = expires_at;
194            }
195            if let Some(plan) = update.plan {
196                account.plan = plan;
197            }
198            if let Some(auth_failed_at) = update.auth_failed_at {
199                account.auth_failed_at = auth_failed_at;
200            }
201            account.updated_at = Utc::now();
202            Ok(account.clone())
203        })
204    }
205
206    fn delete_provider_account(&self, id: Uuid) -> StoreFuture<'_, bool> {
207        Box::pin(async move {
208            let mut state = self.state.write().await;
209            if state.provider_accounts.remove(&id).is_none() {
210                return Ok(false);
211            }
212            state
213                .provider_account_windows
214                .retain(|(account_id, _, _), _| *account_id != id);
215            state.provider_account_usage.retain(|u| u.account_id != id);
216            for step in state.steps.values_mut() {
217                if step.account_id == Some(id) {
218                    step.account_id = None;
219                }
220            }
221            Ok(true)
222        })
223    }
224
225    fn list_provider_account_windows(
226        &self,
227        ids: Vec<Uuid>,
228    ) -> StoreFuture<'_, Vec<ProviderAccountWindow>> {
229        Box::pin(async move {
230            let state = self.state.read().await;
231            Ok(windows_of(&state, &ids))
232        })
233    }
234
235    fn list_provider_account_usage(
236        &self,
237        id: Uuid,
238        since: DateTime<Utc>,
239    ) -> StoreFuture<'_, Vec<ProviderAccountUsagePoint>> {
240        Box::pin(async move {
241            let state = self.state.read().await;
242            let mut points: Vec<ProviderAccountUsagePoint> = state
243                .provider_account_usage
244                .iter()
245                .filter(|u| u.account_id == id && u.observed_at >= since)
246                .cloned()
247                .collect();
248            points.sort_by_key(|u| u.observed_at);
249            Ok(points)
250        })
251    }
252
253    fn record_provider_account_observation(
254        &self,
255        id: Uuid,
256        observation: NewProviderAccountObservation,
257    ) -> StoreFuture<'_, Vec<ProviderAccountWindow>> {
258        Box::pin(async move {
259            let mut state = self.state.write().await;
260            if !state.provider_accounts.contains_key(&id) {
261                return Err(StoreError::ProviderAccountNotFound(id));
262            }
263            let recorded = !observation.windows.is_empty();
264            for window in observation.windows {
265                let key = (
266                    id,
267                    window.window.clone(),
268                    window.model_scope.clone().unwrap_or_default(),
269                );
270                let newer_stored = state
271                    .provider_account_windows
272                    .get(&key)
273                    .is_some_and(|stored| stored.observed_at > window.observed_at);
274                if !newer_stored {
275                    state.provider_account_windows.insert(
276                        key,
277                        ProviderAccountWindow {
278                            account_id: id,
279                            window: window.window.clone(),
280                            utilization: window.utilization,
281                            resets_at: window.resets_at,
282                            status: window.status,
283                            model_scope: window.model_scope.clone(),
284                            observed_at: window.observed_at,
285                        },
286                    );
287                }
288                state
289                    .provider_account_usage
290                    .push(ProviderAccountUsagePoint {
291                        id: Uuid::now_v7(),
292                        account_id: id,
293                        window: window.window,
294                        utilization: window.utilization,
295                        resets_at: window.resets_at,
296                        status: window.status,
297                        model_scope: window.model_scope,
298                        observed_at: window.observed_at,
299                    });
300            }
301            let now = Utc::now();
302            if let Some(account) = state.provider_accounts.get_mut(&id) {
303                if observation.auth_failed {
304                    account.auth_failed_at = Some(now);
305                    account.updated_at = now;
306                } else if recorded {
307                    account.auth_failed_at = None;
308                }
309            }
310            Ok(windows_of(&state, &[id]))
311        })
312    }
313
314    fn purge_provider_account_usage(&self, before: DateTime<Utc>) -> StoreFuture<'_, u64> {
315        Box::pin(async move {
316            let mut state = self.state.write().await;
317            let len_before = state.provider_account_usage.len();
318            state
319                .provider_account_usage
320                .retain(|u| u.observed_at >= before);
321            Ok((len_before - state.provider_account_usage.len()) as u64)
322        })
323    }
324
325    fn list_provider_account_candidates(
326        &self,
327        kind: String,
328    ) -> StoreFuture<'_, Vec<ProviderAccountCandidate>> {
329        Box::pin(async move {
330            let state = self.state.read().await;
331            let now = Utc::now();
332            let mut accounts: Vec<ProviderAccount> = state
333                .provider_accounts
334                .values()
335                .filter(|a| {
336                    a.kind == kind && a.enabled && a.auth_failed_at.is_none() && a.expires_at > now
337                })
338                .cloned()
339                .collect();
340            accounts.sort_by(by_priority_then_name);
341            let running = running_steps(&state, now);
342            Ok(accounts
343                .into_iter()
344                .map(|account| {
345                    let windows = windows_of(&state, &[account.id]);
346                    let running_steps = running.get(&account.id).copied().unwrap_or(0);
347                    ProviderAccountCandidate {
348                        account,
349                        windows,
350                        running_steps,
351                    }
352                })
353                .collect())
354        })
355    }
356}