use crate::Result;
use crate::HttpError;
use crate::{
auth::middleware::AuthUser,
research::coordinator::ResearchCoordinator,
types::{Claims, ResearchRequest, ResearchResponse, Source},
utils::toml_config::{AresConfig, WorkflowConfig},
};
use std::sync::Arc;
use cordis::Context;
use axum::{
extract::{Extension, State},
Json,
};
use std::time::Instant;
async fn jwt_tenant_tier(ctx: &Arc<Context>, tenant_id: &str) -> ares_types::models::TenantTier {
let Some(db) = ctx.get::<ares_store::TenantDb>() else {
return ares_types::models::TenantTier::Free;
};
match db.get_tenant(tenant_id).await {
Ok(Some(tenant)) => tenant.tier,
_ => ares_types::models::TenantTier::Free,
}
}
async fn intercept_jwt_tenant(
ctx: Arc<Context>,
claims: &Claims,
tenant_ext: Option<Extension<ares_types::models::TenantContext>>,
) -> Arc<Context> {
if let Some(Extension(tc)) = tenant_ext {
return ares_agent::request_tenant_ctx(&ctx, tc);
}
let Some(tenant_id) = claims.tenant_id.clone() else {
return ares_agent::request_user_scope(&ctx, &claims.sub);
};
let tier = jwt_tenant_tier(&ctx, &tenant_id).await;
ares_agent::request_tenant_ctx(
&ctx,
ares_types::models::TenantContext::new(tenant_id, tier),
)
}
fn resolve_research_limits(
payload: &ResearchRequest,
workflow: Option<&WorkflowConfig>,
) -> (u8, u8) {
if let Some(workflow) = workflow {
(
payload.depth.unwrap_or(workflow.max_depth),
payload.max_iterations.unwrap_or(workflow.max_iterations),
)
} else {
(
payload.depth.unwrap_or(2),
payload.max_iterations.unwrap_or(5),
)
}
}
fn orchestrator_model_name(config: &AresConfig) -> &str {
config
.get_agent("orchestrator")
.map(|a| a.model.as_str())
.unwrap_or("powerful")
}
pub(crate) fn plan_research_run<'a>(
config: &'a AresConfig,
payload: &ResearchRequest,
) -> (u8, u8, &'a str) {
let (depth, max_iterations) = resolve_research_limits(payload, config.get_workflow("research"));
let model_name = orchestrator_model_name(config);
(depth, max_iterations, model_name)
}
fn research_emergency_stop_message() -> &'static str {
"All agents are currently under human review. Please try again later."
}
fn ensure_research_emergency_stop_inactive(
stop: &ares_agent::EmergencyStop,
) -> Result<()> {
if stop.is_active() {
return Err(HttpError::from(ares_types::types::AppError::Unavailable(
research_emergency_stop_message().to_string(),
)));
}
Ok(())
}
pub(crate) fn finalize_research_response(
findings: String,
sources: Vec<Source>,
duration: std::time::Duration,
) -> ResearchResponse {
ResearchResponse {
findings,
sources,
duration_ms: duration.as_millis() as u64,
}
}
pub async fn deep_research(
State(ctx): State<Arc<Context>>,
AuthUser(claims): AuthUser,
tenant_ctx: Option<Extension<ares_types::models::TenantContext>>,
Json(payload): Json<ResearchRequest>,
) -> Result<Json<ResearchResponse>> {
let ctx = intercept_jwt_tenant(ctx, &claims, tenant_ctx).await;
let start = Instant::now();
ensure_research_emergency_stop_inactive(&ctx.get::<ares_agent::EmergencyStop>().expect("not provided"))?;
let config = ctx.get::<crate::overlay::AresConfigManager>().expect("not provided").config();
let (depth, max_iterations, model_name) = plan_research_run(&config, &payload);
let llm = ctx.get::<ares_llm::Llm>().ok_or_else(|| {
HttpError::from(ares_types::types::AppError::Configuration(
"Llm service is not provided on the request context".to_string(),
))
})?;
let model_ctx = ctx.with_intercept(ares_llm::ModelOverride {
model: model_name.to_string(),
});
let llm_client = llm
.get_client_boxed(&model_ctx, ares_llm::CapabilityRequirements::default())
.await?;
let coordinator = ResearchCoordinator::new(llm_client, depth, max_iterations);
let (findings, sources) = coordinator.research(&payload.query).await?;
let duration = start.elapsed();
Ok(Json(finalize_research_response(
findings, sources, duration,
)))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::overlay::{
AgentConfig, BillingConfig, DatabaseConfig, DynamicConfigPaths, RagConfig,
};
use crate::config::{AuthConfig, ServerConfig};
use std::collections::HashMap;
use std::time::Duration;
fn minimal_overlay_config(agents: HashMap<String, AgentConfig>) -> AresConfig {
AresConfig {
server: ServerConfig::default(),
auth: AuthConfig::default(),
database: DatabaseConfig::default(),
nvidia: None,
config: DynamicConfigPaths::default(),
providers: HashMap::new(),
models: HashMap::new(),
tools: HashMap::new(),
agents,
workflows: HashMap::new(),
rag: RagConfig::default(),
billing: BillingConfig::default(),
skills: None,
}
}
fn sample_workflow() -> WorkflowConfig {
WorkflowConfig {
entry_agent: "orchestrator".to_string(),
fallback_agent: None,
max_depth: 4,
max_iterations: 12,
parallel_subagents: true,
}
}
fn config_with_research_workflow(workflow: WorkflowConfig) -> AresConfig {
let mut config = minimal_overlay_config(HashMap::new());
config.workflows.insert("research".to_string(), workflow);
config
}
fn config_with_orchestrator(model: &str) -> AresConfig {
let mut agents = HashMap::new();
agents.insert(
"orchestrator".to_string(),
AgentConfig {
model: model.to_string(),
system_prompt: None,
tools: vec![],
allowed_tools: None,
max_tool_iterations: 5,
parallel_tools: false,
extra: HashMap::new(),
compaction_enabled: None,
},
);
minimal_overlay_config(agents)
}
fn sample_request(query: &str) -> ResearchRequest {
ResearchRequest {
query: query.to_string(),
depth: None,
max_iterations: None,
}
}
#[test]
fn resolve_research_limits_uses_workflow_defaults() {
let payload = sample_request("topic");
let workflow = sample_workflow();
assert_eq!(resolve_research_limits(&payload, Some(&workflow)), (4, 12));
}
#[test]
fn resolve_research_limits_honors_payload_overrides() {
let payload = ResearchRequest {
query: "topic".to_string(),
depth: Some(1),
max_iterations: Some(3),
};
assert_eq!(
resolve_research_limits(&payload, Some(&sample_workflow())),
(1, 3)
);
}
#[test]
fn resolve_research_limits_falls_back_when_workflow_missing() {
let payload = sample_request("topic");
assert_eq!(resolve_research_limits(&payload, None), (2, 5));
}
#[test]
fn resolve_research_limits_partial_payload_override_depth_only() {
let payload = ResearchRequest {
query: "topic".to_string(),
depth: Some(6),
max_iterations: None,
};
assert_eq!(
resolve_research_limits(&payload, Some(&sample_workflow())),
(6, 12)
);
}
#[test]
fn resolve_research_limits_partial_payload_override_iterations_only() {
let payload = ResearchRequest {
query: "topic".to_string(),
depth: None,
max_iterations: Some(9),
};
assert_eq!(
resolve_research_limits(&payload, Some(&sample_workflow())),
(4, 9)
);
}
#[test]
fn resolve_research_limits_zero_payload_overrides() {
let payload = ResearchRequest {
query: "topic".to_string(),
depth: Some(0),
max_iterations: Some(0),
};
assert_eq!(
resolve_research_limits(&payload, Some(&sample_workflow())),
(0, 0)
);
assert_eq!(resolve_research_limits(&payload, None), (0, 0));
}
#[test]
fn resolve_research_limits_u8_max_overrides() {
let payload = ResearchRequest {
query: "topic".to_string(),
depth: Some(u8::MAX),
max_iterations: Some(u8::MAX),
};
assert_eq!(
resolve_research_limits(&payload, Some(&sample_workflow())),
(u8::MAX, u8::MAX)
);
}
#[test]
fn resolve_research_limits_respects_workflow_zero_defaults() {
let workflow = WorkflowConfig {
max_depth: 0,
max_iterations: 0,
..sample_workflow()
};
let payload = sample_request("topic");
assert_eq!(resolve_research_limits(&payload, Some(&workflow)), (0, 0));
}
#[test]
fn orchestrator_model_name_reads_configured_agent() {
let config = config_with_orchestrator("claude-research");
assert_eq!(orchestrator_model_name(&config), "claude-research");
}
#[test]
fn orchestrator_model_name_defaults_without_agent() {
let config = minimal_overlay_config(HashMap::new());
assert_eq!(orchestrator_model_name(&config), "powerful");
}
#[test]
fn orchestrator_model_name_ignores_non_orchestrator_agents() {
let mut agents = HashMap::new();
agents.insert(
"researcher".to_string(),
AgentConfig {
model: "other-model".to_string(),
system_prompt: None,
tools: vec![],
allowed_tools: None,
max_tool_iterations: 1,
parallel_tools: false,
extra: HashMap::new(),
compaction_enabled: None,
},
);
let config = minimal_overlay_config(agents);
assert_eq!(orchestrator_model_name(&config), "powerful");
}
#[test]
fn orchestrator_model_name_case_sensitive_agent_key() {
let mut agents = HashMap::new();
agents.insert(
"Orchestrator".to_string(),
AgentConfig {
model: "wrong-case".to_string(),
system_prompt: None,
tools: vec![],
allowed_tools: None,
max_tool_iterations: 1,
parallel_tools: false,
extra: HashMap::new(),
compaction_enabled: None,
},
);
let config = minimal_overlay_config(agents);
assert_eq!(orchestrator_model_name(&config), "powerful");
}
#[test]
fn orchestrator_model_name_preserves_empty_model_string() {
let config = config_with_orchestrator("");
assert_eq!(orchestrator_model_name(&config), "");
}
#[test]
fn orchestrator_model_name_does_not_read_llm_env() {
std::env::remove_var("OPENAI_API_KEY");
std::env::remove_var("ANTHROPIC_API_KEY");
std::env::remove_var("OLLAMA_URL");
let config = config_with_orchestrator("local-model");
assert_eq!(orchestrator_model_name(&config), "local-model");
}
#[test]
fn plan_research_run_uses_research_workflow_from_config() {
let config = config_with_research_workflow(sample_workflow());
let payload = sample_request("climate");
assert_eq!(plan_research_run(&config, &payload), (4, 12, "powerful"));
}
#[test]
fn plan_research_run_combines_overrides_and_orchestrator_model() {
let mut config = config_with_orchestrator("gpt-research");
config
.workflows
.insert("research".to_string(), sample_workflow());
let payload = ResearchRequest {
query: "topic".to_string(),
depth: Some(2),
max_iterations: Some(8),
};
assert_eq!(plan_research_run(&config, &payload), (2, 8, "gpt-research"));
}
#[test]
fn plan_research_run_without_workflow_uses_hardcoded_defaults() {
let config = config_with_orchestrator("fast");
let payload = sample_request("topic");
assert_eq!(plan_research_run(&config, &payload), (2, 5, "fast"));
}
#[test]
fn ensure_research_emergency_stop_inactive_rejects_active_stop() {
let active = ares_agent::EmergencyStop::new(true);
let err = ensure_research_emergency_stop_inactive(&active).unwrap_err();
match err.0 {
ares_types::types::AppError::Unavailable(msg) => {
assert_eq!(msg, research_emergency_stop_message())
}
other => panic!("expected Unavailable, got {other:?}"),
}
}
#[test]
fn ensure_research_emergency_stop_inactive_allows_clear_stop() {
let active = ares_agent::EmergencyStop::new(false);
assert!(ensure_research_emergency_stop_inactive(&active).is_ok());
}
#[test]
fn finalize_research_response_maps_duration_to_millis() {
let response =
finalize_research_response("done".to_string(), vec![], Duration::from_millis(1500));
assert_eq!(response.findings, "done");
assert!(response.sources.is_empty());
assert_eq!(response.duration_ms, 1500);
}
#[test]
fn finalize_research_response_truncates_sub_millisecond_duration() {
let response =
finalize_research_response("fast".to_string(), vec![], Duration::from_nanos(999_999));
assert_eq!(response.duration_ms, 0);
}
#[test]
fn finalize_research_response_preserves_sources() {
let sources = vec![Source {
title: "Doc".to_string(),
url: None,
relevance_score: 0.5,
}];
let response = finalize_research_response("x".into(), sources.clone(), Duration::ZERO);
assert_eq!(response.sources.len(), 1);
assert_eq!(response.sources[0].title, "Doc");
assert!(response.sources[0].url.is_none());
}
#[test]
fn research_response_serde_roundtrip() {
let response = ResearchResponse {
findings: "Summary".to_string(),
sources: vec![Source {
title: "Paper".to_string(),
url: Some("https://example.com".to_string()),
relevance_score: 0.9,
}],
duration_ms: 42,
};
let json = serde_json::to_string(&response).unwrap();
let parsed: ResearchResponse = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.findings, "Summary");
assert_eq!(parsed.sources.len(), 1);
assert_eq!(parsed.duration_ms, 42);
}
#[test]
fn research_response_serde_empty_sources() {
let response = ResearchResponse {
findings: String::new(),
sources: vec![],
duration_ms: 0,
};
let json = serde_json::to_string(&response).unwrap();
let parsed: ResearchResponse = serde_json::from_str(&json).unwrap();
assert!(parsed.sources.is_empty());
assert_eq!(parsed.findings, "");
}
#[test]
fn research_request_serde_roundtrip() {
let req = ResearchRequest {
query: "rust async".to_string(),
depth: Some(2),
max_iterations: Some(10),
};
let json = serde_json::to_string(&req).unwrap();
let parsed: ResearchRequest = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.query, req.query);
assert_eq!(parsed.depth, req.depth);
assert_eq!(parsed.max_iterations, req.max_iterations);
}
#[test]
fn research_request_debug_includes_query() {
let req = sample_request("debug me");
let debug = format!("{req:?}");
assert!(debug.contains("debug me"));
}
#[test]
fn research_response_debug_includes_duration() {
let response = ResearchResponse {
findings: "f".into(),
sources: vec![],
duration_ms: 99,
};
let debug = format!("{response:?}");
assert!(debug.contains("99"));
}
#[test]
fn research_request_deserializes_query_only() {
let req: ResearchRequest =
serde_json::from_str(r#"{"query":"quantum computing"}"#).unwrap();
assert_eq!(req.query, "quantum computing");
assert!(req.depth.is_none());
assert!(req.max_iterations.is_none());
}
#[test]
fn research_request_deserializes_with_overrides() {
let req: ResearchRequest =
serde_json::from_str(r#"{"query":"ai safety","depth":3,"max_iterations":7}"#).unwrap();
assert_eq!(req.depth, Some(3));
assert_eq!(req.max_iterations, Some(7));
}
#[test]
fn research_request_deserializes_explicit_null_overrides() {
let req: ResearchRequest =
serde_json::from_str(r#"{"query":"topic","depth":null,"max_iterations":null}"#)
.unwrap();
assert!(req.depth.is_none());
assert!(req.max_iterations.is_none());
}
#[test]
fn research_request_rejects_missing_query() {
let err = serde_json::from_str::<ResearchRequest>(r#"{}"#).unwrap_err();
assert!(err.to_string().contains("query"));
}
fn sample_claims(sub: &str, tenant_id: Option<&str>) -> Claims {
Claims {
sub: sub.into(),
email: "u@example.com".into(),
exp: 0,
iat: 0,
jti: String::new(),
tenant_id: tenant_id.map(str::to_string),
}
}
#[tokio::test]
async fn intercept_jwt_tenant_user_isolate_skips_dummy_tenant_context() {
let ctx = Context::new_root();
let scoped = intercept_jwt_tenant(ctx, &sample_claims("user-1", None), None).await;
assert!(scoped.get::<ares_types::models::TenantContext>().is_none());
assert_eq!(
scoped
.isolate_label(std::any::TypeId::of::<ares_agent::Execute>())
.as_deref(),
None
);
}
#[cfg(feature = "postgres")]
#[tokio::test]
async fn intercept_jwt_tenant_opens_realm_then_intercepts() {
let ctx = Context::new_root();
ctx.provide(ares_store::TenantRealms::new(
std::any::TypeId::of::<ares_tools::Tools>(),
std::any::TypeId::of::<ares_agent::Execute>(),
));
let tc = ares_types::models::TenantContext::new(
"acme".into(),
ares_types::models::TenantTier::Pro,
);
let scoped = intercept_jwt_tenant(
ctx.clone(),
&sample_claims("user-1", Some("acme")),
Some(Extension(tc)),
)
.await;
assert_eq!(
scoped
.get::<ares_types::models::TenantContext>()
.expect("intercept")
.tenant_id,
"acme"
);
let realm = ctx
.get::<ares_store::TenantRealms>()
.expect("TenantRealms")
.open(&ctx, "acme");
assert!(realm.get::<ares_types::models::TenantContext>().is_none());
assert_eq!(
scoped
.isolate_label(std::any::TypeId::of::<ares_tools::Tools>())
.as_deref(),
Some("acme")
);
}
}