Skip to main content

ares_llm/
provider_registry.rs

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