use std::collections::HashSet;
use std::future::Future;
use std::sync::Arc;
use crate::nvidia_catalog::NvidiaCatalogCache;
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use cordis::{Context, CordisError, EventsService, Service};
use parking_lot::RwLock;
use crate::capabilities::CapabilityRequirements;
use crate::client::{GenerationHints, LLMClient, LLMResponse};
use crate::config::ProviderConfig;
use crate::pool::ClientPool;
use crate::provider_registry::{
ConfigBasedLLMFactory, ModelInfo, ProviderRegistry, RuntimeProviderEntry,
};
use ares_types::types::{AppError, ToolDefinition};
#[derive(Debug, Clone)]
pub struct ModelOverride {
pub model: String,
}
impl Service for ModelOverride {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TenantModelPolicy {
tenant_id: String,
allowed_models: HashSet<String>,
}
impl TenantModelPolicy {
pub fn new<I, S>(tenant_id: impl Into<String>, allowed_models: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
Self {
tenant_id: tenant_id.into(),
allowed_models: allowed_models.into_iter().map(Into::into).collect(),
}
}
pub fn tenant_id(&self) -> &str {
&self.tenant_id
}
pub fn allows(&self, model: &str) -> bool {
self.allowed_models.contains(model)
}
pub fn denial_message(tenant_id: &str, model: &str) -> String {
format!(
"Model '{}' is not allowed for tenant '{}'",
model, tenant_id
)
}
pub fn denial_error(tenant_id: &str, model: &str) -> AppError {
AppError::Auth(Self::denial_message(tenant_id, model))
}
pub fn authorize(&self, model: &str) -> Result<(), AppError> {
if self.allows(model) {
Ok(())
} else {
Err(Self::denial_error(&self.tenant_id, model))
}
}
}
impl Service for TenantModelPolicy {}
#[derive(Debug, Clone, Default)]
pub enum Breaker {
#[default]
Closed,
Open { until: DateTime<Utc> },
HalfOpen,
}
impl Breaker {
pub const FAILURE_THRESHOLD: u32 = 5;
pub const COOLDOWN_SECS: i64 = 30;
pub fn check(&self) -> bool {
match self {
Breaker::Closed => true,
Breaker::HalfOpen => true,
Breaker::Open { until } => {
Utc::now() >= *until
}
}
}
pub fn is_closed(&self) -> bool {
matches!(self, Breaker::Closed)
}
pub fn transition_on_failure(&self) -> Breaker {
let now = Utc::now();
let cooldown = chrono::Duration::seconds(Self::COOLDOWN_SECS);
match self {
Breaker::Closed => Breaker::Open {
until: now + cooldown,
},
Breaker::HalfOpen => Breaker::Open {
until: now + cooldown,
},
Breaker::Open { .. } => Breaker::Open {
until: now + cooldown,
},
}
}
pub fn transition_on_failure_with_count(&self, failures: u32) -> Breaker {
if failures >= Self::FAILURE_THRESHOLD {
let now = Utc::now();
Breaker::Open {
until: now + chrono::Duration::seconds(Self::COOLDOWN_SECS),
}
} else {
Breaker::Closed
}
}
}
pub struct Llm {
pub(crate) provider_registry: Arc<ProviderRegistry>,
pub(crate) catalog: Option<Arc<NvidiaCatalogCache>>,
pub(crate) pool: Arc<ClientPool>,
pub(crate) factory: Option<Arc<ConfigBasedLLMFactory>>,
breaker: RwLock<Breaker>,
failures: RwLock<u32>,
test_client: Option<Arc<dyn LLMClient>>,
}
impl Llm {
pub fn new(
provider_registry: Arc<ProviderRegistry>,
pool: Arc<ClientPool>,
catalog: Option<Arc<NvidiaCatalogCache>>,
) -> Self {
Self {
provider_registry,
catalog,
pool,
factory: None,
breaker: RwLock::new(Breaker::Closed),
failures: RwLock::new(0),
test_client: None,
}
}
pub fn with_factory(mut self, factory: Arc<ConfigBasedLLMFactory>) -> Self {
self.factory = Some(factory);
self
}
pub fn from_client(client: Arc<dyn LLMClient>) -> Self {
let mut llm = Self::new(
Arc::new(ProviderRegistry::new()),
Arc::new(ClientPool::with_defaults()),
None,
);
llm.test_client = Some(client);
llm
}
#[cfg(test)]
pub(crate) fn for_test(client: Arc<dyn LLMClient>) -> Self {
Self::from_client(client)
}
pub(crate) fn provider_registry(&self) -> Arc<ProviderRegistry> {
Arc::clone(&self.provider_registry)
}
pub fn registry(&self) -> Arc<ProviderRegistry> {
self.provider_registry()
}
pub fn with_breaker(
provider_registry: Arc<ProviderRegistry>,
catalog: Option<Arc<NvidiaCatalogCache>>,
pool: Arc<ClientPool>,
breaker: Breaker,
) -> Self {
Self {
provider_registry,
catalog,
pool,
factory: None,
breaker: RwLock::new(breaker),
failures: RwLock::new(0),
test_client: None,
}
}
pub fn with_catalog(
provider_registry: Arc<ProviderRegistry>,
catalog: Arc<NvidiaCatalogCache>,
pool: Arc<ClientPool>,
) -> Self {
Self::new(provider_registry, pool, Some(catalog))
}
pub fn breaker(&self) -> Breaker {
self.breaker.read().clone()
}
pub fn trip(&self, until: DateTime<Utc>) {
*self.breaker.write() = Breaker::Open { until };
}
pub fn half_open(&self) {
*self.breaker.write() = Breaker::HalfOpen;
}
pub fn reset(&self) {
*self.breaker.write() = Breaker::Closed;
*self.failures.write() = 0;
}
pub fn record_success(&self) {
*self.breaker.write() = Breaker::Closed;
*self.failures.write() = 0;
}
pub fn record_failure(&self) {
let mut failures = self.failures.write();
*failures = failures.saturating_add(1);
let count = *failures;
drop(failures);
let mut b = self.breaker.write();
match &*b {
Breaker::HalfOpen => {
*b = b.transition_on_failure();
}
Breaker::Closed => {
if count >= Breaker::FAILURE_THRESHOLD {
*b = Breaker::Open {
until: Utc::now() + chrono::Duration::seconds(Breaker::COOLDOWN_SECS),
};
}
}
Breaker::Open { .. } => {
*b = b.transition_on_failure();
}
}
}
pub fn validate_model_override(&self, ctx: &Arc<Context>) -> Result<(), AppError> {
if let (Some(policy), Some(override_model)) =
(ctx.get::<TenantModelPolicy>(), ctx.get::<ModelOverride>())
{
policy.authorize(&override_model.model)?;
}
Ok(())
}
pub async fn get_client(
&self,
ctx: &Arc<Context>,
capability: CapabilityRequirements,
) -> Result<Arc<dyn LLMClient>, AppError> {
let Some(events) = ctx.get::<EventsService>() else {
return self.get_client_inner(ctx, capability).await;
};
let payload = serde_json::to_value(cordis::LlmGetClientPayload {
capability: format!("{capability:?}"),
deny: None,
model: None,
})
.unwrap_or(serde_json::Value::Null);
let result = events
.waterfall_around(
cordis::events_catalog::ev::LLM_GET_CLIENT.to_string(),
payload,
|payload| async move { Ok(payload) },
)
.await
.map_err(map_cordis)?;
if result.get("deny").and_then(|v| v.as_bool()) == Some(true) {
return Err(AppError::InvalidInput("llm.get_client denied".into()));
}
if let Some(model) = result
.get("model")
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
{
if ctx.get::<ModelOverride>().is_none() {
let intercepted = ctx.with_intercept(ModelOverride {
model: model.to_string(),
});
return self.get_client_inner(&intercepted, capability).await;
}
}
self.get_client_inner(ctx, capability).await
}
async fn get_client_inner(
&self,
ctx: &Arc<Context>,
capability: CapabilityRequirements,
) -> Result<Arc<dyn LLMClient>, AppError> {
if let Some(c) = &self.test_client {
return Ok(Arc::clone(c));
}
self.validate_model_override(ctx)?;
if let Some(ov) = ctx.get::<ModelOverride>() {
if let Ok(guard) = self.pool.try_get(&ov.model).await {
let boxed = guard.take();
return Ok(Arc::from(boxed));
}
if let Ok(client) = self
.provider_registry
.create_client_for_model_ctx(ctx, &ov.model)
.await
{
return Ok(Arc::from(client));
}
}
if let Some(catalog) = &self.catalog {
let _snap = catalog.snapshot(); if let Some(best) = self.provider_registry.find_best_model(&capability) {
if let Ok(client) = self
.provider_registry
.create_client_for_model_ctx(ctx, &best.name)
.await
{
return Ok(Arc::from(client));
}
}
} else if let Some(best) = self.provider_registry.find_best_model(&capability) {
if let Ok(client) = self
.provider_registry
.create_client_for_model_ctx(ctx, &best.name)
.await
{
return Ok(Arc::from(client));
}
}
let client = self
.provider_registry
.resolve_with_capability_fallback(Some(capability))
.await?;
Ok(Arc::from(client))
}
pub async fn get_client_boxed(
&self,
ctx: &Arc<Context>,
capability: CapabilityRequirements,
) -> Result<Box<dyn LLMClient>, AppError> {
let client = self.get_client(ctx, capability).await?;
Ok(Box::new(BoxedArcClient(client)))
}
pub async fn complete(&self, ctx: &Arc<Context>, prompt: &str) -> Result<String, AppError> {
let client = self
.get_client(ctx, CapabilityRequirements::default())
.await?;
let Some(events) = ctx.get::<EventsService>() else {
return client.generate(prompt).await;
};
let payload = serde_json::to_value(cordis::LlmCompleteRequest {
prompt: prompt.to_string(),
})
.unwrap_or(serde_json::Value::Null);
let out = events
.waterfall_around(
cordis::events_catalog::ev::LLM_COMPLETE.to_string(),
payload,
move |payload| {
let client = Arc::clone(&client);
async move {
let prompt = payload
.get("prompt")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string();
let text = client
.generate(&prompt)
.await
.map_err(|e| CordisError::Fiber(e.to_string()))?;
Ok(serde_json::json!({ "prompt": prompt, "content": text }))
}
},
)
.await
.map_err(map_cordis)?;
Ok(out
.get("content")
.and_then(|v| v.as_str())
.unwrap_or("")
.to_string())
}
pub fn find_model_stub(&self, _capability: &str) -> Option<String> {
None
}
pub fn list_models(&self) -> Vec<ModelInfo> {
self.provider_registry.list_models()
}
pub fn has_provider_for_tenant(&self, name: &str, tenant_id: Option<&str>) -> bool {
self.provider_registry
.has_provider_for_tenant(name, tenant_id)
}
pub fn get_provider_for_ctx(&self, ctx: &Arc<Context>, name: &str) -> Option<ProviderConfig> {
self.provider_registry.get_provider_for_ctx(ctx, name)
}
pub fn reload_runtime_providers(
&self,
providers: Vec<RuntimeProviderEntry>,
names: Vec<String>,
) {
self.provider_registry
.reload_runtime_providers(providers, names);
}
}
fn map_cordis(err: CordisError) -> AppError {
AppError::Internal(err.to_string())
}
struct BoxedArcClient(Arc<dyn LLMClient>);
#[async_trait]
impl LLMClient for BoxedArcClient {
async fn generate(&self, prompt: &str) -> ares_types::types::Result<String> {
self.0.generate(prompt).await
}
async fn generate_with_system(
&self,
system: &str,
prompt: &str,
) -> ares_types::types::Result<String> {
self.0.generate_with_system(system, prompt).await
}
async fn generate_with_history(
&self,
messages: &[(String, String)],
) -> ares_types::types::Result<LLMResponse> {
self.0.generate_with_history(messages).await
}
async fn generate_with_tools(
&self,
prompt: &str,
tools: &[ToolDefinition],
) -> ares_types::types::Result<LLMResponse> {
self.0.generate_with_tools(prompt, tools).await
}
async fn generate_with_tools_and_history(
&self,
messages: &[crate::coordinator::ConversationMessage],
tools: &[ToolDefinition],
) -> ares_types::types::Result<LLMResponse> {
self.0
.generate_with_tools_and_history(messages, tools)
.await
}
async fn stream(
&self,
prompt: &str,
) -> ares_types::types::Result<
Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
> {
self.0.stream(prompt).await
}
async fn stream_with_system(
&self,
system: &str,
prompt: &str,
) -> ares_types::types::Result<
Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
> {
self.0.stream_with_system(system, prompt).await
}
async fn stream_with_history(
&self,
messages: &[(String, String)],
) -> ares_types::types::Result<
Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
> {
self.0.stream_with_history(messages).await
}
fn model_name(&self) -> &str {
self.0.model_name()
}
fn supports_hints(&self) -> bool {
self.0.supports_hints()
}
fn set_hints(&self, hints: GenerationHints) {
self.0.set_hints(hints)
}
}
impl Service for Llm {
fn name(&self) -> &'static str {
"Llm"
}
fn init(
&self,
_ctx: &Arc<Context>,
) -> std::pin::Pin<
Box<
dyn Future<Output = Result<Option<Box<dyn cordis::Disposable>>, CordisError>>
+ Send
+ '_,
>,
> {
Box::pin(async { Ok(None) })
}
fn check(&self) -> bool {
self.breaker.read().check()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::capabilities::CapabilityRequirements;
use crate::provider_registry::RuntimeProviderEntry;
use ares_types::models::{TenantContext, TenantTier};
use chrono::Duration;
use cordis::Context;
use std::collections::HashMap;
#[test]
fn breaker_closed_allows() {
assert!(Breaker::Closed.check());
}
#[test]
fn breaker_half_open_allows() {
assert!(Breaker::HalfOpen.check());
}
#[test]
fn breaker_open_future_denies() {
let until = Utc::now() + Duration::seconds(60);
assert!(!Breaker::Open { until }.check());
}
#[test]
fn breaker_open_past_allows() {
let until = Utc::now() - Duration::seconds(1);
assert!(Breaker::Open { until }.check());
}
#[test]
fn breaker_failure_threshold_opens() {
let b = Breaker::Closed;
let next = b.transition_on_failure_with_count(5);
assert!(matches!(next, Breaker::Open { .. }));
let still_closed = b.transition_on_failure_with_count(3);
assert!(matches!(still_closed, Breaker::Closed));
}
#[test]
fn breaker_constants_exist() {
assert_eq!(Breaker::FAILURE_THRESHOLD, 5);
assert_eq!(Breaker::COOLDOWN_SECS, 30);
}
#[test]
fn provider_registry_and_factory_accessors() {
let registry = Arc::new(ProviderRegistry::new());
let pool = Arc::new(ClientPool::with_defaults());
let factory = Arc::new(
ConfigBasedLLMFactory::from_config(HashMap::new(), HashMap::new(), None)
.expect("empty factory config"),
);
let llm = Llm::new(Arc::clone(®istry), pool, None).with_factory(Arc::clone(&factory));
assert!(Arc::ptr_eq(&llm.provider_registry(), ®istry));
}
#[tokio::test]
async fn llm_model_override_via_context() {
let registry = Arc::new(ProviderRegistry::new());
let pool = Arc::new(ClientPool::with_defaults());
let svc = Arc::new(Llm::new(registry, pool, None));
let root = Context::new_root();
root.provide::<Llm>(Llm::new(
Arc::new(ProviderRegistry::new()),
Arc::new(ClientPool::with_defaults()),
None,
));
let req_ctx = root.intercept(ModelOverride {
model: "gpt-4o-mini".into(),
});
assert!(req_ctx.get::<ModelOverride>().is_some());
assert_eq!(req_ctx.get::<ModelOverride>().unwrap().model, "gpt-4o-mini");
assert!(!Arc::as_ptr(&svc.provider_registry).is_null());
let _ = svc.catalog.clone();
let _ = svc.pool.provider_names();
assert!(svc.check());
}
#[tokio::test]
async fn record_failure_threshold_opens_breaker() {
let svc = Llm::new(
Arc::new(ProviderRegistry::new()),
Arc::new(ClientPool::with_defaults()),
None,
);
for _ in 0..Breaker::FAILURE_THRESHOLD {
svc.record_failure();
}
assert!(!svc.check());
svc.record_success();
assert!(svc.check());
}
#[test]
fn tenant_model_policy_allows_and_composes_with_model_override() {
let root = Context::new_root();
let tenant_ctx = root.intercept(TenantModelPolicy::new(
"tenant-a",
["gpt-4o-mini".to_string()],
));
let request = tenant_ctx.intercept(ModelOverride {
model: "gpt-4o-mini".into(),
});
let policy = request
.get::<TenantModelPolicy>()
.expect("policy should be inherited by request context");
let override_model = request
.get::<ModelOverride>()
.expect("model override should be visible in request context");
policy
.authorize(&override_model.model)
.expect("allowed model override should pass policy");
let svc = Llm::new(
Arc::new(ProviderRegistry::new()),
Arc::new(ClientPool::with_defaults()),
None,
);
svc.validate_model_override(&request)
.expect("allowed model override should pass LLM validation");
assert!(root.get::<ModelOverride>().is_none());
assert!(root.get::<TenantModelPolicy>().is_none());
}
#[tokio::test]
async fn disallowed_model_override_is_rejected_before_provider_execution() {
let registry = Arc::new(ProviderRegistry::new());
let svc = Arc::new(Llm::new(
registry,
Arc::new(ClientPool::with_defaults()),
None,
));
let root = Context::new_root();
root.provide_arc(svc.clone());
let tenant_ctx = root.intercept(TenantModelPolicy::new("tenant-a", ["gpt-4o".to_string()]));
let request = tenant_ctx.intercept(ModelOverride {
model: "not-allowed".into(),
});
let err = match svc
.get_client(&request, CapabilityRequirements::default())
.await
{
Ok(_) => panic!("disallowed override must fail before provider lookup"),
Err(err) => err,
};
assert!(matches!(err, AppError::Auth(_)));
assert!(err.to_string().contains("not-allowed"));
assert!(root.get::<ModelOverride>().is_none());
assert!(root.get::<TenantModelPolicy>().is_none());
assert!(matches!(
root.get::<Llm>().expect("global service").breaker(),
Breaker::Closed
));
}
#[tokio::test]
async fn get_client_uses_override_when_catalog_absent() {
let registry = Arc::new(ProviderRegistry::new());
let pool = Arc::new(ClientPool::with_defaults());
let svc = Llm::new(registry, pool, None);
let ctx = Context::new_root();
let req_ctx = ctx.intercept(ModelOverride {
model: "nonexistent-model-xyz".into(),
});
let req = CapabilityRequirements::default();
let res = svc.get_client(&req_ctx, req).await;
assert!(res.is_err());
}
#[tokio::test]
async fn get_client_override_uses_tenant_context_intercept() {
let mut registry = ProviderRegistry::new();
registry.register_model(
"pinned-model",
crate::config::ModelConfig {
provider: "shared-runtime".into(),
model: "tenant-model".into(),
temperature: 0.7,
max_tokens: 512,
},
);
let global = RuntimeProviderEntry {
tenant_id: None,
display_name: "Global Shared".to_string(),
provider_type: "openai-compatible".to_string(),
api_base: "https://global.example.com/v1".to_string(),
auth_type: "api_key".to_string(),
default_model: Some("global-model".to_string()),
headers: HashMap::new(),
api_key: Some("global-key".to_string()),
enabled: true,
};
let tenant = RuntimeProviderEntry {
tenant_id: Some("tenant-a".to_string()),
display_name: "Tenant Shared".to_string(),
provider_type: "openai-compatible".to_string(),
api_base: "https://tenant.example.com/v1".to_string(),
auth_type: "api_key".to_string(),
default_model: Some("tenant-model".to_string()),
headers: HashMap::new(),
api_key: Some("tenant-key".to_string()),
enabled: true,
};
registry.reload_runtime_providers(
vec![global, tenant],
vec!["shared-runtime".to_string(), "shared-runtime".to_string()],
);
let registry = Arc::new(registry);
let svc = Llm::new(registry, Arc::new(ClientPool::with_defaults()), None);
let root = Context::new_root();
let ctx = root
.with_intercept(TenantContext::new("tenant-a".into(), TenantTier::Pro))
.intercept(ModelOverride {
model: "pinned-model".into(),
});
let tenant_client = svc
.get_client(&ctx, CapabilityRequirements::default())
.await;
assert!(
tenant_client.is_ok(),
"tenant intercept should construct a client from the tenant runtime entry: {:?}",
tenant_client.as_ref().err()
);
let unlabeled = root.intercept(ModelOverride {
model: "pinned-model".into(),
});
let fleet_client = svc
.get_client(&unlabeled, CapabilityRequirements::default())
.await;
assert!(
fleet_client.is_ok(),
"unlabeled root with ModelOverride should construct a client from the fleet global runtime entry: {:?}",
fleet_client.as_ref().err()
);
}
struct EchoClient {
generated: std::sync::Arc<std::sync::atomic::AtomicBool>,
}
impl EchoClient {
fn new() -> (Self, std::sync::Arc<std::sync::atomic::AtomicBool>) {
let generated = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false));
(
Self {
generated: std::sync::Arc::clone(&generated),
},
generated,
)
}
}
#[async_trait]
impl LLMClient for EchoClient {
async fn generate(&self, prompt: &str) -> ares_types::types::Result<String> {
self.generated
.store(true, std::sync::atomic::Ordering::SeqCst);
Ok(format!("echo:{prompt}"))
}
async fn generate_with_system(
&self,
_system: &str,
prompt: &str,
) -> ares_types::types::Result<String> {
self.generate(prompt).await
}
async fn generate_with_history(
&self,
_messages: &[(String, String)],
) -> ares_types::types::Result<LLMResponse> {
Ok(LLMResponse {
content: String::new(),
tool_calls: vec![],
finish_reason: "stop".into(),
usage: None,
})
}
async fn generate_with_tools(
&self,
_prompt: &str,
_tools: &[ToolDefinition],
) -> ares_types::types::Result<LLMResponse> {
Ok(LLMResponse {
content: String::new(),
tool_calls: vec![],
finish_reason: "stop".into(),
usage: None,
})
}
async fn generate_with_tools_and_history(
&self,
_messages: &[crate::coordinator::ConversationMessage],
_tools: &[ToolDefinition],
) -> ares_types::types::Result<LLMResponse> {
Ok(LLMResponse {
content: String::new(),
tool_calls: vec![],
finish_reason: "stop".into(),
usage: None,
})
}
async fn stream(
&self,
_prompt: &str,
) -> ares_types::types::Result<
Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
> {
Err(AppError::Internal("echo stream not implemented".into()))
}
async fn stream_with_system(
&self,
_system: &str,
_prompt: &str,
) -> ares_types::types::Result<
Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
> {
Err(AppError::Internal("echo stream not implemented".into()))
}
async fn stream_with_history(
&self,
_messages: &[(String, String)],
) -> ares_types::types::Result<
Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
> {
Err(AppError::Internal("echo stream not implemented".into()))
}
fn model_name(&self) -> &str {
"echo"
}
}
#[tokio::test]
async fn llm_complete_runs_generate_without_events() {
let (client, generated) = EchoClient::new();
let llm = Llm::for_test(std::sync::Arc::new(client));
let ctx = Context::new_root();
let out = llm.complete(&ctx, "hi").await.expect("complete");
assert_eq!(out, "echo:hi");
assert!(generated.load(std::sync::atomic::Ordering::SeqCst));
}
#[tokio::test]
async fn llm_complete_waterfall_rewrites_prompt() {
let (client, _) = EchoClient::new();
let llm = Llm::for_test(std::sync::Arc::new(client));
let ctx = Context::new_root();
let events = ctx.provide(EventsService::new());
events.on_waterfall(
cordis::events_catalog::ev::LLM_COMPLETE.to_string(),
|mut payload, next| async move {
if let Some(p) = payload.get("prompt").and_then(|v| v.as_str()) {
payload["prompt"] = serde_json::json!(format!("WRAP:{p}"));
}
next(payload).await
},
);
let out = llm.complete(&ctx, "hi").await.expect("complete");
assert_eq!(out, "echo:WRAP:hi");
}
#[tokio::test]
async fn llm_complete_short_circuit_skips_generate() {
let (client, generated) = EchoClient::new();
let llm = Llm::for_test(std::sync::Arc::new(client));
let ctx = Context::new_root();
let events = ctx.provide(EventsService::new());
events.on_waterfall(
cordis::events_catalog::ev::LLM_COMPLETE.to_string(),
|_payload, _next| async move { Ok(serde_json::json!({ "content": "cached" })) },
);
let out = llm.complete(&ctx, "hi").await.expect("complete");
assert_eq!(out, "cached");
assert!(
!generated.load(std::sync::atomic::Ordering::SeqCst),
"dummy generate must stay false when handler skips next"
);
}
#[tokio::test]
async fn llm_get_client_waterfall_deny() {
let (client, _) = EchoClient::new();
let llm = Llm::for_test(std::sync::Arc::new(client));
let ctx = Context::new_root();
let events = ctx.provide(EventsService::new());
events.on_waterfall(
cordis::events_catalog::ev::LLM_GET_CLIENT.to_string(),
|_payload, _next| async move { Ok(serde_json::json!({ "deny": true })) },
);
let err = match llm
.get_client(&ctx, CapabilityRequirements::default())
.await
{
Ok(_) => panic!("deny"),
Err(err) => err,
};
assert!(matches!(err, AppError::InvalidInput(msg) if msg == "llm.get_client denied"));
}
#[test]
fn llm_list_models_exposes_registry_models() {
let mut registry = ProviderRegistry::new();
registry.register_model(
"stub-model",
crate::config::ModelConfig {
provider: "stub".into(),
model: "stub-model".into(),
temperature: 0.7,
max_tokens: 512,
},
);
let llm = Llm::new(
Arc::new(registry),
Arc::new(ClientPool::with_defaults()),
None,
);
let models = llm.list_models();
assert!(
models
.iter()
.any(|m| m.name == "stub-model" && m.provider == "stub"),
"Llm::list_models should expose registry models: {models:?}"
);
}
#[derive(Default)]
struct HintRecordingClient {
hints: parking_lot::Mutex<Vec<GenerationHints>>,
supports: bool,
}
#[async_trait]
impl LLMClient for HintRecordingClient {
async fn generate(&self, _prompt: &str) -> ares_types::types::Result<String> {
Err(AppError::Internal("unused".into()))
}
async fn generate_with_system(
&self,
_system: &str,
_prompt: &str,
) -> ares_types::types::Result<String> {
Err(AppError::Internal("unused".into()))
}
async fn generate_with_history(
&self,
_messages: &[(String, String)],
) -> ares_types::types::Result<LLMResponse> {
Err(AppError::Internal("unused".into()))
}
async fn generate_with_tools(
&self,
_prompt: &str,
_tools: &[ToolDefinition],
) -> ares_types::types::Result<LLMResponse> {
Err(AppError::Internal("unused".into()))
}
async fn generate_with_tools_and_history(
&self,
_messages: &[crate::coordinator::ConversationMessage],
_tools: &[ToolDefinition],
) -> ares_types::types::Result<LLMResponse> {
Err(AppError::Internal("unused".into()))
}
async fn stream(
&self,
_prompt: &str,
) -> ares_types::types::Result<
Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
> {
Err(AppError::Internal("unused".into()))
}
async fn stream_with_system(
&self,
_system: &str,
_prompt: &str,
) -> ares_types::types::Result<
Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
> {
Err(AppError::Internal("unused".into()))
}
async fn stream_with_history(
&self,
_messages: &[(String, String)],
) -> ares_types::types::Result<
Box<dyn futures::Stream<Item = ares_types::types::Result<String>> + Send + Unpin>,
> {
Err(AppError::Internal("unused".into()))
}
fn model_name(&self) -> &str {
"hint-recording-mock"
}
fn supports_hints(&self) -> bool {
self.supports
}
fn set_hints(&self, hints: GenerationHints) {
self.hints.lock().push(hints);
}
}
#[test]
fn boxed_arc_client_forwards_hints_to_inner_client() {
let concrete = Arc::new(HintRecordingClient {
supports: true,
hints: parking_lot::Mutex::new(Vec::new()),
});
let recorder_handle = Arc::clone(&concrete);
let inner: Arc<dyn LLMClient> = concrete;
let boxed = BoxedArcClient(Arc::clone(&inner));
assert!(boxed.supports_hints());
boxed.set_hints(GenerationHints {
json_mode: true,
suppress_reasoning: false,
max_tokens: Some(256),
guided_grammar: None,
});
boxed.set_hints(GenerationHints::default());
let recorded = recorder_handle.hints.lock();
assert_eq!(
recorded.len(),
2,
"both set_hints calls must reach the inner client"
);
assert!(recorded[0].json_mode && recorded[0].max_tokens == Some(256));
assert_eq!(recorded[1], GenerationHints::default());
}
}