1use crate::capabilities::{CapabilityRequirements, ModelCapabilities, ModelWithCapabilities};
24use crate::client::{LLMClient, ModelParams, Provider};
25use crate::config::{ModelConfig, ProviderConfig};
26use crate::nvidia_catalog::{NvidiaCatalogCache, NvidiaConfig};
27use arc_swap::ArcSwap;
28use ares_types::types::{AppError, Result};
29use std::any::TypeId;
30use std::collections::HashMap;
31use std::sync::Arc;
32pub use cordis::ReflectService;
34
35#[derive(Debug, Clone)]
37pub struct RuntimeProviderEntry {
38 pub tenant_id: Option<String>,
40 pub display_name: String,
42 pub provider_type: String,
44 pub api_base: String,
46 pub auth_type: String,
48 pub default_model: Option<String>,
50 pub headers: HashMap<String, String>,
52 pub api_key: Option<String>,
54 pub enabled: bool,
56}
57
58#[derive(Debug, Clone)]
60pub struct ResolvedProviderConfig {
61 pub provider_name: String,
62 pub model_name: String,
63 pub provider_config: ProviderConfig,
64 pub params: ModelParams,
65 pub tenant_id: Option<String>,
66}
67
68pub struct ProviderRegistry {
73 providers: HashMap<String, ProviderConfig>,
75 models: HashMap<String, ModelConfig>,
77 catalog: Option<Arc<NvidiaCatalogCache>>,
79 default_model: Option<String>,
81 runtime_providers: Arc<ArcSwap<HashMap<String, Vec<RuntimeProviderEntry>>>>,
83}
84
85impl ProviderRegistry {
86 pub fn new() -> Self {
88 Self {
89 providers: HashMap::new(),
90 models: HashMap::new(),
91 catalog: None,
92 default_model: None,
93 runtime_providers: Arc::new(ArcSwap::from_pointee(HashMap::new())),
94 }
95 }
96
97 pub fn from_config(
99 providers: std::collections::HashMap<String, ProviderConfig>,
100 models: std::collections::HashMap<String, ModelConfig>,
101 nvidia: Option<&NvidiaConfig>,
102 ) -> Self {
103 let mut providers = providers;
104
105 if providers.is_empty() {
107 let nvidia = nvidia.cloned().unwrap_or_default();
108 let _ = std::env::var(&nvidia.api_key_env); providers.insert(
110 "nvidia".to_string(),
111 ProviderConfig::OpenAI {
112 api_key_env: nvidia.api_key_env.clone(),
113 api_base: nvidia.api_base.clone(),
114 default_model: nvidia.default_model.clone(),
115 },
116 );
117 }
118
119 #[cfg(feature = "bedrock")]
120 providers
121 .entry("bedrock".to_string())
122 .or_insert_with(Self::default_bedrock_provider_config);
123
124 #[cfg(feature = "azure")]
125 providers
126 .entry("azure".to_string())
127 .or_insert_with(Self::default_azure_provider_config);
128
129 let default_model = nvidia
130 .map(|n| n.default_model.clone())
131 .or_else(|| models.keys().next().cloned());
132
133 Self {
134 providers,
135 models,
136 catalog: None,
137 default_model,
138 runtime_providers: Arc::new(ArcSwap::from_pointee(HashMap::new())),
139 }
140 }
141
142 pub fn reload_runtime_providers(
145 &self,
146 providers: Vec<RuntimeProviderEntry>,
147 names: Vec<String>,
148 ) {
149 let mut map = HashMap::new();
150 for (entry, name) in providers.into_iter().zip(names) {
151 if entry.enabled {
152 map.entry(name).or_insert_with(Vec::new).push(entry);
153 }
154 }
155 self.runtime_providers.store(Arc::new(map));
156 }
157
158 pub fn with_catalog(mut self, catalog: Arc<NvidiaCatalogCache>) -> Self {
160 self.catalog = Some(catalog);
161 self
162 }
163
164 pub fn set_default_model(&mut self, model_name: &str) {
166 self.default_model = Some(model_name.to_string());
167 }
168
169 pub fn register_provider(&mut self, name: &str, config: ProviderConfig) {
171 self.providers.insert(name.to_string(), config);
172 }
173
174 pub fn register_model(&mut self, name: &str, config: ModelConfig) {
176 self.models.insert(name.to_string(), config);
177 }
178
179 pub fn unregister_provider(&mut self, name: &str) -> Option<ProviderConfig> {
181 self.providers.remove(name)
182 }
183
184 pub fn unregister_model(&mut self, name: &str) -> Option<ModelConfig> {
186 self.models.remove(name)
187 }
188
189 pub fn get_provider(&self, name: &str) -> Option<ProviderConfig> {
197 self.provider_for_tenant(name, None)
198 }
199
200 pub(crate) fn provider_for_tenant(
206 &self,
207 name: &str,
208 tenant_id: Option<&str>,
209 ) -> Option<ProviderConfig> {
210 if let Some(entry) = self.runtime_provider_entry_for_tenant(name, tenant_id) {
211 return Some(Self::synthesize_provider_config(&entry));
212 }
213 self.providers.get(name).cloned()
214 }
215
216 pub fn get_provider_for_ctx(
222 &self,
223 ctx: &std::sync::Arc<cordis::Context>,
224 name: &str,
225 ) -> Option<ProviderConfig> {
226 let tenant = tenant_from_ctx(ctx);
227 self.provider_for_tenant(name, tenant.as_deref())
228 }
229
230 fn runtime_provider_entry_for_tenant(
231 &self,
232 name: &str,
233 tenant_id: Option<&str>,
234 ) -> Option<RuntimeProviderEntry> {
235 let runtime = self.runtime_providers.load();
236 let entries = runtime.get(name)?;
237 if let Some(requester) = tenant_id {
238 if let Some(entry) = entries
239 .iter()
240 .find(|entry| entry.tenant_id.as_deref() == Some(requester))
241 {
242 return Some(entry.clone());
243 }
244 }
245 entries
246 .iter()
247 .find(|entry| entry.tenant_id.is_none())
248 .cloned()
249 }
250
251 fn synthesize_provider_config(entry: &RuntimeProviderEntry) -> ProviderConfig {
253 match entry.provider_type.as_str() {
254 "anthropic-compatible" => ProviderConfig::Anthropic {
255 api_key_env: entry
256 .api_key
257 .clone()
258 .unwrap_or_else(|| "ANTHROPIC_API_KEY".to_string()),
259 default_model: entry.default_model.clone().unwrap_or_default(),
260 },
261 "bedrock" | "bedrock-compatible" => ProviderConfig::Bedrock {
262 api_key_env: "AWS_BEARER_TOKEN_BEDROCK".to_string(),
263 region_env: entry
264 .headers
265 .get("region_env")
266 .cloned()
267 .unwrap_or_else(|| "AWS_REGION".to_string()),
268 default_model: entry.default_model.clone().unwrap_or_default(),
269 },
270 "azure" | "azure-compatible" => ProviderConfig::Azure {
271 api_key_env: "AZURE_FOUNDRY_API_KEY".to_string(),
272 base_url_env: "AZURE_FOUNDRY_BASE_URL".to_string(),
273 default_model: entry.default_model.clone().unwrap_or_default(),
274 },
275 _ => ProviderConfig::OpenAI {
276 api_key_env: entry
277 .api_key
278 .clone()
279 .unwrap_or_else(|| "OPENAI_API_KEY".to_string()),
280 api_base: entry.api_base.clone(),
281 default_model: entry.default_model.clone().unwrap_or_default(),
282 },
283 }
284 }
285
286 #[allow(dead_code)]
287 fn runtime_api_key(provider_name: &str, entry: &RuntimeProviderEntry) -> Result<String> {
288 entry
289 .api_key
290 .as_ref()
291 .filter(|api_key| !api_key.is_empty())
292 .cloned()
293 .ok_or_else(|| {
294 AppError::Configuration(format!(
295 "Runtime provider '{}' API key is not resolved",
296 provider_name
297 ))
298 })
299 }
300
301 fn provider_from_runtime_entry(
302 provider_name: &str,
303 entry: &RuntimeProviderEntry,
304 ) -> Result<Provider> {
305 Self::provider_from_runtime_entry_with_params(
306 provider_name,
307 entry,
308 entry.default_model.as_deref(),
309 ModelParams::default(),
310 )
311 }
312
313 #[allow(dead_code)]
314 fn provider_default_model(config: &ProviderConfig) -> &str {
315 match config {
316 ProviderConfig::OpenAI { default_model, .. }
317 | ProviderConfig::Anthropic { default_model, .. }
318 | ProviderConfig::Azure { default_model, .. }
319 | ProviderConfig::Bedrock { default_model, .. }
320 | ProviderConfig::Ollama { default_model, .. } => default_model,
321 }
322 }
323
324 fn default_bedrock_provider_config() -> ProviderConfig {
325 ProviderConfig::Bedrock {
326 api_key_env: "AWS_BEARER_TOKEN_BEDROCK".to_string(),
327 region_env: "AWS_REGION".to_string(),
328 default_model: "us.anthropic.claude-haiku-4-5-20251001-v1:0".to_string(),
329 }
330 }
331
332 fn default_azure_provider_config() -> ProviderConfig {
333 ProviderConfig::Azure {
334 api_key_env: "AZURE_FOUNDRY_API_KEY".to_string(),
335 base_url_env: "AZURE_FOUNDRY_BASE_URL".to_string(),
336 default_model: "DeepSeek-V4-Flash".to_string(),
337 }
338 }
339
340 fn bedrock_model_id_from_name(model_name: &str) -> Option<&str> {
341 let trimmed = model_name.trim();
342 if let Some(model_id) = trimmed.strip_prefix("bedrock/") {
343 return (!model_id.trim().is_empty()).then_some(model_id.trim());
344 }
345 if trimmed.starts_with("us.anthropic.") || trimmed.starts_with("anthropic.claude") {
346 return Some(trimmed);
347 }
348 None
349 }
350
351 fn azure_model_id_from_name(model_name: &str) -> Option<&str> {
352 let trimmed = model_name.trim();
353 if let Some(model_id) = trimmed.strip_prefix("azure/") {
354 return (!model_id.trim().is_empty()).then_some(model_id.trim());
355 }
356 None
357 }
358
359 fn bedrock_model_config(model_id: &str) -> ModelConfig {
360 ModelConfig {
361 provider: "bedrock".to_string(),
362 model: model_id.to_string(),
363 temperature: 0.7,
364 max_tokens: 4096,
365 }
366 }
367
368 fn azure_model_config(model_id: &str) -> ModelConfig {
369 ModelConfig {
370 provider: "azure".to_string(),
371 model: model_id.to_string(),
372 temperature: 0.7,
373 max_tokens: 4096,
374 }
375 }
376
377 #[allow(dead_code)]
378 fn runtime_bedrock_region(provider_name: &str, entry: &RuntimeProviderEntry) -> Result<String> {
379 entry
380 .headers
381 .get("region")
382 .cloned()
383 .or_else(|| {
384 entry
385 .headers
386 .get("region_env")
387 .and_then(|env| std::env::var(env).ok())
388 })
389 .or_else(|| std::env::var("AWS_REGION").ok())
390 .filter(|region| !region.is_empty())
391 .ok_or_else(|| {
392 AppError::Configuration(format!(
393 "Runtime Bedrock provider '{}' must define headers.region or AWS_REGION",
394 provider_name
395 ))
396 })
397 }
398
399 #[allow(unused_variables)]
400 fn provider_from_runtime_entry_with_params(
401 provider_name: &str,
402 entry: &RuntimeProviderEntry,
403 model_override: Option<&str>,
404 params: ModelParams,
405 ) -> Result<Provider> {
406 let model = model_override
407 .map(String::from)
408 .or_else(|| entry.default_model.clone())
409 .unwrap_or_default();
410 match entry.provider_type.as_str() {
411 "openai-compatible" | "custom" => {
412 #[cfg(feature = "openai")]
413 {
414 Ok(Provider::from_runtime_openai(
415 Self::runtime_api_key(provider_name, entry)?,
416 entry.api_base.clone(),
417 model,
418 params,
419 entry.headers.clone(),
420 ))
421 }
422
423 #[cfg(not(feature = "openai"))]
424 {
425 Err(AppError::Configuration(
426 "OpenAI provider configured but the corresponding feature is not enabled in this build".into(),
427 ))
428 }
429 }
430 "anthropic-compatible" => {
431 #[cfg(feature = "anthropic")]
432 {
433 Ok(Provider::Anthropic {
434 api_key: Self::runtime_api_key(provider_name, entry)?,
435 model,
436 params,
437 })
438 }
439
440 #[cfg(not(feature = "anthropic"))]
441 {
442 Err(AppError::Configuration(
443 "Anthropic provider configured but the corresponding feature is not enabled in this build".into(),
444 ))
445 }
446 }
447 "bedrock" | "bedrock-compatible" => {
448 #[cfg(feature = "bedrock")]
449 {
450 Ok(Provider::from_runtime_bedrock(
451 Self::runtime_api_key(provider_name, entry)?,
452 Self::runtime_bedrock_region(provider_name, entry)?,
453 model,
454 params,
455 ))
456 }
457
458 #[cfg(not(feature = "bedrock"))]
459 {
460 Err(AppError::Configuration(
461 "Bedrock provider configured but the corresponding feature is not enabled in this build".into(),
462 ))
463 }
464 }
465 "azure" | "azure-compatible" => {
466 #[cfg(feature = "azure")]
467 {
468 Ok(Provider::Azure {
469 api_key: Self::runtime_api_key(provider_name, entry)?,
470 api_base: entry.api_base.clone(),
471 model,
472 params,
473 })
474 }
475
476 #[cfg(not(feature = "azure"))]
477 {
478 Err(AppError::Configuration(
479 "Azure provider configured but the corresponding feature is not enabled in this build".into(),
480 ))
481 }
482 }
483 provider_type => Err(AppError::Configuration(format!(
484 "Runtime provider '{}' has unsupported provider_type '{}'",
485 provider_name, provider_type
486 ))),
487 }
488 }
489
490 pub fn get_model(&self, name: &str) -> Option<ModelConfig> {
493 if let Some(cfg) = self.models.get(name) {
495 return Some(cfg.clone());
496 }
497 if let Some(model_id) = Self::bedrock_model_id_from_name(name) {
499 return Some(Self::bedrock_model_config(model_id));
500 }
501 if let Some(model_id) = Self::azure_model_id_from_name(name) {
503 return Some(Self::azure_model_config(model_id));
504 }
505 if let Some(catalog) = &self.catalog {
507 let snapshot = catalog.snapshot();
508 if snapshot.iter().any(|e| e.id == name) {
509 return Some(ModelConfig {
510 provider: "nvidia".to_string(),
511 model: name.to_string(),
512 temperature: 0.7,
513 max_tokens: 512,
514 });
515 }
516 }
517 None
518 }
519
520 pub fn provider_names(&self) -> Vec<String> {
522 let mut names: Vec<String> = self.providers.keys().cloned().collect();
523 let runtime = self.runtime_providers.load();
524 for (name, entries) in runtime.iter() {
525 if entries.iter().any(|entry| entry.tenant_id.is_none()) && !names.contains(name) {
526 names.push(name.clone());
527 }
528 }
529 names
530 }
531
532 pub fn model_names(&self) -> Vec<String> {
534 let mut names: Vec<String> = self.models.keys().cloned().collect();
535 if let Some(ProviderConfig::Bedrock { default_model, .. }) = self.get_provider("bedrock") {
536 let name = format!("bedrock/{default_model}");
537 if !default_model.is_empty() && !names.contains(&name) {
538 names.push(name);
539 }
540 }
541 if let Some(ProviderConfig::Azure { default_model, .. }) = self.get_provider("azure") {
542 let name = format!("azure/{default_model}");
543 if !default_model.is_empty() && !names.contains(&name) {
544 names.push(name);
545 }
546 }
547 if let Some(catalog) = &self.catalog {
548 for entry in catalog.snapshot() {
549 names.push(entry.id.clone());
550 }
551 }
552 names
553 }
554
555 pub async fn create_client_for_model(&self, model_name: &str) -> Result<Box<dyn LLMClient>> {
559 self.create_client_for_model_inner(model_name, None).await
560 }
561
562 pub async fn create_client_for_model_ctx(
566 &self,
567 ctx: &std::sync::Arc<cordis::Context>,
568 model_name: &str,
569 ) -> Result<Box<dyn LLMClient>> {
570 let tenant = tenant_from_ctx(ctx);
571 self.create_client_for_model_inner(model_name, tenant.as_deref())
572 .await
573 }
574
575 async fn create_client_for_model_inner(
576 &self,
577 model_name: &str,
578 tenant: Option<&str>,
579 ) -> Result<Box<dyn LLMClient>> {
580 if let Some(model_config) = self.models.get(model_name) {
582 let runtime_entry =
583 self.runtime_provider_entry_for_tenant(&model_config.provider, tenant);
584 if let Some(entry) = runtime_entry {
585 let provider = Self::provider_from_runtime_entry_with_params(
586 &model_config.provider,
587 &entry,
588 Some(&model_config.model),
589 ModelParams::from_model_config(model_config),
590 )?;
591 return provider.create_client().await;
592 }
593
594 let provider_config = self
595 .providers
596 .get(&model_config.provider)
597 .cloned()
598 .ok_or_else(|| {
599 AppError::Configuration(format!(
600 "Provider '{}' referenced by model '{}' not found",
601 model_config.provider, model_name
602 ))
603 })?;
604 let provider = Provider::from_model_config(model_config, &provider_config)?;
605 return provider.create_client().await;
606 }
607
608 if let Some(model_id) = Self::bedrock_model_id_from_name(model_name) {
610 let model_config = Self::bedrock_model_config(model_id);
611 let provider_config = self
612 .provider_for_tenant("bedrock", tenant)
613 .unwrap_or_else(Self::default_bedrock_provider_config);
614 let provider = Provider::from_model_config(&model_config, &provider_config)?;
615 return provider.create_client().await;
616 }
617
618 if let Some(model_id) = Self::azure_model_id_from_name(model_name) {
620 let model_config = Self::azure_model_config(model_id);
621 let provider_config = self
622 .provider_for_tenant("azure", tenant)
623 .unwrap_or_else(Self::default_azure_provider_config);
624 let provider = Provider::from_model_config(&model_config, &provider_config)?;
625 return provider.create_client().await;
626 }
627
628 if let Some(catalog) = &self.catalog {
630 let snapshot = catalog.snapshot();
631 if snapshot.iter().any(|e| e.id == model_name) {
632 let nvidia_cfg = self.nvidia_config_from_providers();
633 let provider_config = ProviderConfig::OpenAI {
634 api_key_env: nvidia_cfg.api_key_env,
635 api_base: nvidia_cfg.api_base,
636 default_model: model_name.to_string(),
637 };
638 let provider = Provider::from_config(&provider_config, Some(model_name))?;
639 return provider.create_client().await;
640 }
641 }
642
643 Err(AppError::Configuration(format!(
644 "Model '{}' not found in configuration",
645 model_name
646 )))
647 }
648
649 pub async fn create_client_for_provider(
651 &self,
652 provider_name: &str,
653 ) -> Result<Box<dyn LLMClient>> {
654 let runtime_entry = self.runtime_provider_entry_for_tenant(provider_name, None);
656 if let Some(entry) = runtime_entry {
657 let provider = Self::provider_from_runtime_entry(provider_name, &entry)?;
658 return provider.create_client().await;
659 }
660
661 let provider_config = self.providers.get(provider_name).ok_or_else(|| {
662 AppError::Configuration(format!(
663 "Provider '{}' not found in configuration",
664 provider_name
665 ))
666 })?;
667
668 let provider = Provider::from_config(provider_config, None)?;
669 provider.create_client().await
670 }
671
672 pub async fn create_client_for_resolved_provider(
674 &self,
675 resolved: &ResolvedProviderConfig,
676 ) -> Result<Box<dyn LLMClient>> {
677 let runtime_entry = self.runtime_provider_entry_for_tenant(
678 &resolved.provider_name,
679 resolved.tenant_id.as_deref(),
680 );
681 if let Some(entry) = runtime_entry {
682 let provider = Self::provider_from_runtime_entry_with_params(
683 &resolved.provider_name,
684 &entry,
685 Some(&resolved.model_name),
686 resolved.params.clone(),
687 )?;
688 return provider.create_client().await;
689 }
690
691 let provider = Provider::from_config_with_params(
692 &resolved.provider_config,
693 Some(&resolved.model_name),
694 resolved.params.clone(),
695 )?;
696 provider.create_client().await
697 }
698
699 pub async fn create_default_client(&self) -> Result<Box<dyn LLMClient>> {
701 let model_name = self
702 .default_model
703 .as_ref()
704 .ok_or_else(|| AppError::Configuration("No default model configured".into()))?;
705
706 self.create_client_for_model(model_name).await
707 }
708
709 pub fn has_model(&self, name: &str) -> bool {
711 self.models.contains_key(name)
712 || Self::bedrock_model_id_from_name(name).is_some()
713 || Self::azure_model_id_from_name(name).is_some()
714 || self
715 .catalog
716 .as_ref()
717 .map(|c| c.snapshot().iter().any(|e| e.id == name))
718 .unwrap_or(false)
719 }
720
721 pub fn has_provider(&self, name: &str) -> bool {
723 self.has_provider_for_tenant(name, None)
724 }
725
726 pub fn has_provider_for_tenant(&self, name: &str, tenant_id: Option<&str>) -> bool {
727 self.providers.contains_key(name)
728 || self
729 .runtime_provider_entry_for_tenant(name, tenant_id)
730 .is_some()
731 }
732
733 pub fn get_model_capabilities(&self, model_name: &str) -> Option<ModelCapabilities> {
737 if let Some(model_id) = Self::bedrock_model_id_from_name(model_name) {
738 let mut caps = ModelCapabilities::for_model(model_id);
739 caps.is_local = false;
740 return Some(caps);
741 }
742 if let Some(model_id) = Self::azure_model_id_from_name(model_name) {
743 let mut caps = ModelCapabilities::for_model(model_id);
744 caps.is_local = false;
745 return Some(caps);
746 }
747
748 if let Some(model_config) = self.models.get(model_name) {
750 let provider_config = self.get_provider(&model_config.provider)?;
751 let mut caps = ModelCapabilities::for_model(&model_config.model);
752 if matches!(
753 provider_config,
754 ProviderConfig::OpenAI { .. }
755 | ProviderConfig::Azure { .. }
756 | ProviderConfig::Bedrock { .. }
757 ) {
758 caps.is_local = false;
759 }
760 return Some(caps);
761 }
762
763 if let Some(catalog) = &self.catalog {
765 let snapshot = catalog.snapshot();
766 if snapshot.iter().any(|e| e.id == model_name) {
767 let mut caps = ModelCapabilities::for_model(model_name);
768 caps.is_local = false;
769 return Some(caps);
770 }
771 }
772
773 None
774 }
775
776 pub fn models_with_capabilities(&self) -> Vec<ModelWithCapabilities> {
778 let mut result = Vec::new();
779
780 for (name, config) in &self.models {
782 if let Some(caps) = self.get_model_capabilities(name) {
783 result.push(ModelWithCapabilities {
784 name: name.clone(),
785 provider: config.provider.clone(),
786 model_id: config.model.clone(),
787 capabilities: caps,
788 });
789 }
790 }
791
792 if let Some(ProviderConfig::Bedrock { default_model, .. }) = self.get_provider("bedrock") {
793 let name = format!("bedrock/{default_model}");
794 if !default_model.is_empty() && !result.iter().any(|model| model.name == name) {
795 let mut caps = ModelCapabilities::for_model(&default_model);
796 caps.is_local = false;
797 result.push(ModelWithCapabilities {
798 name,
799 provider: "bedrock".to_string(),
800 model_id: default_model,
801 capabilities: caps,
802 });
803 }
804 }
805
806 if let Some(ProviderConfig::Azure { default_model, .. }) = self.get_provider("azure") {
807 let name = format!("azure/{default_model}");
808 if !default_model.is_empty() && !result.iter().any(|model| model.name == name) {
809 let mut caps = ModelCapabilities::for_model(&default_model);
810 caps.is_local = false;
811 result.push(ModelWithCapabilities {
812 name,
813 provider: "azure".to_string(),
814 model_id: default_model,
815 capabilities: caps,
816 });
817 }
818 }
819
820 if let Some(catalog) = &self.catalog {
822 for entry in catalog.snapshot() {
823 let caps = self.get_model_capabilities(&entry.id).unwrap_or_else(|| {
824 let mut c = ModelCapabilities::for_model(&entry.id);
825 c.is_local = false;
826 c
827 });
828 result.push(ModelWithCapabilities {
829 name: entry.id.clone(),
830 provider: "nvidia".to_string(),
831 model_id: entry.id.clone(),
832 capabilities: caps,
833 });
834 }
835 }
836
837 result
838 }
839
840 pub fn find_models(&self, requirements: &CapabilityRequirements) -> Vec<ModelWithCapabilities> {
842 let mut matches: Vec<_> = self
843 .models_with_capabilities()
844 .into_iter()
845 .filter(|m| m.capabilities.satisfies(requirements))
846 .collect();
847
848 matches.sort_by(|a, b| {
850 let score_a = a.capabilities.score(requirements);
851 let score_b = b.capabilities.score(requirements);
852 score_b.cmp(&score_a)
853 });
854
855 matches
856 }
857
858 pub fn find_best_model(
860 &self,
861 requirements: &CapabilityRequirements,
862 ) -> Option<ModelWithCapabilities> {
863 self.find_models(requirements).into_iter().next()
864 }
865
866 pub async fn create_client_for_requirements(
868 &self,
869 requirements: &CapabilityRequirements,
870 ) -> Result<Box<dyn LLMClient>> {
871 let model = self.find_best_model(requirements).ok_or_else(|| {
872 AppError::Configuration(format!(
873 "No model found matching requirements: {:?}",
874 requirements
875 ))
876 })?;
877
878 self.create_client_for_model(&model.name).await
879 }
880
881 pub fn find_agent_models(&self) -> Vec<ModelWithCapabilities> {
883 self.find_models(&CapabilityRequirements::for_agent())
884 }
885
886 pub fn find_vision_models(&self) -> Vec<ModelWithCapabilities> {
888 self.find_models(&CapabilityRequirements::for_vision())
889 }
890
891 pub fn find_coding_models(&self) -> Vec<ModelWithCapabilities> {
893 self.find_models(&CapabilityRequirements::for_coding())
894 }
895
896 pub fn find_local_models(&self) -> Vec<ModelWithCapabilities> {
898 self.find_models(&CapabilityRequirements::for_local())
899 }
900
901 pub fn list_models(&self) -> Vec<ModelInfo> {
903 let mut models = Vec::new();
904
905 for (name, config) in &self.models {
907 let capabilities = ModelCapabilities::for_model(&config.model);
908 models.push(ModelInfo {
909 name: name.clone(),
910 provider: config.provider.clone(),
911 model: config.model.clone(),
912 owned_by: config.provider.clone(),
913 quality_score: 75,
914 is_chat: true,
915 supports_reasoning: capabilities.supports_reasoning,
916 supports_streaming: capabilities.supports_streaming,
917 });
918 }
919
920 if let Some(ProviderConfig::Bedrock { default_model, .. }) = self.get_provider("bedrock") {
921 let name = format!("bedrock/{default_model}");
922 if !default_model.is_empty() && !models.iter().any(|model| model.name == name) {
923 let capabilities = ModelCapabilities::for_model(&default_model);
924 models.push(ModelInfo {
925 name,
926 provider: "bedrock".to_string(),
927 model: default_model,
928 owned_by: "aws-bedrock".to_string(),
929 quality_score: 85,
930 is_chat: true,
931 supports_reasoning: capabilities.supports_reasoning,
932 supports_streaming: capabilities.supports_streaming,
933 });
934 }
935 }
936
937 if let Some(ProviderConfig::Azure { default_model, .. }) = self.get_provider("azure") {
938 let name = format!("azure/{default_model}");
939 if !default_model.is_empty() && !models.iter().any(|model| model.name == name) {
940 let capabilities = ModelCapabilities::for_model(&default_model);
941 models.push(ModelInfo {
942 name,
943 provider: "azure".to_string(),
944 model: default_model,
945 owned_by: "azure-foundry".to_string(),
946 quality_score: 80,
947 is_chat: true,
948 supports_reasoning: capabilities.supports_reasoning,
949 supports_streaming: capabilities.supports_streaming,
950 });
951 }
952 }
953
954 if let Some(catalog) = &self.catalog {
956 let snapshot = catalog.snapshot();
957 if !snapshot.is_empty() {
958 for entry in snapshot {
959 let capabilities = self
960 .get_model_capabilities(&entry.id)
961 .unwrap_or_else(|| ModelCapabilities::for_model(&entry.id));
962 models.push(ModelInfo {
963 name: entry.id.clone(),
964 provider: "nvidia".to_string(),
965 model: entry.id.clone(),
966 owned_by: entry.owned_by.clone(),
967 quality_score: entry.quality_score,
968 is_chat: true,
969 supports_reasoning: capabilities.supports_reasoning,
970 supports_streaming: capabilities.supports_streaming,
971 });
972 }
973 } else if let Some(default) = &self.default_model {
974 let capabilities = ModelCapabilities::for_model(default);
976 models.push(ModelInfo {
977 name: default.clone(),
978 provider: "nvidia".to_string(),
979 model: default.clone(),
980 owned_by: "unknown".to_string(),
981 quality_score: 75,
982 is_chat: true,
983 supports_reasoning: capabilities.supports_reasoning,
984 supports_streaming: capabilities.supports_streaming,
985 });
986 }
987 } else if let Some(default) = &self.default_model {
988 let capabilities = ModelCapabilities::for_model(default);
990 models.push(ModelInfo {
991 name: default.clone(),
992 provider: "nvidia".to_string(),
993 model: default.clone(),
994 owned_by: "unknown".to_string(),
995 quality_score: 75,
996 is_chat: true,
997 supports_reasoning: capabilities.supports_reasoning,
998 supports_streaming: capabilities.supports_streaming,
999 });
1000 }
1001
1002 models
1003 }
1004
1005 pub async fn resolve_with_capability_fallback(
1012 &self,
1013 requirements: Option<CapabilityRequirements>,
1014 ) -> Result<Box<dyn crate::client::LLMClient>> {
1015 if let Some(req) = requirements {
1016 if let Ok(client) = self.create_client_for_requirements(&req).await {
1017 return Ok(client);
1018 }
1019 }
1020 self.create_default_client().await
1021 }
1022
1023 #[cfg(not(feature = "postgres"))]
1026 pub async fn resolve_with_fallback(
1027 &self,
1028 requirements: Option<CapabilityRequirements>,
1029 ) -> Result<Box<dyn crate::client::LLMClient>> {
1030 self.resolve_with_capability_fallback(requirements).await
1031 }
1032
1033 #[cfg(feature = "postgres")]
1046 pub async fn resolve_with_fallback(
1047 &self,
1048 tier_or_model: &str,
1049 tenant_id: &str,
1050 pool: &sqlx::PgPool,
1051 fleet_secrets: &ares_store::FleetSecrets,
1052 ) -> Result<Vec<ResolvedProviderConfig>> {
1053 use ares_store::tenant_model_tiers::TenantModelTierStore;
1054 use std::collections::HashSet;
1055
1056 let store = TenantModelTierStore::new(pool);
1057 let primary = match store.get(tenant_id, tier_or_model).await {
1058 Ok(Some(tier)) => {
1059 let provider_config = self
1060 .provider_for_tenant(&tier.provider_name, Some(tenant_id))
1061 .ok_or_else(|| {
1062 AppError::Configuration(format!(
1063 "Provider '{}' configured for tenant '{}' tier '{}' not found",
1064 tier.provider_name, tenant_id, tier_or_model
1065 ))
1066 })?;
1067 ResolvedProviderConfig {
1068 provider_name: tier.provider_name,
1069 model_name: tier.model_name,
1070 provider_config,
1071 params: ModelParams::default(),
1072 tenant_id: Some(tenant_id.to_string()),
1073 }
1074 }
1075 Ok(None) => {
1076 if let Some(model_cfg) = self.get_model(tier_or_model) {
1077 let provider_config = self
1078 .provider_for_tenant(&model_cfg.provider, Some(tenant_id))
1079 .ok_or_else(|| {
1080 AppError::Configuration(format!(
1081 "Provider '{}' referenced by model/tier '{}' not found",
1082 model_cfg.provider, tier_or_model
1083 ))
1084 })?;
1085 ResolvedProviderConfig {
1086 provider_name: model_cfg.provider.clone(),
1087 model_name: model_cfg.model.clone(),
1088 provider_config,
1089 params: ModelParams::from_model_config(&model_cfg),
1090 tenant_id: Some(tenant_id.to_string()),
1091 }
1092 } else if let Some(provider_config) =
1093 self.provider_for_tenant(tier_or_model, Some(tenant_id))
1094 {
1095 let model_name = Self::provider_default_model(&provider_config).to_string();
1096 if model_name.is_empty() {
1097 return Err(AppError::Configuration(format!(
1098 "Provider '{}' has no concrete default model configured",
1099 tier_or_model
1100 )));
1101 }
1102 ResolvedProviderConfig {
1103 provider_name: tier_or_model.to_string(),
1104 model_name,
1105 provider_config,
1106 params: ModelParams::default(),
1107 tenant_id: Some(tenant_id.to_string()),
1108 }
1109 } else {
1110 return Err(AppError::Configuration(format!(
1111 "No provider or model/tier '{}' found for tenant '{}'",
1112 tier_or_model, tenant_id
1113 )));
1114 }
1115 }
1116 Err(e) => {
1117 return Err(AppError::Database(format!(
1118 "Failed to resolve tenant '{}' model tier '{}': {}",
1119 tenant_id, tier_or_model, e
1120 )));
1121 }
1122 };
1123
1124 let primary_provider = primary.provider_name.clone();
1125 let mut result = vec![primary];
1126 let mut seen = HashSet::new();
1127 seen.insert(primary_provider.clone());
1128
1129 if let Some(override_) = fleet_secrets.get(&primary_provider) {
1130 for fallback_name in &override_.fallback_providers {
1131 if seen.contains(fallback_name) {
1132 continue;
1133 }
1134 let provider_config = self
1135 .provider_for_tenant(fallback_name, Some(tenant_id))
1136 .ok_or_else(|| {
1137 AppError::Configuration(format!(
1138 "Fallback provider '{}' configured for primary provider '{}' not found",
1139 fallback_name, primary_provider
1140 ))
1141 })?;
1142 let model_name = Self::provider_default_model(&provider_config).to_string();
1143 if model_name.is_empty() {
1144 return Err(AppError::Configuration(format!(
1145 "Fallback provider '{}' configured for primary provider '{}' has no concrete default model configured",
1146 fallback_name, primary_provider
1147 )));
1148 }
1149 seen.insert(fallback_name.clone());
1150 result.push(ResolvedProviderConfig {
1151 provider_name: fallback_name.clone(),
1152 model_name,
1153 provider_config,
1154 params: ModelParams::default(),
1155 tenant_id: Some(tenant_id.to_string()),
1156 });
1157 }
1158 }
1159
1160 Ok(result)
1161 }
1162
1163 fn nvidia_config_from_providers(&self) -> NvidiaConfig {
1165 if let Some(ProviderConfig::OpenAI {
1166 api_key_env,
1167 api_base,
1168 default_model,
1169 }) = self.providers.get("nvidia")
1170 {
1171 NvidiaConfig {
1172 api_key_env: api_key_env.clone(),
1173 api_base: api_base.clone(),
1174 models_url: format!("{}/models", api_base.trim_end_matches('/')),
1175 catalog_refresh_seconds: 3600,
1176 default_model: default_model.clone(),
1177 }
1178 } else {
1179 NvidiaConfig::default()
1180 }
1181 }
1182}
1183
1184#[derive(Debug, Clone, serde::Serialize)]
1186pub struct ModelInfo {
1187 pub name: String,
1188 pub provider: String,
1189 pub model: String,
1190 #[serde(default)]
1191 pub owned_by: String,
1192 #[serde(default)]
1193 pub quality_score: u8,
1194 #[serde(default)]
1195 pub is_chat: bool,
1196 #[serde(default)]
1197 pub supports_reasoning: bool,
1198 #[serde(default)]
1199 pub supports_streaming: bool,
1200}
1201
1202impl Default for ProviderRegistry {
1203 fn default() -> Self {
1204 Self::new()
1205 }
1206}
1207
1208pub struct ConfigBasedLLMFactory {
1210 registry: Arc<ProviderRegistry>,
1211 default_model: String,
1212}
1213
1214impl ConfigBasedLLMFactory {
1215 pub fn new(registry: Arc<ProviderRegistry>, default_model: &str) -> Self {
1217 Self {
1218 registry,
1219 default_model: default_model.to_string(),
1220 }
1221 }
1222
1223 pub fn from_config(
1225 providers: std::collections::HashMap<String, ProviderConfig>,
1226 models: std::collections::HashMap<String, ModelConfig>,
1227 nvidia: Option<&NvidiaConfig>,
1228 ) -> Result<Self> {
1229 let registry = ProviderRegistry::from_config(providers, models.clone(), nvidia);
1230
1231 let default_model = nvidia
1232 .map(|n| n.default_model.clone())
1233 .or_else(|| models.keys().next().cloned())
1234 .unwrap_or_else(|| "nvidia/nemotron-3-ultra-550b-a55b".to_string());
1235
1236 Ok(Self {
1237 registry: Arc::new(registry),
1238 default_model,
1239 })
1240 }
1241
1242 pub fn registry(&self) -> &Arc<ProviderRegistry> {
1244 &self.registry
1245 }
1246
1247 pub async fn create_for_model(&self, model_name: &str) -> Result<Box<dyn LLMClient>> {
1249 self.registry.create_client_for_model(model_name).await
1250 }
1251
1252 pub async fn create_default(&self) -> Result<Box<dyn LLMClient>> {
1254 self.registry
1255 .create_client_for_model(&self.default_model)
1256 .await
1257 }
1258
1259 pub fn default_model(&self) -> &str {
1261 &self.default_model
1262 }
1263
1264 pub fn set_default_model(&mut self, model_name: &str) {
1266 self.default_model = model_name.to_string();
1267 }
1268}
1269
1270pub fn reflect_notify_stub(ctx: &Arc<cordis::Context>) {
1275 let _ = ctx.get::<cordis::loader::Loader>();
1277 let tid = TypeId::of::<ProviderRegistry>();
1278 if let Some(reflect) = ctx.get::<ReflectService>() {
1280 let _rx = reflect.ensure_notifier(tid);
1281 reflect.register_dependent(tid, 43);
1282 reflect.notify(tid);
1283 }
1284 let _ = tid;
1285}
1286
1287#[cfg(test)]
1288mod tests {
1289 use super::*;
1290 use crate::capabilities::CapabilityRequirements;
1291
1292 use crate::config::{ModelConfig, ProviderConfig};
1293 use std::collections::HashMap;
1294
1295 fn sample_openai_provider() -> ProviderConfig {
1296 ProviderConfig::OpenAI {
1297 api_key_env: "TEST_KEY".to_string(),
1298 api_base: "https://test.example.com/v1".to_string(),
1299 default_model: "test-model".to_string(),
1300 }
1301 }
1302
1303 fn sample_model_config(provider: &str, model: &str) -> ModelConfig {
1304 ModelConfig {
1305 provider: provider.to_string(),
1306 model: model.to_string(),
1307 temperature: 0.7,
1308 max_tokens: 512,
1309 }
1310 }
1311
1312 fn from_maps(
1313 providers: HashMap<String, ProviderConfig>,
1314 models: HashMap<String, ModelConfig>,
1315 ) -> crate::provider_registry::ProviderRegistry {
1316 ProviderRegistry::from_config(providers, models, None)
1317 }
1318
1319 fn assert_configuration_error<T>(result: Result<T>, expected_substring: &str) {
1320 match result {
1321 Err(AppError::Configuration(msg)) => {
1322 assert!(
1323 msg.contains(expected_substring),
1324 "expected message containing {expected_substring:?}, got {msg:?}"
1325 );
1326 }
1327 Err(other) => panic!("expected Configuration error, got: {other:?}"),
1328 Ok(_) => {
1329 panic!("expected Configuration error containing {expected_substring:?}, got Ok")
1330 }
1331 }
1332 }
1333
1334 #[test]
1335 fn test_empty_registry() {
1336 let registry = ProviderRegistry::new();
1337 assert!(registry.provider_names().is_empty());
1338 assert!(registry.model_names().is_empty());
1339 }
1340
1341 #[test]
1342 fn test_register_provider() {
1343 let mut registry = ProviderRegistry::new();
1344 registry.register_provider(
1345 "nvidia",
1346 ProviderConfig::OpenAI {
1347 api_key_env: "TEST_KEY".to_string(),
1348 api_base: "https://test.example.com/v1".to_string(),
1349 default_model: "test-model".to_string(),
1350 },
1351 );
1352
1353 assert!(registry.has_provider("nvidia"));
1354 assert!(!registry.has_provider("nonexistent"));
1355 }
1356
1357 #[test]
1358 fn test_register_model() {
1359 let mut registry = ProviderRegistry::new();
1360 registry.register_provider(
1361 "nvidia",
1362 ProviderConfig::OpenAI {
1363 api_key_env: "TEST_KEY".to_string(),
1364 api_base: "https://test.example.com/v1".to_string(),
1365 default_model: "test-model".to_string(),
1366 },
1367 );
1368 registry.register_model(
1369 "fast",
1370 ModelConfig {
1371 provider: "nvidia".to_string(),
1372 model: "test-model".to_string(),
1373 temperature: 0.7,
1374 max_tokens: 256,
1375 },
1376 );
1377
1378 assert!(registry.has_model("fast"));
1379 assert!(!registry.has_model("nonexistent"));
1380 }
1381
1382 fn create_test_registry() -> ProviderRegistry {
1385 let mut registry = ProviderRegistry::new();
1386
1387 registry.register_provider(
1388 "nvidia",
1389 ProviderConfig::OpenAI {
1390 api_key_env: "TEST_KEY".to_string(),
1391 api_base: "https://integrate.api.nvidia.com/v1".to_string(),
1392 default_model: "nvidia/nemotron-3-ultra-550b-a55b".to_string(),
1393 },
1394 );
1395
1396 registry.register_model(
1397 "fast-local",
1398 ModelConfig {
1399 provider: "nvidia".to_string(),
1400 model: "nvidia/nemotron-3-ultra-550b-a55b".to_string(),
1401 temperature: 0.7,
1402 max_tokens: 512,
1403 },
1404 );
1405
1406 registry.register_model(
1407 "powerful-local",
1408 ModelConfig {
1409 provider: "nvidia".to_string(),
1410 model: "nvidia/nemotron-3-ultra-550b-a55b".to_string(),
1411 temperature: 0.7,
1412 max_tokens: 2048,
1413 },
1414 );
1415
1416 registry.register_model(
1417 "qwen",
1418 ModelConfig {
1419 provider: "nvidia".to_string(),
1420 model: "qwen/qwen-32b".to_string(),
1421 temperature: 0.7,
1422 max_tokens: 4096,
1423 },
1424 );
1425
1426 registry
1427 }
1428
1429 #[test]
1430 fn test_get_model_capabilities() {
1431 let registry = create_test_registry();
1432
1433 let fast_caps = registry.get_model_capabilities("fast-local").unwrap();
1434 assert!(!fast_caps.is_local);
1435 assert!(fast_caps.supports_tools);
1436 }
1437
1438 #[test]
1439 fn test_models_with_capabilities() {
1440 let registry = create_test_registry();
1441 let models = registry.models_with_capabilities();
1442
1443 assert_eq!(models.len(), 3);
1444
1445 for model in &models {
1446 assert!(!model.name.is_empty());
1447 assert!(!model.provider.is_empty());
1448 assert!(model.capabilities.supports_tools);
1449 }
1450 }
1451
1452 #[test]
1453 fn test_find_local_models() {
1454 let registry = create_test_registry();
1455 let local_models = registry.find_local_models();
1456 assert!(local_models.is_empty());
1458 }
1459
1460 #[test]
1461 fn test_find_vision_models() {
1462 let registry = create_test_registry();
1463 let vision_models = registry.find_vision_models();
1464 assert!(vision_models.is_empty());
1466 }
1467
1468 #[test]
1469 fn test_find_best_model_for_agent() {
1470 let registry = create_test_registry();
1471
1472 let requirements = CapabilityRequirements::for_agent();
1473 let best = registry.find_best_model(&requirements);
1474
1475 assert!(best.is_some());
1476 let best = best.unwrap();
1477 assert!(best.capabilities.supports_tools);
1478 assert!(best.capabilities.production_ready);
1479 }
1480
1481 #[test]
1482 fn test_find_best_model_with_context_window() {
1483 let registry = create_test_registry();
1484
1485 let requirements = CapabilityRequirements::builder()
1486 .min_context_window(100_000)
1487 .build();
1488
1489 let matches = registry.find_models(&requirements);
1490
1491 assert!(matches.len() >= 2);
1492 for model in &matches {
1493 assert!(model.capabilities.context_window >= 100_000);
1494 }
1495 }
1496
1497 #[test]
1498 fn test_find_best_model_prefers_cheaper() {
1499 let registry = create_test_registry();
1500
1501 let requirements = CapabilityRequirements::builder().requires_tools().build();
1502
1503 let best = registry.find_best_model(&requirements).unwrap();
1504
1505 assert_eq!(best.capabilities.cost_tier, "free");
1507 }
1508
1509 #[test]
1510 fn test_no_model_matches_impossible_requirements() {
1511 let registry = create_test_registry();
1512
1513 let requirements = CapabilityRequirements::builder()
1514 .requires_local()
1515 .requires_vision()
1516 .build();
1517
1518 let matches = registry.find_models(&requirements);
1519 assert!(matches.is_empty());
1520 }
1521
1522 #[test]
1523 fn test_find_coding_models() {
1524 let registry = create_test_registry();
1525 let coding_models = registry.find_coding_models();
1526
1527 for model in &coding_models {
1528 assert!(model.capabilities.supports_tools);
1529 assert!(model.capabilities.supports_reasoning);
1530 assert!(model.capabilities.context_window >= 32_000);
1531 }
1532 }
1533
1534 #[test]
1535 fn test_unregister_provider() {
1536 let mut registry = ProviderRegistry::new();
1537 registry.register_provider(
1538 "nvidia",
1539 ProviderConfig::OpenAI {
1540 api_key_env: "TEST_KEY".to_string(),
1541 api_base: "https://test.example.com/v1".to_string(),
1542 default_model: "test-model".to_string(),
1543 },
1544 );
1545 assert!(registry.has_provider("nvidia"));
1546 let removed = registry.unregister_provider("nvidia").unwrap();
1547 assert!(matches!(removed, ProviderConfig::OpenAI { .. }));
1548 assert!(!registry.has_provider("nvidia"));
1549 }
1550
1551 #[test]
1552 fn test_unregister_model() {
1553 let mut registry = create_test_registry();
1554 assert!(registry.has_model("fast-local"));
1555 registry.unregister_model("fast-local");
1556 assert!(!registry.has_model("fast-local"));
1557 }
1558
1559 #[test]
1560 fn test_lookup_provider_by_name() {
1561 let registry = create_test_registry();
1562 let provider = registry.get_provider("nvidia").unwrap();
1563 assert!(matches!(provider, ProviderConfig::OpenAI { .. }));
1564 assert!(registry.get_provider("missing").is_none());
1565 }
1566
1567 #[cfg(feature = "openai")]
1568 #[test]
1569 fn test_runtime_openai_provider_preserves_key_and_headers() {
1570 let mut headers = HashMap::new();
1571 headers.insert("X-Test-Header".to_string(), "runtime-value".to_string());
1572 let entry = RuntimeProviderEntry {
1573 tenant_id: None,
1574 display_name: "Runtime OpenAI".to_string(),
1575 provider_type: "openai-compatible".to_string(),
1576 api_base: "https://runtime.example.com/v1".to_string(),
1577 auth_type: "api_key".to_string(),
1578 default_model: Some("runtime-model".to_string()),
1579 headers,
1580 api_key: Some("resolved-runtime-key".to_string()),
1581 enabled: true,
1582 };
1583
1584 let provider = ProviderRegistry::provider_from_runtime_entry("runtime-openai", &entry)
1585 .expect("runtime provider should resolve");
1586 match provider {
1587 Provider::RuntimeOpenAI {
1588 api_key,
1589 api_base,
1590 model,
1591 headers,
1592 ..
1593 } => {
1594 assert_eq!(api_key, "resolved-runtime-key");
1595 assert_eq!(api_base, "https://runtime.example.com/v1");
1596 assert_eq!(model, "runtime-model");
1597 assert_eq!(
1598 headers.get("X-Test-Header").map(String::as_str),
1599 Some("runtime-value")
1600 );
1601 }
1602 _ => panic!("expected RuntimeOpenAI provider"),
1603 }
1604 }
1605
1606 #[cfg(feature = "openai")]
1607 #[test]
1608 fn test_runtime_provider_requires_resolved_api_key() {
1609 let entry = RuntimeProviderEntry {
1610 tenant_id: None,
1611 display_name: "Runtime OpenAI".to_string(),
1612 provider_type: "openai-compatible".to_string(),
1613 api_base: "https://runtime.example.com/v1".to_string(),
1614 auth_type: "api_key".to_string(),
1615 default_model: Some("runtime-model".to_string()),
1616 headers: HashMap::new(),
1617 api_key: None,
1618 enabled: true,
1619 };
1620
1621 assert_configuration_error(
1622 ProviderRegistry::provider_from_runtime_entry("runtime-openai", &entry),
1623 "Runtime provider 'runtime-openai' API key is not resolved",
1624 );
1625 }
1626
1627 #[test]
1628 fn runtime_provider_visibility_respects_tenant_scope() {
1629 let registry = ProviderRegistry::new();
1630 let global = RuntimeProviderEntry {
1631 tenant_id: None,
1632 display_name: "Global Runtime".to_string(),
1633 provider_type: "openai-compatible".to_string(),
1634 api_base: "https://global.example.com/v1".to_string(),
1635 auth_type: "api_key".to_string(),
1636 default_model: Some("global-model".to_string()),
1637 headers: HashMap::new(),
1638 api_key: Some("global-key".to_string()),
1639 enabled: true,
1640 };
1641 let scoped = RuntimeProviderEntry {
1642 tenant_id: Some("tenant-a".to_string()),
1643 display_name: "Scoped Runtime".to_string(),
1644 provider_type: "openai-compatible".to_string(),
1645 api_base: "https://tenant.example.com/v1".to_string(),
1646 auth_type: "api_key".to_string(),
1647 default_model: Some("tenant-model".to_string()),
1648 headers: HashMap::new(),
1649 api_key: Some("tenant-key".to_string()),
1650 enabled: true,
1651 };
1652 registry.reload_runtime_providers(
1653 vec![global, scoped],
1654 vec!["global-runtime".to_string(), "tenant-runtime".to_string()],
1655 );
1656
1657 assert!(registry.has_provider("global-runtime"));
1658 assert!(!registry.has_provider("tenant-runtime"));
1659 assert!(registry.has_provider_for_tenant("tenant-runtime", Some("tenant-a")));
1660 assert!(!registry.has_provider_for_tenant("tenant-runtime", Some("tenant-b")));
1661 assert!(
1662 registry
1663 .provider_for_tenant("tenant-runtime", Some("tenant-a"))
1664 .is_some()
1665 );
1666 assert!(
1667 registry
1668 .provider_for_tenant("tenant-runtime", Some("tenant-b"))
1669 .is_none()
1670 );
1671 assert_eq!(
1672 registry.provider_names(),
1673 vec!["global-runtime".to_string()]
1674 );
1675 }
1676
1677 #[test]
1678 fn runtime_provider_lookup_prefers_tenant_override_same_name() {
1679 let registry = ProviderRegistry::new();
1680 let global = RuntimeProviderEntry {
1681 tenant_id: None,
1682 display_name: "Global Shared".to_string(),
1683 provider_type: "openai-compatible".to_string(),
1684 api_base: "https://global.example.com/v1".to_string(),
1685 auth_type: "api_key".to_string(),
1686 default_model: Some("global-model".to_string()),
1687 headers: HashMap::new(),
1688 api_key: Some("global-key".to_string()),
1689 enabled: true,
1690 };
1691 let tenant = RuntimeProviderEntry {
1692 tenant_id: Some("tenant-a".to_string()),
1693 display_name: "Tenant Shared".to_string(),
1694 provider_type: "openai-compatible".to_string(),
1695 api_base: "https://tenant.example.com/v1".to_string(),
1696 auth_type: "api_key".to_string(),
1697 default_model: Some("tenant-model".to_string()),
1698 headers: HashMap::new(),
1699 api_key: Some("tenant-key".to_string()),
1700 enabled: true,
1701 };
1702 registry.reload_runtime_providers(
1703 vec![global, tenant],
1704 vec!["shared-runtime".to_string(), "shared-runtime".to_string()],
1705 );
1706
1707 assert!(registry.has_provider("shared-runtime"));
1708 assert!(registry.has_provider_for_tenant("shared-runtime", Some("tenant-a")));
1709 assert!(registry.has_provider_for_tenant("shared-runtime", Some("tenant-b")));
1710
1711 let tenant_provider = registry
1712 .provider_for_tenant("shared-runtime", Some("tenant-a"))
1713 .expect("tenant provider");
1714 let global_provider = registry
1715 .provider_for_tenant("shared-runtime", Some("tenant-b"))
1716 .expect("global fallback provider");
1717 assert_eq!(
1718 ProviderRegistry::provider_default_model(&tenant_provider),
1719 "tenant-model"
1720 );
1721 assert_eq!(
1722 ProviderRegistry::provider_default_model(&global_provider),
1723 "global-model"
1724 );
1725 assert_eq!(
1726 registry.provider_names(),
1727 vec!["shared-runtime".to_string()]
1728 );
1729 }
1730
1731 #[test]
1732 fn test_default_registry() {
1733 let registry = ProviderRegistry::default();
1734 assert!(registry.provider_names().is_empty());
1735 assert!(registry.model_names().is_empty());
1736 }
1737
1738 #[test]
1739 fn test_register_provider_overwrites_existing() {
1740 let mut registry = ProviderRegistry::new();
1741 registry.register_provider(
1742 "nvidia",
1743 ProviderConfig::OpenAI {
1744 api_key_env: "TEST_KEY".to_string(),
1745 api_base: "https://old.example.com/v1".to_string(),
1746 default_model: "old-model".to_string(),
1747 },
1748 );
1749 registry.register_provider(
1750 "nvidia",
1751 ProviderConfig::OpenAI {
1752 api_key_env: "TEST_KEY".to_string(),
1753 api_base: "https://new.example.com/v1".to_string(),
1754 default_model: "new-model".to_string(),
1755 },
1756 );
1757
1758 let provider = registry.get_provider("nvidia").unwrap();
1759 if let ProviderConfig::OpenAI { default_model, .. } = provider {
1760 assert_eq!(default_model, "new-model");
1761 } else {
1762 panic!("expected OpenAI provider");
1763 }
1764 }
1765
1766 #[test]
1767 fn test_provider_and_model_name_iteration() {
1768 let mut registry = ProviderRegistry::new();
1769 registry.register_provider("alpha", sample_openai_provider());
1770 registry.register_provider("beta", sample_openai_provider());
1771 registry.register_model("m1", sample_model_config("alpha", "model-a"));
1772 registry.register_model("m2", sample_model_config("beta", "model-b"));
1773
1774 let mut provider_names = registry.provider_names();
1775 provider_names.sort_unstable();
1776 assert_eq!(provider_names, vec!["alpha", "beta"]);
1777
1778 let mut model_names = registry.model_names();
1779 model_names.sort_unstable();
1780 assert_eq!(model_names, vec!["m1", "m2"]);
1781 }
1782
1783 #[test]
1784 fn test_lookup_model_by_name() {
1785 let mut registry = ProviderRegistry::new();
1786 registry.register_provider("nvidia", sample_openai_provider());
1787 registry.register_model("fast", sample_model_config("nvidia", "test-model"));
1788
1789 let model = registry.get_model("fast").unwrap();
1790 assert_eq!(model.provider, "nvidia");
1791 assert_eq!(model.model, "test-model");
1792 assert!(registry.get_model("missing").is_none());
1793 }
1794
1795 #[test]
1796 fn test_list_models_returns_registered_entries() {
1797 let registry = create_test_registry();
1798 let models = registry.list_models();
1799
1800 assert_eq!(models.len(), 3);
1801 let fast = models
1802 .iter()
1803 .find(|m| m.name == "fast-local" && m.provider == "nvidia")
1804 .expect("fast model");
1805 assert!(fast.supports_reasoning);
1806 assert!(fast.supports_streaming);
1807 }
1808
1809 #[test]
1810 fn test_from_config_loads_providers_and_models() {
1811 let mut providers = HashMap::new();
1812 providers.insert("nvidia".to_string(), sample_openai_provider());
1813 let mut models = HashMap::new();
1814 models.insert(
1815 "fast".to_string(),
1816 sample_model_config("nvidia", "test-model"),
1817 );
1818
1819 let registry = from_maps(providers, models);
1820
1821 assert!(registry.has_provider("nvidia"));
1822 assert!(registry.has_model("fast"));
1823
1824 #[cfg(feature = "bedrock")]
1825 let expected = vec![
1826 "fast",
1827 "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
1828 ];
1829 #[cfg(not(feature = "bedrock"))]
1830 let expected = vec!["fast"];
1831 assert_eq!(registry.model_names(), expected);
1832 }
1833
1834 #[tokio::test]
1835 async fn test_set_default_model() {
1836 let mut registry = create_test_registry();
1837 registry.set_default_model("powerful-local");
1838 let saved_key = std::env::var("TEST_KEY").ok();
1841 std::env::remove_var("TEST_KEY");
1842 let result = registry.create_default_client().await;
1843 if let Some(key) = saved_key {
1844 std::env::set_var("TEST_KEY", key);
1845 }
1846 match result {
1847 Err(AppError::Configuration(msg)) => {
1848 assert!(!msg.contains("No default model configured"), "got: {msg}");
1849 }
1850 Err(other) => panic!("expected Configuration error, got: {other:?}"),
1851 Ok(_) => panic!("expected Configuration error, but client creation succeeded"),
1852 }
1853 }
1854
1855 #[test]
1856 fn test_get_model_capabilities_unknown_model() {
1857 let registry = create_test_registry();
1858 assert!(registry.get_model_capabilities("missing").is_none());
1859 }
1860
1861 #[test]
1862 fn test_get_model_capabilities_missing_provider() {
1863 let mut registry = ProviderRegistry::new();
1864 registry.register_model(
1865 "orphan",
1866 sample_model_config("missing-provider", "some-model"),
1867 );
1868 assert!(registry.get_model_capabilities("orphan").is_none());
1869 }
1870
1871 #[test]
1872 fn test_unregister_provider_missing_returns_none() {
1873 let mut registry = ProviderRegistry::new();
1874 assert!(registry.unregister_provider("missing").is_none());
1875 }
1876
1877 #[test]
1878 fn test_unregister_model_returns_removed_config() {
1879 let mut registry = create_test_registry();
1880 let removed = registry.unregister_model("fast-local").unwrap();
1881 assert_eq!(removed.provider, "nvidia");
1882 assert_eq!(removed.model, "nvidia/nemotron-3-ultra-550b-a55b");
1883 assert!(registry.unregister_model("fast-local").is_none());
1884 }
1885
1886 #[test]
1887 fn test_provider_config_serde_roundtrip() {
1888 let configs = [ProviderConfig::OpenAI {
1889 api_key_env: "OPENAI_API_KEY".to_string(),
1890 api_base: "https://api.openai.com/v1".to_string(),
1891 default_model: "gpt-4o".to_string(),
1892 }];
1893
1894 for original in configs {
1895 let json = serde_json::to_string(&original).unwrap();
1896 let decoded: ProviderConfig = serde_json::from_str(&json).unwrap();
1897 assert_eq!(original.type_name(), decoded.type_name());
1898 }
1899 }
1900
1901 #[test]
1902 fn test_model_config_serde_roundtrip() {
1903 let original = sample_model_config("nvidia", "test-model");
1904 let json = serde_json::to_string(&original).unwrap();
1905 let decoded: ModelConfig = serde_json::from_str(&json).unwrap();
1906 assert_eq!(decoded.provider, original.provider);
1907 assert_eq!(decoded.model, original.model);
1908 assert_eq!(decoded.temperature, original.temperature);
1909 assert_eq!(decoded.max_tokens, original.max_tokens);
1910 }
1911
1912 #[test]
1913 fn test_config_factory_from_config() {
1914 let mut providers = HashMap::new();
1915 providers.insert("nvidia".to_string(), sample_openai_provider());
1916 let mut models = HashMap::new();
1917 models.insert(
1918 "fast".to_string(),
1919 sample_model_config("nvidia", "test-model"),
1920 );
1921
1922 let factory = ConfigBasedLLMFactory::from_config(providers, models, None).unwrap();
1923 assert_eq!(factory.default_model(), "fast");
1924 assert!(factory.registry().has_model("fast"));
1925 }
1926
1927 #[test]
1928 fn test_config_factory_from_config_no_models() {
1929 let factory =
1930 ConfigBasedLLMFactory::from_config(HashMap::new(), HashMap::new(), None).unwrap();
1931 assert_eq!(factory.default_model(), "nvidia/nemotron-3-ultra-550b-a55b");
1932 }
1933
1934 #[tokio::test]
1935 async fn test_create_client_for_model_not_found() {
1936 let registry = ProviderRegistry::new();
1937 assert_configuration_error(
1938 registry.create_client_for_model("missing").await,
1939 "Model 'missing' not found in configuration",
1940 );
1941 }
1942
1943 #[tokio::test]
1944 async fn test_create_client_for_model_missing_provider() {
1945 let mut registry = ProviderRegistry::new();
1946 registry.register_model(
1947 "orphan",
1948 sample_model_config("missing-provider", "some-model"),
1949 );
1950 assert_configuration_error(
1951 registry.create_client_for_model("orphan").await,
1952 "Provider 'missing-provider' referenced by model 'orphan' not found",
1953 );
1954 }
1955
1956 #[tokio::test]
1957 async fn test_create_client_for_provider_not_found() {
1958 let registry = ProviderRegistry::new();
1959 assert_configuration_error(
1960 registry.create_client_for_provider("missing").await,
1961 "Provider 'missing' not found in configuration",
1962 );
1963 }
1964
1965 #[tokio::test]
1966 async fn test_create_default_client_without_default_model() {
1967 let registry = ProviderRegistry::new();
1968 assert_configuration_error(
1969 registry.create_default_client().await,
1970 "No default model configured",
1971 );
1972 }
1973
1974 #[tokio::test]
1975 async fn test_create_client_for_requirements_no_match() {
1976 let registry = create_test_registry();
1977 let requirements = CapabilityRequirements::builder()
1978 .requires_local()
1979 .requires_vision()
1980 .build();
1981 assert_configuration_error(
1982 registry.create_client_for_requirements(&requirements).await,
1983 "No model found matching requirements",
1984 );
1985 }
1986
1987 #[test]
1988 fn get_provider_for_ctx_derives_tenant_from_isolate_namespace() {
1989 let registry = ProviderRegistry::new();
1997 let global = RuntimeProviderEntry {
1998 tenant_id: None,
1999 display_name: "Global Shared".to_string(),
2000 provider_type: "openai-compatible".to_string(),
2001 api_base: "https://global.example.com/v1".to_string(),
2002 auth_type: "api_key".to_string(),
2003 default_model: Some("global-model".to_string()),
2004 headers: HashMap::new(),
2005 api_key: Some("global-key".to_string()),
2006 enabled: true,
2007 };
2008 let tenant = RuntimeProviderEntry {
2009 tenant_id: Some("tenant-a".to_string()),
2010 display_name: "Tenant Shared".to_string(),
2011 provider_type: "openai-compatible".to_string(),
2012 api_base: "https://tenant.example.com/v1".to_string(),
2013 auth_type: "api_key".to_string(),
2014 default_model: Some("tenant-model".to_string()),
2015 headers: HashMap::new(),
2016 api_key: Some("tenant-key".to_string()),
2017 enabled: true,
2018 };
2019 registry.reload_runtime_providers(
2020 vec![global, tenant],
2021 vec!["shared-runtime".to_string(), "shared-runtime".to_string()],
2022 );
2023
2024 let ctx: Arc<cordis::Context> = cordis::Context::new_root();
2025
2026 let fleet = registry.get_provider_for_ctx(&ctx, "shared-runtime");
2029 assert!(
2030 fleet.is_some(),
2031 "untagged ctx should resolve the fleet provider"
2032 );
2033
2034 let tenant_ctx = ctx.isolate::<crate::Llm>("tenant:tenant-a");
2038 let tenant_provider = registry.get_provider_for_ctx(&tenant_ctx, "shared-runtime");
2039 assert!(
2040 tenant_provider.is_some(),
2041 "tenant:tenant-a isolated ctx should resolve a provider"
2042 );
2043 }
2044
2045 #[test]
2046 fn tenant_from_ctx_reads_tenant_context_intercept() {
2047 let ctx: Arc<cordis::Context> = cordis::Context::new_root();
2051 assert_eq!(tenant_from_ctx(&ctx), None);
2052
2053 let intercepted = ctx.with_intercept(ares_types::models::TenantContext::new(
2054 "tenant-a".into(),
2055 ares_types::models::TenantTier::Pro,
2056 ));
2057 assert_eq!(tenant_from_ctx(&intercepted), Some("tenant-a".to_string()));
2058
2059 let registry = ProviderRegistry::new();
2061 let global = RuntimeProviderEntry {
2062 tenant_id: None,
2063 display_name: "Global Shared".to_string(),
2064 provider_type: "openai-compatible".to_string(),
2065 api_base: "https://global.example.com/v1".to_string(),
2066 auth_type: "api_key".to_string(),
2067 default_model: Some("global-model".to_string()),
2068 headers: HashMap::new(),
2069 api_key: Some("global-key".to_string()),
2070 enabled: true,
2071 };
2072 let tenant = RuntimeProviderEntry {
2073 tenant_id: Some("tenant-a".to_string()),
2074 display_name: "Tenant Shared".to_string(),
2075 provider_type: "openai-compatible".to_string(),
2076 api_base: "https://tenant.example.com/v1".to_string(),
2077 auth_type: "api_key".to_string(),
2078 default_model: Some("tenant-model".to_string()),
2079 headers: HashMap::new(),
2080 api_key: Some("tenant-key".to_string()),
2081 enabled: true,
2082 };
2083 registry.reload_runtime_providers(
2084 vec![global, tenant],
2085 vec!["shared-runtime".to_string(), "shared-runtime".to_string()],
2086 );
2087 let provider = registry
2088 .get_provider_for_ctx(&intercepted, "shared-runtime")
2089 .expect("intercept-only ctx should resolve the tenant-a provider");
2090 assert_eq!(
2091 ProviderRegistry::provider_default_model(&provider),
2092 "tenant-model"
2093 );
2094 }
2095
2096 #[test]
2097 fn tenant_from_ctx_isolate_label_wins_over_intercept() {
2098 let ctx: Arc<cordis::Context> = cordis::Context::new_root();
2102 let intercepted = ctx.with_intercept(ares_types::models::TenantContext::new(
2103 "from-intercept".into(),
2104 ares_types::models::TenantTier::Pro,
2105 ));
2106 let isolated = intercepted.isolate::<crate::Llm>("tenant:from-isolate");
2107 assert_eq!(tenant_from_ctx(&isolated), Some("from-isolate".to_string()));
2108 }
2109}
2110
2111pub(crate) fn tenant_from_ctx(ctx: &std::sync::Arc<cordis::Context>) -> Option<String> {
2118 ctx.isolate_label(std::any::TypeId::of::<crate::Llm>())
2119 .and_then(|label| {
2120 label
2121 .strip_prefix("tenant:")
2122 .or_else(|| label.strip_prefix("user:"))
2123 .map(|s| s.to_string())
2124 .filter(|s| !s.is_empty())
2125 })
2126 .or_else(|| {
2127 ctx.get::<ares_types::models::TenantContext>()
2128 .map(|tc| tc.tenant_id.clone())
2129 .filter(|s| !s.is_empty())
2130 })
2131}
2132
2133impl cordis::Service for ProviderRegistry {
2136 fn name(&self) -> &'static str {
2137 "provider_registry"
2138 }
2139 fn init(&self, _ctx: &std::sync::Arc<cordis::Context>) -> cordis::ServiceInitFuture<'_> {
2140 Box::pin(async { Ok(None) })
2141 }
2142 fn check(&self) -> bool {
2143 true
2144 }
2145}
2146
2147impl cordis::Service for ConfigBasedLLMFactory {
2149 fn name(&self) -> &'static str {
2150 "llm_factory"
2151 }
2152 fn init(&self, _ctx: &std::sync::Arc<cordis::Context>) -> cordis::ServiceInitFuture<'_> {
2153 Box::pin(async { Ok(None) })
2154 }
2155 fn check(&self) -> bool {
2156 true
2157 }
2158}