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;
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 find_provider_account_by_name(
102        &self,
103        name: &str,
104    ) -> StoreFuture<'_, Option<ProviderAccount>> {
105        let name = name.to_string();
106        Box::pin(async move {
107            let state = self.state.read().await;
108            Ok(state
109                .provider_accounts
110                .values()
111                .find(|a| a.name == name)
112                .cloned())
113        })
114    }
115
116    fn list_provider_accounts(
117        &self,
118        kind: Option<String>,
119        page: u32,
120        per_page: u32,
121    ) -> StoreFuture<'_, Page<ProviderAccount>> {
122        Box::pin(async move {
123            let state = self.state.read().await;
124            let mut all: Vec<ProviderAccount> = state
125                .provider_accounts
126                .values()
127                .filter(|a| kind.as_ref().is_none_or(|k| &a.kind == k))
128                .cloned()
129                .collect();
130            all.sort_by(by_priority_then_name);
131            let total = all.len() as u64;
132            let start = (page.saturating_sub(1) as usize) * (per_page as usize);
133            let items = all
134                .into_iter()
135                .skip(start)
136                .take(per_page as usize)
137                .collect();
138            Ok(Page {
139                items,
140                total,
141                page,
142                per_page,
143            })
144        })
145    }
146
147    fn update_provider_account(
148        &self,
149        id: Uuid,
150        update: ProviderAccountUpdate,
151    ) -> StoreFuture<'_, ProviderAccount> {
152        Box::pin(async move {
153            let mut state = self.state.write().await;
154            let account = state
155                .provider_accounts
156                .get_mut(&id)
157                .ok_or(StoreError::ProviderAccountNotFound(id))?;
158            if let Some(display_name) = update.display_name {
159                account.display_name = display_name;
160            }
161            if let Some(enabled) = update.enabled {
162                account.enabled = enabled;
163            }
164            if let Some(priority) = update.priority {
165                account.priority = priority;
166            }
167            if let Some(tags) = update.tags {
168                account.tags = tags;
169            }
170            if let Some(max_concurrency) = update.max_concurrency {
171                account.max_concurrency = max_concurrency;
172            }
173            if let Some(alert_threshold) = update.alert_threshold {
174                account.alert_threshold = alert_threshold;
175            }
176            if let Some(expires_at) = update.expires_at {
177                account.expires_at = expires_at;
178            }
179            if let Some(plan) = update.plan {
180                account.plan = plan;
181            }
182            if let Some(auth_failed_at) = update.auth_failed_at {
183                account.auth_failed_at = auth_failed_at;
184            }
185            account.updated_at = Utc::now();
186            Ok(account.clone())
187        })
188    }
189
190    fn delete_provider_account(&self, id: Uuid) -> StoreFuture<'_, bool> {
191        Box::pin(async move {
192            let mut state = self.state.write().await;
193            if state.provider_accounts.remove(&id).is_none() {
194                return Ok(false);
195            }
196            state
197                .provider_account_windows
198                .retain(|(account_id, _, _), _| *account_id != id);
199            state.provider_account_usage.retain(|u| u.account_id != id);
200            for step in state.steps.values_mut() {
201                if step.account_id == Some(id) {
202                    step.account_id = None;
203                }
204            }
205            Ok(true)
206        })
207    }
208
209    fn list_provider_account_windows(
210        &self,
211        ids: Vec<Uuid>,
212    ) -> StoreFuture<'_, Vec<ProviderAccountWindow>> {
213        Box::pin(async move {
214            let state = self.state.read().await;
215            Ok(windows_of(&state, &ids))
216        })
217    }
218
219    fn list_provider_account_usage(
220        &self,
221        id: Uuid,
222        since: DateTime<Utc>,
223    ) -> StoreFuture<'_, Vec<ProviderAccountUsagePoint>> {
224        Box::pin(async move {
225            let state = self.state.read().await;
226            let mut points: Vec<ProviderAccountUsagePoint> = state
227                .provider_account_usage
228                .iter()
229                .filter(|u| u.account_id == id && u.observed_at >= since)
230                .cloned()
231                .collect();
232            points.sort_by_key(|u| u.observed_at);
233            Ok(points)
234        })
235    }
236
237    fn record_provider_account_observation(
238        &self,
239        id: Uuid,
240        observation: NewProviderAccountObservation,
241    ) -> StoreFuture<'_, Vec<ProviderAccountWindow>> {
242        Box::pin(async move {
243            let mut state = self.state.write().await;
244            if !state.provider_accounts.contains_key(&id) {
245                return Err(StoreError::ProviderAccountNotFound(id));
246            }
247            let recorded = !observation.windows.is_empty();
248            for window in observation.windows {
249                let key = (
250                    id,
251                    window.window.clone(),
252                    window.model_scope.clone().unwrap_or_default(),
253                );
254                let newer_stored = state
255                    .provider_account_windows
256                    .get(&key)
257                    .is_some_and(|stored| stored.observed_at > window.observed_at);
258                if !newer_stored {
259                    state.provider_account_windows.insert(
260                        key,
261                        ProviderAccountWindow {
262                            account_id: id,
263                            window: window.window.clone(),
264                            utilization: window.utilization,
265                            resets_at: window.resets_at,
266                            status: window.status,
267                            model_scope: window.model_scope.clone(),
268                            observed_at: window.observed_at,
269                        },
270                    );
271                }
272                state
273                    .provider_account_usage
274                    .push(ProviderAccountUsagePoint {
275                        id: Uuid::now_v7(),
276                        account_id: id,
277                        window: window.window,
278                        utilization: window.utilization,
279                        resets_at: window.resets_at,
280                        status: window.status,
281                        model_scope: window.model_scope,
282                        observed_at: window.observed_at,
283                    });
284            }
285            let now = Utc::now();
286            if let Some(account) = state.provider_accounts.get_mut(&id) {
287                if observation.auth_failed {
288                    account.auth_failed_at = Some(now);
289                    account.updated_at = now;
290                } else if recorded {
291                    account.auth_failed_at = None;
292                }
293            }
294            Ok(windows_of(&state, &[id]))
295        })
296    }
297
298    fn purge_provider_account_usage(&self, before: DateTime<Utc>) -> StoreFuture<'_, u64> {
299        Box::pin(async move {
300            let mut state = self.state.write().await;
301            let len_before = state.provider_account_usage.len();
302            state
303                .provider_account_usage
304                .retain(|u| u.observed_at >= before);
305            Ok((len_before - state.provider_account_usage.len()) as u64)
306        })
307    }
308
309    fn list_provider_account_candidates(
310        &self,
311        kind: String,
312    ) -> StoreFuture<'_, Vec<ProviderAccountCandidate>> {
313        Box::pin(async move {
314            let state = self.state.read().await;
315            let now = Utc::now();
316            let mut accounts: Vec<ProviderAccount> = state
317                .provider_accounts
318                .values()
319                .filter(|a| {
320                    a.kind == kind && a.enabled && a.auth_failed_at.is_none() && a.expires_at > now
321                })
322                .cloned()
323                .collect();
324            accounts.sort_by(by_priority_then_name);
325            let running = running_steps(&state, now);
326            Ok(accounts
327                .into_iter()
328                .map(|account| {
329                    let windows = windows_of(&state, &[account.id]);
330                    let running_steps = running.get(&account.id).copied().unwrap_or(0);
331                    ProviderAccountCandidate {
332                        account,
333                        windows,
334                        running_steps,
335                    }
336                })
337                .collect())
338        })
339    }
340}