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::{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;
33// Phase 3 unified hot-reload: re-export ReflectService from core (single source)
34pub use cordis::ReflectService;
35
36/// Runtime provider entry, synthesized from the DB `runtime_providers` table.
37#[derive(Debug, Clone)]
38pub struct RuntimeProviderEntry {
39    /// Optional tenant owner. `None` means fleet-wide.
40    pub tenant_id: Option<String>,
41    /// Display name for UI purposes.
42    pub display_name: String,
43    /// Provider compatibility type: "openai-compatible", "anthropic-compatible", "bedrock", "custom".
44    pub provider_type: String,
45    /// Base URL for the API.
46    pub api_base: String,
47    /// Authentication type: "api_key", "oauth2", "aws_sigv4".
48    pub auth_type: String,
49    /// Default model when none is specified.
50    pub default_model: Option<String>,
51    /// Extra HTTP headers.
52    pub headers: HashMap<String, String>,
53    /// Resolved API key (populated by the reload path).
54    pub api_key: Option<String>,
55    /// Whether this runtime provider is enabled.
56    pub enabled: bool,
57}
58
59/// Resolved provider plus the concrete model id that must be sent to that provider.
60#[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
69/// Registry for managing multiple named LLM providers
70///
71/// The ProviderRegistry holds references to provider configurations and allows
72/// creating LLM clients for specific models or providers by name.
73pub struct ProviderRegistry {
74    /// Provider configurations keyed by name (legacy, kept for backward compat).
75    providers: HashMap<String, ProviderConfig>,
76    /// Explicit model configurations keyed by name (legacy, kept for backward compat).
77    models: HashMap<String, ModelConfig>,
78    /// Live NVIDIA catalog cache.
79    catalog: Option<Arc<NvidiaCatalogCache>>,
80    /// Default model name to use when none specified.
81    default_model: Option<String>,
82    /// Runtime providers loaded from the DB (hot-swapped).
83    runtime_providers: Arc<ArcSwap<HashMap<String, Vec<RuntimeProviderEntry>>>>,
84}
85
86impl ProviderRegistry {
87    /// Create a new empty provider registry
88    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    /// Create a provider registry from TOML configuration
99    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 no legacy providers are configured, synthesize a single NVIDIA provider.
107        if providers.is_empty() {
108            let nvidia = nvidia.cloned().unwrap_or_default();
109            let _ = std::env::var(&nvidia.api_key_env); // we don't error here; refresh will report it
110            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    /// Hot-swap the runtime provider map. Called by admin endpoints after
142    /// mutating the DB so the new providers are visible immediately.
143    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    /// Attach a live catalog cache (used after construction for background refresh).
158    pub fn with_catalog(mut self, catalog: Arc<NvidiaCatalogCache>) -> Self {
159        self.catalog = Some(catalog);
160        self
161    }
162
163    /// Set the default model name
164    pub fn set_default_model(&mut self, model_name: &str) {
165        self.default_model = Some(model_name.to_string());
166    }
167
168    /// Register a provider configuration (legacy no-op if providers are already managed).
169    pub fn register_provider(&mut self, name: &str, config: ProviderConfig) {
170        self.providers.insert(name.to_string(), config);
171    }
172
173    /// Register a model configuration (legacy backward-compat).
174    pub fn register_model(&mut self, name: &str, config: ModelConfig) {
175        self.models.insert(name.to_string(), config);
176    }
177
178    /// Remove a provider by name.
179    pub fn unregister_provider(&mut self, name: &str) -> Option<ProviderConfig> {
180        self.providers.remove(name)
181    }
182
183    /// Remove a model by name.
184    pub fn unregister_model(&mut self, name: &str) -> Option<ModelConfig> {
185        self.models.remove(name)
186    }
187
188    /// Get a provider configuration by name.
189    ///
190    /// Runtime providers are checked first and synthesized into a [`ProviderConfig`]
191    /// on the fly; legacy static configs are checked second. The return value is
192    /// cloned so that the caller owns it — this is required because runtime
193    /// providers are materialised from the arc-swapped map rather than stored as
194    /// [`ProviderConfig`] internally.
195    pub fn get_provider(&self, name: &str) -> Option<ProviderConfig> {
196        self.provider_for_tenant(name, None)
197    }
198
199    /// Crate-private tenant lookup. The former public tenant-provider getter
200    /// is gone; callers use `Llm::get_client` (or `get_provider_for_ctx` inside
201    /// this crate). Tenant-scoped runtime providers are only visible to their
202    /// owning tenant; fleet-wide runtime providers and static providers remain
203    /// visible to every tenant.
204    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    /// Resolve a provider visible to the tenant derived from the context's isolate
216    /// namespace. Reads `ctx.isolate_label(TypeId::of::<Llm>())` and
217    /// strips a leading `tenant:`/`user:` prefix (mirroring
218    /// `ares_agent::resolver::user_id_from_ctx`), delegating to
219    /// crate-private tenant lookup with the derived tenant (`None` when unlabeled).
220    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    /// Synthesize a legacy [`ProviderConfig`] from a runtime provider entry.
251    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    /// Get a model configuration by name.
477    /// Checks explicit legacy models first, then falls back to the live catalog.
478    pub fn get_model(&self, name: &str) -> Option<ModelConfig> {
479        // 1. explicit legacy models
480        if let Some(cfg) = self.models.get(name) {
481            return Some(cfg.clone());
482        }
483        // 2. direct Bedrock model ids (`bedrock/<model-id>` or Bedrock Anthropic ids)
484        if let Some(model_id) = Self::bedrock_model_id_from_name(name) {
485            return Some(Self::bedrock_model_config(model_id));
486        }
487        // 3. direct Azure model ids (`azure/<model-id>`)
488        if let Some(model_id) = Self::azure_model_id_from_name(name) {
489            return Some(Self::azure_model_config(model_id));
490        }
491        // 4. catalog lookup – synthesize a ModelConfig on the fly
492        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    /// Get all provider names (legacy + runtime).
507    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    /// Get all model names (legacy + catalog ids)
519    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    /// Create an LLM client for a specific model by name.
542    ///
543    /// Fleet-wide tenant resolution (tenant `None`).
544    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    /// Create an LLM client for a specific model by name, deriving the tenant
549    /// from the context's isolate namespace so tenant-scoped runtime providers
550    /// are used when the caller holds a tenant-isolated context.
551    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        // 1. Try legacy explicit models first
567        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        // 2. Try direct Bedrock model routing (`bedrock/<model-id>`).
595        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        // 3. Try direct Azure model routing (`azure/<model-id>`).
605        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        // 4. Try catalog lookup
615        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    /// Create an LLM client for a specific provider by name
636    pub async fn create_client_for_provider(
637        &self,
638        provider_name: &str,
639    ) -> Result<Box<dyn LLMClient>> {
640        // Check runtime providers first so resolved API keys and custom headers are preserved.
641        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    /// Create an LLM client for an already-resolved provider/model pair.
659    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    /// Create an LLM client using the default model
686    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    /// Check if a model exists in the registry
696    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    /// Check if a provider exists in the registry (legacy or runtime).
708    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    // ================== Capability-Based Model Selection (DIR-43) ==================
720
721    /// Get capabilities for a registered model.
722    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 it's a legacy model, use the explicit config
735        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 it's in the catalog, use the catalog id directly
750        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    /// Get all models with their capabilities.
763    pub fn models_with_capabilities(&self) -> Vec<ModelWithCapabilities> {
764        let mut result = Vec::new();
765
766        // Legacy models
767        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        // Catalog models
807        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    /// Find models that satisfy the given capability requirements.
827    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        // Sort by score (highest first)
835        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    /// Find the best model for the given requirements.
845    pub fn find_best_model(
846        &self,
847        requirements: &CapabilityRequirements,
848    ) -> Option<ModelWithCapabilities> {
849        self.find_models(requirements).into_iter().next()
850    }
851
852    /// Create an LLM client for the best model matching requirements.
853    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    /// Find models suitable for agent tasks (tool calling required).
868    pub fn find_agent_models(&self) -> Vec<ModelWithCapabilities> {
869        self.find_models(&CapabilityRequirements::for_agent())
870    }
871
872    /// Find models suitable for vision tasks.
873    pub fn find_vision_models(&self) -> Vec<ModelWithCapabilities> {
874        self.find_models(&CapabilityRequirements::for_vision())
875    }
876
877    /// Find models suitable for coding tasks.
878    pub fn find_coding_models(&self) -> Vec<ModelWithCapabilities> {
879        self.find_models(&CapabilityRequirements::for_coding())
880    }
881
882    /// Find local-only models.
883    pub fn find_local_models(&self) -> Vec<ModelWithCapabilities> {
884        self.find_models(&CapabilityRequirements::for_local())
885    }
886
887    /// List all registered models with their provider info.
888    pub fn list_models(&self) -> Vec<ModelInfo> {
889        let mut models = Vec::new();
890
891        // Legacy explicit models
892        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        // Catalog entries
941        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                // Fallback when catalog is empty: expose the default model so the UI is never blank
961                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            // No catalog at all – still expose the default model
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
988        models
989    }
990
991    /// Capability-aware fallback chain for `Llm`.
992    ///
993    /// Tries `create_client_for_requirements` for the given `requirements` if
994    /// `Some`, then falls back to `create_default_client`. This reuses the
995    /// existing `find_best_model` → `create_client_for_model` chain and the
996    /// coordinator's fallback semantics without requiring a database.
997    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    /// Alias satisfying the `Llm` spec's `resolve_with_fallback` name
1010    /// when the postgres-gated tier resolver is not active.
1011    #[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    // ============================================================
1020    // Helpers
1021    // ============================================================
1022
1023    /// Resolve a model tier or model name to concrete provider/model entries,
1024    /// following the fallback chain stored in `fleet_secrets` for the primary provider.
1025    ///
1026    /// 1. Looks up `tenant_model_tiers` for the tenant + tier.
1027    /// 2. Falls back to the registry's configured models.
1028    /// 3. Falls back to treating `tier_or_model` as a provider name.
1029    /// 4. Loads the primary provider's `fallback_providers` from fleet secrets
1030    ///    and appends each resolved provider with its own concrete default model.
1031    #[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    /// Extract NVIDIA config from the synthetic provider we inserted.
1150    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/// Model info for listing available models via API.
1171#[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
1194/// Configuration-based LLM client factory using the provider registry
1195pub struct ConfigBasedLLMFactory {
1196    registry: Arc<ProviderRegistry>,
1197    default_model: String,
1198}
1199
1200impl ConfigBasedLLMFactory {
1201    /// Create a new factory from a provider registry
1202    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    /// Create a factory from TOML configuration
1210    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    /// Get the provider registry
1229    pub fn registry(&self) -> &Arc<ProviderRegistry> {
1230        &self.registry
1231    }
1232
1233    /// Create an LLM client for a specific model
1234    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    /// Create an LLM client using the default model
1239    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    /// Get the default model name
1246    pub fn default_model(&self) -> &str {
1247        &self.default_model
1248    }
1249
1250    /// Set the default model name
1251    pub fn set_default_model(&mut self, model_name: &str) {
1252        self.default_model = model_name.to_string();
1253    }
1254}
1255
1256// 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.
1257
1258/// Phase 3 unified hot-reload demo — watch channel creation on provide.
1259/// Compile-time proof that notifiers/dependents insertion compiles via ReflectService.
1260pub fn reflect_notify_stub(ctx: &Arc<cordis::Context>) {
1261    // Prove Loader integration still compiles
1262    let _ = ctx.get::<cordis::loader::Loader>();
1263    let tid = TypeId::of::<ProviderRegistry>();
1264    // Prove ReflectService watch channel creation on provide + dependents insertion + BFS notify compiles
1265    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
1273/// Derive the tenant id from the context's isolate namespace for [`Llm`],
1274/// stripping a leading `tenant:`/`user:` prefix.
1275///
1276/// Isolate labels win. When unlabeled for `Llm`, falls back to a
1277/// [`ares_types::models::TenantContext`] intercept (`tenant_id` if non-empty).
1278/// Empty labels/ids yield `None` (fleet-wide resolution).
1279pub(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
1295// Cordis Service impl — allows ctx.get::<ProviderRegistry>() for crate wiring.
1296// Per-tenant provider-secret isolate labels key on TypeId::of::<Llm>(), not this type.
1297impl 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
1309// Cordis Service impl — allows direct ctx.get::<ConfigBasedLLMFactory>() without wrapper
1310impl 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    // ================== DIR-43: Capability Tests ==================
1418
1419    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        // NVIDIA models are not local
1492        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        // No explicit vision models in test registry
1500        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        // NVIDIA models are "free" tier in our heuristic
1541        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        // The fixture's provider reads TEST_KEY; other suites export it, so
1862        // clear it here to keep the missing-key error deterministic.
1863        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        // Cordis design (§4): per-tenant provider resolution should be driven by
2013        // the context's isolate namespace, not a throwaway method param.
2014        // get_provider_for_ctx(ctx, name) reads
2015        // ctx.isolate_label(TypeId::of::<Llm>()) and strips a
2016        // leading 'tenant:' prefix (mirroring resolver::user_id_from_ctx) so a
2017        // tenant-isolated context resolves the tenant-scoped provider.
2018
2019        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        // Untagged context -> no isolate label -> falls back to the fleet-wide
2050        // provider (the shared "shared-runtime" entry's tenant is None).
2051        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        // A tenant-isolated context must drive resolution: the isolate label
2058        // 'tenant:tenant-a' derives tenant 'tenant-a', so the tenant-scoped
2059        // provider wins.
2060        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        // Unlabeled root has no isolate label and no intercept, so tenant is
2071        // None (fleet-wide). A TenantContext intercept is the fallback when
2072        // isolate_label is missing.
2073        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        // Intercept-only ctx should resolve the tenant-a runtime provider.
2083        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        // Isolate label is the primary source even when a TenantContext
2122        // intercept is present. Intercept "from-intercept", then isolate
2123        // as tenant:from-isolate, must yield the isolate id.
2124        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}