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
1287pub(crate) fn tenant_from_ctx(ctx: &std::sync::Arc<cordis::Context>) -> Option<String> {
1294 ctx.isolate_label(std::any::TypeId::of::<crate::Llm>())
1295 .and_then(|label| {
1296 label
1297 .strip_prefix("tenant:")
1298 .or_else(|| label.strip_prefix("user:"))
1299 .map(|s| s.to_string())
1300 .filter(|s| !s.is_empty())
1301 })
1302 .or_else(|| {
1303 ctx.get::<ares_types::models::TenantContext>()
1304 .map(|tc| tc.tenant_id.clone())
1305 .filter(|s| !s.is_empty())
1306 })
1307}
1308
1309impl cordis::Service for ProviderRegistry {
1312 fn name(&self) -> &'static str {
1313 "provider_registry"
1314 }
1315 fn init(&self, _ctx: &std::sync::Arc<cordis::Context>) -> cordis::ServiceInitFuture<'_> {
1316 Box::pin(async { Ok(None) })
1317 }
1318 fn check(&self) -> bool {
1319 true
1320 }
1321}
1322
1323impl cordis::Service for ConfigBasedLLMFactory {
1325 fn name(&self) -> &'static str {
1326 "llm_factory"
1327 }
1328 fn init(&self, _ctx: &std::sync::Arc<cordis::Context>) -> cordis::ServiceInitFuture<'_> {
1329 Box::pin(async { Ok(None) })
1330 }
1331 fn check(&self) -> bool {
1332 true
1333 }
1334}
1335
1336#[cfg(test)]
1337mod tests {
1338 use super::*;
1339 use crate::capabilities::CapabilityRequirements;
1340
1341 use crate::config::{ModelConfig, ProviderConfig};
1342 use std::collections::HashMap;
1343
1344 fn sample_openai_provider() -> ProviderConfig {
1345 ProviderConfig::OpenAI {
1346 api_key_env: "TEST_KEY".to_string(),
1347 api_base: "https://test.example.com/v1".to_string(),
1348 default_model: "test-model".to_string(),
1349 }
1350 }
1351
1352 fn sample_model_config(provider: &str, model: &str) -> ModelConfig {
1353 ModelConfig {
1354 provider: provider.to_string(),
1355 model: model.to_string(),
1356 temperature: 0.7,
1357 max_tokens: 512,
1358 }
1359 }
1360
1361 fn from_maps(
1362 providers: HashMap<String, ProviderConfig>,
1363 models: HashMap<String, ModelConfig>,
1364 ) -> crate::provider_registry::ProviderRegistry {
1365 ProviderRegistry::from_config(providers, models, None)
1366 }
1367
1368 fn assert_configuration_error<T>(result: Result<T>, expected_substring: &str) {
1369 match result {
1370 Err(AppError::Configuration(msg)) => {
1371 assert!(
1372 msg.contains(expected_substring),
1373 "expected message containing {expected_substring:?}, got {msg:?}"
1374 );
1375 }
1376 Err(other) => panic!("expected Configuration error, got: {other:?}"),
1377 Ok(_) => {
1378 panic!("expected Configuration error containing {expected_substring:?}, got Ok")
1379 }
1380 }
1381 }
1382
1383 #[test]
1384 fn test_empty_registry() {
1385 let registry = ProviderRegistry::new();
1386 assert!(registry.provider_names().is_empty());
1387 assert!(registry.model_names().is_empty());
1388 }
1389
1390 #[test]
1391 fn test_register_provider() {
1392 let mut registry = ProviderRegistry::new();
1393 registry.register_provider(
1394 "nvidia",
1395 ProviderConfig::OpenAI {
1396 api_key_env: "TEST_KEY".to_string(),
1397 api_base: "https://test.example.com/v1".to_string(),
1398 default_model: "test-model".to_string(),
1399 },
1400 );
1401
1402 assert!(registry.has_provider("nvidia"));
1403 assert!(!registry.has_provider("nonexistent"));
1404 }
1405
1406 #[test]
1407 fn test_register_model() {
1408 let mut registry = ProviderRegistry::new();
1409 registry.register_provider(
1410 "nvidia",
1411 ProviderConfig::OpenAI {
1412 api_key_env: "TEST_KEY".to_string(),
1413 api_base: "https://test.example.com/v1".to_string(),
1414 default_model: "test-model".to_string(),
1415 },
1416 );
1417 registry.register_model(
1418 "fast",
1419 ModelConfig {
1420 provider: "nvidia".to_string(),
1421 model: "test-model".to_string(),
1422 temperature: 0.7,
1423 max_tokens: 256,
1424 },
1425 );
1426
1427 assert!(registry.has_model("fast"));
1428 assert!(!registry.has_model("nonexistent"));
1429 }
1430
1431 fn create_test_registry() -> ProviderRegistry {
1434 let mut registry = ProviderRegistry::new();
1435
1436 registry.register_provider(
1437 "nvidia",
1438 ProviderConfig::OpenAI {
1439 api_key_env: "TEST_KEY".to_string(),
1440 api_base: "https://integrate.api.nvidia.com/v1".to_string(),
1441 default_model: "nvidia/nemotron-3-ultra-550b-a55b".to_string(),
1442 },
1443 );
1444
1445 registry.register_model(
1446 "fast-local",
1447 ModelConfig {
1448 provider: "nvidia".to_string(),
1449 model: "nvidia/nemotron-3-ultra-550b-a55b".to_string(),
1450 temperature: 0.7,
1451 max_tokens: 512,
1452 },
1453 );
1454
1455 registry.register_model(
1456 "powerful-local",
1457 ModelConfig {
1458 provider: "nvidia".to_string(),
1459 model: "nvidia/nemotron-3-ultra-550b-a55b".to_string(),
1460 temperature: 0.7,
1461 max_tokens: 2048,
1462 },
1463 );
1464
1465 registry.register_model(
1466 "qwen",
1467 ModelConfig {
1468 provider: "nvidia".to_string(),
1469 model: "qwen/qwen-32b".to_string(),
1470 temperature: 0.7,
1471 max_tokens: 4096,
1472 },
1473 );
1474
1475 registry
1476 }
1477
1478 #[test]
1479 fn test_get_model_capabilities() {
1480 let registry = create_test_registry();
1481
1482 let fast_caps = registry.get_model_capabilities("fast-local").unwrap();
1483 assert!(!fast_caps.is_local);
1484 assert!(fast_caps.supports_tools);
1485 }
1486
1487 #[test]
1488 fn test_models_with_capabilities() {
1489 let registry = create_test_registry();
1490 let models = registry.models_with_capabilities();
1491
1492 assert_eq!(models.len(), 3);
1493
1494 for model in &models {
1495 assert!(!model.name.is_empty());
1496 assert!(!model.provider.is_empty());
1497 assert!(model.capabilities.supports_tools);
1498 }
1499 }
1500
1501 #[test]
1502 fn test_find_local_models() {
1503 let registry = create_test_registry();
1504 let local_models = registry.find_local_models();
1505 assert!(local_models.is_empty());
1507 }
1508
1509 #[test]
1510 fn test_find_vision_models() {
1511 let registry = create_test_registry();
1512 let vision_models = registry.find_vision_models();
1513 assert!(vision_models.is_empty());
1515 }
1516
1517 #[test]
1518 fn test_find_best_model_for_agent() {
1519 let registry = create_test_registry();
1520
1521 let requirements = CapabilityRequirements::for_agent();
1522 let best = registry.find_best_model(&requirements);
1523
1524 assert!(best.is_some());
1525 let best = best.unwrap();
1526 assert!(best.capabilities.supports_tools);
1527 assert!(best.capabilities.production_ready);
1528 }
1529
1530 #[test]
1531 fn test_find_best_model_with_context_window() {
1532 let registry = create_test_registry();
1533
1534 let requirements = CapabilityRequirements::builder()
1535 .min_context_window(100_000)
1536 .build();
1537
1538 let matches = registry.find_models(&requirements);
1539
1540 assert!(matches.len() >= 2);
1541 for model in &matches {
1542 assert!(model.capabilities.context_window >= 100_000);
1543 }
1544 }
1545
1546 #[test]
1547 fn test_find_best_model_prefers_cheaper() {
1548 let registry = create_test_registry();
1549
1550 let requirements = CapabilityRequirements::builder().requires_tools().build();
1551
1552 let best = registry.find_best_model(&requirements).unwrap();
1553
1554 assert_eq!(best.capabilities.cost_tier, "free");
1556 }
1557
1558 #[test]
1559 fn test_no_model_matches_impossible_requirements() {
1560 let registry = create_test_registry();
1561
1562 let requirements = CapabilityRequirements::builder()
1563 .requires_local()
1564 .requires_vision()
1565 .build();
1566
1567 let matches = registry.find_models(&requirements);
1568 assert!(matches.is_empty());
1569 }
1570
1571 #[test]
1572 fn test_find_coding_models() {
1573 let registry = create_test_registry();
1574 let coding_models = registry.find_coding_models();
1575
1576 for model in &coding_models {
1577 assert!(model.capabilities.supports_tools);
1578 assert!(model.capabilities.supports_reasoning);
1579 assert!(model.capabilities.context_window >= 32_000);
1580 }
1581 }
1582
1583 #[test]
1584 fn test_unregister_provider() {
1585 let mut registry = ProviderRegistry::new();
1586 registry.register_provider(
1587 "nvidia",
1588 ProviderConfig::OpenAI {
1589 api_key_env: "TEST_KEY".to_string(),
1590 api_base: "https://test.example.com/v1".to_string(),
1591 default_model: "test-model".to_string(),
1592 },
1593 );
1594 assert!(registry.has_provider("nvidia"));
1595 let removed = registry.unregister_provider("nvidia").unwrap();
1596 assert!(matches!(removed, ProviderConfig::OpenAI { .. }));
1597 assert!(!registry.has_provider("nvidia"));
1598 }
1599
1600 #[test]
1601 fn test_unregister_model() {
1602 let mut registry = create_test_registry();
1603 assert!(registry.has_model("fast-local"));
1604 registry.unregister_model("fast-local");
1605 assert!(!registry.has_model("fast-local"));
1606 }
1607
1608 #[test]
1609 fn test_lookup_provider_by_name() {
1610 let registry = create_test_registry();
1611 let provider = registry.get_provider("nvidia").unwrap();
1612 assert!(matches!(provider, ProviderConfig::OpenAI { .. }));
1613 assert!(registry.get_provider("missing").is_none());
1614 }
1615
1616 #[cfg(feature = "openai")]
1617 #[test]
1618 fn test_runtime_openai_provider_preserves_key_and_headers() {
1619 let mut headers = HashMap::new();
1620 headers.insert("X-Test-Header".to_string(), "runtime-value".to_string());
1621 let entry = RuntimeProviderEntry {
1622 tenant_id: None,
1623 display_name: "Runtime OpenAI".to_string(),
1624 provider_type: "openai-compatible".to_string(),
1625 api_base: "https://runtime.example.com/v1".to_string(),
1626 auth_type: "api_key".to_string(),
1627 default_model: Some("runtime-model".to_string()),
1628 headers,
1629 api_key: Some("resolved-runtime-key".to_string()),
1630 enabled: true,
1631 };
1632
1633 let provider = ProviderRegistry::provider_from_runtime_entry("runtime-openai", &entry)
1634 .expect("runtime provider should resolve");
1635 match provider {
1636 Provider::RuntimeOpenAI {
1637 api_key,
1638 api_base,
1639 model,
1640 headers,
1641 ..
1642 } => {
1643 assert_eq!(api_key, "resolved-runtime-key");
1644 assert_eq!(api_base, "https://runtime.example.com/v1");
1645 assert_eq!(model, "runtime-model");
1646 assert_eq!(
1647 headers.get("X-Test-Header").map(String::as_str),
1648 Some("runtime-value")
1649 );
1650 }
1651 _ => panic!("expected RuntimeOpenAI provider"),
1652 }
1653 }
1654
1655 #[cfg(feature = "openai")]
1656 #[test]
1657 fn test_runtime_provider_requires_resolved_api_key() {
1658 let entry = RuntimeProviderEntry {
1659 tenant_id: None,
1660 display_name: "Runtime OpenAI".to_string(),
1661 provider_type: "openai-compatible".to_string(),
1662 api_base: "https://runtime.example.com/v1".to_string(),
1663 auth_type: "api_key".to_string(),
1664 default_model: Some("runtime-model".to_string()),
1665 headers: HashMap::new(),
1666 api_key: None,
1667 enabled: true,
1668 };
1669
1670 assert_configuration_error(
1671 ProviderRegistry::provider_from_runtime_entry("runtime-openai", &entry),
1672 "Runtime provider 'runtime-openai' API key is not resolved",
1673 );
1674 }
1675
1676 #[test]
1677 fn runtime_provider_visibility_respects_tenant_scope() {
1678 let registry = ProviderRegistry::new();
1679 let global = RuntimeProviderEntry {
1680 tenant_id: None,
1681 display_name: "Global Runtime".to_string(),
1682 provider_type: "openai-compatible".to_string(),
1683 api_base: "https://global.example.com/v1".to_string(),
1684 auth_type: "api_key".to_string(),
1685 default_model: Some("global-model".to_string()),
1686 headers: HashMap::new(),
1687 api_key: Some("global-key".to_string()),
1688 enabled: true,
1689 };
1690 let scoped = RuntimeProviderEntry {
1691 tenant_id: Some("tenant-a".to_string()),
1692 display_name: "Scoped Runtime".to_string(),
1693 provider_type: "openai-compatible".to_string(),
1694 api_base: "https://tenant.example.com/v1".to_string(),
1695 auth_type: "api_key".to_string(),
1696 default_model: Some("tenant-model".to_string()),
1697 headers: HashMap::new(),
1698 api_key: Some("tenant-key".to_string()),
1699 enabled: true,
1700 };
1701 registry.reload_runtime_providers(
1702 vec![global, scoped],
1703 vec!["global-runtime".to_string(), "tenant-runtime".to_string()],
1704 );
1705
1706 assert!(registry.has_provider("global-runtime"));
1707 assert!(!registry.has_provider("tenant-runtime"));
1708 assert!(registry.has_provider_for_tenant("tenant-runtime", Some("tenant-a")));
1709 assert!(!registry.has_provider_for_tenant("tenant-runtime", Some("tenant-b")));
1710 assert!(
1711 registry
1712 .provider_for_tenant("tenant-runtime", Some("tenant-a"))
1713 .is_some()
1714 );
1715 assert!(
1716 registry
1717 .provider_for_tenant("tenant-runtime", Some("tenant-b"))
1718 .is_none()
1719 );
1720 assert_eq!(
1721 registry.provider_names(),
1722 vec!["global-runtime".to_string()]
1723 );
1724 }
1725
1726 #[test]
1727 fn runtime_provider_lookup_prefers_tenant_override_same_name() {
1728 let registry = ProviderRegistry::new();
1729 let global = RuntimeProviderEntry {
1730 tenant_id: None,
1731 display_name: "Global Shared".to_string(),
1732 provider_type: "openai-compatible".to_string(),
1733 api_base: "https://global.example.com/v1".to_string(),
1734 auth_type: "api_key".to_string(),
1735 default_model: Some("global-model".to_string()),
1736 headers: HashMap::new(),
1737 api_key: Some("global-key".to_string()),
1738 enabled: true,
1739 };
1740 let tenant = RuntimeProviderEntry {
1741 tenant_id: Some("tenant-a".to_string()),
1742 display_name: "Tenant Shared".to_string(),
1743 provider_type: "openai-compatible".to_string(),
1744 api_base: "https://tenant.example.com/v1".to_string(),
1745 auth_type: "api_key".to_string(),
1746 default_model: Some("tenant-model".to_string()),
1747 headers: HashMap::new(),
1748 api_key: Some("tenant-key".to_string()),
1749 enabled: true,
1750 };
1751 registry.reload_runtime_providers(
1752 vec![global, tenant],
1753 vec!["shared-runtime".to_string(), "shared-runtime".to_string()],
1754 );
1755
1756 assert!(registry.has_provider("shared-runtime"));
1757 assert!(registry.has_provider_for_tenant("shared-runtime", Some("tenant-a")));
1758 assert!(registry.has_provider_for_tenant("shared-runtime", Some("tenant-b")));
1759
1760 let tenant_provider = registry
1761 .provider_for_tenant("shared-runtime", Some("tenant-a"))
1762 .expect("tenant provider");
1763 let global_provider = registry
1764 .provider_for_tenant("shared-runtime", Some("tenant-b"))
1765 .expect("global fallback provider");
1766 assert_eq!(
1767 ProviderRegistry::provider_default_model(&tenant_provider),
1768 "tenant-model"
1769 );
1770 assert_eq!(
1771 ProviderRegistry::provider_default_model(&global_provider),
1772 "global-model"
1773 );
1774 assert_eq!(
1775 registry.provider_names(),
1776 vec!["shared-runtime".to_string()]
1777 );
1778 }
1779
1780 #[test]
1781 fn test_default_registry() {
1782 let registry = ProviderRegistry::default();
1783 assert!(registry.provider_names().is_empty());
1784 assert!(registry.model_names().is_empty());
1785 }
1786
1787 #[test]
1788 fn test_register_provider_overwrites_existing() {
1789 let mut registry = ProviderRegistry::new();
1790 registry.register_provider(
1791 "nvidia",
1792 ProviderConfig::OpenAI {
1793 api_key_env: "TEST_KEY".to_string(),
1794 api_base: "https://old.example.com/v1".to_string(),
1795 default_model: "old-model".to_string(),
1796 },
1797 );
1798 registry.register_provider(
1799 "nvidia",
1800 ProviderConfig::OpenAI {
1801 api_key_env: "TEST_KEY".to_string(),
1802 api_base: "https://new.example.com/v1".to_string(),
1803 default_model: "new-model".to_string(),
1804 },
1805 );
1806
1807 let provider = registry.get_provider("nvidia").unwrap();
1808 if let ProviderConfig::OpenAI { default_model, .. } = provider {
1809 assert_eq!(default_model, "new-model");
1810 } else {
1811 panic!("expected OpenAI provider");
1812 }
1813 }
1814
1815 #[test]
1816 fn test_provider_and_model_name_iteration() {
1817 let mut registry = ProviderRegistry::new();
1818 registry.register_provider("alpha", sample_openai_provider());
1819 registry.register_provider("beta", sample_openai_provider());
1820 registry.register_model("m1", sample_model_config("alpha", "model-a"));
1821 registry.register_model("m2", sample_model_config("beta", "model-b"));
1822
1823 let mut provider_names = registry.provider_names();
1824 provider_names.sort_unstable();
1825 assert_eq!(provider_names, vec!["alpha", "beta"]);
1826
1827 let mut model_names = registry.model_names();
1828 model_names.sort_unstable();
1829 assert_eq!(model_names, vec!["m1", "m2"]);
1830 }
1831
1832 #[test]
1833 fn test_lookup_model_by_name() {
1834 let mut registry = ProviderRegistry::new();
1835 registry.register_provider("nvidia", sample_openai_provider());
1836 registry.register_model("fast", sample_model_config("nvidia", "test-model"));
1837
1838 let model = registry.get_model("fast").unwrap();
1839 assert_eq!(model.provider, "nvidia");
1840 assert_eq!(model.model, "test-model");
1841 assert!(registry.get_model("missing").is_none());
1842 }
1843
1844 #[test]
1845 fn test_list_models_returns_registered_entries() {
1846 let registry = create_test_registry();
1847 let models = registry.list_models();
1848
1849 assert_eq!(models.len(), 3);
1850 let fast = models
1851 .iter()
1852 .find(|m| m.name == "fast-local" && m.provider == "nvidia")
1853 .expect("fast model");
1854 assert!(fast.supports_reasoning);
1855 assert!(fast.supports_streaming);
1856 }
1857
1858 #[test]
1859 fn test_from_config_loads_providers_and_models() {
1860 let mut providers = HashMap::new();
1861 providers.insert("nvidia".to_string(), sample_openai_provider());
1862 let mut models = HashMap::new();
1863 models.insert(
1864 "fast".to_string(),
1865 sample_model_config("nvidia", "test-model"),
1866 );
1867
1868 let registry = from_maps(providers, models);
1869
1870 assert!(registry.has_provider("nvidia"));
1871 assert!(registry.has_model("fast"));
1872
1873 #[cfg(feature = "bedrock")]
1874 let expected = vec![
1875 "fast",
1876 "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
1877 ];
1878 #[cfg(not(feature = "bedrock"))]
1879 let expected = vec!["fast"];
1880 assert_eq!(registry.model_names(), expected);
1881 }
1882
1883 #[tokio::test]
1884 async fn test_set_default_model() {
1885 let mut registry = create_test_registry();
1886 registry.set_default_model("powerful-local");
1887 let saved_key = std::env::var("TEST_KEY").ok();
1890 std::env::remove_var("TEST_KEY");
1891 let result = registry.create_default_client().await;
1892 if let Some(key) = saved_key {
1893 std::env::set_var("TEST_KEY", key);
1894 }
1895 match result {
1896 Err(AppError::Configuration(msg)) => {
1897 assert!(!msg.contains("No default model configured"), "got: {msg}");
1898 }
1899 Err(other) => panic!("expected Configuration error, got: {other:?}"),
1900 Ok(_) => panic!("expected Configuration error, but client creation succeeded"),
1901 }
1902 }
1903
1904 #[test]
1905 fn test_get_model_capabilities_unknown_model() {
1906 let registry = create_test_registry();
1907 assert!(registry.get_model_capabilities("missing").is_none());
1908 }
1909
1910 #[test]
1911 fn test_get_model_capabilities_missing_provider() {
1912 let mut registry = ProviderRegistry::new();
1913 registry.register_model(
1914 "orphan",
1915 sample_model_config("missing-provider", "some-model"),
1916 );
1917 assert!(registry.get_model_capabilities("orphan").is_none());
1918 }
1919
1920 #[test]
1921 fn test_unregister_provider_missing_returns_none() {
1922 let mut registry = ProviderRegistry::new();
1923 assert!(registry.unregister_provider("missing").is_none());
1924 }
1925
1926 #[test]
1927 fn test_unregister_model_returns_removed_config() {
1928 let mut registry = create_test_registry();
1929 let removed = registry.unregister_model("fast-local").unwrap();
1930 assert_eq!(removed.provider, "nvidia");
1931 assert_eq!(removed.model, "nvidia/nemotron-3-ultra-550b-a55b");
1932 assert!(registry.unregister_model("fast-local").is_none());
1933 }
1934
1935 #[test]
1936 fn test_provider_config_serde_roundtrip() {
1937 let configs = [ProviderConfig::OpenAI {
1938 api_key_env: "OPENAI_API_KEY".to_string(),
1939 api_base: "https://api.openai.com/v1".to_string(),
1940 default_model: "gpt-4o".to_string(),
1941 }];
1942
1943 for original in configs {
1944 let json = serde_json::to_string(&original).unwrap();
1945 let decoded: ProviderConfig = serde_json::from_str(&json).unwrap();
1946 assert_eq!(original.type_name(), decoded.type_name());
1947 }
1948 }
1949
1950 #[test]
1951 fn test_model_config_serde_roundtrip() {
1952 let original = sample_model_config("nvidia", "test-model");
1953 let json = serde_json::to_string(&original).unwrap();
1954 let decoded: ModelConfig = serde_json::from_str(&json).unwrap();
1955 assert_eq!(decoded.provider, original.provider);
1956 assert_eq!(decoded.model, original.model);
1957 assert_eq!(decoded.temperature, original.temperature);
1958 assert_eq!(decoded.max_tokens, original.max_tokens);
1959 }
1960
1961 #[test]
1962 fn test_config_factory_from_config() {
1963 let mut providers = HashMap::new();
1964 providers.insert("nvidia".to_string(), sample_openai_provider());
1965 let mut models = HashMap::new();
1966 models.insert(
1967 "fast".to_string(),
1968 sample_model_config("nvidia", "test-model"),
1969 );
1970
1971 let factory = ConfigBasedLLMFactory::from_config(providers, models, None).unwrap();
1972 assert_eq!(factory.default_model(), "fast");
1973 assert!(factory.registry().has_model("fast"));
1974 }
1975
1976 #[test]
1977 fn test_config_factory_from_config_no_models() {
1978 let factory =
1979 ConfigBasedLLMFactory::from_config(HashMap::new(), HashMap::new(), None).unwrap();
1980 assert_eq!(factory.default_model(), "nvidia/nemotron-3-ultra-550b-a55b");
1981 }
1982
1983 #[tokio::test]
1984 async fn test_create_client_for_model_not_found() {
1985 let registry = ProviderRegistry::new();
1986 assert_configuration_error(
1987 registry.create_client_for_model("missing").await,
1988 "Model 'missing' not found in configuration",
1989 );
1990 }
1991
1992 #[tokio::test]
1993 async fn test_create_client_for_model_missing_provider() {
1994 let mut registry = ProviderRegistry::new();
1995 registry.register_model(
1996 "orphan",
1997 sample_model_config("missing-provider", "some-model"),
1998 );
1999 assert_configuration_error(
2000 registry.create_client_for_model("orphan").await,
2001 "Provider 'missing-provider' referenced by model 'orphan' not found",
2002 );
2003 }
2004
2005 #[tokio::test]
2006 async fn test_create_client_for_provider_not_found() {
2007 let registry = ProviderRegistry::new();
2008 assert_configuration_error(
2009 registry.create_client_for_provider("missing").await,
2010 "Provider 'missing' not found in configuration",
2011 );
2012 }
2013
2014 #[tokio::test]
2015 async fn test_create_default_client_without_default_model() {
2016 let registry = ProviderRegistry::new();
2017 assert_configuration_error(
2018 registry.create_default_client().await,
2019 "No default model configured",
2020 );
2021 }
2022
2023 #[tokio::test]
2024 async fn test_create_client_for_requirements_no_match() {
2025 let registry = create_test_registry();
2026 let requirements = CapabilityRequirements::builder()
2027 .requires_local()
2028 .requires_vision()
2029 .build();
2030 assert_configuration_error(
2031 registry.create_client_for_requirements(&requirements).await,
2032 "No model found matching requirements",
2033 );
2034 }
2035
2036 #[test]
2037 fn get_provider_for_ctx_derives_tenant_from_isolate_namespace() {
2038 let registry = ProviderRegistry::new();
2046 let global = RuntimeProviderEntry {
2047 tenant_id: None,
2048 display_name: "Global Shared".to_string(),
2049 provider_type: "openai-compatible".to_string(),
2050 api_base: "https://global.example.com/v1".to_string(),
2051 auth_type: "api_key".to_string(),
2052 default_model: Some("global-model".to_string()),
2053 headers: HashMap::new(),
2054 api_key: Some("global-key".to_string()),
2055 enabled: true,
2056 };
2057 let tenant = RuntimeProviderEntry {
2058 tenant_id: Some("tenant-a".to_string()),
2059 display_name: "Tenant Shared".to_string(),
2060 provider_type: "openai-compatible".to_string(),
2061 api_base: "https://tenant.example.com/v1".to_string(),
2062 auth_type: "api_key".to_string(),
2063 default_model: Some("tenant-model".to_string()),
2064 headers: HashMap::new(),
2065 api_key: Some("tenant-key".to_string()),
2066 enabled: true,
2067 };
2068 registry.reload_runtime_providers(
2069 vec![global, tenant],
2070 vec!["shared-runtime".to_string(), "shared-runtime".to_string()],
2071 );
2072
2073 let ctx: Arc<cordis::Context> = cordis::Context::new_root();
2074
2075 let fleet = registry.get_provider_for_ctx(&ctx, "shared-runtime");
2078 assert!(
2079 fleet.is_some(),
2080 "untagged ctx should resolve the fleet provider"
2081 );
2082
2083 let tenant_ctx = ctx.isolate::<crate::Llm>("tenant:tenant-a");
2087 let tenant_provider = registry.get_provider_for_ctx(&tenant_ctx, "shared-runtime");
2088 assert!(
2089 tenant_provider.is_some(),
2090 "tenant:tenant-a isolated ctx should resolve a provider"
2091 );
2092 }
2093
2094 #[test]
2095 fn tenant_from_ctx_reads_tenant_context_intercept() {
2096 let ctx: Arc<cordis::Context> = cordis::Context::new_root();
2100 assert_eq!(tenant_from_ctx(&ctx), None);
2101
2102 let intercepted = ctx.with_intercept(ares_types::models::TenantContext::new(
2103 "tenant-a".into(),
2104 ares_types::models::TenantTier::Pro,
2105 ));
2106 assert_eq!(tenant_from_ctx(&intercepted), Some("tenant-a".to_string()));
2107
2108 let registry = ProviderRegistry::new();
2110 let global = RuntimeProviderEntry {
2111 tenant_id: None,
2112 display_name: "Global Shared".to_string(),
2113 provider_type: "openai-compatible".to_string(),
2114 api_base: "https://global.example.com/v1".to_string(),
2115 auth_type: "api_key".to_string(),
2116 default_model: Some("global-model".to_string()),
2117 headers: HashMap::new(),
2118 api_key: Some("global-key".to_string()),
2119 enabled: true,
2120 };
2121 let tenant = RuntimeProviderEntry {
2122 tenant_id: Some("tenant-a".to_string()),
2123 display_name: "Tenant Shared".to_string(),
2124 provider_type: "openai-compatible".to_string(),
2125 api_base: "https://tenant.example.com/v1".to_string(),
2126 auth_type: "api_key".to_string(),
2127 default_model: Some("tenant-model".to_string()),
2128 headers: HashMap::new(),
2129 api_key: Some("tenant-key".to_string()),
2130 enabled: true,
2131 };
2132 registry.reload_runtime_providers(
2133 vec![global, tenant],
2134 vec!["shared-runtime".to_string(), "shared-runtime".to_string()],
2135 );
2136 let provider = registry
2137 .get_provider_for_ctx(&intercepted, "shared-runtime")
2138 .expect("intercept-only ctx should resolve the tenant-a provider");
2139 assert_eq!(
2140 ProviderRegistry::provider_default_model(&provider),
2141 "tenant-model"
2142 );
2143 }
2144
2145 #[test]
2146 fn tenant_from_ctx_isolate_label_wins_over_intercept() {
2147 let ctx: Arc<cordis::Context> = cordis::Context::new_root();
2151 let intercepted = ctx.with_intercept(ares_types::models::TenantContext::new(
2152 "from-intercept".into(),
2153 ares_types::models::TenantTier::Pro,
2154 ));
2155 let isolated = intercepted.isolate::<crate::Llm>("tenant:from-isolate");
2156 assert_eq!(tenant_from_ctx(&isolated), Some("from-isolate".to_string()));
2157 }
2158}