use std::collections::{HashMap, HashSet};
use std::fmt::Write as _;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Instant;
use crate::core::agent::AgentDefinition;
use crate::core::crew::{CrewProfile, CrewSpec, CrewState, CrewStatus};
use crate::core::task::{ProcessMode, Task, TaskId, TaskResult, TaskStatus};
use crate::llm::{
AuditChain, CostTracker, HooshClient, InferenceRequest, Message, ProviderType, ResponseCache,
Role, cache_key,
};
use crate::server::sse::CrewEvent;
use tokio::sync::Semaphore;
use tokio::sync::broadcast;
use tracing::{debug, info, warn};
use crate::orchestrator::scoring;
pub struct CrewRunner {
spec: CrewSpec,
event_tx: Option<broadcast::Sender<CrewEvent>>,
llm: Option<Arc<HooshClient>>,
cache: Arc<ResponseCache>,
cost_tracker: Arc<CostTracker>,
cancelled: Arc<AtomicBool>,
audit: Option<Arc<AuditChain>>,
}
impl CrewRunner {
pub fn new(spec: CrewSpec) -> Self {
Self {
spec,
event_tx: None,
llm: None,
cache: Arc::new(ResponseCache::new(Default::default())),
cost_tracker: Arc::new(CostTracker::new()),
cancelled: Arc::new(AtomicBool::new(false)),
audit: None,
}
}
pub fn with_cache(mut self, cache: Arc<ResponseCache>) -> Self {
self.cache = cache;
self
}
pub fn with_cost_tracker(mut self, tracker: Arc<CostTracker>) -> Self {
self.cost_tracker = tracker;
self
}
pub fn with_llm(mut self, client: Arc<HooshClient>) -> Self {
self.llm = Some(client);
self
}
pub fn with_cancel_token(mut self, token: Arc<AtomicBool>) -> Self {
self.cancelled = token;
self
}
pub fn with_audit(mut self, audit: Arc<AuditChain>) -> Self {
self.audit = Some(audit);
self
}
pub fn with_events(mut self, tx: broadcast::Sender<CrewEvent>) -> Self {
self.event_tx = Some(tx);
self
}
#[inline]
fn is_cancelled(&self) -> bool {
self.cancelled.load(Ordering::Acquire)
}
fn audit_record(&self, event: &str, level: &str, message: &str, metadata: serde_json::Value) {
if let Some(ref audit) = self.audit {
audit.record(event, level, message, None, None, Some(metadata));
}
}
fn emit(&self, event_type: &str, data: serde_json::Value) {
if let Some(ref tx) = self.event_tx {
let _ = tx.send(CrewEvent {
crew_id: self.spec.id.to_string(),
event_type: event_type.to_string(),
data,
});
}
}
#[tracing::instrument(skip(self), fields(crew_id = %self.spec.id, process = ?self.spec.process))]
pub async fn run(&mut self) -> crate::core::Result<CrewState> {
let crew_start = Instant::now();
info!(crew_id = %self.spec.id, name = %self.spec.name, "starting crew run");
self.emit(
"crew_started",
serde_json::json!({
"name": self.spec.name,
"task_count": self.spec.tasks.len(),
}),
);
let results = match self.spec.process {
ProcessMode::Sequential => self.run_sequential().await?,
ProcessMode::Parallel { max_concurrency } => self.run_parallel(max_concurrency).await?,
ProcessMode::Dag => self.run_dag().await?,
ProcessMode::Hierarchical { .. } => {
warn!("hierarchical mode not yet implemented, falling back to sequential");
self.run_sequential().await?
}
};
let wall_ms = crew_start.elapsed().as_millis() as u64;
let status = if self.is_cancelled() {
CrewStatus::Cancelled
} else if results.iter().all(|r| r.status == TaskStatus::Completed) {
CrewStatus::Completed
} else {
CrewStatus::Failed
};
let task_ms: HashMap<TaskId, u64> = results
.iter()
.filter_map(|r| {
r.metadata
.get("task_duration_ms")
.and_then(|v| v.as_u64())
.map(|ms| (r.task_id, ms))
})
.collect();
let task_cost_usd: HashMap<TaskId, f64> = results
.iter()
.filter_map(|r| {
r.metadata
.get("cost_usd")
.and_then(|v| v.as_f64())
.filter(|&c| c > 0.0)
.map(|c| (r.task_id, c))
})
.collect();
let cost_usd: f64 = task_cost_usd.values().sum();
let mut agent_cost_usd: HashMap<String, f64> = HashMap::new();
for r in &results {
if let Some(cost) = r.metadata.get("cost_usd").and_then(|v| v.as_f64())
&& cost > 0.0
&& let Some(agent_key) = r.metadata.get("agent").and_then(|v| v.as_str())
{
*agent_cost_usd.entry(agent_key.to_string()).or_default() += cost;
}
}
#[cfg(feature = "kavach")]
let sandbox_strength = {
let sandbox_policy = crate::sandbox::policy::SandboxPolicy::process();
Some(crate::sandbox::kavach_bridge::strength_for_policy(&sandbox_policy).value())
};
#[cfg(not(feature = "kavach"))]
let sandbox_strength: Option<u8> = None;
let profile = CrewProfile {
wall_ms,
task_count: results.len(),
task_ms,
cost_usd,
agent_cost_usd,
task_cost_usd,
sandbox_strength,
};
info!(
crew_id = %self.spec.id,
?status,
wall_ms,
"crew run finished"
);
self.emit(
"crew_completed",
serde_json::json!({
"status": format!("{status:?}"),
"task_count": results.len(),
"wall_ms": wall_ms,
}),
);
Ok(CrewState {
crew_id: self.spec.id,
status,
results,
profile: Some(profile),
})
}
#[tracing::instrument(skip(self), fields(crew_id = %self.spec.id, task_count = self.spec.tasks.len()))]
async fn run_sequential(&mut self) -> crate::core::Result<Vec<TaskResult>> {
let mut results = Vec::with_capacity(self.spec.tasks.len());
for i in 0..self.spec.tasks.len() {
if self.is_cancelled() {
info!(crew_id = %self.spec.id, "crew cancelled — stopping sequential execution");
break;
}
let agent = pick_best_agent(&self.spec.agents, &self.spec.tasks[i]);
self.spec.tasks[i].status = TaskStatus::Queued;
let agent_key = agent.as_ref().map(|a| a.agent_key.clone());
if let Some(ref a) = agent {
debug!(task_id = %self.spec.tasks[i].id, agent = %a.agent_key, "assigned");
}
self.emit(
"task_started",
serde_json::json!({
"task_id": self.spec.tasks[i].id.to_string(),
"description": self.spec.tasks[i].description,
"agent": agent_key,
}),
);
self.spec.tasks[i].status = TaskStatus::Running;
let result = execute_task(
&self.spec.tasks[i],
agent.as_ref(),
self.llm.as_ref(),
&self.cache,
&self.cost_tracker,
self.event_tx.as_ref(),
)
.await;
self.spec.tasks[i].status = result.status;
self.emit(
"task_completed",
serde_json::json!({
"task_id": result.task_id.to_string(),
"status": format!("{:?}", result.status),
}),
);
let task_level = if result.status == TaskStatus::Completed {
"info"
} else {
"error"
};
self.audit_record(
"task_completed",
task_level,
&self.spec.tasks[i].description,
serde_json::json!({
"crew_id": self.spec.id.to_string(),
"task_id": result.task_id.to_string(),
"status": format!("{:?}", result.status),
"agent": agent_key,
}),
);
results.push(result);
}
Ok(results)
}
#[tracing::instrument(skip(self), fields(crew_id = %self.spec.id, max_concurrency))]
async fn run_parallel(
&mut self,
max_concurrency: usize,
) -> crate::core::Result<Vec<TaskResult>> {
let semaphore = std::sync::Arc::new(Semaphore::new(max_concurrency));
for task in &mut self.spec.tasks {
task.status = TaskStatus::Queued;
}
let task_snapshots: Vec<(Task, Option<AgentDefinition>)> = self
.spec
.tasks
.iter()
.map(|t| {
let agent = pick_best_agent(&self.spec.agents, t);
(t.clone(), agent)
})
.collect();
for (task, agent) in &task_snapshots {
self.emit(
"task_started",
serde_json::json!({
"task_id": task.id.to_string(),
"description": task.description,
"agent": agent.as_ref().map(|a| &a.agent_key),
}),
);
}
let mut join_set = tokio::task::JoinSet::new();
for (task, agent) in task_snapshots {
let permit = semaphore.clone();
let llm = self.llm.clone();
let cache = Arc::clone(&self.cache);
let cost_tracker = Arc::clone(&self.cost_tracker);
let cancel = Arc::clone(&self.cancelled);
join_set.spawn(async move {
if cancel.load(Ordering::Acquire) {
return TaskResult {
task_id: task.id,
status: TaskStatus::Failed,
output: "crew cancelled".into(),
metadata: Default::default(),
};
}
let _permit = match permit.acquire().await {
Ok(p) => p,
Err(_) => {
return TaskResult {
task_id: task.id,
status: TaskStatus::Failed,
output: "internal error: concurrency semaphore closed".into(),
metadata: Default::default(),
};
}
};
if cancel.load(Ordering::Acquire) {
return TaskResult {
task_id: task.id,
status: TaskStatus::Failed,
output: "crew cancelled".into(),
metadata: Default::default(),
};
}
execute_task(
&task,
agent.as_ref(),
llm.as_ref(),
&cache,
&cost_tracker,
None,
)
.await
});
}
let mut results = Vec::with_capacity(self.spec.tasks.len());
while let Some(res) = join_set.join_next().await {
match res {
Ok(task_result) => {
self.emit(
"task_completed",
serde_json::json!({
"task_id": task_result.task_id.to_string(),
"status": format!("{:?}", task_result.status),
}),
);
results.push(task_result);
}
Err(e) => {
warn!(error = %e, "task join error (task panicked)");
results.push(TaskResult {
task_id: uuid::Uuid::nil(),
output: format!("task panicked: {e}"),
status: TaskStatus::Failed,
metadata: Default::default(),
});
}
}
}
let status_map: HashMap<TaskId, TaskStatus> =
results.iter().map(|r| (r.task_id, r.status)).collect();
for task in &mut self.spec.tasks {
if let Some(&s) = status_map.get(&task.id) {
task.status = s;
}
}
Ok(results)
}
#[tracing::instrument(skip(self), fields(crew_id = %self.spec.id, task_count = self.spec.tasks.len()))]
async fn run_dag(&mut self) -> crate::core::Result<Vec<TaskResult>> {
let order = topological_sort(&self.spec.tasks)?;
let dep_sets: HashMap<TaskId, HashSet<TaskId>> = self
.spec
.tasks
.iter()
.map(|t| (t.id, t.dependencies.iter().copied().collect()))
.collect();
let task_map: HashMap<TaskId, usize> = self
.spec
.tasks
.iter()
.enumerate()
.map(|(i, t)| (t.id, i))
.collect();
let mut completed: HashSet<TaskId> = HashSet::new();
let mut results: Vec<TaskResult> = Vec::with_capacity(self.spec.tasks.len());
let mut remaining: Vec<TaskId> = order;
while !remaining.is_empty() {
if self.is_cancelled() {
info!(crew_id = %self.spec.id, "crew cancelled — stopping DAG execution");
break;
}
let (ready, not_ready): (Vec<TaskId>, Vec<TaskId>) =
remaining.into_iter().partition(|id| {
dep_sets
.get(id)
.is_none_or(|deps| deps.is_subset(&completed))
});
if ready.is_empty() {
return Err(crate::core::AgnosaiError::Scheduling(
"DAG deadlock: no ready tasks but remaining exist".into(),
));
}
let mut join_set = tokio::task::JoinSet::new();
for id in &ready {
let idx = task_map[id];
self.spec.tasks[idx].status = TaskStatus::Queued;
let agent = pick_best_agent(&self.spec.agents, &self.spec.tasks[idx]);
self.emit(
"task_started",
serde_json::json!({
"task_id": self.spec.tasks[idx].id.to_string(),
"description": self.spec.tasks[idx].description,
"agent": agent.as_ref().map(|a| &a.agent_key),
}),
);
let task_snap = self.spec.tasks[idx].clone();
let llm = self.llm.clone();
let cache = Arc::clone(&self.cache);
let cost_tracker = Arc::clone(&self.cost_tracker);
join_set.spawn(async move {
execute_task(
&task_snap,
agent.as_ref(),
llm.as_ref(),
&cache,
&cost_tracker,
None,
)
.await
});
}
let mut wave_failed = false;
while let Some(res) = join_set.join_next().await {
match res {
Ok(tr) => {
if let Some(&idx) = task_map.get(&tr.task_id) {
self.spec.tasks[idx].status = tr.status;
}
self.emit(
"task_completed",
serde_json::json!({
"task_id": tr.task_id.to_string(),
"status": format!("{:?}", tr.status),
}),
);
if tr.status == TaskStatus::Failed {
wave_failed = true;
warn!(task_id = %tr.task_id, "DAG task failed — downstream tasks will be skipped");
} else {
completed.insert(tr.task_id);
}
results.push(tr);
}
Err(e) => {
warn!(error = %e, "DAG task join error");
wave_failed = true;
}
}
}
remaining = not_ready;
if wave_failed && !remaining.is_empty() {
let any_runnable = remaining.iter().any(|id| {
dep_sets
.get(id)
.is_none_or(|deps| deps.is_subset(&completed))
});
if !any_runnable {
warn!("DAG execution halted: no runnable tasks after failure");
break;
}
}
}
Ok(results)
}
}
fn pick_best_agent(agents: &[AgentDefinition], task: &Task) -> Option<AgentDefinition> {
if agents.is_empty() {
return None;
}
let mut ranked: Vec<(&AgentDefinition, f64)> = agents
.iter()
.map(|a| (a, scoring::score_agent(a, task)))
.collect();
ranked.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
if let Some((agent, score)) = ranked.first() {
debug!(
agent_key = %agent.agent_key,
score,
task_id = %task.id,
"picked best agent for task"
);
}
ranked.first().map(|(a, _)| (*a).clone())
}
fn infer_provider(model: &str) -> ProviderType {
let m = model.to_lowercase();
if m.starts_with("gpt-") || m.starts_with("o1") || m.starts_with("o3") {
ProviderType::OpenAi
} else if m.starts_with("claude") {
ProviderType::Anthropic
} else if m.starts_with("deepseek") {
ProviderType::DeepSeek
} else if m.starts_with("gemini") {
ProviderType::Google
} else if m.starts_with("grok") {
ProviderType::Grok
} else if m.starts_with("mistral") || m.starts_with("mixtral") {
ProviderType::Mistral
} else {
ProviderType::Ollama
}
}
#[tracing::instrument(skip(task, agent, llm, cache, cost_tracker, event_tx), fields(task_id = %task.id, agent = agent.map(|a| a.agent_key.as_str()).unwrap_or("none")))]
async fn execute_task(
task: &Task,
agent: Option<&AgentDefinition>,
llm: Option<&Arc<HooshClient>>,
cache: &Arc<ResponseCache>,
cost_tracker: &Arc<CostTracker>,
event_tx: Option<&broadcast::Sender<CrewEvent>>,
) -> TaskResult {
let task_start = Instant::now();
let agent_label = agent.map(|a| a.agent_key.as_str()).unwrap_or("unassigned");
let Some(client) = llm else {
debug!(task_id = %task.id, agent = agent_label, "executing task (placeholder — no LLM client)");
tokio::task::yield_now().await;
let mut metadata = HashMap::new();
metadata.insert(
"task_duration_ms".into(),
serde_json::json!(task_start.elapsed().as_millis() as u64),
);
return TaskResult {
task_id: task.id,
output: task.description.clone(),
status: TaskStatus::Completed,
metadata,
};
};
debug!(task_id = %task.id, agent = agent_label, "executing task via LLM");
let raw_system = build_system_prompt(agent);
let system_prompt = crate::server::prompt_guard::wrap_system_prompt(&raw_system);
let model = select_model(agent);
let mut messages = Vec::new();
if !task.context.is_empty() {
let ctx_json = serde_json::to_string_pretty(&task.context).unwrap_or_default();
let sanitized_ctx = crate::server::prompt_guard::sanitize(&ctx_json, "context");
messages.push(Message::new(Role::User, sanitized_ctx));
messages.push(Message::new(
Role::Assistant,
"Understood, I have the context.",
));
}
let mut user_msg = crate::server::prompt_guard::sanitize(&task.description, "task_description");
if let Some(ref expected) = task.expected_output {
let sanitized_expected = crate::server::prompt_guard::sanitize(expected, "expected_output");
let _ = write!(user_msg, "\n\n{sanitized_expected}");
}
let mut temperature = 0.7;
if let Some(agent) = agent
&& let Some(ref profile) = agent.personality
{
temperature = mood_adjusted_temperature(profile, temperature);
}
let request = InferenceRequest {
model: model.to_string(),
prompt: user_msg,
system: Some(system_prompt),
messages,
max_tokens: Some(4096),
temperature: Some(temperature),
..Default::default()
};
let ck = cache_key(&request.model, &request.messages);
if let Some(cached) = cache.get(&ck) {
let task_duration_ms = task_start.elapsed().as_millis() as u64;
let mut metadata = HashMap::new();
metadata.insert(
"model".into(),
serde_json::Value::String(request.model.clone()),
);
metadata.insert("cached".into(), serde_json::json!(true));
metadata.insert(
"task_duration_ms".into(),
serde_json::json!(task_duration_ms),
);
debug!(
task_id = %task.id,
agent = agent_label,
model = %request.model,
"task completed from cache"
);
return TaskResult {
task_id: task.id,
output: (*cached).clone(),
status: TaskStatus::Completed,
metadata,
};
}
let mut current_request = request;
let original_prompt = current_request.prompt.clone();
let retry_config = crate::llm::retry::RetryConfig::default();
let task_id_str = task.id.to_string();
let mut final_response = {
let req = ¤t_request;
crate::llm::retry::with_retry(&retry_config, &task_id_str, || client.infer(req)).await
};
if let Some(ref schema) = task.output_schema {
for attempt in 1..=crate::orchestrator::output_validation::MAX_VALIDATION_RETRIES {
let Ok(ref resp) = final_response else { break };
let (extracted, result) =
crate::orchestrator::output_validation::extract_and_validate(&resp.text, schema);
match result {
crate::orchestrator::output_validation::ValidationResult::Valid => break,
crate::orchestrator::output_validation::ValidationResult::Invalid(err) => {
crate::orchestrator::output_validation::log_retry(
&task.id.to_string(),
attempt,
&err,
);
let sanitized_output =
crate::server::prompt_guard::sanitize(&extracted, "failed_output");
current_request.prompt =
crate::orchestrator::output_validation::build_retry_prompt(
&original_prompt,
&sanitized_output,
&err,
schema,
);
current_request.temperature = Some(0.1);
final_response = {
let req = ¤t_request;
crate::llm::retry::with_retry(&retry_config, &task_id_str, || {
client.infer(req)
})
.await
};
}
}
}
}
match final_response {
Ok(response) => {
cache.insert(ck, response.text.clone());
if let Some(tx) = event_tx {
let _ = tx.send(CrewEvent {
crew_id: task
.context
.get("crew_id")
.and_then(|v| v.as_str())
.unwrap_or("unknown")
.to_string(),
event_type: "token".into(),
data: serde_json::json!({
"task_id": task.id.to_string(),
"token": response.text,
"complete": true,
}),
});
}
let task_duration_ms = task_start.elapsed().as_millis() as u64;
let provider = infer_provider(&response.model);
let cost_usd = cost_tracker.record(provider, "hoosh", &response.model, &response.usage);
crate::llm::llm_metrics::record_request(
&provider.to_string(),
&response.model,
"success",
task_duration_ms as f64 / 1000.0,
response.usage.prompt_tokens,
response.usage.completion_tokens,
);
let mut metadata = HashMap::new();
metadata.insert(
"model".into(),
serde_json::Value::String(response.model.clone()),
);
metadata.insert(
"provider".into(),
serde_json::Value::String(response.provider.clone()),
);
metadata.insert("latency_ms".into(), serde_json::json!(response.latency_ms));
metadata.insert(
"tokens".into(),
serde_json::json!({
"prompt": response.usage.prompt_tokens,
"completion": response.usage.completion_tokens,
"total": response.usage.total_tokens,
}),
);
metadata.insert("cost_usd".into(), serde_json::json!(cost_usd));
metadata.insert(
"task_duration_ms".into(),
serde_json::json!(task_duration_ms),
);
info!(
task_id = %task.id,
agent = agent_label,
model = %response.model,
latency_ms = response.latency_ms,
task_duration_ms,
tokens = response.usage.total_tokens,
cost_usd,
"task completed via LLM"
);
TaskResult {
task_id: task.id,
output: response.text,
status: TaskStatus::Completed,
metadata,
}
}
Err(e) => {
let task_duration_ms = task_start.elapsed().as_millis() as u64;
crate::llm::llm_metrics::record_request(
"hoosh",
¤t_request.model,
"error",
task_duration_ms as f64 / 1000.0,
0,
0,
);
warn!(
task_id = %task.id,
agent = agent_label,
task_duration_ms,
error = %e,
"LLM inference failed"
);
let mut metadata = HashMap::new();
metadata.insert("error".into(), serde_json::Value::String(e.to_string()));
metadata.insert(
"task_duration_ms".into(),
serde_json::json!(task_duration_ms),
);
let mut output = String::from("LLM error: ");
let _ = write!(output, "{e}");
TaskResult {
task_id: task.id,
output,
status: TaskStatus::Failed,
metadata,
}
}
}
}
fn build_system_prompt(agent: Option<&AgentDefinition>) -> String {
let Some(agent) = agent else {
return "You are a helpful AI assistant executing tasks within a crew.".into();
};
let mut prompt = format!(
"You are {name}, a {role}.\n\nGoal: {goal}",
name = agent.name,
role = agent.role,
goal = agent.goal,
);
if let Some(ref backstory) = agent.backstory {
let _ = write!(prompt, "\n\nBackstory: {backstory}");
}
if let Some(ref domain) = agent.domain {
let _ = write!(prompt, "\n\nDomain expertise: {domain}");
}
if !agent.tools.is_empty() {
let _ = write!(prompt, "\n\nAvailable tools: {}", agent.tools.join(", "));
}
if let Some(ref profile) = agent.personality {
let disposition = profile.compose_prompt();
if !disposition.is_empty() {
prompt.push('\n');
prompt.push_str(&disposition);
}
}
prompt
}
fn strip_provider_prefix(model: &str) -> &str {
if let Some(idx) = model.find('/') {
let prefix = &model[..idx];
match prefix {
"ollama" | "openai" | "anthropic" | "groq" | "deepseek" | "mistral" | "together"
| "fireworks" | "anyscale" | "perplexity" | "bedrock" | "azure" => &model[idx + 1..],
_ => model,
}
} else {
model
}
}
fn select_model(agent: Option<&AgentDefinition>) -> &str {
if let Some(agent) = agent {
if let Some(ref model) = agent.llm_model {
return strip_provider_prefix(model.as_str());
}
let complexity = crate::llm::parse_complexity(&agent.complexity);
let profile = crate::llm::TaskProfile {
task_type: crate::llm::TaskType::Reason,
complexity,
};
let tier = crate::llm::router::route(&profile);
return crate::llm::default_model(tier);
}
crate::llm::default_model(crate::llm::ModelTier::Capable)
}
fn mood_adjusted_temperature(profile: &bhava::traits::PersonalityProfile, base: f64) -> f64 {
use bhava::traits::TraitKind;
let creativity = profile.get_trait(TraitKind::Creativity).normalized() as f64;
let curiosity = profile.get_trait(TraitKind::Curiosity).normalized() as f64;
let precision = profile.get_trait(TraitKind::Precision).normalized() as f64;
let risk = profile.get_trait(TraitKind::RiskTolerance).normalized() as f64;
let confidence = profile.get_trait(TraitKind::Confidence).normalized() as f64;
let delta = (creativity * 0.15) + (curiosity * 0.1) + (risk * 0.1)
- (precision * 0.15)
- (confidence * 0.05);
(base + delta).clamp(0.1, 1.5)
}
fn topological_sort(tasks: &[Task]) -> crate::core::Result<Vec<TaskId>> {
crate::orchestrator::scheduler::topological_sort_tasks(tasks)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::agent::AgentDefinition;
use crate::core::crew::CrewSpec;
use crate::core::task::{ProcessMode, Task};
use uuid::Uuid;
fn test_agent(key: &str) -> AgentDefinition {
AgentDefinition {
agent_key: key.into(),
name: key.into(),
role: "tester".into(),
goal: "test things".into(),
backstory: None,
domain: None,
tools: vec![],
complexity: "medium".into(),
llm_model: None,
gpu_required: false,
gpu_preferred: false,
gpu_memory_min_mb: None,
hardware: None,
personality: None,
}
}
fn test_task(desc: &str) -> Task {
Task::new(desc)
}
fn test_spec(tasks: Vec<Task>, process: ProcessMode) -> CrewSpec {
CrewSpec {
id: Uuid::new_v4(),
name: "test-crew".into(),
agents: vec![test_agent("agent-a"), test_agent("agent-b")],
tasks,
process,
metadata: Default::default(),
trust_level: "basic".into(),
}
}
#[tokio::test]
async fn test_sequential_execution() {
let tasks = vec![
test_task("step one"),
test_task("step two"),
test_task("step three"),
];
let spec = test_spec(tasks, ProcessMode::Sequential);
let mut runner = CrewRunner::new(spec);
let state = runner.run().await.unwrap();
assert_eq!(state.status, CrewStatus::Completed);
assert_eq!(state.results.len(), 3);
assert_eq!(state.results[0].output, "step one");
assert_eq!(state.results[1].output, "step two");
assert_eq!(state.results[2].output, "step three");
for r in &state.results {
assert_eq!(r.status, TaskStatus::Completed);
}
}
#[tokio::test]
async fn test_sequential_empty() {
let spec = test_spec(vec![], ProcessMode::Sequential);
let mut runner = CrewRunner::new(spec);
let state = runner.run().await.unwrap();
assert_eq!(state.status, CrewStatus::Completed);
assert!(state.results.is_empty());
}
#[tokio::test]
async fn test_parallel_execution() {
let tasks = vec![
test_task("par one"),
test_task("par two"),
test_task("par three"),
test_task("par four"),
];
let spec = test_spec(tasks, ProcessMode::Parallel { max_concurrency: 2 });
let mut runner = CrewRunner::new(spec);
let state = runner.run().await.unwrap();
assert_eq!(state.status, CrewStatus::Completed);
assert_eq!(state.results.len(), 4);
let outputs: HashSet<String> = state.results.iter().map(|r| r.output.clone()).collect();
assert!(outputs.contains("par one"));
assert!(outputs.contains("par two"));
assert!(outputs.contains("par three"));
assert!(outputs.contains("par four"));
}
#[tokio::test]
async fn test_parallel_single_concurrency() {
let tasks = vec![test_task("a"), test_task("b")];
let spec = test_spec(tasks, ProcessMode::Parallel { max_concurrency: 1 });
let mut runner = CrewRunner::new(spec);
let state = runner.run().await.unwrap();
assert_eq!(state.status, CrewStatus::Completed);
assert_eq!(state.results.len(), 2);
}
#[tokio::test]
async fn test_dag_execution_with_dependencies() {
let task_a = test_task("task A");
let mut task_b = test_task("task B");
let mut task_c = test_task("task C");
task_b.dependencies.push(task_a.id);
task_c.dependencies.push(task_b.id);
let spec = test_spec(vec![task_a, task_b, task_c], ProcessMode::Dag);
let mut runner = CrewRunner::new(spec);
let state = runner.run().await.unwrap();
assert_eq!(state.status, CrewStatus::Completed);
assert_eq!(state.results.len(), 3);
let pos = |desc: &str| state.results.iter().position(|r| r.output == desc).unwrap();
assert!(pos("task A") < pos("task B"));
assert!(pos("task B") < pos("task C"));
}
#[tokio::test]
async fn test_dag_diamond() {
let a = test_task("A");
let mut b = test_task("B");
let mut c = test_task("C");
let mut d = test_task("D");
b.dependencies.push(a.id);
c.dependencies.push(a.id);
d.dependencies.push(b.id);
d.dependencies.push(c.id);
let spec = test_spec(vec![a, b, c, d], ProcessMode::Dag);
let mut runner = CrewRunner::new(spec);
let state = runner.run().await.unwrap();
assert_eq!(state.status, CrewStatus::Completed);
assert_eq!(state.results.len(), 4);
let pos = |desc: &str| state.results.iter().position(|r| r.output == desc).unwrap();
assert!(pos("A") < pos("B"));
assert!(pos("A") < pos("C"));
assert!(pos("B") < pos("D"));
assert!(pos("C") < pos("D"));
}
#[tokio::test]
async fn test_dag_no_deps_runs_all() {
let tasks = vec![test_task("x"), test_task("y"), test_task("z")];
let spec = test_spec(tasks, ProcessMode::Dag);
let mut runner = CrewRunner::new(spec);
let state = runner.run().await.unwrap();
assert_eq!(state.status, CrewStatus::Completed);
assert_eq!(state.results.len(), 3);
}
#[test]
fn test_topo_sort_detects_cycle() {
let mut a = test_task("a");
let mut b = test_task("b");
a.dependencies.push(b.id);
b.dependencies.push(a.id);
let err = topological_sort(&[a, b]);
assert!(err.is_err());
}
#[test]
fn test_pick_best_agent_empty_roster() {
let task = test_task("something");
assert!(pick_best_agent(&[], &task).is_none());
}
#[test]
fn test_pick_best_agent_returns_some() {
let task = test_task("something");
let agents = vec![test_agent("a1")];
let picked = pick_best_agent(&agents, &task);
assert!(picked.is_some());
assert_eq!(picked.unwrap().agent_key, "a1");
}
#[tokio::test]
async fn test_hierarchical_falls_back_to_sequential() {
let tasks = vec![test_task("h1"), test_task("h2")];
let spec = test_spec(
tasks,
ProcessMode::Hierarchical {
manager: Uuid::new_v4(),
},
);
let mut runner = CrewRunner::new(spec);
let state = runner.run().await.unwrap();
assert_eq!(state.status, CrewStatus::Completed);
assert_eq!(state.results.len(), 2);
assert_eq!(state.results[0].output, "h1");
assert_eq!(state.results[1].output, "h2");
}
#[test]
fn test_build_system_prompt_no_agent() {
let prompt = build_system_prompt(None);
assert!(prompt.contains("helpful AI assistant"));
}
#[test]
fn test_build_system_prompt_full_agent() {
let mut agent = test_agent("qa");
agent.name = "QA Lead".into();
agent.role = "quality assurance".into();
agent.goal = "ensure zero defects".into();
agent.backstory = Some("10 years in QA".into());
agent.domain = Some("testing".into());
agent.tools = vec!["selenium".into(), "pytest".into()];
let prompt = build_system_prompt(Some(&agent));
assert!(prompt.contains("QA Lead"));
assert!(prompt.contains("quality assurance"));
assert!(prompt.contains("ensure zero defects"));
assert!(prompt.contains("10 years in QA"));
assert!(prompt.contains("testing"));
assert!(prompt.contains("selenium, pytest"));
}
#[test]
fn test_build_system_prompt_minimal_agent() {
let agent = test_agent("min");
let prompt = build_system_prompt(Some(&agent));
assert!(prompt.contains("min")); assert!(prompt.contains("tester")); assert!(prompt.contains("test things")); assert!(!prompt.contains("Backstory"));
assert!(!prompt.contains("Domain"));
assert!(!prompt.contains("Available tools"));
}
#[test]
fn test_select_model_no_agent() {
let model = select_model(None);
assert_eq!(model, "llama3:70b");
}
#[test]
fn test_select_model_agent_override() {
let mut agent = test_agent("a");
agent.llm_model = Some("gpt-4o".into());
assert_eq!(select_model(Some(&agent)), "gpt-4o");
}
#[test]
fn test_select_model_strips_provider_prefix() {
let mut agent = test_agent("a");
agent.llm_model = Some("ollama/llama3.2:1b".into());
assert_eq!(select_model(Some(&agent)), "llama3.2:1b");
agent.llm_model = Some("openai/gpt-4o".into());
assert_eq!(select_model(Some(&agent)), "gpt-4o");
agent.llm_model = Some("anthropic/claude-sonnet-4-20250514".into());
assert_eq!(select_model(Some(&agent)), "claude-sonnet-4-20250514");
}
#[test]
fn test_strip_provider_prefix_preserves_unknown() {
assert_eq!(strip_provider_prefix("custom/model"), "custom/model");
assert_eq!(strip_provider_prefix("llama3:70b"), "llama3:70b");
}
#[test]
fn test_select_model_routes_by_complexity() {
let mut low = test_agent("low");
low.complexity = "low".into();
assert_eq!(select_model(Some(&low)), "llama3:70b");
let mut high = test_agent("high");
high.complexity = "high".into();
assert_eq!(select_model(Some(&high)), "llama3:405b");
}
#[tokio::test]
async fn test_execute_task_placeholder_when_no_llm() {
let task = test_task("do something");
let agent = test_agent("a");
let cache = Arc::new(ResponseCache::new(Default::default()));
let cost_tracker = Arc::new(CostTracker::new());
let result = execute_task(&task, Some(&agent), None, &cache, &cost_tracker, None).await;
assert_eq!(result.status, TaskStatus::Completed);
assert_eq!(result.output, "do something");
}
#[tokio::test]
async fn test_events_emitted_during_sequential_run() {
let tasks = vec![test_task("step one"), test_task("step two")];
let spec = test_spec(tasks, ProcessMode::Sequential);
let (tx, mut rx) = broadcast::channel::<CrewEvent>(64);
let mut runner = CrewRunner::new(spec).with_events(tx);
let state = runner.run().await.unwrap();
assert_eq!(state.status, CrewStatus::Completed);
let mut events = Vec::new();
while let Ok(ev) = rx.try_recv() {
events.push(ev);
}
let types: Vec<&str> = events.iter().map(|e| e.event_type.as_str()).collect();
assert!(types.contains(&"crew_started"));
assert!(types.contains(&"crew_completed"));
assert_eq!(types.iter().filter(|&&t| t == "task_started").count(), 2);
assert_eq!(types.iter().filter(|&&t| t == "task_completed").count(), 2);
}
#[test]
fn test_infer_provider_openai() {
assert_eq!(infer_provider("gpt-4o"), ProviderType::OpenAi);
assert_eq!(infer_provider("gpt-4o-mini"), ProviderType::OpenAi);
assert_eq!(infer_provider("o1"), ProviderType::OpenAi);
assert_eq!(infer_provider("o3-mini"), ProviderType::OpenAi);
}
#[test]
fn test_infer_provider_anthropic() {
assert_eq!(infer_provider("claude-sonnet-4"), ProviderType::Anthropic);
assert_eq!(
infer_provider("claude-sonnet-4-20250514"),
ProviderType::Anthropic
);
assert_eq!(infer_provider("claude-opus-4"), ProviderType::Anthropic);
}
#[test]
fn test_infer_provider_deepseek() {
assert_eq!(infer_provider("deepseek-chat"), ProviderType::DeepSeek);
assert_eq!(infer_provider("deepseek-coder"), ProviderType::DeepSeek);
}
#[test]
fn test_infer_provider_local_models() {
assert_eq!(infer_provider("llama3"), ProviderType::Ollama);
assert_eq!(infer_provider("llama3:70b"), ProviderType::Ollama);
assert_eq!(infer_provider("phi3"), ProviderType::Ollama);
assert_eq!(infer_provider("unknown-model"), ProviderType::Ollama);
}
#[test]
fn test_infer_provider_case_insensitive() {
assert_eq!(infer_provider("GPT-4o"), ProviderType::OpenAi);
assert_eq!(infer_provider("Claude-Opus-4"), ProviderType::Anthropic);
assert_eq!(infer_provider("DEEPSEEK-CHAT"), ProviderType::DeepSeek);
}
#[tokio::test]
async fn test_execute_task_placeholder_has_duration() {
let task = test_task("check duration");
let agent = test_agent("a");
let cache = Arc::new(ResponseCache::new(Default::default()));
let cost_tracker = Arc::new(CostTracker::new());
let result = execute_task(&task, Some(&agent), None, &cache, &cost_tracker, None).await;
assert!(result.metadata.contains_key("task_duration_ms"));
assert!(cost_tracker.total_cost() == 0.0);
}
#[tokio::test]
async fn test_crew_profile_includes_cost() {
let tasks = vec![test_task("task a")];
let spec = test_spec(tasks, ProcessMode::Sequential);
let mut runner = CrewRunner::new(spec);
let state = runner.run().await.unwrap();
let profile = state.profile.unwrap();
assert_eq!(profile.cost_usd, 0.0);
assert_eq!(profile.task_count, 1);
assert!(profile.wall_ms < 1000); }
#[tokio::test]
async fn test_parallel_tasks_have_duration_metadata() {
let tasks = vec![test_task("par a"), test_task("par b"), test_task("par c")];
let spec = test_spec(tasks, ProcessMode::Parallel { max_concurrency: 3 });
let mut runner = CrewRunner::new(spec);
let state = runner.run().await.unwrap();
assert_eq!(state.status, CrewStatus::Completed);
let profile = state.profile.unwrap();
assert_eq!(profile.task_count, 3);
assert_eq!(profile.task_ms.len(), 3);
}
#[tokio::test]
async fn test_dag_tasks_have_duration_metadata() {
let t1 = test_task("dag root");
let mut t2 = test_task("dag child");
t2.dependencies.push(t1.id);
let tasks = vec![t1, t2];
let spec = test_spec(tasks, ProcessMode::Dag);
let mut runner = CrewRunner::new(spec);
let state = runner.run().await.unwrap();
assert_eq!(state.status, CrewStatus::Completed);
let profile = state.profile.unwrap();
assert_eq!(profile.task_count, 2);
assert_eq!(profile.task_ms.len(), 2);
}
}