1use 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
41fn 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
63fn 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}