1use crate::error::{ProviderError, Result};
4use crate::runtime::op_context::OpContext;
5use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
6use std::sync::{Arc, Mutex};
7use std::time::{Duration, Instant};
8
9#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
11pub enum PermitKind {
12 ModelLoad,
13 LocalStt,
14 LocalTts,
15 Remote,
16 Blocking,
17}
18
19#[derive(Debug, Clone)]
21pub struct GovernorConfig {
22 pub max_model_loads: usize,
23 pub max_local_stt: usize,
24 pub max_local_tts: usize,
25 pub max_remote: usize,
26 pub max_blocking: usize,
27 pub max_cpu_threads: usize,
29 pub max_memory_bytes: u64,
31 pub queue_timeout: Duration,
33 pub fail_fast: bool,
35}
36
37impl Default for GovernorConfig {
38 fn default() -> Self {
39 let cpus = std::thread::available_parallelism()
40 .map(|n| n.get())
41 .unwrap_or(4)
42 .clamp(1, 16);
43 Self {
44 max_model_loads: 2,
45 max_local_stt: 2,
46 max_local_tts: 2,
47 max_remote: 4,
48 max_blocking: 4,
49 max_cpu_threads: cpus,
50 max_memory_bytes: 2 * 1024 * 1024 * 1024, queue_timeout: Duration::from_secs(30),
52 fail_fast: false,
53 }
54 }
55}
56
57impl GovernorConfig {
58 pub fn mobile() -> Self {
60 Self {
61 max_model_loads: 1,
62 max_local_stt: 1,
63 max_local_tts: 1,
64 max_remote: 2,
65 max_blocking: 2,
66 max_cpu_threads: 2,
67 max_memory_bytes: 512 * 1024 * 1024,
68 queue_timeout: Duration::from_secs(15),
69 fail_fast: false,
70 }
71 }
72
73 pub fn server() -> Self {
75 let cpus = std::thread::available_parallelism()
76 .map(|n| n.get())
77 .unwrap_or(8);
78 Self {
79 max_model_loads: 4,
80 max_local_stt: cpus.max(4),
81 max_local_tts: 4,
82 max_remote: 16,
83 max_blocking: cpus.max(8),
84 max_cpu_threads: cpus,
85 max_memory_bytes: 8 * 1024 * 1024 * 1024,
86 queue_timeout: Duration::from_secs(60),
87 fail_fast: false,
88 }
89 }
90
91 pub fn validate(&self) -> Result<()> {
93 if self.max_model_loads == 0
94 || self.max_local_stt == 0
95 || self.max_local_tts == 0
96 || self.max_remote == 0
97 || self.max_blocking == 0
98 || self.max_cpu_threads == 0
99 {
100 return Err(crate::error::UserError::InvalidConfig {
101 reason: "governor permit and CPU budgets must be >= 1".into(),
102 }
103 .into());
104 }
105 Ok(())
106 }
107}
108
109struct CounterPool {
110 max: usize,
111 in_use: AtomicUsize,
112}
113
114impl CounterPool {
115 fn new(max: usize) -> Self {
116 Self {
117 max: max.max(1),
118 in_use: AtomicUsize::new(0),
119 }
120 }
121
122 fn try_acquire(&self) -> bool {
123 loop {
124 let cur = self.in_use.load(Ordering::SeqCst);
125 if cur >= self.max {
126 return false;
127 }
128 if self
129 .in_use
130 .compare_exchange(cur, cur + 1, Ordering::SeqCst, Ordering::SeqCst)
131 .is_ok()
132 {
133 return true;
134 }
135 }
136 }
137
138 fn release(&self) {
139 let prev = self.in_use.fetch_sub(1, Ordering::SeqCst);
140 debug_assert!(prev > 0, "permit released more times than acquired");
141 }
142
143 fn in_use(&self) -> usize {
144 self.in_use.load(Ordering::SeqCst)
145 }
146
147 fn max(&self) -> usize {
148 self.max
149 }
150}
151
152pub struct ResourceGovernor {
154 config: GovernorConfig,
155 model_loads: CounterPool,
156 local_stt: CounterPool,
157 local_tts: CounterPool,
158 remote: CounterPool,
159 blocking: CounterPool,
160 cpu_threads_in_use: AtomicUsize,
162 memory_reserved: AtomicU64,
163 acquire_lock: Mutex<()>,
165}
166
167impl Default for ResourceGovernor {
168 fn default() -> Self {
169 Self::new(GovernorConfig::default())
170 }
171}
172
173impl ResourceGovernor {
174 pub fn new(config: GovernorConfig) -> Self {
175 Self {
176 model_loads: CounterPool::new(config.max_model_loads),
177 local_stt: CounterPool::new(config.max_local_stt),
178 local_tts: CounterPool::new(config.max_local_tts),
179 remote: CounterPool::new(config.max_remote),
180 blocking: CounterPool::new(config.max_blocking),
181 cpu_threads_in_use: AtomicUsize::new(0),
182 memory_reserved: AtomicU64::new(0),
183 acquire_lock: Mutex::new(()),
184 config,
185 }
186 }
187
188 pub fn config(&self) -> &GovernorConfig {
189 &self.config
190 }
191
192 pub fn process_global() -> Arc<Self> {
196 use once_cell::sync::Lazy;
197 static G: Lazy<Arc<ResourceGovernor>> = Lazy::new(|| Arc::new(ResourceGovernor::default()));
198 Arc::clone(&G)
199 }
200
201 fn pool(&self, kind: PermitKind) -> &CounterPool {
202 match kind {
203 PermitKind::ModelLoad => &self.model_loads,
204 PermitKind::LocalStt => &self.local_stt,
205 PermitKind::LocalTts => &self.local_tts,
206 PermitKind::Remote => &self.remote,
207 PermitKind::Blocking => &self.blocking,
208 }
209 }
210
211 pub fn acquire(
213 self: &Arc<Self>,
214 kind: PermitKind,
215 ctx: Option<&OpContext>,
216 ) -> Result<ResourcePermit> {
217 let timeout = self.wait_budget(ctx);
218 let deadline = Instant::now() + timeout;
219 loop {
220 if let Some(c) = ctx {
221 c.check()?;
222 }
223 if self.pool(kind).try_acquire() {
224 return Ok(ResourcePermit {
225 governor: Arc::clone(self),
226 kind,
227 cpu_threads: 0,
228 memory: 0,
229 holds_blocking: false,
230 released: false,
231 });
232 }
233 if self.config.fail_fast || Instant::now() >= deadline {
234 return Err(ProviderError::Overload {
235 reason: format!(
236 "{kind:?} permits exhausted ({}/{})",
237 self.pool(kind).in_use(),
238 self.pool(kind).max()
239 ),
240 }
241 .into());
242 }
243 std::thread::sleep(Duration::from_millis(2));
244 }
245 }
246
247 fn wait_budget(&self, ctx: Option<&OpContext>) -> Duration {
248 if self.config.fail_fast {
249 return Duration::ZERO;
250 }
251 ctx.and_then(|c| c.remaining())
252 .unwrap_or(self.config.queue_timeout)
253 .min(self.config.queue_timeout)
254 }
255
256 pub fn acquire_stt(
262 self: &Arc<Self>,
263 cpu_threads: usize,
264 memory: u64,
265 ctx: Option<&OpContext>,
266 ) -> Result<ResourcePermit> {
267 let timeout = self.wait_budget(ctx);
268 let deadline = Instant::now() + timeout;
269 let want = cpu_threads.max(1);
270
271 loop {
272 if let Some(c) = ctx {
273 c.check()?;
274 }
275
276 {
277 let _order = self.acquire_lock.lock().unwrap_or_else(|e| e.into_inner());
278
279 if memory == 0 || self.try_reserve_memory(memory).is_ok() {
281 if self.local_stt.try_acquire() {
282 if self.blocking.try_acquire() {
283 if self.try_reserve_cpu(want) {
284 return Ok(ResourcePermit {
285 governor: Arc::clone(self),
286 kind: PermitKind::LocalStt,
287 cpu_threads: want,
288 memory,
289 holds_blocking: true,
290 released: false,
291 });
292 }
293 self.blocking.release();
294 }
295 self.local_stt.release();
296 }
297 self.release_memory(memory);
298 }
299 }
300
301 if self.config.fail_fast || Instant::now() >= deadline {
302 return Err(ProviderError::Overload {
303 reason: format!(
304 "STT resources unavailable (cpu want {want}, budget {})",
305 self.config.max_cpu_threads
306 ),
307 }
308 .into());
309 }
310 std::thread::sleep(Duration::from_millis(2));
311 }
312 }
313
314 pub fn acquire_tts(
316 self: &Arc<Self>,
317 memory: u64,
318 ctx: Option<&OpContext>,
319 ) -> Result<ResourcePermit> {
320 let timeout = self.wait_budget(ctx);
321 let deadline = Instant::now() + timeout;
322
323 loop {
324 if let Some(c) = ctx {
325 c.check()?;
326 }
327
328 {
329 let _order = self.acquire_lock.lock().unwrap_or_else(|e| e.into_inner());
330
331 if memory == 0 || self.try_reserve_memory(memory).is_ok() {
332 if self.local_tts.try_acquire() {
333 if self.blocking.try_acquire() {
334 return Ok(ResourcePermit {
335 governor: Arc::clone(self),
336 kind: PermitKind::LocalTts,
337 cpu_threads: 0,
338 memory,
339 holds_blocking: true,
340 released: false,
341 });
342 }
343 self.local_tts.release();
344 }
345 self.release_memory(memory);
346 }
347 }
348
349 if self.config.fail_fast || Instant::now() >= deadline {
350 return Err(ProviderError::Overload {
351 reason: format!(
352 "LocalTts permits exhausted ({}/{})",
353 self.local_tts.in_use(),
354 self.local_tts.max()
355 ),
356 }
357 .into());
358 }
359 std::thread::sleep(Duration::from_millis(2));
360 }
361 }
362
363 fn try_reserve_memory(&self, bytes: u64) -> Result<()> {
364 if bytes == 0 {
365 return Ok(());
366 }
367 loop {
368 let cur = self.memory_reserved.load(Ordering::SeqCst);
369 let new = cur.saturating_add(bytes);
370 if new > self.config.max_memory_bytes {
371 return Err(ProviderError::Overload {
372 reason: format!(
373 "memory reservation {bytes} would exceed budget {} (in use {cur})",
374 self.config.max_memory_bytes
375 ),
376 }
377 .into());
378 }
379 if self
380 .memory_reserved
381 .compare_exchange(cur, new, Ordering::SeqCst, Ordering::SeqCst)
382 .is_ok()
383 {
384 return Ok(());
385 }
386 }
387 }
388
389 fn release_memory(&self, bytes: u64) {
390 if bytes > 0 {
391 self.memory_reserved.fetch_sub(bytes, Ordering::SeqCst);
392 }
393 }
394
395 fn try_reserve_cpu(&self, n: usize) -> bool {
396 loop {
397 let cur = self.cpu_threads_in_use.load(Ordering::SeqCst);
398 if cur + n > self.config.max_cpu_threads {
399 return false;
400 }
401 if self
402 .cpu_threads_in_use
403 .compare_exchange(cur, cur + n, Ordering::SeqCst, Ordering::SeqCst)
404 .is_ok()
405 {
406 return true;
407 }
408 }
409 }
410
411 fn release_cpu(&self, n: usize) {
412 if n > 0 {
413 self.cpu_threads_in_use.fetch_sub(n, Ordering::SeqCst);
414 }
415 }
416
417 pub fn recommend_stt_threads(&self) -> usize {
419 let used = self.cpu_threads_in_use.load(Ordering::SeqCst);
420 let rem = self.config.max_cpu_threads.saturating_sub(used).max(1);
421 let fair = (self.config.max_cpu_threads / self.config.max_local_stt.max(1)).max(1);
423 fair.min(rem).clamp(1, 8)
424 }
425
426 pub fn stats(&self) -> GovernorStats {
427 GovernorStats {
428 model_loads: self.model_loads.in_use(),
429 local_stt: self.local_stt.in_use(),
430 local_tts: self.local_tts.in_use(),
431 remote: self.remote.in_use(),
432 blocking: self.blocking.in_use(),
433 cpu_threads: self.cpu_threads_in_use.load(Ordering::SeqCst),
434 memory_reserved: self.memory_reserved.load(Ordering::SeqCst),
435 max_cpu_threads: self.config.max_cpu_threads,
436 max_memory_bytes: self.config.max_memory_bytes,
437 }
438 }
439}
440
441#[derive(Debug, Clone)]
443pub struct GovernorStats {
444 pub model_loads: usize,
445 pub local_stt: usize,
446 pub local_tts: usize,
447 pub remote: usize,
448 pub blocking: usize,
449 pub cpu_threads: usize,
450 pub memory_reserved: u64,
451 pub max_cpu_threads: usize,
452 pub max_memory_bytes: u64,
453}
454
455pub struct ResourcePermit {
459 governor: Arc<ResourceGovernor>,
460 kind: PermitKind,
461 cpu_threads: usize,
462 memory: u64,
463 holds_blocking: bool,
464 released: bool,
465}
466
467impl std::fmt::Debug for ResourcePermit {
468 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
469 f.debug_struct("ResourcePermit")
470 .field("kind", &self.kind)
471 .field("cpu_threads", &self.cpu_threads)
472 .field("memory", &self.memory)
473 .field("holds_blocking", &self.holds_blocking)
474 .finish()
475 }
476}
477
478impl ResourcePermit {
479 pub fn kind(&self) -> PermitKind {
480 self.kind
481 }
482
483 pub fn cpu_threads(&self) -> usize {
484 self.cpu_threads
485 }
486
487 fn release_inner(&mut self) {
488 if self.released {
489 return;
490 }
491 self.released = true;
492 if self.cpu_threads > 0 {
493 self.governor.release_cpu(self.cpu_threads);
494 self.cpu_threads = 0;
495 }
496 if self.holds_blocking {
497 self.governor.blocking.release();
498 self.holds_blocking = false;
499 }
500 if self.memory > 0 {
501 self.governor.release_memory(self.memory);
502 self.memory = 0;
503 }
504 self.governor.pool(self.kind).release();
505 }
506}
507
508impl Drop for ResourcePermit {
509 fn drop(&mut self) {
510 self.release_inner();
511 }
512}
513
514#[cfg(test)]
515mod tests {
516 use super::*;
517
518 #[test]
519 fn permits_cap() {
520 let g = Arc::new(ResourceGovernor::new(GovernorConfig {
521 max_local_stt: 1,
522 fail_fast: true,
523 ..GovernorConfig::default()
524 }));
525 let a = g.acquire(PermitKind::LocalStt, None).unwrap();
526 let err = g.acquire(PermitKind::LocalStt, None).unwrap_err();
527 assert!(matches!(
528 err,
529 crate::error::TranscriptionError::Provider(ProviderError::Overload { .. })
530 ));
531 drop(a);
532 let _b = g.acquire(PermitKind::LocalStt, None).unwrap();
533 }
534
535 #[test]
536 fn memory_budget() {
537 let g = Arc::new(ResourceGovernor::new(GovernorConfig {
538 max_memory_bytes: 1000,
539 fail_fast: true,
540 ..GovernorConfig::default()
541 }));
542 assert!(g.try_reserve_memory(600).is_ok());
543 assert!(g.try_reserve_memory(600).is_err());
544 g.release_memory(600);
545 assert!(g.try_reserve_memory(600).is_ok());
546 }
547
548 #[test]
549 fn cpu_budget() {
550 let g = Arc::new(ResourceGovernor::new(GovernorConfig {
551 max_cpu_threads: 4,
552 max_local_stt: 4,
553 max_blocking: 4,
554 fail_fast: true,
555 ..GovernorConfig::default()
556 }));
557 let p = g.acquire_stt(3, 0, None).unwrap();
558 assert_eq!(p.cpu_threads(), 3);
559 let err = g.acquire_stt(3, 0, None).unwrap_err();
560 assert!(err.to_string().contains("CPU") || err.to_string().contains("overload"));
561 drop(p);
562 let _p2 = g.acquire_stt(2, 0, None).unwrap();
563 }
564
565 #[test]
566 fn tts_releases_blocking() {
567 let g = Arc::new(ResourceGovernor::new(GovernorConfig {
568 max_local_tts: 1,
569 max_blocking: 1,
570 fail_fast: true,
571 ..GovernorConfig::default()
572 }));
573 let p = g.acquire_tts(0, None).unwrap();
574 let err = g.acquire_tts(0, None).unwrap_err();
575 assert!(matches!(
576 err,
577 crate::error::TranscriptionError::Provider(ProviderError::Overload { .. })
578 ));
579 drop(p);
580 let _p2 = g.acquire_tts(0, None).unwrap();
581 }
582
583 #[test]
584 fn cancel_aborts_queue_wait() {
585 let g = Arc::new(ResourceGovernor::new(GovernorConfig {
586 max_local_stt: 1,
587 fail_fast: false,
588 queue_timeout: Duration::from_secs(5),
589 ..GovernorConfig::default()
590 }));
591 let _hold = g.acquire(PermitKind::LocalStt, None).unwrap();
592 let ctx = OpContext::new();
593 ctx.cancel.cancel();
594 let err = g.acquire(PermitKind::LocalStt, Some(&ctx)).unwrap_err();
595 assert!(matches!(
596 err,
597 crate::error::TranscriptionError::Provider(ProviderError::Cancelled)
598 ));
599 }
600
601 #[test]
602 fn config_validate_rejects_zero() {
603 let c = GovernorConfig {
604 max_cpu_threads: 0,
605 ..GovernorConfig::default()
606 };
607 assert!(c.validate().is_err());
608 }
609}