1use std::collections::HashSet;
16use std::future::Future;
17use std::sync::Arc;
18
19use crate::nvidia_catalog::NvidiaCatalogCache;
20use async_trait::async_trait;
21use chrono::{DateTime, Utc};
22use cordis::{Context, CordisError, EventsService, Service};
23use parking_lot::RwLock;
24
25use crate::capabilities::CapabilityRequirements;
26use crate::client::{GenerationHints, LLMClient, LLMResponse};
27use crate::config::ProviderConfig;
28use crate::pool::ClientPool;
29use crate::provider_registry::{
30 ConfigBasedLLMFactory, ModelInfo, ProviderRegistry, RuntimeProviderEntry,
31};
32use ares_types::types::{AppError, ToolDefinition};
33
34#[derive(Debug, Clone)]
43pub struct ModelOverride {
44 pub model: String,
46}
47
48impl Service for ModelOverride {}
52
53#[derive(Debug, Clone, PartialEq, Eq)]
60pub struct TenantModelPolicy {
61 tenant_id: String,
62 allowed_models: HashSet<String>,
63}
64
65impl TenantModelPolicy {
66 pub fn new<I, S>(tenant_id: impl Into<String>, allowed_models: I) -> Self
68 where
69 I: IntoIterator<Item = S>,
70 S: Into<String>,
71 {
72 Self {
73 tenant_id: tenant_id.into(),
74 allowed_models: allowed_models.into_iter().map(Into::into).collect(),
75 }
76 }
77
78 pub fn tenant_id(&self) -> &str {
80 &self.tenant_id
81 }
82
83 pub fn allows(&self, model: &str) -> bool {
85 self.allowed_models.contains(model)
86 }
87
88 pub fn denial_message(tenant_id: &str, model: &str) -> String {
90 format!(
91 "Model '{}' is not allowed for tenant '{}'",
92 model, tenant_id
93 )
94 }
95
96 pub fn denial_error(tenant_id: &str, model: &str) -> AppError {
98 AppError::Auth(Self::denial_message(tenant_id, model))
99 }
100
101 pub fn authorize(&self, model: &str) -> Result<(), AppError> {
103 if self.allows(model) {
104 Ok(())
105 } else {
106 Err(Self::denial_error(&self.tenant_id, model))
107 }
108 }
109}
110
111impl Service for TenantModelPolicy {}
112
113#[derive(Debug, Clone, Default)]
119pub enum Breaker {
120 #[default]
122 Closed,
123 Open { until: DateTime<Utc> },
125 HalfOpen,
127}
128
129impl Breaker {
130 pub const FAILURE_THRESHOLD: u32 = 5;
132 pub const COOLDOWN_SECS: i64 = 30;
134
135 pub fn check(&self) -> bool {
143 match self {
144 Breaker::Closed => true,
145 Breaker::HalfOpen => true,
146 Breaker::Open { until } => {
147 Utc::now() >= *until
150 }
151 }
152 }
153
154 pub fn is_closed(&self) -> bool {
156 matches!(self, Breaker::Closed)
157 }
158
159 pub fn transition_on_failure(&self) -> Breaker {
169 let now = Utc::now();
170 let cooldown = chrono::Duration::seconds(Self::COOLDOWN_SECS);
171 match self {
172 Breaker::Closed => Breaker::Open {
173 until: now + cooldown,
174 },
175 Breaker::HalfOpen => Breaker::Open {
176 until: now + cooldown,
177 },
178 Breaker::Open { .. } => Breaker::Open {
179 until: now + cooldown,
180 },
181 }
182 }
183
184 pub fn transition_on_failure_with_count(&self, failures: u32) -> Breaker {
186 if failures >= Self::FAILURE_THRESHOLD {
187 let now = Utc::now();
188 Breaker::Open {
189 until: now + chrono::Duration::seconds(Self::COOLDOWN_SECS),
190 }
191 } else {
192 Breaker::Closed
193 }
194 }
195}
196
197pub struct Llm {
206 pub(crate) provider_registry: Arc<ProviderRegistry>,
208 pub(crate) catalog: Option<Arc<NvidiaCatalogCache>>,
211 pub(crate) pool: Arc<ClientPool>,
213 pub(crate) factory: Option<Arc<ConfigBasedLLMFactory>>,
215 breaker: RwLock<Breaker>,
217 failures: RwLock<u32>,
219 test_client: Option<Arc<dyn LLMClient>>,
223}
224
225impl Llm {
226 pub fn new(
228 provider_registry: Arc<ProviderRegistry>,
229 pool: Arc<ClientPool>,
230 catalog: Option<Arc<NvidiaCatalogCache>>,
231 ) -> Self {
232 Self {
233 provider_registry,
234 catalog,
235 pool,
236 factory: None,
237 breaker: RwLock::new(Breaker::Closed),
238 failures: RwLock::new(0),
239 test_client: None,
240 }
241 }
242
243 pub fn with_factory(mut self, factory: Arc<ConfigBasedLLMFactory>) -> Self {
245 self.factory = Some(factory);
246 self
247 }
248
249 pub fn from_client(client: Arc<dyn LLMClient>) -> Self {
253 let mut llm = Self::new(
254 Arc::new(ProviderRegistry::new()),
255 Arc::new(ClientPool::with_defaults()),
256 None,
257 );
258 llm.test_client = Some(client);
259 llm
260 }
261
262 #[cfg(test)]
264 pub(crate) fn for_test(client: Arc<dyn LLMClient>) -> Self {
265 Self::from_client(client)
266 }
267
268 pub(crate) fn provider_registry(&self) -> Arc<ProviderRegistry> {
270 Arc::clone(&self.provider_registry)
271 }
272
273 pub fn registry(&self) -> Arc<ProviderRegistry> {
277 self.provider_registry()
278 }
279
280 pub fn with_breaker(
282 provider_registry: Arc<ProviderRegistry>,
283 catalog: Option<Arc<NvidiaCatalogCache>>,
284 pool: Arc<ClientPool>,
285 breaker: Breaker,
286 ) -> Self {
287 Self {
288 provider_registry,
289 catalog,
290 pool,
291 factory: None,
292 breaker: RwLock::new(breaker),
293 failures: RwLock::new(0),
294 test_client: None,
295 }
296 }
297
298 pub fn with_catalog(
300 provider_registry: Arc<ProviderRegistry>,
301 catalog: Arc<NvidiaCatalogCache>,
302 pool: Arc<ClientPool>,
303 ) -> Self {
304 Self::new(provider_registry, pool, Some(catalog))
305 }
306
307 pub fn breaker(&self) -> Breaker {
309 self.breaker.read().clone()
310 }
311
312 pub fn trip(&self, until: DateTime<Utc>) {
314 *self.breaker.write() = Breaker::Open { until };
315 }
316
317 pub fn half_open(&self) {
319 *self.breaker.write() = Breaker::HalfOpen;
320 }
321
322 pub fn reset(&self) {
324 *self.breaker.write() = Breaker::Closed;
325 *self.failures.write() = 0;
326 }
327
328 pub fn record_success(&self) {
330 *self.breaker.write() = Breaker::Closed;
331 *self.failures.write() = 0;
332 }
333
334 pub fn record_failure(&self) {
339 let mut failures = self.failures.write();
340 *failures = failures.saturating_add(1);
341 let count = *failures;
342 drop(failures);
343 let mut b = self.breaker.write();
344 match &*b {
346 Breaker::HalfOpen => {
347 *b = b.transition_on_failure();
348 }
349 Breaker::Closed => {
350 if count >= Breaker::FAILURE_THRESHOLD {
351 *b = Breaker::Open {
352 until: Utc::now() + chrono::Duration::seconds(Breaker::COOLDOWN_SECS),
353 };
354 }
355 }
356 Breaker::Open { .. } => {
357 *b = b.transition_on_failure();
359 }
360 }
361 }
362
363 pub fn validate_model_override(&self, ctx: &Arc<Context>) -> Result<(), AppError> {
369 if let (Some(policy), Some(override_model)) =
370 (ctx.get::<TenantModelPolicy>(), ctx.get::<ModelOverride>())
371 {
372 policy.authorize(&override_model.model)?;
373 }
374 Ok(())
375 }
376
377 pub async fn get_client(
386 &self,
387 ctx: &Arc<Context>,
388 capability: CapabilityRequirements,
389 ) -> Result<Arc<dyn LLMClient>, AppError> {
390 let Some(events) = ctx.get::<EventsService>() else {
391 return self.get_client_inner(ctx, capability).await;
392 };
393 let payload = serde_json::to_value(cordis::LlmGetClientPayload {
394 capability: format!("{capability:?}"),
395 deny: None,
396 model: None,
397 })
398 .unwrap_or(serde_json::Value::Null);
399 let result = events
400 .waterfall_around(
401 cordis::events_catalog::ev::LLM_GET_CLIENT.to_string(),
402 payload,
403 |payload| async move { Ok(payload) },
404 )
405 .await
406 .map_err(map_cordis)?;
407 if result.get("deny").and_then(|v| v.as_bool()) == Some(true) {
408 return Err(AppError::InvalidInput("llm.get_client denied".into()));
409 }
410 if let Some(model) = result
411 .get("model")
412 .and_then(|v| v.as_str())
413 .filter(|s| !s.is_empty())
414 {
415 if ctx.get::<ModelOverride>().is_none() {
416 let intercepted = ctx.with_intercept(ModelOverride {
417 model: model.to_string(),
418 });
419 return self.get_client_inner(&intercepted, capability).await;
420 }
421 }
422 self.get_client_inner(ctx, capability).await
423 }
424
425 async fn get_client_inner(
427 &self,
428 ctx: &Arc<Context>,
429 capability: CapabilityRequirements,
430 ) -> Result<Arc<dyn LLMClient>, AppError> {
431 if let Some(c) = &self.test_client {
432 return Ok(Arc::clone(c));
433 }
434 self.validate_model_override(ctx)?;
436 if let Some(ov) = ctx.get::<ModelOverride>() {
438 if let Ok(guard) = self.pool.try_get(&ov.model).await {
440 let boxed = guard.take();
441 return Ok(Arc::from(boxed));
442 }
443 if let Ok(client) = self
444 .provider_registry
445 .create_client_for_model_ctx(ctx, &ov.model)
446 .await
447 {
448 return Ok(Arc::from(client));
449 }
450 }
452
453 if let Some(catalog) = &self.catalog {
455 let _snap = catalog.snapshot(); if let Some(best) = self.provider_registry.find_best_model(&capability) {
457 if let Ok(client) = self
458 .provider_registry
459 .create_client_for_model_ctx(ctx, &best.name)
460 .await
461 {
462 return Ok(Arc::from(client));
463 }
464 }
465 } else if let Some(best) = self.provider_registry.find_best_model(&capability) {
466 if let Ok(client) = self
467 .provider_registry
468 .create_client_for_model_ctx(ctx, &best.name)
469 .await
470 {
471 return Ok(Arc::from(client));
472 }
473 }
474
475 let client = self
481 .provider_registry
482 .resolve_with_capability_fallback(Some(capability))
483 .await?;
484 Ok(Arc::from(client))
485 }
486
487 pub async fn get_client_boxed(
489 &self,
490 ctx: &Arc<Context>,
491 capability: CapabilityRequirements,
492 ) -> Result<Box<dyn LLMClient>, AppError> {
493 let client = self.get_client(ctx, capability).await?;
494 Ok(Box::new(BoxedArcClient(client)))
495 }
496
497 pub async fn complete(&self, ctx: &Arc<Context>, prompt: &str) -> Result<String, AppError> {
503 let client = self
504 .get_client(ctx, CapabilityRequirements::default())
505 .await?;
506 let Some(events) = ctx.get::<EventsService>() else {
507 return client.generate(prompt).await;
508 };
509 let payload = serde_json::to_value(cordis::LlmCompleteRequest {
510 prompt: prompt.to_string(),
511 })
512 .unwrap_or(serde_json::Value::Null);
513 let out = events
514 .waterfall_around(
515 cordis::events_catalog::ev::LLM_COMPLETE.to_string(),
516 payload,
517 move |payload| {
518 let client = Arc::clone(&client);
519 async move {
520 let prompt = payload
521 .get("prompt")
522 .and_then(|v| v.as_str())
523 .unwrap_or("")
524 .to_string();
525 let text = client
526 .generate(&prompt)
527 .await
528 .map_err(|e| CordisError::Fiber(e.to_string()))?;
529 Ok(serde_json::json!({ "prompt": prompt, "content": text }))
530 }
531 },
532 )
533 .await
534 .map_err(map_cordis)?;
535 Ok(out
536 .get("content")
537 .and_then(|v| v.as_str())
538 .unwrap_or("")
539 .to_string())
540 }
541
542 pub fn find_model_stub(&self, _capability: &str) -> Option<String> {
544 None
545 }
546
547 pub fn list_models(&self) -> Vec<ModelInfo> {
549 self.provider_registry.list_models()
550 }
551
552 pub fn has_provider_for_tenant(&self, name: &str, tenant_id: Option<&str>) -> bool {
554 self.provider_registry
555 .has_provider_for_tenant(name, tenant_id)
556 }
557
558 pub fn get_provider_for_ctx(&self, ctx: &Arc<Context>, name: &str) -> Option<ProviderConfig> {
560 self.provider_registry.get_provider_for_ctx(ctx, name)
561 }
562
563 pub fn reload_runtime_providers(
565 &self,
566 providers: Vec<RuntimeProviderEntry>,
567 names: Vec<String>,
568 ) {
569 self.provider_registry
570 .reload_runtime_providers(providers, names);
571 }
572}
573
574fn map_cordis(err: CordisError) -> AppError {
575 AppError::Internal(err.to_string())
576}
577
578struct BoxedArcClient(Arc<dyn LLMClient>);
580
581#[async_trait]
582impl LLMClient for BoxedArcClient {
583 async fn generate(&self, prompt: &str) -> ares_types::types::Result<String> {
584 self.0.generate(prompt).await
585 }
586
587 async fn generate_with_system(
588 &self,
589 system: &str,
590 prompt: &str,
591 ) -> ares_types::types::Result<String> {
592 self.0.generate_with_system(system, prompt).await
593 }
594
595 async fn generate_with_history(
596 &self,
597 messages: &[(String, String)],
598 ) -> ares_types::types::Result<LLMResponse> {
599 self.0.generate_with_history(messages).await
600 }
601
602 async fn generate_with_tools(
603 &self,
604 prompt: &str,
605 tools: &[ToolDefinition],
606 ) -> ares_types::types::Result<LLMResponse> {
607 self.0.generate_with_tools(prompt, tools).await
608 }
609
610 async fn generate_with_tools_and_history(
611 &self,
612 messages: &[crate::coordinator::ConversationMessage],
613 tools: &[ToolDefinition],
614 ) -> ares_types::types::Result<LLMResponse> {
615 self.0
616 .generate_with_tools_and_history(messages, tools)
617 .await
618 }
619
620 async fn stream(
621 &self,
622 prompt: &str,
623 ) -> ares_types::types::Result<
624 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
625 > {
626 self.0.stream(prompt).await
627 }
628
629 async fn stream_with_system(
630 &self,
631 system: &str,
632 prompt: &str,
633 ) -> ares_types::types::Result<
634 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
635 > {
636 self.0.stream_with_system(system, prompt).await
637 }
638
639 async fn stream_with_history(
640 &self,
641 messages: &[(String, String)],
642 ) -> ares_types::types::Result<
643 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
644 > {
645 self.0.stream_with_history(messages).await
646 }
647
648 fn model_name(&self) -> &str {
649 self.0.model_name()
650 }
651 fn supports_hints(&self) -> bool {
652 self.0.supports_hints()
653 }
654 fn set_hints(&self, hints: GenerationHints) {
655 self.0.set_hints(hints)
656 }
657}
658
659impl Service for Llm {
660 fn name(&self) -> &'static str {
661 "Llm"
662 }
663
664 fn init(
665 &self,
666 _ctx: &Arc<Context>,
667 ) -> std::pin::Pin<
668 Box<
669 dyn Future<Output = Result<Option<Box<dyn cordis::Disposable>>, CordisError>>
670 + Send
671 + '_,
672 >,
673 > {
674 Box::pin(async { Ok(None) })
675 }
676
677 fn check(&self) -> bool {
678 self.breaker.read().check()
681 }
682}
683
684#[cfg(test)]
685mod tests {
686 use super::*;
687 use crate::capabilities::CapabilityRequirements;
688 use crate::provider_registry::RuntimeProviderEntry;
689 use ares_types::models::{TenantContext, TenantTier};
690 use chrono::Duration;
691 use cordis::Context;
692 use std::collections::HashMap;
693
694 #[test]
695 fn breaker_closed_allows() {
696 assert!(Breaker::Closed.check());
697 }
698
699 #[test]
700 fn breaker_half_open_allows() {
701 assert!(Breaker::HalfOpen.check());
702 }
703
704 #[test]
705 fn breaker_open_future_denies() {
706 let until = Utc::now() + Duration::seconds(60);
707 assert!(!Breaker::Open { until }.check());
708 }
709
710 #[test]
711 fn breaker_open_past_allows() {
712 let until = Utc::now() - Duration::seconds(1);
713 assert!(Breaker::Open { until }.check());
714 }
715
716 #[test]
717 fn breaker_failure_threshold_opens() {
718 let b = Breaker::Closed;
719 let next = b.transition_on_failure_with_count(5);
720 assert!(matches!(next, Breaker::Open { .. }));
721 let still_closed = b.transition_on_failure_with_count(3);
722 assert!(matches!(still_closed, Breaker::Closed));
723 }
724
725 #[test]
726 fn breaker_constants_exist() {
727 assert_eq!(Breaker::FAILURE_THRESHOLD, 5);
728 assert_eq!(Breaker::COOLDOWN_SECS, 30);
729 }
730
731 #[test]
732 fn provider_registry_and_factory_accessors() {
733 let registry = Arc::new(ProviderRegistry::new());
734 let pool = Arc::new(ClientPool::with_defaults());
735 let factory = Arc::new(
736 ConfigBasedLLMFactory::from_config(HashMap::new(), HashMap::new(), None)
737 .expect("empty factory config"),
738 );
739 let llm = Llm::new(Arc::clone(®istry), pool, None).with_factory(Arc::clone(&factory));
740 assert!(Arc::ptr_eq(&llm.provider_registry(), ®istry));
741 }
742
743 #[tokio::test]
744 async fn llm_model_override_via_context() {
745 let registry = Arc::new(ProviderRegistry::new());
746 let pool = Arc::new(ClientPool::with_defaults());
747 let svc = Arc::new(Llm::new(registry, pool, None));
748 let root = Context::new_root();
749 root.provide::<Llm>(Llm::new(
750 Arc::new(ProviderRegistry::new()),
751 Arc::new(ClientPool::with_defaults()),
752 None,
753 ));
754 let req_ctx = root.intercept(ModelOverride {
756 model: "gpt-4o-mini".into(),
757 });
758 assert!(req_ctx.get::<ModelOverride>().is_some());
759 assert_eq!(req_ctx.get::<ModelOverride>().unwrap().model, "gpt-4o-mini");
760 assert!(!Arc::as_ptr(&svc.provider_registry).is_null());
762 let _ = svc.catalog.clone();
763 let _ = svc.pool.provider_names();
764 assert!(svc.check());
766 }
767
768 #[tokio::test]
769 async fn record_failure_threshold_opens_breaker() {
770 let svc = Llm::new(
771 Arc::new(ProviderRegistry::new()),
772 Arc::new(ClientPool::with_defaults()),
773 None,
774 );
775 for _ in 0..Breaker::FAILURE_THRESHOLD {
776 svc.record_failure();
777 }
778 assert!(!svc.check());
780 svc.record_success();
781 assert!(svc.check());
782 }
783
784 #[test]
785 fn tenant_model_policy_allows_and_composes_with_model_override() {
786 let root = Context::new_root();
787 let tenant_ctx = root.intercept(TenantModelPolicy::new(
788 "tenant-a",
789 ["gpt-4o-mini".to_string()],
790 ));
791 let request = tenant_ctx.intercept(ModelOverride {
792 model: "gpt-4o-mini".into(),
793 });
794 let policy = request
795 .get::<TenantModelPolicy>()
796 .expect("policy should be inherited by request context");
797 let override_model = request
798 .get::<ModelOverride>()
799 .expect("model override should be visible in request context");
800 policy
801 .authorize(&override_model.model)
802 .expect("allowed model override should pass policy");
803 let svc = Llm::new(
804 Arc::new(ProviderRegistry::new()),
805 Arc::new(ClientPool::with_defaults()),
806 None,
807 );
808 svc.validate_model_override(&request)
809 .expect("allowed model override should pass LLM validation");
810 assert!(root.get::<ModelOverride>().is_none());
811 assert!(root.get::<TenantModelPolicy>().is_none());
812 }
813
814 #[tokio::test]
815 async fn disallowed_model_override_is_rejected_before_provider_execution() {
816 let registry = Arc::new(ProviderRegistry::new());
817 let svc = Arc::new(Llm::new(
818 registry,
819 Arc::new(ClientPool::with_defaults()),
820 None,
821 ));
822 let root = Context::new_root();
823 root.provide_arc(svc.clone());
824 let tenant_ctx = root.intercept(TenantModelPolicy::new("tenant-a", ["gpt-4o".to_string()]));
825 let request = tenant_ctx.intercept(ModelOverride {
826 model: "not-allowed".into(),
827 });
828 let err = match svc
829 .get_client(&request, CapabilityRequirements::default())
830 .await
831 {
832 Ok(_) => panic!("disallowed override must fail before provider lookup"),
833 Err(err) => err,
834 };
835 assert!(matches!(err, AppError::Auth(_)));
836 assert!(err.to_string().contains("not-allowed"));
837 assert!(root.get::<ModelOverride>().is_none());
838 assert!(root.get::<TenantModelPolicy>().is_none());
839 assert!(matches!(
840 root.get::<Llm>().expect("global service").breaker(),
841 Breaker::Closed
842 ));
843 }
844
845 #[tokio::test]
846 async fn get_client_uses_override_when_catalog_absent() {
847 let registry = Arc::new(ProviderRegistry::new());
848 let pool = Arc::new(ClientPool::with_defaults());
849 let svc = Llm::new(registry, pool, None);
850 let ctx = Context::new_root();
851 let req_ctx = ctx.intercept(ModelOverride {
852 model: "nonexistent-model-xyz".into(),
853 });
854 let req = CapabilityRequirements::default();
855 let res = svc.get_client(&req_ctx, req).await;
857 assert!(res.is_err());
858 }
859
860 #[tokio::test]
861 async fn get_client_override_uses_tenant_context_intercept() {
862 let mut registry = ProviderRegistry::new();
863 registry.register_model(
864 "pinned-model",
865 crate::config::ModelConfig {
866 provider: "shared-runtime".into(),
867 model: "tenant-model".into(),
868 temperature: 0.7,
869 max_tokens: 512,
870 },
871 );
872 let global = RuntimeProviderEntry {
873 tenant_id: None,
874 display_name: "Global Shared".to_string(),
875 provider_type: "openai-compatible".to_string(),
876 api_base: "https://global.example.com/v1".to_string(),
877 auth_type: "api_key".to_string(),
878 default_model: Some("global-model".to_string()),
879 headers: HashMap::new(),
880 api_key: Some("global-key".to_string()),
881 enabled: true,
882 };
883 let tenant = RuntimeProviderEntry {
884 tenant_id: Some("tenant-a".to_string()),
885 display_name: "Tenant Shared".to_string(),
886 provider_type: "openai-compatible".to_string(),
887 api_base: "https://tenant.example.com/v1".to_string(),
888 auth_type: "api_key".to_string(),
889 default_model: Some("tenant-model".to_string()),
890 headers: HashMap::new(),
891 api_key: Some("tenant-key".to_string()),
892 enabled: true,
893 };
894 registry.reload_runtime_providers(
895 vec![global, tenant],
896 vec!["shared-runtime".to_string(), "shared-runtime".to_string()],
897 );
898 let registry = Arc::new(registry);
899 let svc = Llm::new(registry, Arc::new(ClientPool::with_defaults()), None);
900 let root = Context::new_root();
901 let ctx = root
902 .with_intercept(TenantContext::new("tenant-a".into(), TenantTier::Pro))
903 .intercept(ModelOverride {
904 model: "pinned-model".into(),
905 });
906 let tenant_client = svc
907 .get_client(&ctx, CapabilityRequirements::default())
908 .await;
909 assert!(
910 tenant_client.is_ok(),
911 "tenant intercept should construct a client from the tenant runtime entry: {:?}",
912 tenant_client.as_ref().err()
913 );
914
915 let unlabeled = root.intercept(ModelOverride {
916 model: "pinned-model".into(),
917 });
918 let fleet_client = svc
919 .get_client(&unlabeled, CapabilityRequirements::default())
920 .await;
921 assert!(
922 fleet_client.is_ok(),
923 "unlabeled root with ModelOverride should construct a client from the fleet global runtime entry: {:?}",
924 fleet_client.as_ref().err()
925 );
926 }
927
928 struct EchoClient {
929 generated: std::sync::Arc<std::sync::atomic::AtomicBool>,
930 }
931
932 impl EchoClient {
933 fn new() -> (Self, std::sync::Arc<std::sync::atomic::AtomicBool>) {
934 let generated = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
935 (
936 Self {
937 generated: std::sync::Arc::clone(&generated),
938 },
939 generated,
940 )
941 }
942 }
943
944 #[async_trait]
945 impl LLMClient for EchoClient {
946 async fn generate(&self, prompt: &str) -> ares_types::types::Result<String> {
947 self.generated
948 .store(true, std::sync::atomic::Ordering::SeqCst);
949 Ok(format!("echo:{prompt}"))
950 }
951 async fn generate_with_system(
952 &self,
953 _system: &str,
954 prompt: &str,
955 ) -> ares_types::types::Result<String> {
956 self.generate(prompt).await
957 }
958 async fn generate_with_history(
959 &self,
960 _messages: &[(String, String)],
961 ) -> ares_types::types::Result<LLMResponse> {
962 Ok(LLMResponse {
963 content: String::new(),
964 tool_calls: vec![],
965 finish_reason: "stop".into(),
966 usage: None,
967 })
968 }
969 async fn generate_with_tools(
970 &self,
971 _prompt: &str,
972 _tools: &[ToolDefinition],
973 ) -> ares_types::types::Result<LLMResponse> {
974 Ok(LLMResponse {
975 content: String::new(),
976 tool_calls: vec![],
977 finish_reason: "stop".into(),
978 usage: None,
979 })
980 }
981 async fn generate_with_tools_and_history(
982 &self,
983 _messages: &[crate::coordinator::ConversationMessage],
984 _tools: &[ToolDefinition],
985 ) -> ares_types::types::Result<LLMResponse> {
986 Ok(LLMResponse {
987 content: String::new(),
988 tool_calls: vec![],
989 finish_reason: "stop".into(),
990 usage: None,
991 })
992 }
993 async fn stream(
994 &self,
995 _prompt: &str,
996 ) -> ares_types::types::Result<
997 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
998 > {
999 Err(AppError::Internal("echo stream not implemented".into()))
1000 }
1001 async fn stream_with_system(
1002 &self,
1003 _system: &str,
1004 _prompt: &str,
1005 ) -> ares_types::types::Result<
1006 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1007 > {
1008 Err(AppError::Internal("echo stream not implemented".into()))
1009 }
1010 async fn stream_with_history(
1011 &self,
1012 _messages: &[(String, String)],
1013 ) -> ares_types::types::Result<
1014 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1015 > {
1016 Err(AppError::Internal("echo stream not implemented".into()))
1017 }
1018 fn model_name(&self) -> &str {
1019 "echo"
1020 }
1021 }
1022
1023 #[tokio::test]
1024 async fn llm_complete_runs_generate_without_events() {
1025 let (client, generated) = EchoClient::new();
1026 let llm = Llm::for_test(std::sync::Arc::new(client));
1027 let ctx = Context::new_root();
1028 let out = llm.complete(&ctx, "hi").await.expect("complete");
1029 assert_eq!(out, "echo:hi");
1030 assert!(generated.load(std::sync::atomic::Ordering::SeqCst));
1031 }
1032
1033 #[tokio::test]
1034 async fn llm_complete_waterfall_rewrites_prompt() {
1035 let (client, _) = EchoClient::new();
1036 let llm = Llm::for_test(std::sync::Arc::new(client));
1037 let ctx = Context::new_root();
1038 let events = ctx.provide(EventsService::new());
1039 events.on_waterfall(
1040 cordis::events_catalog::ev::LLM_COMPLETE.to_string(),
1041 |mut payload, next| async move {
1042 if let Some(p) = payload.get("prompt").and_then(|v| v.as_str()) {
1043 payload["prompt"] = serde_json::json!(format!("WRAP:{p}"));
1044 }
1045 next(payload).await
1046 },
1047 );
1048 let out = llm.complete(&ctx, "hi").await.expect("complete");
1049 assert_eq!(out, "echo:WRAP:hi");
1050 }
1051
1052 #[tokio::test]
1053 async fn llm_complete_short_circuit_skips_generate() {
1054 let (client, generated) = EchoClient::new();
1055 let llm = Llm::for_test(std::sync::Arc::new(client));
1056 let ctx = Context::new_root();
1057 let events = ctx.provide(EventsService::new());
1058 events.on_waterfall(
1059 cordis::events_catalog::ev::LLM_COMPLETE.to_string(),
1060 |_payload, _next| async move { Ok(serde_json::json!({ "content": "cached" })) },
1061 );
1062 let out = llm.complete(&ctx, "hi").await.expect("complete");
1063 assert_eq!(out, "cached");
1064 assert!(
1065 !generated.load(std::sync::atomic::Ordering::SeqCst),
1066 "dummy generate must stay false when handler skips next"
1067 );
1068 }
1069
1070 #[tokio::test]
1071 async fn llm_get_client_waterfall_deny() {
1072 let (client, _) = EchoClient::new();
1073 let llm = Llm::for_test(std::sync::Arc::new(client));
1074 let ctx = Context::new_root();
1075 let events = ctx.provide(EventsService::new());
1076 events.on_waterfall(
1077 cordis::events_catalog::ev::LLM_GET_CLIENT.to_string(),
1078 |_payload, _next| async move { Ok(serde_json::json!({ "deny": true })) },
1079 );
1080 let err = match llm
1081 .get_client(&ctx, CapabilityRequirements::default())
1082 .await
1083 {
1084 Ok(_) => panic!("deny"),
1085 Err(err) => err,
1086 };
1087 assert!(matches!(err, AppError::InvalidInput(msg) if msg == "llm.get_client denied"));
1088 }
1089
1090 #[test]
1091 fn llm_list_models_exposes_registry_models() {
1092 let mut registry = ProviderRegistry::new();
1093 registry.register_model(
1094 "stub-model",
1095 crate::config::ModelConfig {
1096 provider: "stub".into(),
1097 model: "stub-model".into(),
1098 temperature: 0.7,
1099 max_tokens: 512,
1100 },
1101 );
1102 let llm = Llm::new(
1103 Arc::new(registry),
1104 Arc::new(ClientPool::with_defaults()),
1105 None,
1106 );
1107 let models = llm.list_models();
1108 assert!(
1109 models
1110 .iter()
1111 .any(|m| m.name == "stub-model" && m.provider == "stub"),
1112 "Llm::list_models should expose registry models: {models:?}"
1113 );
1114 }
1115
1116 #[derive(Default)]
1118 struct HintRecordingClient {
1119 hints: parking_lot::Mutex<Vec<GenerationHints>>,
1120 supports: bool,
1121 }
1122
1123 #[async_trait]
1124 impl LLMClient for HintRecordingClient {
1125 async fn generate(&self, _prompt: &str) -> ares_types::types::Result<String> {
1126 Err(AppError::Internal("unused".into()))
1127 }
1128
1129 async fn generate_with_system(
1130 &self,
1131 _system: &str,
1132 _prompt: &str,
1133 ) -> ares_types::types::Result<String> {
1134 Err(AppError::Internal("unused".into()))
1135 }
1136
1137 async fn generate_with_history(
1138 &self,
1139 _messages: &[(String, String)],
1140 ) -> ares_types::types::Result<LLMResponse> {
1141 Err(AppError::Internal("unused".into()))
1142 }
1143
1144 async fn generate_with_tools(
1145 &self,
1146 _prompt: &str,
1147 _tools: &[ToolDefinition],
1148 ) -> ares_types::types::Result<LLMResponse> {
1149 Err(AppError::Internal("unused".into()))
1150 }
1151
1152 async fn generate_with_tools_and_history(
1153 &self,
1154 _messages: &[crate::coordinator::ConversationMessage],
1155 _tools: &[ToolDefinition],
1156 ) -> ares_types::types::Result<LLMResponse> {
1157 Err(AppError::Internal("unused".into()))
1158 }
1159
1160 async fn stream(
1161 &self,
1162 _prompt: &str,
1163 ) -> ares_types::types::Result<
1164 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1165 > {
1166 Err(AppError::Internal("unused".into()))
1167 }
1168
1169 async fn stream_with_system(
1170 &self,
1171 _system: &str,
1172 _prompt: &str,
1173 ) -> ares_types::types::Result<
1174 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1175 > {
1176 Err(AppError::Internal("unused".into()))
1177 }
1178
1179 async fn stream_with_history(
1180 &self,
1181 _messages: &[(String, String)],
1182 ) -> ares_types::types::Result<
1183 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1184 > {
1185 Err(AppError::Internal("unused".into()))
1186 }
1187
1188 fn model_name(&self) -> &str {
1189 "hint-recording-mock"
1190 }
1191
1192 fn supports_hints(&self) -> bool {
1193 self.supports
1194 }
1195
1196 fn set_hints(&self, hints: GenerationHints) {
1197 self.hints.lock().push(hints);
1198 }
1199 }
1200
1201 #[test]
1202 fn boxed_arc_client_forwards_hints_to_inner_client() {
1203 let concrete = Arc::new(HintRecordingClient {
1204 supports: true,
1205 hints: parking_lot::Mutex::new(Vec::new()),
1206 });
1207 let recorder_handle = Arc::clone(&concrete);
1208 let inner: Arc<dyn LLMClient> = concrete;
1209 let boxed = BoxedArcClient(Arc::clone(&inner));
1210
1211 assert!(boxed.supports_hints());
1212 boxed.set_hints(GenerationHints {
1213 json_mode: true,
1214 suppress_reasoning: false,
1215 max_tokens: Some(256),
1216 guided_grammar: None,
1217 });
1218 boxed.set_hints(GenerationHints::default());
1219
1220 let recorded = recorder_handle.hints.lock();
1224 assert_eq!(
1225 recorded.len(),
1226 2,
1227 "both set_hints calls must reach the inner client"
1228 );
1229 assert!(recorded[0].json_mode && recorded[0].max_tokens == Some(256));
1230 assert_eq!(recorded[1], GenerationHints::default());
1231 }
1232}