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