1use std::sync::Arc;
17use std::time::Duration;
18
19use bevy_ecs::entity::Entity;
20use leviath_providers::{InferenceRequest, InferenceResponse, Provider, ProviderError};
21use tokio::sync::Notify;
22use tokio::sync::mpsc::UnboundedSender;
23
24use crate::inference_pool::InferencePermit;
25
26#[derive(Debug, Clone, Copy)]
33pub struct RetryPolicy {
34 pub max_attempts: u32,
36 pub base_delay: Duration,
39 pub job_timeout: Duration,
52}
53
54impl Default for RetryPolicy {
55 fn default() -> Self {
56 Self {
57 max_attempts: 4,
58 base_delay: Duration::from_secs(1),
59 job_timeout: Duration::from_secs(leviath_providers::DEFAULT_INFERENCE_TIMEOUT_SECS),
65 }
66 }
67}
68
69pub struct InferenceJob {
71 pub entity: Entity,
73 pub provider: Arc<dyn Provider>,
75 pub request: InferenceRequest,
77 pub permit: InferencePermit,
80 pub exact_token_counting: bool,
86}
87
88fn flatten_request_text(request: &InferenceRequest) -> String {
94 let mut parts: Vec<String> = Vec::new();
95 for block in &request.system {
96 parts.push(block.text.clone());
97 }
98 for msg in &request.messages {
99 parts.push(msg.content.as_text());
100 }
101 for tool in &request.tools {
102 parts.push(tool.name.clone());
103 parts.push(tool.description.clone());
104 parts.push(tool.parameters.to_string());
105 }
106 parts.join("\n")
107}
108
109pub struct InferenceOutcome {
112 pub entity: Entity,
114 pub result: Result<InferenceResponse, ProviderError>,
116 pub latency: std::time::Duration,
120}
121
122pub async fn run_inference_job(
129 job: InferenceJob,
130 results: UnboundedSender<InferenceOutcome>,
131 wake: Arc<Notify>,
132 retry: RetryPolicy,
133 cancel: crate::cancel::CancelToken,
134) {
135 let InferenceJob {
136 entity,
137 provider,
138 request,
139 permit,
140 exact_token_counting,
141 } = job;
142 let started = std::time::Instant::now();
143 if exact_token_counting {
148 let text = flatten_request_text(&request);
149 let used = provider.count_tokens(&text, &request.model).await;
150 let max = provider.max_context_tokens(&request.model);
151 if used.saturating_add(request.max_tokens) > max {
152 drop(permit);
153 let _ = results.send(InferenceOutcome {
154 entity,
155 result: Err(ProviderError::TokenLimitExceeded { used, max }),
156 latency: started.elapsed(),
157 });
158 wake.notify_one();
159 return;
160 }
161 }
162 let attempts = async {
167 let mut attempt = 1u32;
168 loop {
169 match provider.infer(request.clone()).await {
170 Ok(response) => break Ok(response),
171 Err(e) if e.is_transient() && attempt < retry.max_attempts => {
172 tokio::time::sleep(retry.base_delay * 2u32.pow(attempt - 1)).await;
173 attempt += 1;
174 }
175 Err(e) => break Err(e),
176 }
177 }
178 };
179 let result = tokio::select! {
190 biased;
191 _ = cancel.cancelled() => {
192 drop(permit);
193 return;
194 }
195 outcome = tokio::time::timeout(retry.job_timeout, attempts) => match outcome {
196 Ok(result) => result,
197 Err(_elapsed) => Err(leviath_providers::ProviderError::Other(format!(
198 "inference exceeded the {}s job timeout and was aborted to free the \
199 pool slot (a stalled or never-completing response)",
200 retry.job_timeout.as_secs()
201 ))),
202 },
203 };
204 drop(permit); let _ = results.send(InferenceOutcome {
206 entity,
207 result,
208 latency: started.elapsed(),
209 });
210 wake.notify_one();
211}
212
213#[cfg(test)]
214mod tests {
215 use super::*;
216 use crate::inference_pool::{InferencePoolConfig, InferencePools};
217 use tokio::sync::mpsc;
218
219 fn test_request() -> InferenceRequest {
220 InferenceRequest {
221 system: vec![],
222 messages: vec![],
223 model: "m".to_string(),
224 max_tokens: 100,
225 temperature: 0.0,
226 tools: vec![],
227 extra: serde_json::Value::Null,
228 request_timeout_secs: None,
229 }
230 }
231
232 fn response(text: &str) -> InferenceResponse {
233 InferenceResponse {
234 content: text.to_string(),
235 tool_calls: vec![],
236 tokens_used: leviath_providers::TokenUsage {
237 prompt_tokens: 1,
238 completion_tokens: 1,
239 total_tokens: 2,
240 cached_tokens: 0,
241 cache_write_tokens: 0,
242 },
243 finish_reason: leviath_providers::FinishReason::Complete,
244 }
245 }
246
247 enum Fixed {
250 Ok(InferenceResponse),
251 Err(String),
252 }
253
254 #[async_trait::async_trait]
255 impl Provider for Fixed {
256 async fn infer(
257 &self,
258 _req: InferenceRequest,
259 ) -> leviath_providers::Result<InferenceResponse> {
260 match self {
261 Fixed::Ok(r) => Ok(r.clone()),
262 Fixed::Err(m) => Err(ProviderError::Other(m.clone())),
263 }
264 }
265 async fn count_tokens(&self, _text: &str, _model: &str) -> usize {
266 1
267 }
268 fn max_context_tokens(&self, _model: &str) -> usize {
269 100_000
270 }
271 fn name(&self) -> &str {
272 "fixed"
273 }
274 fn capabilities(&self, _model: &str) -> leviath_providers::ModelCapabilities {
275 leviath_providers::ModelCapabilities::default()
276 }
277 }
278
279 fn job(provider: Arc<dyn Provider>) -> InferenceJob {
280 let pools = InferencePools::new(InferencePoolConfig::new());
281 InferenceJob {
282 entity: Entity::from_raw_u32(7)
283 .expect("a small literal index is always a valid entity id"),
284 provider,
285 request: test_request(),
286 permit: pools.try_acquire("m").expect("free pool"),
287 exact_token_counting: false,
288 }
289 }
290
291 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
295 async fn a_cancelled_job_frees_its_pool_slot_without_reporting() {
296 let mut cfg = InferencePoolConfig::new();
297 cfg.set_limit("m", 1);
298 let pools = InferencePools::new(cfg);
299 let permit = pools.try_acquire("m").expect("free pool");
300 assert!(pools.try_acquire("m").is_none(), "pool should be full");
301
302 let provider = Arc::new(Scripted {
305 steps: std::sync::Mutex::new(vec![Step::Hang].into()),
306 calls: std::sync::Mutex::new(0),
307 });
308 let job = InferenceJob {
309 entity: Entity::from_raw_u32(7)
310 .expect("a small literal index is always a valid entity id"),
311 provider,
312 request: test_request(),
313 permit,
314 exact_token_counting: false,
315 };
316 let (tx, mut rx) = mpsc::unbounded_channel();
317 let cancel = crate::cancel::CancelToken::new();
318 let running = tokio::spawn(run_inference_job(
319 job,
320 tx,
321 Arc::new(Notify::new()),
322 RetryPolicy {
324 max_attempts: 1,
325 base_delay: Duration::ZERO,
326 job_timeout: Duration::from_secs(3600),
327 },
328 cancel.clone(),
329 ));
330 tokio::task::yield_now().await;
331 cancel.cancel();
332
333 tokio::time::timeout(Duration::from_secs(5), running)
334 .await
335 .expect("the cancel ended the job")
336 .unwrap();
337 assert!(
338 pools.try_acquire("m").is_some(),
339 "the pool slot is free for the next agent"
340 );
341 assert!(
342 rx.try_recv().is_err(),
343 "and no outcome is reported for a cancelled run"
344 );
345 }
346
347 #[tokio::test]
348 async fn run_job_aborts_a_hung_call_and_frees_the_pool_slot() {
349 let mut cfg = InferencePoolConfig::new();
351 cfg.set_limit("m", 1);
352 let pools = InferencePools::new(cfg);
353 let permit = pools.try_acquire("m").expect("free pool");
354 assert!(pools.try_acquire("m").is_none(), "pool should be full");
355
356 let provider = Arc::new(Scripted {
357 steps: std::sync::Mutex::new(vec![Step::Hang].into()),
358 calls: std::sync::Mutex::new(0),
359 });
360 let job = InferenceJob {
361 entity: Entity::from_raw_u32(7)
362 .expect("a small literal index is always a valid entity id"),
363 provider,
364 request: test_request(),
365 permit,
366 exact_token_counting: false,
367 };
368 let (tx, mut rx) = mpsc::unbounded_channel();
369 let policy = RetryPolicy {
370 max_attempts: 1,
371 base_delay: Duration::ZERO,
372 job_timeout: Duration::from_millis(50),
373 };
374 run_inference_job(
375 job,
376 tx,
377 Arc::new(Notify::new()),
378 policy,
379 crate::cancel::CancelToken::new(),
380 )
381 .await;
382
383 let outcome = rx.try_recv().expect("outcome sent");
385 let err = outcome.result.expect_err("hung call should error");
386 assert!(err.to_string().contains("job timeout"), "got: {err}");
387 assert!(
389 pools.try_acquire("m").is_some(),
390 "the slot must be released after the timeout"
391 );
392 }
393
394 #[tokio::test]
395 async fn run_job_reports_ok_and_wakes() {
396 let (tx, mut rx) = mpsc::unbounded_channel();
397 let wake = Arc::new(Notify::new());
398 run_inference_job(
399 job(Arc::new(Fixed::Ok(response("hi")))),
400 tx,
401 wake.clone(),
402 RetryPolicy::default(),
403 crate::cancel::CancelToken::new(),
404 )
405 .await;
406
407 let outcome = rx.try_recv().expect("outcome sent");
408 assert_eq!(
409 outcome.entity,
410 Entity::from_raw_u32(7).expect("a small literal index is always a valid entity id")
411 );
412 assert_eq!(outcome.result.unwrap().content, "hi");
413 wake.notified().await;
415 }
416
417 #[tokio::test]
418 async fn run_job_reports_provider_error() {
419 let (tx, mut rx) = mpsc::unbounded_channel();
420 let wake = Arc::new(Notify::new());
421 let err = Arc::new(Fixed::Err("boom".to_string()));
422 run_inference_job(
423 job(err),
424 tx,
425 wake,
426 RetryPolicy::default(),
427 crate::cancel::CancelToken::new(),
428 )
429 .await;
430
431 let outcome = rx.try_recv().expect("outcome sent");
432 assert!(outcome.result.is_err());
433 }
434
435 struct Counter {
438 count: usize,
439 max: usize,
440 }
441
442 #[async_trait::async_trait]
443 impl Provider for Counter {
444 async fn infer(
445 &self,
446 _req: InferenceRequest,
447 ) -> leviath_providers::Result<InferenceResponse> {
448 Ok(response("ok"))
449 }
450 async fn count_tokens(&self, _text: &str, _model: &str) -> usize {
451 self.count
452 }
453 fn max_context_tokens(&self, _model: &str) -> usize {
454 self.max
455 }
456 fn name(&self) -> &str {
457 "counter"
458 }
459 fn capabilities(&self, _model: &str) -> leviath_providers::ModelCapabilities {
460 leviath_providers::ModelCapabilities::default()
461 }
462 }
463
464 fn counting_job(provider: Arc<dyn Provider>, exact: bool) -> InferenceJob {
465 let pools = InferencePools::new(InferencePoolConfig::new());
466 InferenceJob {
467 entity: Entity::from_raw_u32(7)
468 .expect("a small literal index is always a valid entity id"),
469 provider,
470 request: test_request(), permit: pools.try_acquire("m").expect("free pool"),
472 exact_token_counting: exact,
473 }
474 }
475
476 #[test]
477 fn flatten_request_text_includes_system_messages_and_tools() {
478 use leviath_providers::{SystemBlock, Tool};
479 let req = InferenceRequest {
480 system: vec![SystemBlock {
481 text: "sys".to_string(),
482 cache_hint: leviath_core::CacheHint::Never,
483 }],
484 messages: vec![leviath_providers::Message {
485 role: "user".to_string(),
486 content: "hello".into(),
487 cache_breakpoint: false,
488 }],
489 model: "m".to_string(),
490 max_tokens: 10,
491 temperature: 0.0,
492 tools: vec![Tool {
493 name: "search".to_string(),
494 description: "find things".to_string(),
495 parameters: serde_json::json!({"type": "object"}),
496 }],
497 extra: serde_json::Value::Null,
498 request_timeout_secs: None,
499 };
500 let text = flatten_request_text(&req);
501 assert!(text.contains("sys"));
502 assert!(text.contains("hello"));
503 assert!(text.contains("search"));
504 assert!(text.contains("find things"));
505 assert!(text.contains("object"));
506 }
507
508 #[tokio::test]
509 async fn guard_rejects_request_over_context_window() {
510 let (tx, mut rx) = mpsc::unbounded_channel();
512 let provider = Arc::new(Counter {
513 count: 950,
514 max: 1000,
515 });
516 run_inference_job(
517 counting_job(provider, true),
518 tx,
519 Arc::new(Notify::new()),
520 RetryPolicy::default(),
521 crate::cancel::CancelToken::new(),
522 )
523 .await;
524 let outcome = rx.try_recv().expect("outcome sent");
525 let err = outcome.result.expect_err("should be rejected");
526 assert_eq!(err.to_string(), "Token limit exceeded: 950 > 1000");
529 }
530
531 #[tokio::test]
532 async fn guard_allows_request_within_context_window() {
533 let (tx, mut rx) = mpsc::unbounded_channel();
535 let provider = Arc::new(Counter {
536 count: 800,
537 max: 1000,
538 });
539 run_inference_job(
540 counting_job(provider, true),
541 tx,
542 Arc::new(Notify::new()),
543 RetryPolicy::default(),
544 crate::cancel::CancelToken::new(),
545 )
546 .await;
547 let outcome = rx.try_recv().expect("outcome sent");
548 assert_eq!(outcome.result.expect("should succeed").content, "ok");
549 }
550
551 #[tokio::test]
552 async fn counter_provider_metadata_is_exercised() {
553 let p = Counter { count: 5, max: 10 };
555 assert_eq!(p.name(), "counter");
556 assert_eq!(p.max_context_tokens("m"), 10);
557 assert_eq!(p.count_tokens("t", "m").await, 5);
558 assert!(p.capabilities("m").supports_streaming);
559 }
560
561 #[tokio::test]
562 async fn guard_off_skips_the_count_and_proceeds() {
563 let (tx, mut rx) = mpsc::unbounded_channel();
565 let provider = Arc::new(Counter {
566 count: 1_000_000,
567 max: 1000,
568 });
569 run_inference_job(
570 counting_job(provider, false),
571 tx,
572 Arc::new(Notify::new()),
573 RetryPolicy::default(),
574 crate::cancel::CancelToken::new(),
575 )
576 .await;
577 let outcome = rx.try_recv().expect("outcome sent");
578 assert_eq!(outcome.result.expect("should succeed").content, "ok");
579 }
580
581 #[tokio::test]
582 async fn fixed_provider_metadata_is_exercised() {
583 let p = Fixed::Ok(response("x"));
586 assert_eq!(p.name(), "fixed");
587 assert_eq!(p.count_tokens("t", "m").await, 1);
588 assert_eq!(p.max_context_tokens("m"), 100_000);
589 let _ = p.capabilities("m");
590 }
591
592 #[tokio::test]
593 async fn run_job_survives_dropped_receiver() {
594 let (tx, rx) = mpsc::unbounded_channel();
595 drop(rx); let wake = Arc::new(Notify::new());
597 run_inference_job(
599 job(Arc::new(Fixed::Ok(response("x")))),
600 tx,
601 wake,
602 RetryPolicy::default(),
603 crate::cancel::CancelToken::new(),
604 )
605 .await;
606 }
607
608 enum Step {
611 Ok(String),
612 Transient,
613 Permanent,
614 Hang,
616 }
617
618 struct Scripted {
620 steps: std::sync::Mutex<std::collections::VecDeque<Step>>,
621 calls: std::sync::Mutex<u32>,
622 }
623
624 #[async_trait::async_trait]
625 impl Provider for Scripted {
626 async fn infer(
627 &self,
628 _req: InferenceRequest,
629 ) -> leviath_providers::Result<InferenceResponse> {
630 *self.calls.lock().unwrap() += 1;
631 let step = self.steps.lock().unwrap().pop_front();
634 match step {
635 Some(Step::Ok(t)) => Ok(response(&t)),
636 Some(Step::Transient) => Err(ProviderError::RateLimitExceeded),
637 Some(Step::Permanent) => Err(ProviderError::Other("permanent".to_string())),
638 Some(Step::Hang) => std::future::pending().await,
639 None => Err(ProviderError::Other("exhausted".to_string())),
640 }
641 }
642 async fn count_tokens(&self, _t: &str, _m: &str) -> usize {
643 1
644 }
645 fn max_context_tokens(&self, _m: &str) -> usize {
646 100_000
647 }
648 fn name(&self) -> &str {
649 "scripted"
650 }
651 fn capabilities(&self, _m: &str) -> leviath_providers::ModelCapabilities {
652 leviath_providers::ModelCapabilities::default()
653 }
654 }
655
656 fn no_delay(max_attempts: u32) -> RetryPolicy {
657 RetryPolicy {
658 max_attempts,
659 base_delay: Duration::ZERO,
660 job_timeout: Duration::from_secs(30),
661 }
662 }
663
664 #[tokio::test]
665 async fn run_job_retries_transient_then_succeeds() {
666 let provider = Arc::new(Scripted {
667 steps: std::sync::Mutex::new(
668 vec![
669 Step::Transient,
670 Step::Transient,
671 Step::Ok("done".to_string()),
672 ]
673 .into(),
674 ),
675 calls: std::sync::Mutex::new(0),
676 });
677 let (tx, mut rx) = mpsc::unbounded_channel();
678 run_inference_job(
679 job(provider.clone()),
680 tx,
681 Arc::new(Notify::new()),
682 no_delay(4),
683 crate::cancel::CancelToken::new(),
684 )
685 .await;
686 let outcome = rx.try_recv().expect("outcome sent");
687 assert_eq!(outcome.result.unwrap().content, "done");
688 assert_eq!(*provider.calls.lock().unwrap(), 3); }
690
691 #[tokio::test]
692 async fn run_job_gives_up_after_max_attempts() {
693 let provider = Arc::new(Scripted {
694 steps: std::sync::Mutex::new(
695 vec![
696 Step::Transient,
697 Step::Transient,
698 Step::Transient,
699 Step::Transient,
700 ]
701 .into(),
702 ),
703 calls: std::sync::Mutex::new(0),
704 });
705 let (tx, mut rx) = mpsc::unbounded_channel();
706 run_inference_job(
707 job(provider.clone()),
708 tx,
709 Arc::new(Notify::new()),
710 no_delay(3),
711 crate::cancel::CancelToken::new(),
712 )
713 .await;
714 let outcome = rx.try_recv().expect("outcome sent");
715 assert!(outcome.result.is_err());
716 assert_eq!(*provider.calls.lock().unwrap(), 3); }
718
719 #[tokio::test]
720 async fn run_job_does_not_retry_a_permanent_error() {
721 let provider = Arc::new(Scripted {
722 steps: std::sync::Mutex::new(vec![Step::Permanent, Step::Ok("x".to_string())].into()),
723 calls: std::sync::Mutex::new(0),
724 });
725 let (tx, mut rx) = mpsc::unbounded_channel();
726 run_inference_job(
727 job(provider.clone()),
728 tx,
729 Arc::new(Notify::new()),
730 no_delay(4),
731 crate::cancel::CancelToken::new(),
732 )
733 .await;
734 let outcome = rx.try_recv().expect("outcome sent");
735 assert!(outcome.result.is_err());
736 assert_eq!(*provider.calls.lock().unwrap(), 1); }
738
739 #[tokio::test]
740 async fn scripted_provider_metadata_is_exercised() {
741 let p = Scripted {
742 steps: std::sync::Mutex::new(std::collections::VecDeque::new()),
743 calls: std::sync::Mutex::new(0),
744 };
745 assert_eq!(p.name(), "scripted");
746 assert_eq!(p.count_tokens("t", "m").await, 1);
747 assert_eq!(p.max_context_tokens("m"), 100_000);
748 let _ = p.capabilities("m");
749 }
750}