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