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, LlmStreamEvent};
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 async fn embed(
548 &self,
549 ctx: &Arc<Context>,
550 inputs: &[String],
551 ) -> Result<Vec<Vec<f32>>, AppError> {
552 let client = self
553 .get_client(ctx, CapabilityRequirements::default())
554 .await?;
555 let Some(events) = ctx.get::<EventsService>() else {
556 return client.embed(inputs).await;
557 };
558 let payload = serde_json::to_value(cordis::LlmEmbedRequest {
559 inputs: inputs.to_vec(),
560 })
561 .unwrap_or(serde_json::Value::Null);
562 let out = events
563 .waterfall_around(
564 cordis::events_catalog::ev::LLM_EMBED.to_string(),
565 payload,
566 move |payload| {
567 let client = Arc::clone(&client);
568 async move {
569 let inputs: Vec<String> = payload
570 .get("inputs")
571 .cloned()
572 .map(serde_json::from_value::<Vec<String>>)
573 .transpose()
574 .map_err(|e| CordisError::Fiber(e.to_string()))?
575 .unwrap_or_default();
576 let embeddings = client
577 .embed(&inputs)
578 .await
579 .map_err(|e| CordisError::Fiber(e.to_string()))?;
580 serde_json::to_value(cordis::LlmEmbedResponse { inputs, embeddings })
581 .map_err(|e| CordisError::Fiber(e.to_string()))
582 }
583 },
584 )
585 .await
586 .map_err(map_cordis)?;
587 Ok(match out.get("embeddings") {
588 Some(value) => serde_json::from_value::<Vec<Vec<f32>>>(value.clone())
589 .map_err(|e| AppError::Internal(e.to_string()))?,
590 None => Vec::new(),
591 })
592 }
593
594 pub fn find_model_stub(&self, _capability: &str) -> Option<String> {
596 None
597 }
598
599 pub fn list_models(&self) -> Vec<ModelInfo> {
601 self.provider_registry.list_models()
602 }
603
604 pub fn has_provider_for_tenant(&self, name: &str, tenant_id: Option<&str>) -> bool {
606 self.provider_registry
607 .has_provider_for_tenant(name, tenant_id)
608 }
609
610 pub fn get_provider_for_ctx(&self, ctx: &Arc<Context>, name: &str) -> Option<ProviderConfig> {
612 self.provider_registry.get_provider_for_ctx(ctx, name)
613 }
614
615 pub fn reload_runtime_providers(
617 &self,
618 providers: Vec<RuntimeProviderEntry>,
619 names: Vec<String>,
620 ) {
621 self.provider_registry
622 .reload_runtime_providers(providers, names);
623 }
624}
625
626fn map_cordis(err: CordisError) -> AppError {
627 AppError::Internal(err.to_string())
628}
629
630struct BoxedArcClient(Arc<dyn LLMClient>);
632
633#[async_trait]
634impl LLMClient for BoxedArcClient {
635 async fn generate(&self, prompt: &str) -> ares_types::types::Result<String> {
636 self.0.generate(prompt).await
637 }
638
639 async fn generate_with_system(
640 &self,
641 system: &str,
642 prompt: &str,
643 ) -> ares_types::types::Result<String> {
644 self.0.generate_with_system(system, prompt).await
645 }
646
647 async fn generate_with_history(
648 &self,
649 messages: &[(String, String)],
650 ) -> ares_types::types::Result<LLMResponse> {
651 self.0.generate_with_history(messages).await
652 }
653
654 async fn generate_with_tools(
655 &self,
656 prompt: &str,
657 tools: &[ToolDefinition],
658 ) -> ares_types::types::Result<LLMResponse> {
659 self.0.generate_with_tools(prompt, tools).await
660 }
661
662 async fn generate_with_tools_and_history(
663 &self,
664 messages: &[crate::coordinator::ConversationMessage],
665 tools: &[ToolDefinition],
666 ) -> ares_types::types::Result<LLMResponse> {
667 self.0
668 .generate_with_tools_and_history(messages, tools)
669 .await
670 }
671
672 async fn stream(
673 &self,
674 prompt: &str,
675 ) -> ares_types::types::Result<
676 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
677 > {
678 self.0.stream(prompt).await
679 }
680
681 async fn stream_with_system(
682 &self,
683 system: &str,
684 prompt: &str,
685 ) -> ares_types::types::Result<
686 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
687 > {
688 self.0.stream_with_system(system, prompt).await
689 }
690
691 async fn stream_with_history(
692 &self,
693 messages: &[(String, String)],
694 ) -> ares_types::types::Result<
695 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
696 > {
697 self.0.stream_with_history(messages).await
698 }
699
700 async fn stream_with_tools_and_history(
701 &self,
702 messages: &[crate::coordinator::ConversationMessage],
703 tools: &[ToolDefinition],
704 ) -> ares_types::types::Result<
705 Box<dyn futures::Stream<Item = ares_types::types::Result<LlmStreamEvent>> + Send + Unpin>,
706 > {
707 self.0
708 .stream_with_tools_and_history(messages, tools)
709 .await
710 }
711
712 fn model_name(&self) -> &str {
713 self.0.model_name()
714 }
715 fn supports_hints(&self) -> bool {
716 self.0.supports_hints()
717 }
718 fn set_hints(&self, hints: GenerationHints) {
719 self.0.set_hints(hints)
720 }
721
722 async fn embed(&self, inputs: &[String]) -> ares_types::types::Result<Vec<Vec<f32>>> {
723 self.0.embed(inputs).await
724 }
725
726 fn supports_vision(&self) -> bool {
727 self.0.supports_vision()
728 }
729
730 fn supports_provider_web_search(&self) -> bool {
731 self.0.supports_provider_web_search()
732 }
733}
734
735impl Service for Llm {
736 fn name(&self) -> &'static str {
737 "Llm"
738 }
739
740 fn init(
741 &self,
742 _ctx: &Arc<Context>,
743 ) -> std::pin::Pin<
744 Box<
745 dyn Future<Output = Result<Option<Box<dyn cordis::Disposable>>, CordisError>>
746 + Send
747 + '_,
748 >,
749 > {
750 Box::pin(async { Ok(None) })
751 }
752
753 fn check(&self) -> bool {
754 self.breaker.read().check()
757 }
758}
759
760#[cfg(test)]
761mod tests {
762 use super::*;
763 use crate::capabilities::CapabilityRequirements;
764 #[cfg(feature = "genai")]
765 use crate::provider_registry::RuntimeProviderEntry;
766 #[cfg(feature = "genai")]
767 use ares_types::models::{TenantContext, TenantTier};
768 use chrono::Duration;
769 use cordis::Context;
770 use std::collections::HashMap;
771
772 #[test]
773 fn breaker_closed_allows() {
774 assert!(Breaker::Closed.check());
775 }
776
777 #[test]
778 fn breaker_half_open_allows() {
779 assert!(Breaker::HalfOpen.check());
780 }
781
782 #[test]
783 fn breaker_open_future_denies() {
784 let until = Utc::now() + Duration::seconds(60);
785 assert!(!Breaker::Open { until }.check());
786 }
787
788 #[test]
789 fn breaker_open_past_allows() {
790 let until = Utc::now() - Duration::seconds(1);
791 assert!(Breaker::Open { until }.check());
792 }
793
794 #[test]
795 fn breaker_failure_threshold_opens() {
796 let b = Breaker::Closed;
797 let next = b.transition_on_failure_with_count(5);
798 assert!(matches!(next, Breaker::Open { .. }));
799 let still_closed = b.transition_on_failure_with_count(3);
800 assert!(matches!(still_closed, Breaker::Closed));
801 }
802
803 #[test]
804 fn breaker_constants_exist() {
805 assert_eq!(Breaker::FAILURE_THRESHOLD, 5);
806 assert_eq!(Breaker::COOLDOWN_SECS, 30);
807 }
808
809 #[test]
810 fn provider_registry_and_factory_accessors() {
811 let registry = Arc::new(ProviderRegistry::new());
812 let pool = Arc::new(ClientPool::with_defaults());
813 let factory = Arc::new(
814 ConfigBasedLLMFactory::from_config(HashMap::new(), HashMap::new(), None)
815 .expect("empty factory config"),
816 );
817 let llm = Llm::new(Arc::clone(®istry), pool, None).with_factory(Arc::clone(&factory));
818 assert!(Arc::ptr_eq(&llm.provider_registry(), ®istry));
819 }
820
821 #[tokio::test]
822 async fn llm_model_override_via_context() {
823 let registry = Arc::new(ProviderRegistry::new());
824 let pool = Arc::new(ClientPool::with_defaults());
825 let svc = Arc::new(Llm::new(registry, pool, None));
826 let root = Context::new_root();
827 root.provide::<Llm>(Llm::new(
828 Arc::new(ProviderRegistry::new()),
829 Arc::new(ClientPool::with_defaults()),
830 None,
831 ));
832 let req_ctx = root.intercept(ModelOverride {
834 model: "gpt-4o-mini".into(),
835 });
836 assert!(req_ctx.get::<ModelOverride>().is_some());
837 assert_eq!(req_ctx.get::<ModelOverride>().unwrap().model, "gpt-4o-mini");
838 let _ = Arc::clone(&svc.provider_registry);
840 let _ = svc.catalog.clone();
841 let _ = svc.pool.provider_names();
842 assert!(svc.check());
844 }
845
846 #[tokio::test]
847 async fn record_failure_threshold_opens_breaker() {
848 let svc = Llm::new(
849 Arc::new(ProviderRegistry::new()),
850 Arc::new(ClientPool::with_defaults()),
851 None,
852 );
853 for _ in 0..Breaker::FAILURE_THRESHOLD {
854 svc.record_failure();
855 }
856 assert!(!svc.check());
858 svc.record_success();
859 assert!(svc.check());
860 }
861
862 #[test]
863 fn tenant_model_policy_allows_and_composes_with_model_override() {
864 let root = Context::new_root();
865 let tenant_ctx = root.intercept(TenantModelPolicy::new(
866 "tenant-a",
867 ["gpt-4o-mini".to_string()],
868 ));
869 let request = tenant_ctx.intercept(ModelOverride {
870 model: "gpt-4o-mini".into(),
871 });
872 let policy = request
873 .get::<TenantModelPolicy>()
874 .expect("policy should be inherited by request context");
875 let override_model = request
876 .get::<ModelOverride>()
877 .expect("model override should be visible in request context");
878 policy
879 .authorize(&override_model.model)
880 .expect("allowed model override should pass policy");
881 let svc = Llm::new(
882 Arc::new(ProviderRegistry::new()),
883 Arc::new(ClientPool::with_defaults()),
884 None,
885 );
886 svc.validate_model_override(&request)
887 .expect("allowed model override should pass LLM validation");
888 assert!(root.get::<ModelOverride>().is_none());
889 assert!(root.get::<TenantModelPolicy>().is_none());
890 }
891
892 #[tokio::test]
893 async fn disallowed_model_override_is_rejected_before_provider_execution() {
894 let registry = Arc::new(ProviderRegistry::new());
895 let svc = Arc::new(Llm::new(
896 registry,
897 Arc::new(ClientPool::with_defaults()),
898 None,
899 ));
900 let root = Context::new_root();
901 root.provide_arc(svc.clone());
902 let tenant_ctx = root.intercept(TenantModelPolicy::new("tenant-a", ["gpt-4o".to_string()]));
903 let request = tenant_ctx.intercept(ModelOverride {
904 model: "not-allowed".into(),
905 });
906 let err = match svc
907 .get_client(&request, CapabilityRequirements::default())
908 .await
909 {
910 Ok(_) => panic!("disallowed override must fail before provider lookup"),
911 Err(err) => err,
912 };
913 assert!(matches!(err, AppError::Auth(_)));
914 assert!(err.to_string().contains("not-allowed"));
915 assert!(root.get::<ModelOverride>().is_none());
916 assert!(root.get::<TenantModelPolicy>().is_none());
917 assert!(matches!(
918 root.get::<Llm>().expect("global service").breaker(),
919 Breaker::Closed
920 ));
921 }
922
923 #[tokio::test]
924 async fn get_client_uses_override_when_catalog_absent() {
925 let registry = Arc::new(ProviderRegistry::new());
926 let pool = Arc::new(ClientPool::with_defaults());
927 let svc = Llm::new(registry, pool, None);
928 let ctx = Context::new_root();
929 let req_ctx = ctx.intercept(ModelOverride {
930 model: "nonexistent-model-xyz".into(),
931 });
932 let req = CapabilityRequirements::default();
933 let res = svc.get_client(&req_ctx, req).await;
935 assert!(res.is_err());
936 }
937
938 #[cfg(feature = "genai")]
939 #[tokio::test]
940 async fn get_client_override_uses_tenant_context_intercept() {
941 let mut registry = ProviderRegistry::new();
942 registry.register_model(
943 "pinned-model",
944 crate::config::ModelConfig {
945 provider: "shared-runtime".into(),
946 model: "tenant-model".into(),
947 temperature: 0.7,
948 max_tokens: 512,
949 },
950 );
951 let global = RuntimeProviderEntry {
952 tenant_id: None,
953 display_name: "Global Shared".to_string(),
954 provider_type: "openai-compatible".to_string(),
955 api_base: "https://global.example.com/v1".to_string(),
956 auth_type: "api_key".to_string(),
957 default_model: Some("global-model".to_string()),
958 headers: HashMap::new(),
959 api_key: Some("global-key".to_string()),
960 enabled: true,
961 };
962 let tenant = RuntimeProviderEntry {
963 tenant_id: Some("tenant-a".to_string()),
964 display_name: "Tenant Shared".to_string(),
965 provider_type: "openai-compatible".to_string(),
966 api_base: "https://tenant.example.com/v1".to_string(),
967 auth_type: "api_key".to_string(),
968 default_model: Some("tenant-model".to_string()),
969 headers: HashMap::new(),
970 api_key: Some("tenant-key".to_string()),
971 enabled: true,
972 };
973 registry.reload_runtime_providers(
974 vec![global, tenant],
975 vec!["shared-runtime".to_string(), "shared-runtime".to_string()],
976 );
977 let registry = Arc::new(registry);
978 let svc = Llm::new(registry, Arc::new(ClientPool::with_defaults()), None);
979 let root = Context::new_root();
980 let ctx = root
981 .with_intercept(TenantContext::new("tenant-a".into(), TenantTier::Pro))
982 .intercept(ModelOverride {
983 model: "pinned-model".into(),
984 });
985 let tenant_client = svc
986 .get_client(&ctx, CapabilityRequirements::default())
987 .await;
988 assert!(
989 tenant_client.is_ok(),
990 "tenant intercept should construct a client from the tenant runtime entry: {:?}",
991 tenant_client.as_ref().err()
992 );
993
994 let unlabeled = root.intercept(ModelOverride {
995 model: "pinned-model".into(),
996 });
997 let fleet_client = svc
998 .get_client(&unlabeled, CapabilityRequirements::default())
999 .await;
1000 assert!(
1001 fleet_client.is_ok(),
1002 "unlabeled root with ModelOverride should construct a client from the fleet global runtime entry: {:?}",
1003 fleet_client.as_ref().err()
1004 );
1005 }
1006
1007 struct EchoClient {
1008 generated: std::sync::Arc<std::sync::atomic::AtomicBool>,
1009 }
1010
1011 impl EchoClient {
1012 fn new() -> (Self, std::sync::Arc<std::sync::atomic::AtomicBool>) {
1013 let generated = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
1014 (
1015 Self {
1016 generated: std::sync::Arc::clone(&generated),
1017 },
1018 generated,
1019 )
1020 }
1021 }
1022
1023 #[async_trait]
1024 impl LLMClient for EchoClient {
1025 async fn generate(&self, prompt: &str) -> ares_types::types::Result<String> {
1026 self.generated
1027 .store(true, std::sync::atomic::Ordering::SeqCst);
1028 Ok(format!("echo:{prompt}"))
1029 }
1030 async fn generate_with_system(
1031 &self,
1032 _system: &str,
1033 prompt: &str,
1034 ) -> ares_types::types::Result<String> {
1035 self.generate(prompt).await
1036 }
1037 async fn generate_with_history(
1038 &self,
1039 _messages: &[(String, String)],
1040 ) -> ares_types::types::Result<LLMResponse> {
1041 Ok(LLMResponse {
1042 content: String::new(),
1043 tool_calls: vec![],
1044 finish_reason: "stop".into(),
1045 usage: None,
1046 reasoning_content: None,
1047 response_id: None,
1048 })
1049 }
1050 async fn generate_with_tools(
1051 &self,
1052 _prompt: &str,
1053 _tools: &[ToolDefinition],
1054 ) -> ares_types::types::Result<LLMResponse> {
1055 Ok(LLMResponse {
1056 content: String::new(),
1057 tool_calls: vec![],
1058 finish_reason: "stop".into(),
1059 usage: None,
1060 reasoning_content: None,
1061 response_id: None,
1062 })
1063 }
1064 async fn generate_with_tools_and_history(
1065 &self,
1066 _messages: &[crate::coordinator::ConversationMessage],
1067 _tools: &[ToolDefinition],
1068 ) -> ares_types::types::Result<LLMResponse> {
1069 Ok(LLMResponse {
1070 content: String::new(),
1071 tool_calls: vec![],
1072 finish_reason: "stop".into(),
1073 usage: None,
1074 reasoning_content: None,
1075 response_id: None,
1076 })
1077 }
1078 async fn stream(
1079 &self,
1080 _prompt: &str,
1081 ) -> ares_types::types::Result<
1082 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1083 > {
1084 Err(AppError::Internal("echo stream not implemented".into()))
1085 }
1086 async fn stream_with_system(
1087 &self,
1088 _system: &str,
1089 _prompt: &str,
1090 ) -> ares_types::types::Result<
1091 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1092 > {
1093 Err(AppError::Internal("echo stream not implemented".into()))
1094 }
1095 async fn stream_with_history(
1096 &self,
1097 _messages: &[(String, String)],
1098 ) -> ares_types::types::Result<
1099 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1100 > {
1101 Err(AppError::Internal("echo stream not implemented".into()))
1102 }
1103 fn model_name(&self) -> &str {
1104 "echo"
1105 }
1106 }
1107
1108 struct EmbedClient {
1109 called: std::sync::Arc<std::sync::atomic::AtomicBool>,
1110 }
1111
1112 impl EmbedClient {
1113 fn new() -> (Self, std::sync::Arc<std::sync::atomic::AtomicBool>) {
1114 let called = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
1115 (
1116 Self {
1117 called: std::sync::Arc::clone(&called),
1118 },
1119 called,
1120 )
1121 }
1122 }
1123
1124 #[async_trait]
1125 impl LLMClient for EmbedClient {
1126 async fn generate(&self, _prompt: &str) -> ares_types::types::Result<String> {
1127 Err(AppError::Internal("embed-only mock".into()))
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("embed-only mock".into()))
1135 }
1136 async fn generate_with_history(
1137 &self,
1138 _messages: &[(String, String)],
1139 ) -> ares_types::types::Result<LLMResponse> {
1140 Err(AppError::Internal("embed-only mock".into()))
1141 }
1142 async fn generate_with_tools(
1143 &self,
1144 _prompt: &str,
1145 _tools: &[ToolDefinition],
1146 ) -> ares_types::types::Result<LLMResponse> {
1147 Err(AppError::Internal("embed-only mock".into()))
1148 }
1149 async fn generate_with_tools_and_history(
1150 &self,
1151 _messages: &[crate::coordinator::ConversationMessage],
1152 _tools: &[ToolDefinition],
1153 ) -> ares_types::types::Result<LLMResponse> {
1154 Err(AppError::Internal("embed-only mock".into()))
1155 }
1156 async fn stream(
1157 &self,
1158 _prompt: &str,
1159 ) -> ares_types::types::Result<
1160 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1161 > {
1162 Err(AppError::Internal("embed-only mock".into()))
1163 }
1164 async fn stream_with_system(
1165 &self,
1166 _system: &str,
1167 _prompt: &str,
1168 ) -> ares_types::types::Result<
1169 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1170 > {
1171 Err(AppError::Internal("embed-only mock".into()))
1172 }
1173 async fn stream_with_history(
1174 &self,
1175 _messages: &[(String, String)],
1176 ) -> ares_types::types::Result<
1177 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1178 > {
1179 Err(AppError::Internal("embed-only mock".into()))
1180 }
1181 fn model_name(&self) -> &str {
1182 "embed-mock"
1183 }
1184 async fn embed(&self, inputs: &[String]) -> ares_types::types::Result<Vec<Vec<f32>>> {
1185 self.called.store(true, std::sync::atomic::Ordering::SeqCst);
1186 Ok(inputs.iter().map(|s| vec![s.len() as f32]).collect())
1187 }
1188 }
1189
1190 #[tokio::test]
1191 async fn llm_complete_runs_generate_without_events() {
1192 let (client, generated) = EchoClient::new();
1193 let llm = Llm::for_test(std::sync::Arc::new(client));
1194 let ctx = Context::new_root();
1195 let out = llm.complete(&ctx, "hi").await.expect("complete");
1196 assert_eq!(out, "echo:hi");
1197 assert!(generated.load(std::sync::atomic::Ordering::SeqCst));
1198 }
1199
1200 #[tokio::test]
1201 async fn llm_complete_waterfall_rewrites_prompt() {
1202 let (client, _) = EchoClient::new();
1203 let llm = Llm::for_test(std::sync::Arc::new(client));
1204 let ctx = Context::new_root();
1205 let events = ctx.provide(EventsService::new());
1206 events.on_waterfall(
1207 cordis::events_catalog::ev::LLM_COMPLETE.to_string(),
1208 |mut payload, next| async move {
1209 if let Some(p) = payload.get("prompt").and_then(|v| v.as_str()) {
1210 payload["prompt"] = serde_json::json!(format!("WRAP:{p}"));
1211 }
1212 next(payload).await
1213 },
1214 );
1215 let out = llm.complete(&ctx, "hi").await.expect("complete");
1216 assert_eq!(out, "echo:WRAP:hi");
1217 }
1218
1219 #[tokio::test]
1220 async fn llm_complete_short_circuit_skips_generate() {
1221 let (client, generated) = EchoClient::new();
1222 let llm = Llm::for_test(std::sync::Arc::new(client));
1223 let ctx = Context::new_root();
1224 let events = ctx.provide(EventsService::new());
1225 events.on_waterfall(
1226 cordis::events_catalog::ev::LLM_COMPLETE.to_string(),
1227 |_payload, _next| async move { Ok(serde_json::json!({ "content": "cached" })) },
1228 );
1229 let out = llm.complete(&ctx, "hi").await.expect("complete");
1230 assert_eq!(out, "cached");
1231 assert!(
1232 !generated.load(std::sync::atomic::Ordering::SeqCst),
1233 "dummy generate must stay false when handler skips next"
1234 );
1235 }
1236
1237 #[tokio::test]
1238 async fn llm_get_client_waterfall_deny() {
1239 let (client, _) = EchoClient::new();
1240 let llm = Llm::for_test(std::sync::Arc::new(client));
1241 let ctx = Context::new_root();
1242 let events = ctx.provide(EventsService::new());
1243 events.on_waterfall(
1244 cordis::events_catalog::ev::LLM_GET_CLIENT.to_string(),
1245 |_payload, _next| async move { Ok(serde_json::json!({ "deny": true })) },
1246 );
1247 let err = match llm
1248 .get_client(&ctx, CapabilityRequirements::default())
1249 .await
1250 {
1251 Ok(_) => panic!("deny"),
1252 Err(err) => err,
1253 };
1254 assert!(matches!(err, AppError::InvalidInput(msg) if msg == "llm.get_client denied"));
1255 }
1256
1257 #[tokio::test]
1258 async fn llm_embed_runs_without_events() {
1259 let (client, called) = EmbedClient::new();
1260 let llm = Llm::for_test(std::sync::Arc::new(client));
1261 let ctx = std::sync::Arc::new(Context::new_root());
1262 let out = llm
1263 .embed(&ctx, &["ab".into(), "c".into()])
1264 .await
1265 .expect("embed");
1266 assert_eq!(out, vec![vec![2.0], vec![1.0]]);
1267 assert!(called.load(std::sync::atomic::Ordering::SeqCst));
1268 }
1269
1270 #[tokio::test]
1271 async fn llm_embed_waterfall_rewrites_inputs() {
1272 let (client, _) = EmbedClient::new();
1273 let llm = Llm::for_test(std::sync::Arc::new(client));
1274 let ctx = Context::new_root();
1275 let events = ctx.provide(EventsService::new());
1276 events.on_waterfall(
1277 cordis::events_catalog::ev::LLM_EMBED.to_string(),
1278 |mut payload, next| async move {
1279 if let Some(inputs) = payload.get("inputs").and_then(|v| v.as_array()) {
1280 let rewritten: Vec<String> = inputs
1281 .iter()
1282 .filter_map(|v| v.as_str().map(|s| format!("WRAP:{s}")))
1283 .collect();
1284 payload["inputs"] = serde_json::json!(rewritten);
1285 }
1286 next(payload).await
1287 },
1288 );
1289 let out = llm.embed(&ctx, &["hi".into()]).await.expect("embed");
1290 assert_eq!(out, vec![vec![7.0]]);
1291 }
1292
1293 #[tokio::test]
1294 async fn llm_embed_short_circuit_skips_client() {
1295 let (client, called) = EmbedClient::new();
1296 let llm = Llm::for_test(std::sync::Arc::new(client));
1297 let ctx = std::sync::Arc::new(Context::new_root());
1298 let events = ctx.provide(EventsService::new());
1299 events.on_waterfall(
1300 cordis::events_catalog::ev::LLM_EMBED.to_string(),
1301 |_payload, _next| async move { Ok(serde_json::json!({ "embeddings": [[9.0, 8.0]] })) },
1302 );
1303 let out = llm.embed(&ctx, &["hi".into()]).await.expect("embed");
1304 assert_eq!(out, vec![vec![9.0, 8.0]]);
1305 assert!(
1306 !called.load(std::sync::atomic::Ordering::SeqCst),
1307 "client embed must stay false when handler skips next"
1308 );
1309 }
1310
1311 #[test]
1312 fn llm_list_models_exposes_registry_models() {
1313 let mut registry = ProviderRegistry::new();
1314 registry.register_model(
1315 "stub-model",
1316 crate::config::ModelConfig {
1317 provider: "stub".into(),
1318 model: "stub-model".into(),
1319 temperature: 0.7,
1320 max_tokens: 512,
1321 },
1322 );
1323 let llm = Llm::new(
1324 Arc::new(registry),
1325 Arc::new(ClientPool::with_defaults()),
1326 None,
1327 );
1328 let models = llm.list_models();
1329 assert!(
1330 models
1331 .iter()
1332 .any(|m| m.name == "stub-model" && m.provider == "stub"),
1333 "Llm::list_models should expose registry models: {models:?}"
1334 );
1335 }
1336
1337 #[derive(Default)]
1339 struct HintRecordingClient {
1340 hints: parking_lot::Mutex<Vec<GenerationHints>>,
1341 supports: bool,
1342 }
1343
1344 #[async_trait]
1345 impl LLMClient for HintRecordingClient {
1346 async fn generate(&self, _prompt: &str) -> ares_types::types::Result<String> {
1347 Err(AppError::Internal("unused".into()))
1348 }
1349
1350 async fn generate_with_system(
1351 &self,
1352 _system: &str,
1353 _prompt: &str,
1354 ) -> ares_types::types::Result<String> {
1355 Err(AppError::Internal("unused".into()))
1356 }
1357
1358 async fn generate_with_history(
1359 &self,
1360 _messages: &[(String, String)],
1361 ) -> ares_types::types::Result<LLMResponse> {
1362 Err(AppError::Internal("unused".into()))
1363 }
1364
1365 async fn generate_with_tools(
1366 &self,
1367 _prompt: &str,
1368 _tools: &[ToolDefinition],
1369 ) -> ares_types::types::Result<LLMResponse> {
1370 Err(AppError::Internal("unused".into()))
1371 }
1372
1373 async fn generate_with_tools_and_history(
1374 &self,
1375 _messages: &[crate::coordinator::ConversationMessage],
1376 _tools: &[ToolDefinition],
1377 ) -> ares_types::types::Result<LLMResponse> {
1378 Err(AppError::Internal("unused".into()))
1379 }
1380
1381 async fn stream(
1382 &self,
1383 _prompt: &str,
1384 ) -> ares_types::types::Result<
1385 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1386 > {
1387 Err(AppError::Internal("unused".into()))
1388 }
1389
1390 async fn stream_with_system(
1391 &self,
1392 _system: &str,
1393 _prompt: &str,
1394 ) -> ares_types::types::Result<
1395 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1396 > {
1397 Err(AppError::Internal("unused".into()))
1398 }
1399
1400 async fn stream_with_history(
1401 &self,
1402 _messages: &[(String, String)],
1403 ) -> ares_types::types::Result<
1404 Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1405 > {
1406 Err(AppError::Internal("unused".into()))
1407 }
1408
1409 fn model_name(&self) -> &str {
1410 "hint-recording-mock"
1411 }
1412
1413 fn supports_hints(&self) -> bool {
1414 self.supports
1415 }
1416
1417 fn set_hints(&self, hints: GenerationHints) {
1418 self.hints.lock().push(hints);
1419 }
1420 }
1421
1422 #[test]
1423 fn boxed_arc_client_forwards_hints_to_inner_client() {
1424 let concrete = Arc::new(HintRecordingClient {
1425 supports: true,
1426 hints: parking_lot::Mutex::new(Vec::new()),
1427 });
1428 let recorder_handle = Arc::clone(&concrete);
1429 let inner: Arc<dyn LLMClient> = concrete;
1430 let boxed = BoxedArcClient(Arc::clone(&inner));
1431
1432 assert!(boxed.supports_hints());
1433 boxed.set_hints(GenerationHints {
1434 json_mode: true,
1435 suppress_reasoning: false,
1436 max_tokens: Some(256),
1437 guided_grammar: None,
1438 ..Default::default()
1439 });
1440 boxed.set_hints(GenerationHints::default());
1441
1442 let recorded = recorder_handle.hints.lock();
1446 assert_eq!(
1447 recorded.len(),
1448 2,
1449 "both set_hints calls must reach the inner client"
1450 );
1451 assert!(recorded[0].json_mode && recorded[0].max_tokens == Some(256));
1452 assert_eq!(recorded[1], GenerationHints::default());
1453 }
1454}