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