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