Skip to main content

ares_llm/
provider_registry.rs

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