Skip to main content

ares_llm/
llm_service.rs

1//! Unified LLM capability (Cordis Phase 3).
2//!
3//! Wraps `ProviderRegistry` + `NvidiaCatalogCache` + `ClientPool` +
4//! `ConfigBasedLLMFactory` behind a single `Service` injected via
5//! `ctx.get::<Llm>()`.
6//!
7//! Per-request model pinning uses `ctx.intercept(ModelOverride { model })` so
8//! the override is visible to `Llm` via `ctx.get::<ModelOverride>()`
9//! without mutating global state.
10//!
11//! Circuit-breaker (`Breaker`) causes `Service::check` to return `false` when
12//! open, causing dependent fibers to deactivate (guarded withdrawal per Thm 63:
13//! provider does not withdraw until dependents deactivate).
14
15use std::collections::HashSet;
16use std::future::Future;
17use std::sync::Arc;
18
19use crate::nvidia_catalog::NvidiaCatalogCache;
20use async_trait::async_trait;
21use chrono::{DateTime, Utc};
22use cordis::{Context, CordisError, EventsService, Service};
23use parking_lot::RwLock;
24
25use crate::capabilities::CapabilityRequirements;
26use crate::client::{GenerationHints, LLMClient, LLMResponse};
27use crate::config::ProviderConfig;
28use crate::pool::ClientPool;
29use crate::provider_registry::{
30    ConfigBasedLLMFactory, ModelInfo, ProviderRegistry, RuntimeProviderEntry,
31};
32use ares_types::types::{AppError, ToolDefinition};
33
34/// Per-request model override for `ctx.intercept`.
35///
36/// Example:
37/// ```ignore
38/// let req_ctx = root_ctx.intercept(ModelOverride { model: "gpt-4o-mini".into() });
39/// let llm = req_ctx.get::<Llm>().unwrap();
40/// // inside Llm, check `ctx.get::<ModelOverride>()` for pinning
41/// ```
42#[derive(Debug, Clone)]
43pub struct ModelOverride {
44    /// Model id to pin for this request scope.
45    pub model: String,
46}
47
48// ModelOverride itself can be intercepted, not necessarily a Service provider,
49// but we implement Service so `ctx.intercept(ModelOverride)` and `ctx.get`
50// work via the same `Service` type map when needed.
51impl Service for ModelOverride {}
52
53/// Tenant-scoped model allowlist carried by a request context.
54///
55/// This is a snapshot of the tenant policy (normally populated from
56/// `TenantAllowlistStore`) and is intentionally request-scoped.  It lets the
57/// LLM service authorize an intercepted [`ModelOverride`] without changing the
58/// process-wide provider registry or its default model.
59#[derive(Debug, Clone, PartialEq, Eq)]
60pub struct TenantModelPolicy {
61    tenant_id: String,
62    allowed_models: HashSet<String>,
63}
64
65impl TenantModelPolicy {
66    /// Build a policy from the currently enabled model ids for a tenant.
67    pub fn new<I, S>(tenant_id: impl Into<String>, allowed_models: I) -> Self
68    where
69        I: IntoIterator<Item = S>,
70        S: Into<String>,
71    {
72        Self {
73            tenant_id: tenant_id.into(),
74            allowed_models: allowed_models.into_iter().map(Into::into).collect(),
75        }
76    }
77
78    /// Tenant whose allowlist is represented by this policy.
79    pub fn tenant_id(&self) -> &str {
80        &self.tenant_id
81    }
82
83    /// Return whether this tenant may use `model`.
84    pub fn allows(&self, model: &str) -> bool {
85        self.allowed_models.contains(model)
86    }
87
88    /// Build the authorization message for a model denied by a tenant policy.
89    pub fn denial_message(tenant_id: &str, model: &str) -> String {
90        format!(
91            "Model '{}' is not allowed for tenant '{}'",
92            model, tenant_id
93        )
94    }
95
96    /// Build the authorization error for a model denied by a tenant policy.
97    pub fn denial_error(tenant_id: &str, model: &str) -> AppError {
98        AppError::Auth(Self::denial_message(tenant_id, model))
99    }
100
101    /// Authorize a model selected by a request interceptor.
102    pub fn authorize(&self, model: &str) -> Result<(), AppError> {
103        if self.allows(model) {
104            Ok(())
105        } else {
106            Err(Self::denial_error(&self.tenant_id, model))
107        }
108    }
109}
110
111impl Service for TenantModelPolicy {}
112
113/// Circuit-breaker state for `Llm`.
114///
115/// `check()` returns:
116/// - `Closed` / `HalfOpen` → `true` (service advertises healthy, fibers stay Active)
117/// - `Open { until }` → `false` until cooldown expires (fibers deactivate, guarded withdrawal per Thm 63)
118#[derive(Debug, Clone, Default)]
119pub enum Breaker {
120    /// Normal operation.
121    #[default]
122    Closed,
123    /// Provider is failing; do not use until `until`.
124    Open { until: DateTime<Utc> },
125    /// Trial half-open after cooldown.
126    HalfOpen,
127}
128
129impl Breaker {
130    /// Failure threshold before opening the breaker.
131    pub const FAILURE_THRESHOLD: u32 = 5;
132    /// Cooldown duration while open (seconds).
133    pub const COOLDOWN_SECS: i64 = 30;
134
135    /// Returns `true` if the breaker allows requests.
136    ///
137    /// When `Open` and `Utc::now() < until`, the breaker is still
138    /// considered open — `Service::check` returns `false` so dependent
139    /// fibers deactivate via guarded withdrawal (Thm 63). After cooldown
140    /// expires the breaker reports healthy until an external transition
141    /// moves it to `HalfOpen` or `Closed`.
142    pub fn check(&self) -> bool {
143        match self {
144            Breaker::Closed => true,
145            Breaker::HalfOpen => true,
146            Breaker::Open { until } => {
147                // Guarded withdrawal: while open, dependent fibers see `check() == false`
148                // and deactivate gracefully per Cordis Thm 63.
149                Utc::now() >= *until
150            }
151        }
152    }
153
154    /// Convenience: `true` if strictly closed.
155    pub fn is_closed(&self) -> bool {
156        matches!(self, Breaker::Closed)
157    }
158
159    /// Transition on failure with threshold and cooldown.
160    ///
161    /// - `Closed` → `Open{until: now+cooldown}` after `FAILURE_THRESHOLD` is reached,
162    ///   otherwise stays `Closed` (counting is external via `Llm::record_failure`
163    ///   failure counter; this pure transition always opens — the service layer
164    ///   decides when to call it. For `Closed` we open immediately; callers that
165    ///   want threshold counting should use `Llm::record_failure`).
166    /// - `HalfOpen` → `Open`
167    /// - `Open` → `Open` (refresh cooldown)
168    pub fn transition_on_failure(&self) -> Breaker {
169        let now = Utc::now();
170        let cooldown = chrono::Duration::seconds(Self::COOLDOWN_SECS);
171        match self {
172            Breaker::Closed => Breaker::Open {
173                until: now + cooldown,
174            },
175            Breaker::HalfOpen => Breaker::Open {
176                until: now + cooldown,
177            },
178            Breaker::Open { .. } => Breaker::Open {
179                until: now + cooldown,
180            },
181        }
182    }
183
184    /// Transition that respects an explicit failure count.
185    pub fn transition_on_failure_with_count(&self, failures: u32) -> Breaker {
186        if failures >= Self::FAILURE_THRESHOLD {
187            let now = Utc::now();
188            Breaker::Open {
189                until: now + chrono::Duration::seconds(Self::COOLDOWN_SECS),
190            }
191        } else {
192            Breaker::Closed
193        }
194    }
195}
196
197/// Unified LLM capability composing provider registry, catalog, factory, and pool.
198///
199/// Supports `ctx.get::<Llm>()`, `ctx.isolate::<Llm>` per-tenant scoping, and
200/// `ctx.intercept::<ModelOverride>` per-request model pinning.
201///
202/// `Service::check` delegates to the circuit breaker — when the breaker is
203/// `Open`, `check()` returns `false` so dependent fibers (e.g. `Execute`)
204/// deactivate via guarded withdrawal (Thm 63).
205pub struct Llm {
206    /// Named provider registry (OpenAI, Anthropic, etc.).
207    pub(crate) provider_registry: Arc<ProviderRegistry>,
208    /// NVIDIA catalog cache for capability-based model selection (optional so
209    /// `cargo check --no-default-features` remains viable).
210    pub(crate) catalog: Option<Arc<NvidiaCatalogCache>>,
211    /// Pooled LLM clients.
212    pub(crate) pool: Arc<ClientPool>,
213    /// Config-based factory used by crate-internal wiring.
214    pub(crate) factory: Option<Arc<ConfigBasedLLMFactory>>,
215    /// Circuit-breaker state.
216    breaker: RwLock<Breaker>,
217    /// Consecutive failure count for thresholded transition.
218    failures: RwLock<u32>,
219    /// Pinned client used by `complete` / `get_client_inner` when set.
220    ///
221    /// Set by [`Llm::from_client`] for in-process tests and library proofs.
222    test_client: Option<Arc<dyn LLMClient>>,
223}
224
225impl Llm {
226    /// Create a new `Llm`.
227    pub fn new(
228        provider_registry: Arc<ProviderRegistry>,
229        pool: Arc<ClientPool>,
230        catalog: Option<Arc<NvidiaCatalogCache>>,
231    ) -> Self {
232        Self {
233            provider_registry,
234            catalog,
235            pool,
236            factory: None,
237            breaker: RwLock::new(Breaker::Closed),
238            failures: RwLock::new(0),
239            test_client: None,
240        }
241    }
242
243    /// Attach a config-based factory (crate wiring / Execute boot).
244    pub fn with_factory(mut self, factory: Arc<ConfigBasedLLMFactory>) -> Self {
245        self.factory = Some(factory);
246        self
247    }
248
249    /// Pin a client used by [`get_client`](Self::get_client) / [`complete`](Self::complete).
250    ///
251    /// Intended for in-process tests and the library proof (`Execute` with no HTTP).
252    pub fn from_client(client: Arc<dyn LLMClient>) -> Self {
253        let mut llm = Self::new(
254            Arc::new(ProviderRegistry::new()),
255            Arc::new(ClientPool::with_defaults()),
256            None,
257        );
258        llm.test_client = Some(client);
259        llm
260    }
261
262    /// Test helper: pin a client used by `complete` / `get_client_inner`.
263    #[cfg(test)]
264    pub(crate) fn for_test(client: Arc<dyn LLMClient>) -> Self {
265        Self::from_client(client)
266    }
267
268    /// Clone the provider registry for named-provider lookup.
269    pub(crate) fn provider_registry(&self) -> Arc<ProviderRegistry> {
270        Arc::clone(&self.provider_registry)
271    }
272
273    /// Handle for `AgentRegistry` construction in `ares-agent`.
274    ///
275    /// Application code should use [`get_client`](Self::get_client).
276    pub fn registry(&self) -> Arc<ProviderRegistry> {
277        self.provider_registry()
278    }
279
280    /// Create with explicit breaker (e.g. `Open` for testing).
281    pub fn with_breaker(
282        provider_registry: Arc<ProviderRegistry>,
283        catalog: Option<Arc<NvidiaCatalogCache>>,
284        pool: Arc<ClientPool>,
285        breaker: Breaker,
286    ) -> Self {
287        Self {
288            provider_registry,
289            catalog,
290            pool,
291            factory: None,
292            breaker: RwLock::new(breaker),
293            failures: RwLock::new(0),
294            test_client: None,
295        }
296    }
297
298    /// Legacy constructor with non-optional catalog (kept for existing call-sites).
299    pub fn with_catalog(
300        provider_registry: Arc<ProviderRegistry>,
301        catalog: Arc<NvidiaCatalogCache>,
302        pool: Arc<ClientPool>,
303    ) -> Self {
304        Self::new(provider_registry, pool, Some(catalog))
305    }
306
307    /// Get current breaker state.
308    pub fn breaker(&self) -> Breaker {
309        self.breaker.read().clone()
310    }
311
312    /// Transition breaker to `Open` with cooldown.
313    pub fn trip(&self, until: DateTime<Utc>) {
314        *self.breaker.write() = Breaker::Open { until };
315    }
316
317    /// Transition breaker to `HalfOpen`.
318    pub fn half_open(&self) {
319        *self.breaker.write() = Breaker::HalfOpen;
320    }
321
322    /// Reset breaker to `Closed`.
323    pub fn reset(&self) {
324        *self.breaker.write() = Breaker::Closed;
325        *self.failures.write() = 0;
326    }
327
328    /// Record a successful request — closes the breaker and resets failure count.
329    pub fn record_success(&self) {
330        *self.breaker.write() = Breaker::Closed;
331        *self.failures.write() = 0;
332    }
333
334    /// Record a failure — increments counter and opens breaker when threshold reached.
335    ///
336    /// Threshold: `Breaker::FAILURE_THRESHOLD` (5). Cooldown: `Breaker::COOLDOWN_SECS` (30s).
337    /// `HalfOpen` immediately re-opens on failure.
338    pub fn record_failure(&self) {
339        let mut failures = self.failures.write();
340        *failures = failures.saturating_add(1);
341        let count = *failures;
342        drop(failures);
343        let mut b = self.breaker.write();
344        // HalfOpen always transitions to Open on failure; Closed respects threshold
345        match &*b {
346            Breaker::HalfOpen => {
347                *b = b.transition_on_failure();
348            }
349            Breaker::Closed => {
350                if count >= Breaker::FAILURE_THRESHOLD {
351                    *b = Breaker::Open {
352                        until: Utc::now() + chrono::Duration::seconds(Breaker::COOLDOWN_SECS),
353                    };
354                }
355            }
356            Breaker::Open { .. } => {
357                // refresh cooldown
358                *b = b.transition_on_failure();
359            }
360        }
361    }
362
363    /// Validate a request's model override against its tenant policy.
364    ///
365    /// Both values are resolved through the context prototype chain, so callers
366    /// can compose `TenantModelPolicy` and `ModelOverride` on separate child
367    /// contexts. With no policy, legacy override behavior is preserved.
368    pub fn validate_model_override(&self, ctx: &Arc<Context>) -> Result<(), AppError> {
369        if let (Some(policy), Some(override_model)) =
370            (ctx.get::<TenantModelPolicy>(), ctx.get::<ModelOverride>())
371        {
372            policy.authorize(&override_model.model)?;
373        }
374        Ok(())
375    }
376
377    /// Capability-aware client resolution with per-request `ModelOverride` pinning.
378    ///
379    /// Public API: `Llm::get_client(&self, ctx: &Arc<cordis::Context>, capability)`.
380    ///
381    /// When `EventsService` is on `ctx`, runs waterfall `"llm.get_client"` first.
382    /// Core is identity so handlers can set `"deny": true` or `"model"`.
383    /// If the result has a `model` string and no intercept `ModelOverride`,
384    /// `get_client_inner` runs on `ctx.with_intercept(ModelOverride { model })`.
385    pub async fn get_client(
386        &self,
387        ctx: &Arc<Context>,
388        capability: CapabilityRequirements,
389    ) -> Result<Arc<dyn LLMClient>, AppError> {
390        let Some(events) = ctx.get::<EventsService>() else {
391            return self.get_client_inner(ctx, capability).await;
392        };
393        let payload = serde_json::to_value(cordis::LlmGetClientPayload {
394            capability: format!("{capability:?}"),
395            deny: None,
396            model: None,
397        })
398        .unwrap_or(serde_json::Value::Null);
399        let result = events
400            .waterfall_around(
401                cordis::events_catalog::ev::LLM_GET_CLIENT.to_string(),
402                payload,
403                |payload| async move { Ok(payload) },
404            )
405            .await
406            .map_err(map_cordis)?;
407        if result.get("deny").and_then(|v| v.as_bool()) == Some(true) {
408            return Err(AppError::InvalidInput("llm.get_client denied".into()));
409        }
410        if let Some(model) = result
411            .get("model")
412            .and_then(|v| v.as_str())
413            .filter(|s| !s.is_empty())
414        {
415            if ctx.get::<ModelOverride>().is_none() {
416                let intercepted = ctx.with_intercept(ModelOverride {
417                    model: model.to_string(),
418                });
419                return self.get_client_inner(&intercepted, capability).await;
420            }
421        }
422        self.get_client_inner(ctx, capability).await
423    }
424
425    /// Existing get_client body (no events). Extracted so public `get_client` can wrap it.
426    async fn get_client_inner(
427        &self,
428        ctx: &Arc<Context>,
429        capability: CapabilityRequirements,
430    ) -> Result<Arc<dyn LLMClient>, AppError> {
431        if let Some(c) = &self.test_client {
432            return Ok(Arc::clone(c));
433        }
434        // Authorize before touching the pool or provider registry.
435        self.validate_model_override(ctx)?;
436        // 1. Per-request pinning via `ctx.get::<ModelOverride>()` (intercept realm).
437        if let Some(ov) = ctx.get::<ModelOverride>() {
438            // Fast-path: pooled provider named exactly like the override model
439            if let Ok(guard) = self.pool.try_get(&ov.model).await {
440                let boxed = guard.take();
441                return Ok(Arc::from(boxed));
442            }
443            if let Ok(client) = self
444                .provider_registry
445                .create_client_for_model_ctx(ctx, &ov.model)
446                .await
447            {
448                return Ok(Arc::from(client));
449            }
450            // fall through to capability path if override model not found
451        }
452
453        // 2. Catalog-aware capability selection (catalog is Option for no-default-features builds)
454        if let Some(catalog) = &self.catalog {
455            let _snap = catalog.snapshot(); // touch catalog to prove composition
456            if let Some(best) = self.provider_registry.find_best_model(&capability) {
457                if let Ok(client) = self
458                    .provider_registry
459                    .create_client_for_model_ctx(ctx, &best.name)
460                    .await
461                {
462                    return Ok(Arc::from(client));
463                }
464            }
465        } else if let Some(best) = self.provider_registry.find_best_model(&capability) {
466            if let Ok(client) = self
467                .provider_registry
468                .create_client_for_model_ctx(ctx, &best.name)
469                .await
470            {
471                return Ok(Arc::from(client));
472            }
473        }
474
475        // 3. Fallback chain (Coordinator / ProviderRegistry) — uses
476        // `ProviderRegistry::resolve_with_capability_fallback` which delegates
477        // to `create_client_for_requirements` → `find_best_model` and falls
478        // back to `create_default_client`. Also exposed as
479        // `resolve_with_fallback` when postgres feature is off.
480        let client = self
481            .provider_registry
482            .resolve_with_capability_fallback(Some(capability))
483            .await?;
484        Ok(Arc::from(client))
485    }
486
487    /// Box variant for callers that prefer `Box<dyn LLMClient>`.
488    pub async fn get_client_boxed(
489        &self,
490        ctx: &Arc<Context>,
491        capability: CapabilityRequirements,
492    ) -> Result<Box<dyn LLMClient>, AppError> {
493        let client = self.get_client(ctx, capability).await?;
494        Ok(Box::new(BoxedArcClient(client)))
495    }
496
497    /// Generate a completion, optionally through waterfall `"llm.complete"`.
498    ///
499    /// Without `EventsService`, this is `get_client` then `generate`. With events,
500    /// handlers wrap payload `{"prompt"}`; core generates using `payload["prompt"]`
501    /// and returns `{"prompt", "content"}`.
502    pub async fn complete(&self, ctx: &Arc<Context>, prompt: &str) -> Result<String, AppError> {
503        let client = self
504            .get_client(ctx, CapabilityRequirements::default())
505            .await?;
506        let Some(events) = ctx.get::<EventsService>() else {
507            return client.generate(prompt).await;
508        };
509        let payload = serde_json::to_value(cordis::LlmCompleteRequest {
510            prompt: prompt.to_string(),
511        })
512        .unwrap_or(serde_json::Value::Null);
513        let out = events
514            .waterfall_around(
515                cordis::events_catalog::ev::LLM_COMPLETE.to_string(),
516                payload,
517                move |payload| {
518                    let client = Arc::clone(&client);
519                    async move {
520                        let prompt = payload
521                            .get("prompt")
522                            .and_then(|v| v.as_str())
523                            .unwrap_or("")
524                            .to_string();
525                        let text = client
526                            .generate(&prompt)
527                            .await
528                            .map_err(|e| CordisError::Fiber(e.to_string()))?;
529                        Ok(serde_json::json!({ "prompt": prompt, "content": text }))
530                    }
531                },
532            )
533            .await
534            .map_err(map_cordis)?;
535        Ok(out
536            .get("content")
537            .and_then(|v| v.as_str())
538            .unwrap_or("")
539            .to_string())
540    }
541
542    /// Stub for capability-based model selection (delegates to registry).
543    pub fn find_model_stub(&self, _capability: &str) -> Option<String> {
544        None
545    }
546
547    /// List registered models with their provider info.
548    pub fn list_models(&self) -> Vec<ModelInfo> {
549        self.provider_registry.list_models()
550    }
551
552    /// Check if a provider exists for the given tenant (legacy or runtime).
553    pub fn has_provider_for_tenant(&self, name: &str, tenant_id: Option<&str>) -> bool {
554        self.provider_registry
555            .has_provider_for_tenant(name, tenant_id)
556    }
557
558    /// Resolve a provider visible to the tenant derived from `ctx`.
559    pub fn get_provider_for_ctx(&self, ctx: &Arc<Context>, name: &str) -> Option<ProviderConfig> {
560        self.provider_registry.get_provider_for_ctx(ctx, name)
561    }
562
563    /// Hot-swap the runtime provider map.
564    pub fn reload_runtime_providers(
565        &self,
566        providers: Vec<RuntimeProviderEntry>,
567        names: Vec<String>,
568    ) {
569        self.provider_registry
570            .reload_runtime_providers(providers, names);
571    }
572}
573
574fn map_cordis(err: CordisError) -> AppError {
575    AppError::Internal(err.to_string())
576}
577
578/// `Box<dyn LLMClient>` adapter around the Arc returned by [`Llm::get_client`].
579struct BoxedArcClient(Arc<dyn LLMClient>);
580
581#[async_trait]
582impl LLMClient for BoxedArcClient {
583    async fn generate(&self, prompt: &str) -> ares_types::types::Result<String> {
584        self.0.generate(prompt).await
585    }
586
587    async fn generate_with_system(
588        &self,
589        system: &str,
590        prompt: &str,
591    ) -> ares_types::types::Result<String> {
592        self.0.generate_with_system(system, prompt).await
593    }
594
595    async fn generate_with_history(
596        &self,
597        messages: &[(String, String)],
598    ) -> ares_types::types::Result<LLMResponse> {
599        self.0.generate_with_history(messages).await
600    }
601
602    async fn generate_with_tools(
603        &self,
604        prompt: &str,
605        tools: &[ToolDefinition],
606    ) -> ares_types::types::Result<LLMResponse> {
607        self.0.generate_with_tools(prompt, tools).await
608    }
609
610    async fn generate_with_tools_and_history(
611        &self,
612        messages: &[crate::coordinator::ConversationMessage],
613        tools: &[ToolDefinition],
614    ) -> ares_types::types::Result<LLMResponse> {
615        self.0
616            .generate_with_tools_and_history(messages, tools)
617            .await
618    }
619
620    async fn stream(
621        &self,
622        prompt: &str,
623    ) -> ares_types::types::Result<
624        Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
625    > {
626        self.0.stream(prompt).await
627    }
628
629    async fn stream_with_system(
630        &self,
631        system: &str,
632        prompt: &str,
633    ) -> ares_types::types::Result<
634        Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
635    > {
636        self.0.stream_with_system(system, prompt).await
637    }
638
639    async fn stream_with_history(
640        &self,
641        messages: &[(String, String)],
642    ) -> ares_types::types::Result<
643        Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
644    > {
645        self.0.stream_with_history(messages).await
646    }
647
648    fn model_name(&self) -> &str {
649        self.0.model_name()
650    }
651    fn supports_hints(&self) -> bool {
652        self.0.supports_hints()
653    }
654    fn set_hints(&self, hints: GenerationHints) {
655        self.0.set_hints(hints)
656    }
657}
658
659impl Service for Llm {
660    fn name(&self) -> &'static str {
661        "Llm"
662    }
663
664    fn init(
665        &self,
666        _ctx: &Arc<Context>,
667    ) -> std::pin::Pin<
668        Box<
669            dyn Future<Output = Result<Option<Box<dyn cordis::Disposable>>, CordisError>>
670                + Send
671                + '_,
672        >,
673    > {
674        Box::pin(async { Ok(None) })
675    }
676
677    fn check(&self) -> bool {
678        // Circuit-breaker advertisement: when Open (and cooldown not elapsed),
679        // this service is unhealthy → dependent fibers deactivate (guarded withdrawal per Thm 63).
680        self.breaker.read().check()
681    }
682}
683
684#[cfg(test)]
685mod tests {
686    use super::*;
687    use crate::capabilities::CapabilityRequirements;
688    use crate::provider_registry::RuntimeProviderEntry;
689    use ares_types::models::{TenantContext, TenantTier};
690    use chrono::Duration;
691    use cordis::Context;
692    use std::collections::HashMap;
693
694    #[test]
695    fn breaker_closed_allows() {
696        assert!(Breaker::Closed.check());
697    }
698
699    #[test]
700    fn breaker_half_open_allows() {
701        assert!(Breaker::HalfOpen.check());
702    }
703
704    #[test]
705    fn breaker_open_future_denies() {
706        let until = Utc::now() + Duration::seconds(60);
707        assert!(!Breaker::Open { until }.check());
708    }
709
710    #[test]
711    fn breaker_open_past_allows() {
712        let until = Utc::now() - Duration::seconds(1);
713        assert!(Breaker::Open { until }.check());
714    }
715
716    #[test]
717    fn breaker_failure_threshold_opens() {
718        let b = Breaker::Closed;
719        let next = b.transition_on_failure_with_count(5);
720        assert!(matches!(next, Breaker::Open { .. }));
721        let still_closed = b.transition_on_failure_with_count(3);
722        assert!(matches!(still_closed, Breaker::Closed));
723    }
724
725    #[test]
726    fn breaker_constants_exist() {
727        assert_eq!(Breaker::FAILURE_THRESHOLD, 5);
728        assert_eq!(Breaker::COOLDOWN_SECS, 30);
729    }
730
731    #[test]
732    fn provider_registry_and_factory_accessors() {
733        let registry = Arc::new(ProviderRegistry::new());
734        let pool = Arc::new(ClientPool::with_defaults());
735        let factory = Arc::new(
736            ConfigBasedLLMFactory::from_config(HashMap::new(), HashMap::new(), None)
737                .expect("empty factory config"),
738        );
739        let llm = Llm::new(Arc::clone(&registry), pool, None).with_factory(Arc::clone(&factory));
740        assert!(Arc::ptr_eq(&llm.provider_registry(), &registry));
741    }
742
743    #[tokio::test]
744    async fn llm_model_override_via_context() {
745        let registry = Arc::new(ProviderRegistry::new());
746        let pool = Arc::new(ClientPool::with_defaults());
747        let svc = Arc::new(Llm::new(registry, pool, None));
748        let root = Context::new_root();
749        root.provide::<Llm>(Llm::new(
750            Arc::new(ProviderRegistry::new()),
751            Arc::new(ClientPool::with_defaults()),
752            None,
753        ));
754        // Intercept realm carries ModelOverride
755        let req_ctx = root.intercept(ModelOverride {
756            model: "gpt-4o-mini".into(),
757        });
758        assert!(req_ctx.get::<ModelOverride>().is_some());
759        assert_eq!(req_ctx.get::<ModelOverride>().unwrap().model, "gpt-4o-mini");
760        // Llm composition fields exist
761        assert!(!Arc::as_ptr(&svc.provider_registry).is_null());
762        let _ = svc.catalog.clone();
763        let _ = svc.pool.provider_names();
764        // Service check guarded withdrawal comment path
765        assert!(svc.check());
766    }
767
768    #[tokio::test]
769    async fn record_failure_threshold_opens_breaker() {
770        let svc = Llm::new(
771            Arc::new(ProviderRegistry::new()),
772            Arc::new(ClientPool::with_defaults()),
773            None,
774        );
775        for _ in 0..Breaker::FAILURE_THRESHOLD {
776            svc.record_failure();
777        }
778        // After threshold, breaker should be Open and check() denies
779        assert!(!svc.check());
780        svc.record_success();
781        assert!(svc.check());
782    }
783
784    #[test]
785    fn tenant_model_policy_allows_and_composes_with_model_override() {
786        let root = Context::new_root();
787        let tenant_ctx = root.intercept(TenantModelPolicy::new(
788            "tenant-a",
789            ["gpt-4o-mini".to_string()],
790        ));
791        let request = tenant_ctx.intercept(ModelOverride {
792            model: "gpt-4o-mini".into(),
793        });
794        let policy = request
795            .get::<TenantModelPolicy>()
796            .expect("policy should be inherited by request context");
797        let override_model = request
798            .get::<ModelOverride>()
799            .expect("model override should be visible in request context");
800        policy
801            .authorize(&override_model.model)
802            .expect("allowed model override should pass policy");
803        let svc = Llm::new(
804            Arc::new(ProviderRegistry::new()),
805            Arc::new(ClientPool::with_defaults()),
806            None,
807        );
808        svc.validate_model_override(&request)
809            .expect("allowed model override should pass LLM validation");
810        assert!(root.get::<ModelOverride>().is_none());
811        assert!(root.get::<TenantModelPolicy>().is_none());
812    }
813
814    #[tokio::test]
815    async fn disallowed_model_override_is_rejected_before_provider_execution() {
816        let registry = Arc::new(ProviderRegistry::new());
817        let svc = Arc::new(Llm::new(
818            registry,
819            Arc::new(ClientPool::with_defaults()),
820            None,
821        ));
822        let root = Context::new_root();
823        root.provide_arc(svc.clone());
824        let tenant_ctx = root.intercept(TenantModelPolicy::new("tenant-a", ["gpt-4o".to_string()]));
825        let request = tenant_ctx.intercept(ModelOverride {
826            model: "not-allowed".into(),
827        });
828        let err = match svc
829            .get_client(&request, CapabilityRequirements::default())
830            .await
831        {
832            Ok(_) => panic!("disallowed override must fail before provider lookup"),
833            Err(err) => err,
834        };
835        assert!(matches!(err, AppError::Auth(_)));
836        assert!(err.to_string().contains("not-allowed"));
837        assert!(root.get::<ModelOverride>().is_none());
838        assert!(root.get::<TenantModelPolicy>().is_none());
839        assert!(matches!(
840            root.get::<Llm>().expect("global service").breaker(),
841            Breaker::Closed
842        ));
843    }
844
845    #[tokio::test]
846    async fn get_client_uses_override_when_catalog_absent() {
847        let registry = Arc::new(ProviderRegistry::new());
848        let pool = Arc::new(ClientPool::with_defaults());
849        let svc = Llm::new(registry, pool, None);
850        let ctx = Context::new_root();
851        let req_ctx = ctx.intercept(ModelOverride {
852            model: "nonexistent-model-xyz".into(),
853        });
854        let req = CapabilityRequirements::default();
855        // Should attempt override then fallback; fallback will fail because no provider configured
856        let res = svc.get_client(&req_ctx, req).await;
857        assert!(res.is_err());
858    }
859
860    #[tokio::test]
861    async fn get_client_override_uses_tenant_context_intercept() {
862        let mut registry = ProviderRegistry::new();
863        registry.register_model(
864            "pinned-model",
865            crate::config::ModelConfig {
866                provider: "shared-runtime".into(),
867                model: "tenant-model".into(),
868                temperature: 0.7,
869                max_tokens: 512,
870            },
871        );
872        let global = RuntimeProviderEntry {
873            tenant_id: None,
874            display_name: "Global Shared".to_string(),
875            provider_type: "openai-compatible".to_string(),
876            api_base: "https://global.example.com/v1".to_string(),
877            auth_type: "api_key".to_string(),
878            default_model: Some("global-model".to_string()),
879            headers: HashMap::new(),
880            api_key: Some("global-key".to_string()),
881            enabled: true,
882        };
883        let tenant = RuntimeProviderEntry {
884            tenant_id: Some("tenant-a".to_string()),
885            display_name: "Tenant Shared".to_string(),
886            provider_type: "openai-compatible".to_string(),
887            api_base: "https://tenant.example.com/v1".to_string(),
888            auth_type: "api_key".to_string(),
889            default_model: Some("tenant-model".to_string()),
890            headers: HashMap::new(),
891            api_key: Some("tenant-key".to_string()),
892            enabled: true,
893        };
894        registry.reload_runtime_providers(
895            vec![global, tenant],
896            vec!["shared-runtime".to_string(), "shared-runtime".to_string()],
897        );
898        let registry = Arc::new(registry);
899        let svc = Llm::new(registry, Arc::new(ClientPool::with_defaults()), None);
900        let root = Context::new_root();
901        let ctx = root
902            .with_intercept(TenantContext::new("tenant-a".into(), TenantTier::Pro))
903            .intercept(ModelOverride {
904                model: "pinned-model".into(),
905            });
906        let tenant_client = svc
907            .get_client(&ctx, CapabilityRequirements::default())
908            .await;
909        assert!(
910            tenant_client.is_ok(),
911            "tenant intercept should construct a client from the tenant runtime entry: {:?}",
912            tenant_client.as_ref().err()
913        );
914
915        let unlabeled = root.intercept(ModelOverride {
916            model: "pinned-model".into(),
917        });
918        let fleet_client = svc
919            .get_client(&unlabeled, CapabilityRequirements::default())
920            .await;
921        assert!(
922            fleet_client.is_ok(),
923            "unlabeled root with ModelOverride should construct a client from the fleet global runtime entry: {:?}",
924            fleet_client.as_ref().err()
925        );
926    }
927
928    struct EchoClient {
929        generated: std::sync::Arc<std::sync::atomic::AtomicBool>,
930    }
931
932    impl EchoClient {
933        fn new() -> (Self, std::sync::Arc<std::sync::atomic::AtomicBool>) {
934            let generated = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
935            (
936                Self {
937                    generated: std::sync::Arc::clone(&generated),
938                },
939                generated,
940            )
941        }
942    }
943
944    #[async_trait]
945    impl LLMClient for EchoClient {
946        async fn generate(&self, prompt: &str) -> ares_types::types::Result<String> {
947            self.generated
948                .store(true, std::sync::atomic::Ordering::SeqCst);
949            Ok(format!("echo:{prompt}"))
950        }
951        async fn generate_with_system(
952            &self,
953            _system: &str,
954            prompt: &str,
955        ) -> ares_types::types::Result<String> {
956            self.generate(prompt).await
957        }
958        async fn generate_with_history(
959            &self,
960            _messages: &[(String, String)],
961        ) -> ares_types::types::Result<LLMResponse> {
962            Ok(LLMResponse {
963                content: String::new(),
964                tool_calls: vec![],
965                finish_reason: "stop".into(),
966                usage: None,
967            })
968        }
969        async fn generate_with_tools(
970            &self,
971            _prompt: &str,
972            _tools: &[ToolDefinition],
973        ) -> ares_types::types::Result<LLMResponse> {
974            Ok(LLMResponse {
975                content: String::new(),
976                tool_calls: vec![],
977                finish_reason: "stop".into(),
978                usage: None,
979            })
980        }
981        async fn generate_with_tools_and_history(
982            &self,
983            _messages: &[crate::coordinator::ConversationMessage],
984            _tools: &[ToolDefinition],
985        ) -> ares_types::types::Result<LLMResponse> {
986            Ok(LLMResponse {
987                content: String::new(),
988                tool_calls: vec![],
989                finish_reason: "stop".into(),
990                usage: None,
991            })
992        }
993        async fn stream(
994            &self,
995            _prompt: &str,
996        ) -> ares_types::types::Result<
997            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
998        > {
999            Err(AppError::Internal("echo stream not implemented".into()))
1000        }
1001        async fn stream_with_system(
1002            &self,
1003            _system: &str,
1004            _prompt: &str,
1005        ) -> ares_types::types::Result<
1006            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1007        > {
1008            Err(AppError::Internal("echo stream not implemented".into()))
1009        }
1010        async fn stream_with_history(
1011            &self,
1012            _messages: &[(String, String)],
1013        ) -> ares_types::types::Result<
1014            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1015        > {
1016            Err(AppError::Internal("echo stream not implemented".into()))
1017        }
1018        fn model_name(&self) -> &str {
1019            "echo"
1020        }
1021    }
1022
1023    #[tokio::test]
1024    async fn llm_complete_runs_generate_without_events() {
1025        let (client, generated) = EchoClient::new();
1026        let llm = Llm::for_test(std::sync::Arc::new(client));
1027        let ctx = Context::new_root();
1028        let out = llm.complete(&ctx, "hi").await.expect("complete");
1029        assert_eq!(out, "echo:hi");
1030        assert!(generated.load(std::sync::atomic::Ordering::SeqCst));
1031    }
1032
1033    #[tokio::test]
1034    async fn llm_complete_waterfall_rewrites_prompt() {
1035        let (client, _) = EchoClient::new();
1036        let llm = Llm::for_test(std::sync::Arc::new(client));
1037        let ctx = Context::new_root();
1038        let events = ctx.provide(EventsService::new());
1039        events.on_waterfall(
1040            cordis::events_catalog::ev::LLM_COMPLETE.to_string(),
1041            |mut payload, next| async move {
1042                if let Some(p) = payload.get("prompt").and_then(|v| v.as_str()) {
1043                    payload["prompt"] = serde_json::json!(format!("WRAP:{p}"));
1044                }
1045                next(payload).await
1046            },
1047        );
1048        let out = llm.complete(&ctx, "hi").await.expect("complete");
1049        assert_eq!(out, "echo:WRAP:hi");
1050    }
1051
1052    #[tokio::test]
1053    async fn llm_complete_short_circuit_skips_generate() {
1054        let (client, generated) = EchoClient::new();
1055        let llm = Llm::for_test(std::sync::Arc::new(client));
1056        let ctx = Context::new_root();
1057        let events = ctx.provide(EventsService::new());
1058        events.on_waterfall(
1059            cordis::events_catalog::ev::LLM_COMPLETE.to_string(),
1060            |_payload, _next| async move { Ok(serde_json::json!({ "content": "cached" })) },
1061        );
1062        let out = llm.complete(&ctx, "hi").await.expect("complete");
1063        assert_eq!(out, "cached");
1064        assert!(
1065            !generated.load(std::sync::atomic::Ordering::SeqCst),
1066            "dummy generate must stay false when handler skips next"
1067        );
1068    }
1069
1070    #[tokio::test]
1071    async fn llm_get_client_waterfall_deny() {
1072        let (client, _) = EchoClient::new();
1073        let llm = Llm::for_test(std::sync::Arc::new(client));
1074        let ctx = Context::new_root();
1075        let events = ctx.provide(EventsService::new());
1076        events.on_waterfall(
1077            cordis::events_catalog::ev::LLM_GET_CLIENT.to_string(),
1078            |_payload, _next| async move { Ok(serde_json::json!({ "deny": true })) },
1079        );
1080        let err = match llm
1081            .get_client(&ctx, CapabilityRequirements::default())
1082            .await
1083        {
1084            Ok(_) => panic!("deny"),
1085            Err(err) => err,
1086        };
1087        assert!(matches!(err, AppError::InvalidInput(msg) if msg == "llm.get_client denied"));
1088    }
1089
1090    #[test]
1091    fn llm_list_models_exposes_registry_models() {
1092        let mut registry = ProviderRegistry::new();
1093        registry.register_model(
1094            "stub-model",
1095            crate::config::ModelConfig {
1096                provider: "stub".into(),
1097                model: "stub-model".into(),
1098                temperature: 0.7,
1099                max_tokens: 512,
1100            },
1101        );
1102        let llm = Llm::new(
1103            Arc::new(registry),
1104            Arc::new(ClientPool::with_defaults()),
1105            None,
1106        );
1107        let models = llm.list_models();
1108        assert!(
1109            models
1110                .iter()
1111                .any(|m| m.name == "stub-model" && m.provider == "stub"),
1112            "Llm::list_models should expose registry models: {models:?}"
1113        );
1114    }
1115
1116    /// Mock recording every `set_hints` payload routed through it.
1117    #[derive(Default)]
1118    struct HintRecordingClient {
1119        hints: parking_lot::Mutex<Vec<GenerationHints>>,
1120        supports: bool,
1121    }
1122
1123    #[async_trait]
1124    impl LLMClient for HintRecordingClient {
1125        async fn generate(&self, _prompt: &str) -> ares_types::types::Result<String> {
1126            Err(AppError::Internal("unused".into()))
1127        }
1128
1129        async fn generate_with_system(
1130            &self,
1131            _system: &str,
1132            _prompt: &str,
1133        ) -> ares_types::types::Result<String> {
1134            Err(AppError::Internal("unused".into()))
1135        }
1136
1137        async fn generate_with_history(
1138            &self,
1139            _messages: &[(String, String)],
1140        ) -> ares_types::types::Result<LLMResponse> {
1141            Err(AppError::Internal("unused".into()))
1142        }
1143
1144        async fn generate_with_tools(
1145            &self,
1146            _prompt: &str,
1147            _tools: &[ToolDefinition],
1148        ) -> ares_types::types::Result<LLMResponse> {
1149            Err(AppError::Internal("unused".into()))
1150        }
1151
1152        async fn generate_with_tools_and_history(
1153            &self,
1154            _messages: &[crate::coordinator::ConversationMessage],
1155            _tools: &[ToolDefinition],
1156        ) -> ares_types::types::Result<LLMResponse> {
1157            Err(AppError::Internal("unused".into()))
1158        }
1159
1160        async fn stream(
1161            &self,
1162            _prompt: &str,
1163        ) -> ares_types::types::Result<
1164            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1165        > {
1166            Err(AppError::Internal("unused".into()))
1167        }
1168
1169        async fn stream_with_system(
1170            &self,
1171            _system: &str,
1172            _prompt: &str,
1173        ) -> ares_types::types::Result<
1174            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1175        > {
1176            Err(AppError::Internal("unused".into()))
1177        }
1178
1179        async fn stream_with_history(
1180            &self,
1181            _messages: &[(String, String)],
1182        ) -> ares_types::types::Result<
1183            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1184        > {
1185            Err(AppError::Internal("unused".into()))
1186        }
1187
1188        fn model_name(&self) -> &str {
1189            "hint-recording-mock"
1190        }
1191
1192        fn supports_hints(&self) -> bool {
1193            self.supports
1194        }
1195
1196        fn set_hints(&self, hints: GenerationHints) {
1197            self.hints.lock().push(hints);
1198        }
1199    }
1200
1201    #[test]
1202    fn boxed_arc_client_forwards_hints_to_inner_client() {
1203        let concrete = Arc::new(HintRecordingClient {
1204            supports: true,
1205            hints: parking_lot::Mutex::new(Vec::new()),
1206        });
1207        let recorder_handle = Arc::clone(&concrete);
1208        let inner: Arc<dyn LLMClient> = concrete;
1209        let boxed = BoxedArcClient(Arc::clone(&inner));
1210
1211        assert!(boxed.supports_hints());
1212        boxed.set_hints(GenerationHints {
1213            json_mode: true,
1214            suppress_reasoning: false,
1215            max_tokens: Some(256),
1216            guided_grammar: None,
1217        });
1218        boxed.set_hints(GenerationHints::default());
1219
1220        // The adapter must reach the INNER client instead of stopping at the
1221        // trait's no-op defaults on the box itself. `concrete` and `inner`
1222        // share one allocation, so writes via the box are visible here.
1223        let recorded = recorder_handle.hints.lock();
1224        assert_eq!(
1225            recorded.len(),
1226            2,
1227            "both set_hints calls must reach the inner client"
1228        );
1229        assert!(recorded[0].json_mode && recorded[0].max_tokens == Some(256));
1230        assert_eq!(recorded[1], GenerationHints::default());
1231    }
1232}