1use crate::error::{ProviderError, Result};
4use crate::runtime::singleflight::LoadKey;
5use std::collections::HashMap;
6use std::sync::atomic::{AtomicU64, Ordering};
7use std::sync::{Arc, Mutex};
8use std::time::{Duration, Instant};
9
10#[derive(Debug, Clone, Copy)]
12pub struct ResidencyWeight {
13 pub bytes: u64,
14}
15
16#[derive(Debug, Clone)]
18pub struct RegistryConfig {
19 pub max_resident_bytes: u64,
20 pub max_entries: usize,
21 pub idle_ttl: Option<Duration>,
22}
23
24impl Default for RegistryConfig {
25 fn default() -> Self {
26 Self {
27 max_resident_bytes: 3 * 1024 * 1024 * 1024, max_entries: 8,
29 idle_ttl: Some(Duration::from_secs(30 * 60)),
30 }
31 }
32}
33
34#[derive(Debug)]
36pub struct RegistryEntry<T> {
37 pub key: LoadKey,
38 pub value: Arc<T>,
39 pub weight: ResidencyWeight,
40 last_used_ms: AtomicU64,
42 pub active_refs: AtomicU64,
43 pub pinned: bool,
44}
45
46impl<T> RegistryEntry<T> {
47 fn touch(&self) {
48 self.last_used_ms.store(now_ms(), Ordering::Relaxed);
49 }
50
51 fn last_used_instant(&self) -> Instant {
52 let ms = self.last_used_ms.load(Ordering::Relaxed);
53 Instant::now()
54 .checked_sub(Duration::from_millis(now_ms().saturating_sub(ms)))
55 .unwrap_or_else(Instant::now)
56 }
57
58 pub fn last_used_ms(&self) -> u64 {
59 self.last_used_ms.load(Ordering::Relaxed)
60 }
61}
62
63pub struct ModelRegistry<T> {
65 config: RegistryConfig,
66 inner: Mutex<HashMap<LoadKey, Arc<RegistryEntry<T>>>>,
67 total_weight: AtomicU64,
68}
69
70impl<T> ModelRegistry<T> {
71 pub fn new(config: RegistryConfig) -> Self {
72 Self {
73 config,
74 inner: Mutex::new(HashMap::new()),
75 total_weight: AtomicU64::new(0),
76 }
77 }
78
79 pub fn config(&self) -> &RegistryConfig {
80 &self.config
81 }
82
83 pub fn total_weight(&self) -> u64 {
84 self.total_weight.load(Ordering::SeqCst)
85 }
86
87 pub fn len(&self) -> usize {
88 self.inner.lock().map(|g| g.len()).unwrap_or(0)
89 }
90
91 pub fn is_empty(&self) -> bool {
92 self.len() == 0
93 }
94
95 pub fn insert(
100 &self,
101 key: LoadKey,
102 value: Arc<T>,
103 weight: ResidencyWeight,
104 ) -> Result<Arc<RegistryEntry<T>>> {
105 let mut guard = self.inner.lock().unwrap_or_else(|e| e.into_inner());
106 if let Some(existing) = guard.get(&key) {
107 existing.touch();
108 return Ok(Arc::clone(existing));
109 }
110 self.evict_locked(&mut guard, weight.bytes, 1)?;
111 let entry = Arc::new(RegistryEntry {
112 key: key.clone(),
113 value,
114 weight,
115 last_used_ms: AtomicU64::new(now_ms()),
116 active_refs: AtomicU64::new(0),
117 pinned: false,
118 });
119 guard.insert(key, Arc::clone(&entry));
120 self.total_weight.fetch_add(weight.bytes, Ordering::SeqCst);
121 Ok(entry)
122 }
123
124 pub fn insert_and_pin(
127 &self,
128 key: LoadKey,
129 value: Arc<T>,
130 weight: ResidencyWeight,
131 ) -> Result<RegistryPin<T>> {
132 let mut guard = self.inner.lock().unwrap_or_else(|e| e.into_inner());
133 if let Some(existing) = guard.get(&key) {
134 existing.touch();
135 existing.active_refs.fetch_add(1, Ordering::SeqCst);
136 return Ok(RegistryPin {
137 entry: Arc::clone(existing),
138 });
139 }
140 self.evict_locked(&mut guard, weight.bytes, 1)?;
141 let entry = Arc::new(RegistryEntry {
142 key: key.clone(),
143 value,
144 weight,
145 last_used_ms: AtomicU64::new(now_ms()),
146 active_refs: AtomicU64::new(1), pinned: false,
148 });
149 guard.insert(key, Arc::clone(&entry));
150 self.total_weight.fetch_add(weight.bytes, Ordering::SeqCst);
151 Ok(RegistryPin { entry })
152 }
153
154 pub fn get(&self, key: &LoadKey) -> Option<Arc<RegistryEntry<T>>> {
155 let guard = self.inner.lock().unwrap_or_else(|e| e.into_inner());
156 guard.get(key).map(|e| {
157 e.touch();
158 Arc::clone(e)
159 })
160 }
161
162 pub fn pin(&self, key: &LoadKey) -> Option<RegistryPin<T>> {
167 self.get_and_pin(key)
168 }
169
170 pub fn get_and_pin(&self, key: &LoadKey) -> Option<RegistryPin<T>> {
175 let guard = self.inner.lock().unwrap_or_else(|e| e.into_inner());
176 let entry = guard.get(key)?;
177 entry.touch();
178 entry.active_refs.fetch_add(1, Ordering::SeqCst);
179 Some(RegistryPin {
180 entry: Arc::clone(entry),
181 })
182 }
183
184 pub fn clear_idle(&self) -> (usize, usize) {
192 let mut guard = self.inner.lock().unwrap_or_else(|e| e.into_inner());
193 let before = guard.len();
194 let mut removed_weight = 0u64;
195 guard.retain(|_, e| {
196 let active = e.active_refs.load(Ordering::SeqCst) > 0 || e.pinned;
197 if !active {
198 removed_weight = removed_weight.saturating_add(e.weight.bytes);
199 }
200 active
201 });
202 if removed_weight > 0 {
203 self.total_weight
204 .fetch_sub(removed_weight, Ordering::SeqCst);
205 }
206 let retained = guard.len();
207 let removed = before.saturating_sub(retained);
208 (removed, retained)
209 }
210
211 pub fn clear(&self) {
213 let _ = self.clear_idle();
214 }
215
216 pub fn try_unload(&self, key: &LoadKey) -> bool {
218 let mut guard = self.inner.lock().unwrap_or_else(|e| e.into_inner());
219 if let Some(e) = guard.get(key) {
220 if e.active_refs.load(Ordering::SeqCst) > 0 || e.pinned {
221 return false;
222 }
223 let w = e.weight.bytes;
224 guard.remove(key);
225 self.total_weight.fetch_sub(w, Ordering::SeqCst);
226 return true;
227 }
228 false
229 }
230
231 fn evict_locked(
232 &self,
233 guard: &mut HashMap<LoadKey, Arc<RegistryEntry<T>>>,
234 need_bytes: u64,
235 need_slots: usize,
236 ) -> Result<()> {
237 if let Some(ttl) = self.config.idle_ttl {
239 let now = Instant::now();
240 let expired: Vec<LoadKey> = guard
241 .iter()
242 .filter(|(_, e)| {
243 e.active_refs.load(Ordering::SeqCst) == 0
244 && !e.pinned
245 && now.duration_since(e.last_used_instant()) > ttl
246 })
247 .map(|(k, _)| k.clone())
248 .collect();
249 for k in expired {
250 if let Some(e) = guard.remove(&k) {
251 self.total_weight
252 .fetch_sub(e.weight.bytes, Ordering::SeqCst);
253 }
254 }
255 }
256
257 while guard.len() + need_slots > self.config.max_entries
258 || self.total_weight.load(Ordering::SeqCst) + need_bytes
259 > self.config.max_resident_bytes
260 {
261 let victim = guard
263 .iter()
264 .filter(|(_, e)| e.active_refs.load(Ordering::SeqCst) == 0 && !e.pinned)
265 .min_by_key(|(_, e)| e.last_used_ms.load(Ordering::Relaxed))
266 .map(|(k, _)| k.clone());
267 let Some(v) = victim else {
268 if need_bytes > self.config.max_resident_bytes {
269 return Err(ProviderError::Overload {
270 reason: format!(
271 "model weight {need_bytes} exceeds residency budget {}",
272 self.config.max_resident_bytes
273 ),
274 }
275 .into());
276 }
277 return Err(ProviderError::Overload {
278 reason: "model residency budget exhausted (no idle entries to evict)".into(),
279 }
280 .into());
281 };
282 if let Some(e) = guard.remove(&v) {
283 self.total_weight
284 .fetch_sub(e.weight.bytes, Ordering::SeqCst);
285 }
286 }
287 Ok(())
288 }
289
290 pub fn snapshot(&self) -> Vec<RegistrySnapshot> {
291 let guard = self.inner.lock().unwrap_or_else(|e| e.into_inner());
292 let now = now_ms();
293 guard
294 .values()
295 .map(|e| {
296 let last = e.last_used_ms.load(Ordering::Relaxed);
297 RegistrySnapshot {
298 id: e.key.id.clone(),
299 kind: e.key.kind,
300 weight_bytes: e.weight.bytes,
301 active_refs: e.active_refs.load(Ordering::SeqCst),
302 idle_secs: now.saturating_sub(last) / 1000,
303 pinned: e.pinned,
304 }
305 })
306 .collect()
307 }
308}
309
310fn now_ms() -> u64 {
311 use std::time::SystemTime;
312 SystemTime::now()
313 .duration_since(SystemTime::UNIX_EPOCH)
314 .map(|d| d.as_millis() as u64)
315 .unwrap_or(0)
316}
317
318#[derive(Debug, Clone)]
320pub struct RegistrySnapshot {
321 pub id: String,
322 pub kind: &'static str,
323 pub weight_bytes: u64,
324 pub active_refs: u64,
325 pub idle_secs: u64,
326 pub pinned: bool,
327}
328
329pub struct RegistryPin<T> {
331 entry: Arc<RegistryEntry<T>>,
332}
333
334impl<T> RegistryPin<T> {
335 pub fn value(&self) -> &Arc<T> {
336 &self.entry.value
337 }
338
339 pub fn entry(&self) -> &RegistryEntry<T> {
340 &self.entry
341 }
342}
343
344impl<T> Drop for RegistryPin<T> {
345 fn drop(&mut self) {
346 self.entry.active_refs.fetch_sub(1, Ordering::SeqCst);
347 }
348}
349
350#[cfg(test)]
351mod tests {
352 use super::*;
353
354 fn key(id: &str) -> LoadKey {
355 LoadKey::stt(id, format!("/tmp/{id}"))
356 }
357
358 #[test]
359 fn evicts_idle_when_over_entry_cap() {
360 let reg = ModelRegistry::new(RegistryConfig {
361 max_entries: 2,
362 max_resident_bytes: 10_000,
363 idle_ttl: None,
364 });
365 reg.insert(key("a"), Arc::new(1u32), ResidencyWeight { bytes: 100 })
366 .unwrap();
367 reg.insert(key("b"), Arc::new(2u32), ResidencyWeight { bytes: 100 })
368 .unwrap();
369 assert_eq!(reg.len(), 2);
370 reg.insert(key("c"), Arc::new(3u32), ResidencyWeight { bytes: 100 })
371 .unwrap();
372 assert_eq!(reg.len(), 2);
373 assert!(reg.get(&key("c")).is_some());
374 }
375
376 #[test]
377 fn active_not_evicted() {
378 let reg = ModelRegistry::new(RegistryConfig {
379 max_entries: 1,
380 max_resident_bytes: 10_000,
381 idle_ttl: None,
382 });
383 reg.insert(key("a"), Arc::new(1u32), ResidencyWeight { bytes: 100 })
384 .unwrap();
385 let pin = reg.pin(&key("a")).unwrap();
386 let err = reg
387 .insert(key("b"), Arc::new(2u32), ResidencyWeight { bytes: 100 })
388 .unwrap_err();
389 assert!(err.to_string().contains("residency") || err.to_string().contains("overload"));
390 drop(pin);
391 reg.insert(key("b"), Arc::new(2u32), ResidencyWeight { bytes: 100 })
392 .unwrap();
393 assert!(reg.get(&key("b")).is_some());
394 }
395
396 #[test]
397 fn lru_prefers_older() {
398 let reg = ModelRegistry::new(RegistryConfig {
399 max_entries: 2,
400 max_resident_bytes: 10_000,
401 idle_ttl: None,
402 });
403 reg.insert(key("a"), Arc::new(1u32), ResidencyWeight { bytes: 100 })
404 .unwrap();
405 std::thread::sleep(Duration::from_millis(5));
406 reg.insert(key("b"), Arc::new(2u32), ResidencyWeight { bytes: 100 })
407 .unwrap();
408 let _ = reg.get(&key("b"));
410 reg.insert(key("c"), Arc::new(3u32), ResidencyWeight { bytes: 100 })
411 .unwrap();
412 assert!(reg.get(&key("a")).is_none());
413 assert!(reg.get(&key("b")).is_some());
414 assert!(reg.get(&key("c")).is_some());
415 }
416
417 #[test]
418 fn get_and_pin_prevents_eviction() {
419 let reg = ModelRegistry::new(RegistryConfig {
420 max_entries: 1,
421 max_resident_bytes: 10_000,
422 idle_ttl: None,
423 });
424 let pin = reg
425 .insert_and_pin(key("a"), Arc::new(1u32), ResidencyWeight { bytes: 100 })
426 .unwrap();
427 assert_eq!(**pin.value(), 1);
428 let err = reg
429 .insert(key("b"), Arc::new(2u32), ResidencyWeight { bytes: 100 })
430 .unwrap_err();
431 assert!(err.to_string().contains("residency") || err.to_string().contains("overload"));
432 drop(pin);
433 reg.insert(key("b"), Arc::new(2u32), ResidencyWeight { bytes: 100 })
434 .unwrap();
435 assert!(reg.get(&key("a")).is_none());
436 }
437
438 #[test]
439 fn insert_and_pin_existing_increments_active() {
440 let reg = ModelRegistry::new(RegistryConfig::default());
441 reg.insert(key("a"), Arc::new(7u32), ResidencyWeight { bytes: 100 })
442 .unwrap();
443 let p1 = reg
444 .insert_and_pin(key("a"), Arc::new(999u32), ResidencyWeight { bytes: 100 })
445 .unwrap();
446 assert_eq!(**p1.value(), 7);
448 let p2 = reg.get_and_pin(&key("a")).unwrap();
449 assert_eq!(p1.entry().active_refs.load(Ordering::SeqCst), 2);
450 drop(p1);
451 drop(p2);
452 assert_eq!(
453 reg.get(&key("a"))
454 .unwrap()
455 .active_refs
456 .load(Ordering::SeqCst),
457 0
458 );
459 }
460
461 #[test]
462 fn clear_idle_retains_active_and_blocks_reload_multiplication() {
463 let reg = ModelRegistry::new(RegistryConfig {
464 max_entries: 8,
465 max_resident_bytes: 10_000,
466 idle_ttl: None,
467 });
468 reg.insert(key("idle"), Arc::new(1u32), ResidencyWeight { bytes: 100 })
470 .unwrap();
471 let lease = reg
473 .insert_and_pin(
474 key("active"),
475 Arc::new(42u32),
476 ResidencyWeight { bytes: 200 },
477 )
478 .unwrap();
479 assert_eq!(reg.len(), 2);
480 assert_eq!(reg.total_weight(), 300);
481
482 let (removed, retained) = reg.clear_idle();
483 assert_eq!(removed, 1);
484 assert_eq!(retained, 1);
485 assert_eq!(reg.len(), 1);
486 assert_eq!(reg.total_weight(), 200);
487 assert!(reg.get(&key("idle")).is_none());
488 let again = reg.get_and_pin(&key("active")).unwrap();
490 assert_eq!(**again.value(), 42);
491 assert_eq!(lease.entry().active_refs.load(Ordering::SeqCst), 2);
492 drop(again);
493 drop(lease);
494 let (removed, retained) = reg.clear_idle();
496 assert_eq!(removed, 1);
497 assert_eq!(retained, 0);
498 assert!(reg.is_empty());
499 assert_eq!(reg.total_weight(), 0);
500 }
501}