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, ProviderKind, RunStatus, 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
63/// Make every run sleeping on capacity for `kind` due now, so the run waker
64/// resumes it on its next tick: a new, re-enabled or renewed account may have
65/// the capacity it waits for.
66fn wake_capacity_sleepers(state: &mut State, kind: &ProviderKind, now: DateTime<Utc>) {
67    for run in state.runs.values_mut() {
68        if run.status.state == RunStatus::Sleeping && run.capacity_wait_kind.as_ref() == Some(kind)
69        {
70            run.scheduled_at = Some(now);
71            run.updated_at = now;
72        }
73    }
74}
75
76impl ProviderAccountStore for InMemoryStore {
77    fn create_provider_account(&self, req: NewProviderAccount) -> StoreFuture<'_, ProviderAccount> {
78        Box::pin(async move {
79            let mut state = self.state.write().await;
80            if state.provider_accounts.values().any(|a| a.name == req.name) {
81                return Err(StoreError::DuplicateProviderAccount(req.name));
82            }
83            let now = Utc::now();
84            let account = ProviderAccount {
85                id: req.id,
86                name: req.name,
87                display_name: req.display_name,
88                kind: req.kind,
89                secret_key: req.secret_key,
90                enabled: req.enabled,
91                priority: req.priority,
92                tags: req.tags,
93                max_concurrency: req.max_concurrency,
94                alert_threshold: req.alert_threshold,
95                expires_at: req.expires_at,
96                plan: req.plan,
97                auth_failed_at: None,
98                created_by: req.created_by,
99                created_at: now,
100                updated_at: now,
101            };
102            state.provider_accounts.insert(account.id, account.clone());
103            if account.enabled {
104                wake_capacity_sleepers(&mut state, &ProviderKind::new(account.kind.as_str()), now);
105            }
106            Ok(account)
107        })
108    }
109
110    fn get_provider_account(&self, id: Uuid) -> StoreFuture<'_, Option<ProviderAccount>> {
111        Box::pin(async move {
112            let state = self.state.read().await;
113            Ok(state.provider_accounts.get(&id).cloned())
114        })
115    }
116
117    fn list_provider_accounts_by_ids(
118        &self,
119        ids: Vec<Uuid>,
120    ) -> StoreFuture<'_, Vec<ProviderAccount>> {
121        Box::pin(async move {
122            let wanted: HashSet<Uuid> = ids.into_iter().collect();
123            let state = self.state.read().await;
124            Ok(state
125                .provider_accounts
126                .values()
127                .filter(|a| wanted.contains(&a.id))
128                .cloned()
129                .collect())
130        })
131    }
132
133    fn find_provider_account_by_name(
134        &self,
135        name: &str,
136    ) -> StoreFuture<'_, Option<ProviderAccount>> {
137        let name = name.to_string();
138        Box::pin(async move {
139            let state = self.state.read().await;
140            Ok(state
141                .provider_accounts
142                .values()
143                .find(|a| a.name == name)
144                .cloned())
145        })
146    }
147
148    fn list_provider_accounts(
149        &self,
150        kind: Option<String>,
151        page: u32,
152        per_page: u32,
153    ) -> StoreFuture<'_, Page<ProviderAccount>> {
154        Box::pin(async move {
155            let state = self.state.read().await;
156            let mut all: Vec<ProviderAccount> = state
157                .provider_accounts
158                .values()
159                .filter(|a| kind.as_ref().is_none_or(|k| &a.kind == k))
160                .cloned()
161                .collect();
162            all.sort_by(by_priority_then_name);
163            let total = all.len() as u64;
164            let start = (page.saturating_sub(1) as usize) * (per_page as usize);
165            let items = all
166                .into_iter()
167                .skip(start)
168                .take(per_page as usize)
169                .collect();
170            Ok(Page {
171                items,
172                total,
173                page,
174                per_page,
175            })
176        })
177    }
178
179    fn update_provider_account(
180        &self,
181        id: Uuid,
182        update: ProviderAccountUpdate,
183    ) -> StoreFuture<'_, ProviderAccount> {
184        Box::pin(async move {
185            let mut state = self.state.write().await;
186            let account = state
187                .provider_accounts
188                .get_mut(&id)
189                .ok_or(StoreError::ProviderAccountNotFound(id))?;
190            if let Some(display_name) = update.display_name {
191                account.display_name = display_name;
192            }
193            if let Some(enabled) = update.enabled {
194                account.enabled = enabled;
195            }
196            if let Some(priority) = update.priority {
197                account.priority = priority;
198            }
199            if let Some(tags) = update.tags {
200                account.tags = tags;
201            }
202            if let Some(max_concurrency) = update.max_concurrency {
203                account.max_concurrency = max_concurrency;
204            }
205            if let Some(alert_threshold) = update.alert_threshold {
206                account.alert_threshold = alert_threshold;
207            }
208            if let Some(expires_at) = update.expires_at {
209                account.expires_at = expires_at;
210            }
211            if let Some(plan) = update.plan {
212                account.plan = plan;
213            }
214            if let Some(auth_failed_at) = update.auth_failed_at {
215                account.auth_failed_at = auth_failed_at;
216            }
217            let now = Utc::now();
218            account.updated_at = now;
219            let account = account.clone();
220            let renewed = update.enabled == Some(true) || update.auth_failed_at == Some(None);
221            if renewed && account.enabled {
222                wake_capacity_sleepers(&mut state, &ProviderKind::new(account.kind.as_str()), now);
223            }
224            Ok(account)
225        })
226    }
227
228    fn delete_provider_account(&self, id: Uuid) -> StoreFuture<'_, bool> {
229        Box::pin(async move {
230            let mut state = self.state.write().await;
231            if state.provider_accounts.remove(&id).is_none() {
232                return Ok(false);
233            }
234            state
235                .provider_account_windows
236                .retain(|(account_id, _, _), _| *account_id != id);
237            state.provider_account_usage.retain(|u| u.account_id != id);
238            for step in state.steps.values_mut() {
239                if step.account_id == Some(id) {
240                    step.account_id = None;
241                }
242            }
243            Ok(true)
244        })
245    }
246
247    fn list_provider_account_windows(
248        &self,
249        ids: Vec<Uuid>,
250    ) -> StoreFuture<'_, Vec<ProviderAccountWindow>> {
251        Box::pin(async move {
252            let state = self.state.read().await;
253            Ok(windows_of(&state, &ids))
254        })
255    }
256
257    fn list_provider_account_usage(
258        &self,
259        id: Uuid,
260        since: DateTime<Utc>,
261    ) -> StoreFuture<'_, Vec<ProviderAccountUsagePoint>> {
262        Box::pin(async move {
263            let state = self.state.read().await;
264            let mut points: Vec<ProviderAccountUsagePoint> = state
265                .provider_account_usage
266                .iter()
267                .filter(|u| u.account_id == id && u.observed_at >= since)
268                .cloned()
269                .collect();
270            points.sort_by_key(|u| u.observed_at);
271            Ok(points)
272        })
273    }
274
275    fn record_provider_account_observation(
276        &self,
277        id: Uuid,
278        observation: NewProviderAccountObservation,
279    ) -> StoreFuture<'_, Vec<ProviderAccountWindow>> {
280        Box::pin(async move {
281            let mut state = self.state.write().await;
282            if !state.provider_accounts.contains_key(&id) {
283                return Err(StoreError::ProviderAccountNotFound(id));
284            }
285            let recorded = !observation.windows.is_empty();
286            for window in observation.windows {
287                let key = (
288                    id,
289                    window.window.clone(),
290                    window.model_scope.clone().unwrap_or_default(),
291                );
292                let newer_stored = state
293                    .provider_account_windows
294                    .get(&key)
295                    .is_some_and(|stored| stored.observed_at > window.observed_at);
296                if !newer_stored {
297                    state.provider_account_windows.insert(
298                        key,
299                        ProviderAccountWindow {
300                            account_id: id,
301                            window: window.window.clone(),
302                            utilization: window.utilization,
303                            resets_at: window.resets_at,
304                            status: window.status,
305                            model_scope: window.model_scope.clone(),
306                            observed_at: window.observed_at,
307                        },
308                    );
309                }
310                state
311                    .provider_account_usage
312                    .push(ProviderAccountUsagePoint {
313                        id: Uuid::now_v7(),
314                        account_id: id,
315                        window: window.window,
316                        utilization: window.utilization,
317                        resets_at: window.resets_at,
318                        status: window.status,
319                        model_scope: window.model_scope,
320                        observed_at: window.observed_at,
321                    });
322            }
323            let now = Utc::now();
324            if let Some(account) = state.provider_accounts.get_mut(&id) {
325                if observation.auth_failed {
326                    account.auth_failed_at = Some(now);
327                    account.updated_at = now;
328                } else if recorded {
329                    account.auth_failed_at = None;
330                }
331            }
332            Ok(windows_of(&state, &[id]))
333        })
334    }
335
336    fn purge_provider_account_usage(&self, before: DateTime<Utc>) -> StoreFuture<'_, u64> {
337        Box::pin(async move {
338            let mut state = self.state.write().await;
339            let len_before = state.provider_account_usage.len();
340            state
341                .provider_account_usage
342                .retain(|u| u.observed_at >= before);
343            Ok((len_before - state.provider_account_usage.len()) as u64)
344        })
345    }
346
347    fn list_provider_account_candidates(
348        &self,
349        kind: String,
350    ) -> StoreFuture<'_, Vec<ProviderAccountCandidate>> {
351        Box::pin(async move {
352            let state = self.state.read().await;
353            let now = Utc::now();
354            let mut accounts: Vec<ProviderAccount> = state
355                .provider_accounts
356                .values()
357                .filter(|a| {
358                    a.kind == kind && a.enabled && a.auth_failed_at.is_none() && a.expires_at > now
359                })
360                .cloned()
361                .collect();
362            accounts.sort_by(by_priority_then_name);
363            let running = running_steps(&state, now);
364            Ok(accounts
365                .into_iter()
366                .map(|account| {
367                    let windows = windows_of(&state, &[account.id]);
368                    let running_steps = running.get(&account.id).copied().unwrap_or(0);
369                    ProviderAccountCandidate {
370                        account,
371                        windows,
372                        running_steps,
373                    }
374                })
375                .collect())
376        })
377    }
378}