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, LlmStreamEvent};
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    /// Embed `inputs` using the same client resolution as [`Llm::complete`].
543    ///
544    /// Without `EventsService`, this is `get_client` then `LLMClient::embed`.
545    /// With events, handlers wrap payload `{"inputs"}`; core embeds using
546    /// `payload["inputs"]` and returns `{"inputs", "embeddings"}`.
547    pub async fn embed(
548        &self,
549        ctx: &Arc<Context>,
550        inputs: &[String],
551    ) -> Result<Vec<Vec<f32>>, AppError> {
552        let client = self
553            .get_client(ctx, CapabilityRequirements::default())
554            .await?;
555        let Some(events) = ctx.get::<EventsService>() else {
556            return client.embed(inputs).await;
557        };
558        let payload = serde_json::to_value(cordis::LlmEmbedRequest {
559            inputs: inputs.to_vec(),
560        })
561        .unwrap_or(serde_json::Value::Null);
562        let out = events
563            .waterfall_around(
564                cordis::events_catalog::ev::LLM_EMBED.to_string(),
565                payload,
566                move |payload| {
567                    let client = Arc::clone(&client);
568                    async move {
569                        let inputs: Vec<String> = payload
570                            .get("inputs")
571                            .cloned()
572                            .map(serde_json::from_value::<Vec<String>>)
573                            .transpose()
574                            .map_err(|e| CordisError::Fiber(e.to_string()))?
575                            .unwrap_or_default();
576                        let embeddings = client
577                            .embed(&inputs)
578                            .await
579                            .map_err(|e| CordisError::Fiber(e.to_string()))?;
580                        serde_json::to_value(cordis::LlmEmbedResponse { inputs, embeddings })
581                            .map_err(|e| CordisError::Fiber(e.to_string()))
582                    }
583                },
584            )
585            .await
586            .map_err(map_cordis)?;
587        Ok(match out.get("embeddings") {
588            Some(value) => serde_json::from_value::<Vec<Vec<f32>>>(value.clone())
589                .map_err(|e| AppError::Internal(e.to_string()))?,
590            None => Vec::new(),
591        })
592    }
593
594    /// Stub for capability-based model selection (delegates to registry).
595    pub fn find_model_stub(&self, _capability: &str) -> Option<String> {
596        None
597    }
598
599    /// List registered models with their provider info.
600    pub fn list_models(&self) -> Vec<ModelInfo> {
601        self.provider_registry.list_models()
602    }
603
604    /// Check if a provider exists for the given tenant (legacy or runtime).
605    pub fn has_provider_for_tenant(&self, name: &str, tenant_id: Option<&str>) -> bool {
606        self.provider_registry
607            .has_provider_for_tenant(name, tenant_id)
608    }
609
610    /// Resolve a provider visible to the tenant derived from `ctx`.
611    pub fn get_provider_for_ctx(&self, ctx: &Arc<Context>, name: &str) -> Option<ProviderConfig> {
612        self.provider_registry.get_provider_for_ctx(ctx, name)
613    }
614
615    /// Hot-swap the runtime provider map.
616    pub fn reload_runtime_providers(
617        &self,
618        providers: Vec<RuntimeProviderEntry>,
619        names: Vec<String>,
620    ) {
621        self.provider_registry
622            .reload_runtime_providers(providers, names);
623    }
624}
625
626fn map_cordis(err: CordisError) -> AppError {
627    AppError::Internal(err.to_string())
628}
629
630/// `Box<dyn LLMClient>` adapter around the Arc returned by [`Llm::get_client`].
631struct BoxedArcClient(Arc<dyn LLMClient>);
632
633#[async_trait]
634impl LLMClient for BoxedArcClient {
635    async fn generate(&self, prompt: &str) -> ares_types::types::Result<String> {
636        self.0.generate(prompt).await
637    }
638
639    async fn generate_with_system(
640        &self,
641        system: &str,
642        prompt: &str,
643    ) -> ares_types::types::Result<String> {
644        self.0.generate_with_system(system, prompt).await
645    }
646
647    async fn generate_with_history(
648        &self,
649        messages: &[(String, String)],
650    ) -> ares_types::types::Result<LLMResponse> {
651        self.0.generate_with_history(messages).await
652    }
653
654    async fn generate_with_tools(
655        &self,
656        prompt: &str,
657        tools: &[ToolDefinition],
658    ) -> ares_types::types::Result<LLMResponse> {
659        self.0.generate_with_tools(prompt, tools).await
660    }
661
662    async fn generate_with_tools_and_history(
663        &self,
664        messages: &[crate::coordinator::ConversationMessage],
665        tools: &[ToolDefinition],
666    ) -> ares_types::types::Result<LLMResponse> {
667        self.0
668            .generate_with_tools_and_history(messages, tools)
669            .await
670    }
671
672    async fn stream(
673        &self,
674        prompt: &str,
675    ) -> ares_types::types::Result<
676        Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
677    > {
678        self.0.stream(prompt).await
679    }
680
681    async fn stream_with_system(
682        &self,
683        system: &str,
684        prompt: &str,
685    ) -> ares_types::types::Result<
686        Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
687    > {
688        self.0.stream_with_system(system, prompt).await
689    }
690
691    async fn stream_with_history(
692        &self,
693        messages: &[(String, String)],
694    ) -> ares_types::types::Result<
695        Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
696    > {
697        self.0.stream_with_history(messages).await
698    }
699
700    async fn stream_with_tools_and_history(
701        &self,
702        messages: &[crate::coordinator::ConversationMessage],
703        tools: &[ToolDefinition],
704    ) -> ares_types::types::Result<
705        Box<dyn futures::Stream<Item = ares_types::types::Result<LlmStreamEvent>> + Send + Unpin>,
706    > {
707        self.0
708            .stream_with_tools_and_history(messages, tools)
709            .await
710    }
711
712    fn model_name(&self) -> &str {
713        self.0.model_name()
714    }
715    fn supports_hints(&self) -> bool {
716        self.0.supports_hints()
717    }
718    fn set_hints(&self, hints: GenerationHints) {
719        self.0.set_hints(hints)
720    }
721
722    async fn embed(&self, inputs: &[String]) -> ares_types::types::Result<Vec<Vec<f32>>> {
723        self.0.embed(inputs).await
724    }
725
726    fn supports_vision(&self) -> bool {
727        self.0.supports_vision()
728    }
729
730    fn supports_provider_web_search(&self) -> bool {
731        self.0.supports_provider_web_search()
732    }
733}
734
735impl Service for Llm {
736    fn name(&self) -> &'static str {
737        "Llm"
738    }
739
740    fn init(
741        &self,
742        _ctx: &Arc<Context>,
743    ) -> std::pin::Pin<
744        Box<
745            dyn Future<Output = Result<Option<Box<dyn cordis::Disposable>>, CordisError>>
746                + Send
747                + '_,
748        >,
749    > {
750        Box::pin(async { Ok(None) })
751    }
752
753    fn check(&self) -> bool {
754        // Circuit-breaker advertisement: when Open (and cooldown not elapsed),
755        // this service is unhealthy → dependent fibers deactivate (guarded withdrawal per Thm 63).
756        self.breaker.read().check()
757    }
758}
759
760#[cfg(test)]
761mod tests {
762    use super::*;
763    use crate::capabilities::CapabilityRequirements;
764    #[cfg(feature = "genai")]
765    use crate::provider_registry::RuntimeProviderEntry;
766    #[cfg(feature = "genai")]
767    use ares_types::models::{TenantContext, TenantTier};
768    use chrono::Duration;
769    use cordis::Context;
770    use std::collections::HashMap;
771
772    #[test]
773    fn breaker_closed_allows() {
774        assert!(Breaker::Closed.check());
775    }
776
777    #[test]
778    fn breaker_half_open_allows() {
779        assert!(Breaker::HalfOpen.check());
780    }
781
782    #[test]
783    fn breaker_open_future_denies() {
784        let until = Utc::now() + Duration::seconds(60);
785        assert!(!Breaker::Open { until }.check());
786    }
787
788    #[test]
789    fn breaker_open_past_allows() {
790        let until = Utc::now() - Duration::seconds(1);
791        assert!(Breaker::Open { until }.check());
792    }
793
794    #[test]
795    fn breaker_failure_threshold_opens() {
796        let b = Breaker::Closed;
797        let next = b.transition_on_failure_with_count(5);
798        assert!(matches!(next, Breaker::Open { .. }));
799        let still_closed = b.transition_on_failure_with_count(3);
800        assert!(matches!(still_closed, Breaker::Closed));
801    }
802
803    #[test]
804    fn breaker_constants_exist() {
805        assert_eq!(Breaker::FAILURE_THRESHOLD, 5);
806        assert_eq!(Breaker::COOLDOWN_SECS, 30);
807    }
808
809    #[test]
810    fn provider_registry_and_factory_accessors() {
811        let registry = Arc::new(ProviderRegistry::new());
812        let pool = Arc::new(ClientPool::with_defaults());
813        let factory = Arc::new(
814            ConfigBasedLLMFactory::from_config(HashMap::new(), HashMap::new(), None)
815                .expect("empty factory config"),
816        );
817        let llm = Llm::new(Arc::clone(&registry), pool, None).with_factory(Arc::clone(&factory));
818        assert!(Arc::ptr_eq(&llm.provider_registry(), &registry));
819    }
820
821    #[tokio::test]
822    async fn llm_model_override_via_context() {
823        let registry = Arc::new(ProviderRegistry::new());
824        let pool = Arc::new(ClientPool::with_defaults());
825        let svc = Arc::new(Llm::new(registry, pool, None));
826        let root = Context::new_root();
827        root.provide::<Llm>(Llm::new(
828            Arc::new(ProviderRegistry::new()),
829            Arc::new(ClientPool::with_defaults()),
830            None,
831        ));
832        // Intercept realm carries ModelOverride
833        let req_ctx = root.intercept(ModelOverride {
834            model: "gpt-4o-mini".into(),
835        });
836        assert!(req_ctx.get::<ModelOverride>().is_some());
837        assert_eq!(req_ctx.get::<ModelOverride>().unwrap().model, "gpt-4o-mini");
838        // Llm composition fields exist
839        let _ = Arc::clone(&svc.provider_registry);
840        let _ = svc.catalog.clone();
841        let _ = svc.pool.provider_names();
842        // Service check guarded withdrawal comment path
843        assert!(svc.check());
844    }
845
846    #[tokio::test]
847    async fn record_failure_threshold_opens_breaker() {
848        let svc = Llm::new(
849            Arc::new(ProviderRegistry::new()),
850            Arc::new(ClientPool::with_defaults()),
851            None,
852        );
853        for _ in 0..Breaker::FAILURE_THRESHOLD {
854            svc.record_failure();
855        }
856        // After threshold, breaker should be Open and check() denies
857        assert!(!svc.check());
858        svc.record_success();
859        assert!(svc.check());
860    }
861
862    #[test]
863    fn tenant_model_policy_allows_and_composes_with_model_override() {
864        let root = Context::new_root();
865        let tenant_ctx = root.intercept(TenantModelPolicy::new(
866            "tenant-a",
867            ["gpt-4o-mini".to_string()],
868        ));
869        let request = tenant_ctx.intercept(ModelOverride {
870            model: "gpt-4o-mini".into(),
871        });
872        let policy = request
873            .get::<TenantModelPolicy>()
874            .expect("policy should be inherited by request context");
875        let override_model = request
876            .get::<ModelOverride>()
877            .expect("model override should be visible in request context");
878        policy
879            .authorize(&override_model.model)
880            .expect("allowed model override should pass policy");
881        let svc = Llm::new(
882            Arc::new(ProviderRegistry::new()),
883            Arc::new(ClientPool::with_defaults()),
884            None,
885        );
886        svc.validate_model_override(&request)
887            .expect("allowed model override should pass LLM validation");
888        assert!(root.get::<ModelOverride>().is_none());
889        assert!(root.get::<TenantModelPolicy>().is_none());
890    }
891
892    #[tokio::test]
893    async fn disallowed_model_override_is_rejected_before_provider_execution() {
894        let registry = Arc::new(ProviderRegistry::new());
895        let svc = Arc::new(Llm::new(
896            registry,
897            Arc::new(ClientPool::with_defaults()),
898            None,
899        ));
900        let root = Context::new_root();
901        root.provide_arc(svc.clone());
902        let tenant_ctx = root.intercept(TenantModelPolicy::new("tenant-a", ["gpt-4o".to_string()]));
903        let request = tenant_ctx.intercept(ModelOverride {
904            model: "not-allowed".into(),
905        });
906        let err = match svc
907            .get_client(&request, CapabilityRequirements::default())
908            .await
909        {
910            Ok(_) => panic!("disallowed override must fail before provider lookup"),
911            Err(err) => err,
912        };
913        assert!(matches!(err, AppError::Auth(_)));
914        assert!(err.to_string().contains("not-allowed"));
915        assert!(root.get::<ModelOverride>().is_none());
916        assert!(root.get::<TenantModelPolicy>().is_none());
917        assert!(matches!(
918            root.get::<Llm>().expect("global service").breaker(),
919            Breaker::Closed
920        ));
921    }
922
923    #[tokio::test]
924    async fn get_client_uses_override_when_catalog_absent() {
925        let registry = Arc::new(ProviderRegistry::new());
926        let pool = Arc::new(ClientPool::with_defaults());
927        let svc = Llm::new(registry, pool, None);
928        let ctx = Context::new_root();
929        let req_ctx = ctx.intercept(ModelOverride {
930            model: "nonexistent-model-xyz".into(),
931        });
932        let req = CapabilityRequirements::default();
933        // Should attempt override then fallback; fallback will fail because no provider configured
934        let res = svc.get_client(&req_ctx, req).await;
935        assert!(res.is_err());
936    }
937
938    #[cfg(feature = "genai")]
939    #[tokio::test]
940    async fn get_client_override_uses_tenant_context_intercept() {
941        let mut registry = ProviderRegistry::new();
942        registry.register_model(
943            "pinned-model",
944            crate::config::ModelConfig {
945                provider: "shared-runtime".into(),
946                model: "tenant-model".into(),
947                temperature: 0.7,
948                max_tokens: 512,
949            },
950        );
951        let global = RuntimeProviderEntry {
952            tenant_id: None,
953            display_name: "Global Shared".to_string(),
954            provider_type: "openai-compatible".to_string(),
955            api_base: "https://global.example.com/v1".to_string(),
956            auth_type: "api_key".to_string(),
957            default_model: Some("global-model".to_string()),
958            headers: HashMap::new(),
959            api_key: Some("global-key".to_string()),
960            enabled: true,
961        };
962        let tenant = RuntimeProviderEntry {
963            tenant_id: Some("tenant-a".to_string()),
964            display_name: "Tenant Shared".to_string(),
965            provider_type: "openai-compatible".to_string(),
966            api_base: "https://tenant.example.com/v1".to_string(),
967            auth_type: "api_key".to_string(),
968            default_model: Some("tenant-model".to_string()),
969            headers: HashMap::new(),
970            api_key: Some("tenant-key".to_string()),
971            enabled: true,
972        };
973        registry.reload_runtime_providers(
974            vec![global, tenant],
975            vec!["shared-runtime".to_string(), "shared-runtime".to_string()],
976        );
977        let registry = Arc::new(registry);
978        let svc = Llm::new(registry, Arc::new(ClientPool::with_defaults()), None);
979        let root = Context::new_root();
980        let ctx = root
981            .with_intercept(TenantContext::new("tenant-a".into(), TenantTier::Pro))
982            .intercept(ModelOverride {
983                model: "pinned-model".into(),
984            });
985        let tenant_client = svc
986            .get_client(&ctx, CapabilityRequirements::default())
987            .await;
988        assert!(
989            tenant_client.is_ok(),
990            "tenant intercept should construct a client from the tenant runtime entry: {:?}",
991            tenant_client.as_ref().err()
992        );
993
994        let unlabeled = root.intercept(ModelOverride {
995            model: "pinned-model".into(),
996        });
997        let fleet_client = svc
998            .get_client(&unlabeled, CapabilityRequirements::default())
999            .await;
1000        assert!(
1001            fleet_client.is_ok(),
1002            "unlabeled root with ModelOverride should construct a client from the fleet global runtime entry: {:?}",
1003            fleet_client.as_ref().err()
1004        );
1005    }
1006
1007    struct EchoClient {
1008        generated: std::sync::Arc<std::sync::atomic::AtomicBool>,
1009    }
1010
1011    impl EchoClient {
1012        fn new() -> (Self, std::sync::Arc<std::sync::atomic::AtomicBool>) {
1013            let generated = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
1014            (
1015                Self {
1016                    generated: std::sync::Arc::clone(&generated),
1017                },
1018                generated,
1019            )
1020        }
1021    }
1022
1023    #[async_trait]
1024    impl LLMClient for EchoClient {
1025        async fn generate(&self, prompt: &str) -> ares_types::types::Result<String> {
1026            self.generated
1027                .store(true, std::sync::atomic::Ordering::SeqCst);
1028            Ok(format!("echo:{prompt}"))
1029        }
1030        async fn generate_with_system(
1031            &self,
1032            _system: &str,
1033            prompt: &str,
1034        ) -> ares_types::types::Result<String> {
1035            self.generate(prompt).await
1036        }
1037        async fn generate_with_history(
1038            &self,
1039            _messages: &[(String, String)],
1040        ) -> ares_types::types::Result<LLMResponse> {
1041            Ok(LLMResponse {
1042                content: String::new(),
1043                tool_calls: vec![],
1044                finish_reason: "stop".into(),
1045                usage: None,
1046                reasoning_content: None,
1047                response_id: None,
1048            })
1049        }
1050        async fn generate_with_tools(
1051            &self,
1052            _prompt: &str,
1053            _tools: &[ToolDefinition],
1054        ) -> ares_types::types::Result<LLMResponse> {
1055            Ok(LLMResponse {
1056                content: String::new(),
1057                tool_calls: vec![],
1058                finish_reason: "stop".into(),
1059                usage: None,
1060                reasoning_content: None,
1061                response_id: None,
1062            })
1063        }
1064        async fn generate_with_tools_and_history(
1065            &self,
1066            _messages: &[crate::coordinator::ConversationMessage],
1067            _tools: &[ToolDefinition],
1068        ) -> ares_types::types::Result<LLMResponse> {
1069            Ok(LLMResponse {
1070                content: String::new(),
1071                tool_calls: vec![],
1072                finish_reason: "stop".into(),
1073                usage: None,
1074                reasoning_content: None,
1075                response_id: None,
1076            })
1077        }
1078        async fn stream(
1079            &self,
1080            _prompt: &str,
1081        ) -> ares_types::types::Result<
1082            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1083        > {
1084            Err(AppError::Internal("echo stream not implemented".into()))
1085        }
1086        async fn stream_with_system(
1087            &self,
1088            _system: &str,
1089            _prompt: &str,
1090        ) -> ares_types::types::Result<
1091            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1092        > {
1093            Err(AppError::Internal("echo stream not implemented".into()))
1094        }
1095        async fn stream_with_history(
1096            &self,
1097            _messages: &[(String, String)],
1098        ) -> ares_types::types::Result<
1099            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1100        > {
1101            Err(AppError::Internal("echo stream not implemented".into()))
1102        }
1103        fn model_name(&self) -> &str {
1104            "echo"
1105        }
1106    }
1107
1108    struct EmbedClient {
1109        called: std::sync::Arc<std::sync::atomic::AtomicBool>,
1110    }
1111
1112    impl EmbedClient {
1113        fn new() -> (Self, std::sync::Arc<std::sync::atomic::AtomicBool>) {
1114            let called = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
1115            (
1116                Self {
1117                    called: std::sync::Arc::clone(&called),
1118                },
1119                called,
1120            )
1121        }
1122    }
1123
1124    #[async_trait]
1125    impl LLMClient for EmbedClient {
1126        async fn generate(&self, _prompt: &str) -> ares_types::types::Result<String> {
1127            Err(AppError::Internal("embed-only mock".into()))
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("embed-only mock".into()))
1135        }
1136        async fn generate_with_history(
1137            &self,
1138            _messages: &[(String, String)],
1139        ) -> ares_types::types::Result<LLMResponse> {
1140            Err(AppError::Internal("embed-only mock".into()))
1141        }
1142        async fn generate_with_tools(
1143            &self,
1144            _prompt: &str,
1145            _tools: &[ToolDefinition],
1146        ) -> ares_types::types::Result<LLMResponse> {
1147            Err(AppError::Internal("embed-only mock".into()))
1148        }
1149        async fn generate_with_tools_and_history(
1150            &self,
1151            _messages: &[crate::coordinator::ConversationMessage],
1152            _tools: &[ToolDefinition],
1153        ) -> ares_types::types::Result<LLMResponse> {
1154            Err(AppError::Internal("embed-only mock".into()))
1155        }
1156        async fn stream(
1157            &self,
1158            _prompt: &str,
1159        ) -> ares_types::types::Result<
1160            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1161        > {
1162            Err(AppError::Internal("embed-only mock".into()))
1163        }
1164        async fn stream_with_system(
1165            &self,
1166            _system: &str,
1167            _prompt: &str,
1168        ) -> ares_types::types::Result<
1169            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1170        > {
1171            Err(AppError::Internal("embed-only mock".into()))
1172        }
1173        async fn stream_with_history(
1174            &self,
1175            _messages: &[(String, String)],
1176        ) -> ares_types::types::Result<
1177            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1178        > {
1179            Err(AppError::Internal("embed-only mock".into()))
1180        }
1181        fn model_name(&self) -> &str {
1182            "embed-mock"
1183        }
1184        async fn embed(&self, inputs: &[String]) -> ares_types::types::Result<Vec<Vec<f32>>> {
1185            self.called.store(true, std::sync::atomic::Ordering::SeqCst);
1186            Ok(inputs.iter().map(|s| vec![s.len() as f32]).collect())
1187        }
1188    }
1189
1190    #[tokio::test]
1191    async fn llm_complete_runs_generate_without_events() {
1192        let (client, generated) = EchoClient::new();
1193        let llm = Llm::for_test(std::sync::Arc::new(client));
1194        let ctx = Context::new_root();
1195        let out = llm.complete(&ctx, "hi").await.expect("complete");
1196        assert_eq!(out, "echo:hi");
1197        assert!(generated.load(std::sync::atomic::Ordering::SeqCst));
1198    }
1199
1200    #[tokio::test]
1201    async fn llm_complete_waterfall_rewrites_prompt() {
1202        let (client, _) = EchoClient::new();
1203        let llm = Llm::for_test(std::sync::Arc::new(client));
1204        let ctx = Context::new_root();
1205        let events = ctx.provide(EventsService::new());
1206        events.on_waterfall(
1207            cordis::events_catalog::ev::LLM_COMPLETE.to_string(),
1208            |mut payload, next| async move {
1209                if let Some(p) = payload.get("prompt").and_then(|v| v.as_str()) {
1210                    payload["prompt"] = serde_json::json!(format!("WRAP:{p}"));
1211                }
1212                next(payload).await
1213            },
1214        );
1215        let out = llm.complete(&ctx, "hi").await.expect("complete");
1216        assert_eq!(out, "echo:WRAP:hi");
1217    }
1218
1219    #[tokio::test]
1220    async fn llm_complete_short_circuit_skips_generate() {
1221        let (client, generated) = EchoClient::new();
1222        let llm = Llm::for_test(std::sync::Arc::new(client));
1223        let ctx = Context::new_root();
1224        let events = ctx.provide(EventsService::new());
1225        events.on_waterfall(
1226            cordis::events_catalog::ev::LLM_COMPLETE.to_string(),
1227            |_payload, _next| async move { Ok(serde_json::json!({ "content": "cached" })) },
1228        );
1229        let out = llm.complete(&ctx, "hi").await.expect("complete");
1230        assert_eq!(out, "cached");
1231        assert!(
1232            !generated.load(std::sync::atomic::Ordering::SeqCst),
1233            "dummy generate must stay false when handler skips next"
1234        );
1235    }
1236
1237    #[tokio::test]
1238    async fn llm_get_client_waterfall_deny() {
1239        let (client, _) = EchoClient::new();
1240        let llm = Llm::for_test(std::sync::Arc::new(client));
1241        let ctx = Context::new_root();
1242        let events = ctx.provide(EventsService::new());
1243        events.on_waterfall(
1244            cordis::events_catalog::ev::LLM_GET_CLIENT.to_string(),
1245            |_payload, _next| async move { Ok(serde_json::json!({ "deny": true })) },
1246        );
1247        let err = match llm
1248            .get_client(&ctx, CapabilityRequirements::default())
1249            .await
1250        {
1251            Ok(_) => panic!("deny"),
1252            Err(err) => err,
1253        };
1254        assert!(matches!(err, AppError::InvalidInput(msg) if msg == "llm.get_client denied"));
1255    }
1256
1257    #[tokio::test]
1258    async fn llm_embed_runs_without_events() {
1259        let (client, called) = EmbedClient::new();
1260        let llm = Llm::for_test(std::sync::Arc::new(client));
1261        let ctx = std::sync::Arc::new(Context::new_root());
1262        let out = llm
1263            .embed(&ctx, &["ab".into(), "c".into()])
1264            .await
1265            .expect("embed");
1266        assert_eq!(out, vec![vec![2.0], vec![1.0]]);
1267        assert!(called.load(std::sync::atomic::Ordering::SeqCst));
1268    }
1269
1270    #[tokio::test]
1271    async fn llm_embed_waterfall_rewrites_inputs() {
1272        let (client, _) = EmbedClient::new();
1273        let llm = Llm::for_test(std::sync::Arc::new(client));
1274        let ctx = Context::new_root();
1275        let events = ctx.provide(EventsService::new());
1276        events.on_waterfall(
1277            cordis::events_catalog::ev::LLM_EMBED.to_string(),
1278            |mut payload, next| async move {
1279                if let Some(inputs) = payload.get("inputs").and_then(|v| v.as_array()) {
1280                    let rewritten: Vec<String> = inputs
1281                        .iter()
1282                        .filter_map(|v| v.as_str().map(|s| format!("WRAP:{s}")))
1283                        .collect();
1284                    payload["inputs"] = serde_json::json!(rewritten);
1285                }
1286                next(payload).await
1287            },
1288        );
1289        let out = llm.embed(&ctx, &["hi".into()]).await.expect("embed");
1290        assert_eq!(out, vec![vec![7.0]]);
1291    }
1292
1293    #[tokio::test]
1294    async fn llm_embed_short_circuit_skips_client() {
1295        let (client, called) = EmbedClient::new();
1296        let llm = Llm::for_test(std::sync::Arc::new(client));
1297        let ctx = std::sync::Arc::new(Context::new_root());
1298        let events = ctx.provide(EventsService::new());
1299        events.on_waterfall(
1300            cordis::events_catalog::ev::LLM_EMBED.to_string(),
1301            |_payload, _next| async move { Ok(serde_json::json!({ "embeddings": [[9.0, 8.0]] })) },
1302        );
1303        let out = llm.embed(&ctx, &["hi".into()]).await.expect("embed");
1304        assert_eq!(out, vec![vec![9.0, 8.0]]);
1305        assert!(
1306            !called.load(std::sync::atomic::Ordering::SeqCst),
1307            "client embed must stay false when handler skips next"
1308        );
1309    }
1310
1311    #[test]
1312    fn llm_list_models_exposes_registry_models() {
1313        let mut registry = ProviderRegistry::new();
1314        registry.register_model(
1315            "stub-model",
1316            crate::config::ModelConfig {
1317                provider: "stub".into(),
1318                model: "stub-model".into(),
1319                temperature: 0.7,
1320                max_tokens: 512,
1321            },
1322        );
1323        let llm = Llm::new(
1324            Arc::new(registry),
1325            Arc::new(ClientPool::with_defaults()),
1326            None,
1327        );
1328        let models = llm.list_models();
1329        assert!(
1330            models
1331                .iter()
1332                .any(|m| m.name == "stub-model" && m.provider == "stub"),
1333            "Llm::list_models should expose registry models: {models:?}"
1334        );
1335    }
1336
1337    /// Mock recording every `set_hints` payload routed through it.
1338    #[derive(Default)]
1339    struct HintRecordingClient {
1340        hints: parking_lot::Mutex<Vec<GenerationHints>>,
1341        supports: bool,
1342    }
1343
1344    #[async_trait]
1345    impl LLMClient for HintRecordingClient {
1346        async fn generate(&self, _prompt: &str) -> ares_types::types::Result<String> {
1347            Err(AppError::Internal("unused".into()))
1348        }
1349
1350        async fn generate_with_system(
1351            &self,
1352            _system: &str,
1353            _prompt: &str,
1354        ) -> ares_types::types::Result<String> {
1355            Err(AppError::Internal("unused".into()))
1356        }
1357
1358        async fn generate_with_history(
1359            &self,
1360            _messages: &[(String, String)],
1361        ) -> ares_types::types::Result<LLMResponse> {
1362            Err(AppError::Internal("unused".into()))
1363        }
1364
1365        async fn generate_with_tools(
1366            &self,
1367            _prompt: &str,
1368            _tools: &[ToolDefinition],
1369        ) -> ares_types::types::Result<LLMResponse> {
1370            Err(AppError::Internal("unused".into()))
1371        }
1372
1373        async fn generate_with_tools_and_history(
1374            &self,
1375            _messages: &[crate::coordinator::ConversationMessage],
1376            _tools: &[ToolDefinition],
1377        ) -> ares_types::types::Result<LLMResponse> {
1378            Err(AppError::Internal("unused".into()))
1379        }
1380
1381        async fn stream(
1382            &self,
1383            _prompt: &str,
1384        ) -> ares_types::types::Result<
1385            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1386        > {
1387            Err(AppError::Internal("unused".into()))
1388        }
1389
1390        async fn stream_with_system(
1391            &self,
1392            _system: &str,
1393            _prompt: &str,
1394        ) -> ares_types::types::Result<
1395            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1396        > {
1397            Err(AppError::Internal("unused".into()))
1398        }
1399
1400        async fn stream_with_history(
1401            &self,
1402            _messages: &[(String, String)],
1403        ) -> ares_types::types::Result<
1404            Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
1405        > {
1406            Err(AppError::Internal("unused".into()))
1407        }
1408
1409        fn model_name(&self) -> &str {
1410            "hint-recording-mock"
1411        }
1412
1413        fn supports_hints(&self) -> bool {
1414            self.supports
1415        }
1416
1417        fn set_hints(&self, hints: GenerationHints) {
1418            self.hints.lock().push(hints);
1419        }
1420    }
1421
1422    #[test]
1423    fn boxed_arc_client_forwards_hints_to_inner_client() {
1424        let concrete = Arc::new(HintRecordingClient {
1425            supports: true,
1426            hints: parking_lot::Mutex::new(Vec::new()),
1427        });
1428        let recorder_handle = Arc::clone(&concrete);
1429        let inner: Arc<dyn LLMClient> = concrete;
1430        let boxed = BoxedArcClient(Arc::clone(&inner));
1431
1432        assert!(boxed.supports_hints());
1433        boxed.set_hints(GenerationHints {
1434            json_mode: true,
1435            suppress_reasoning: false,
1436            max_tokens: Some(256),
1437            guided_grammar: None,
1438            ..Default::default()
1439        });
1440        boxed.set_hints(GenerationHints::default());
1441
1442        // The adapter must reach the INNER client instead of stopping at the
1443        // trait's no-op defaults on the box itself. `concrete` and `inner`
1444        // share one allocation, so writes via the box are visible here.
1445        let recorded = recorder_handle.hints.lock();
1446        assert_eq!(
1447            recorded.len(),
1448            2,
1449            "both set_hints calls must reach the inner client"
1450        );
1451        assert!(recorded[0].json_mode && recorded[0].max_tokens == Some(256));
1452        assert_eq!(recorded[1], GenerationHints::default());
1453    }
1454}