use axum::{
extract::{Path, Query, State},
response::Json,
};
use serde::Deserialize;
use crate::ai::{orchestrator::TaskOrchestrator, AITask, AITaskStatus, ListAITasksResponse};
use crate::error::DbError;
use crate::server::handlers::AppState;
#[derive(Debug, Deserialize)]
pub struct ListAITasksQuery {
pub contribution_id: Option<String>,
pub status: Option<String>,
pub limit: Option<usize>,
pub offset: Option<usize>,
}
pub async fn list_ai_tasks_handler(
State(state): State<AppState>,
Path(db_name): Path<String>,
Query(query): Query<ListAITasksQuery>,
) -> Result<Json<ListAITasksResponse>, DbError> {
let db = state.storage.get_database(&db_name)?;
if db.get_collection("_ai_tasks").is_err() {
return Ok(Json(ListAITasksResponse {
tasks: Vec::new(),
total: 0,
}));
}
let coll = db.get_collection("_ai_tasks")?;
let mut tasks = Vec::new();
for doc in coll.scan(None) {
let task: AITask = serde_json::from_value(doc.to_value())
.map_err(|_| DbError::InternalError("Corrupted task data".to_string()))?;
if let Some(ref contribution_id) = query.contribution_id {
if task.contribution_id != *contribution_id {
continue;
}
}
if let Some(ref status_filter) = query.status {
let status_str = task.status.to_string();
if status_str != *status_filter {
continue;
}
}
tasks.push(task);
}
tasks.sort_by(|a, b| {
b.priority
.cmp(&a.priority)
.then_with(|| a.created_at.cmp(&b.created_at))
});
let total = tasks.len();
let offset = query.offset.unwrap_or(0);
let limit = query.limit.unwrap_or(100);
let tasks: Vec<AITask> = tasks.into_iter().skip(offset).take(limit).collect();
Ok(Json(ListAITasksResponse { tasks, total }))
}
pub async fn get_ai_task_handler(
State(state): State<AppState>,
Path((db_name, task_id)): Path<(String, String)>,
) -> Result<Json<AITask>, DbError> {
let db = state.storage.get_database(&db_name)?;
let coll = db.get_collection("_ai_tasks")?;
let doc = coll.get(&task_id)?;
let task: AITask = serde_json::from_value(doc.to_value())
.map_err(|_| DbError::InternalError("Corrupted task data".to_string()))?;
Ok(Json(task))
}
#[derive(Debug, Deserialize)]
pub struct ClaimTaskRequest {
pub agent_id: String,
}
pub async fn claim_task_handler(
State(state): State<AppState>,
Path((db_name, task_id)): Path<(String, String)>,
Json(request): Json<ClaimTaskRequest>,
) -> Result<Json<AITask>, DbError> {
let db = state.storage.get_database(&db_name)?;
let coll = db.get_collection("_ai_tasks")?;
let doc = coll.get(&task_id)?;
let mut task: AITask = serde_json::from_value(doc.to_value())
.map_err(|e| DbError::InternalError(format!("Corrupted task data: {}", e)))?;
if task.status != AITaskStatus::Pending {
return Err(DbError::BadRequest(format!(
"Task {} is not pending (current status: {})",
task_id, task.status
)));
}
task.status = AITaskStatus::Running;
task.agent_id = Some(request.agent_id);
task.started_at = Some(chrono::Utc::now());
let doc_value = serde_json::to_value(&task)
.map_err(|e| DbError::InternalError(format!("Serialization error: {}", e)))?;
coll.update(&task_id, doc_value)?;
Ok(Json(task))
}
#[derive(Debug, Deserialize)]
#[serde(untagged)]
pub enum CompleteTaskRequest {
Nested { output: serde_json::Value },
Flat {
summary: Option<String>,
risk_score: Option<f64>,
requires_review: Option<bool>,
affected_files: Option<Vec<String>>,
passed: Option<bool>,
stages: Option<Vec<String>>,
errors: Option<Vec<serde_json::Value>>,
tests_run: Option<u32>,
tests_passed: Option<u32>,
test_failures: Option<serde_json::Value>,
files: Option<Vec<serde_json::Value>>,
},
}
#[derive(Debug, serde::Serialize)]
pub struct CompleteTaskResponse {
pub task: AITask,
pub next_stage: Option<String>,
pub message: String,
}
pub async fn complete_task_handler(
State(state): State<AppState>,
Path((db_name, task_id)): Path<(String, String)>,
Json(request): Json<CompleteTaskRequest>,
) -> Result<Json<CompleteTaskResponse>, DbError> {
let db = state.storage.get_database(&db_name)?;
let tasks_coll = db.get_collection("_ai_tasks")?;
let contribs_coll = db.get_collection("_ai_contributions")?;
let doc = tasks_coll.get(&task_id)?;
let mut task: AITask = serde_json::from_value(doc.to_value())
.map_err(|e| DbError::InternalError(format!("Corrupted task data: {}", e)))?;
if task.status != AITaskStatus::Running {
return Err(DbError::BadRequest(format!(
"Task {} is not in progress (current status: {})",
task_id, task.status
)));
}
let output = match &request {
CompleteTaskRequest::Nested { output } => output.clone(),
CompleteTaskRequest::Flat {
summary,
risk_score,
requires_review,
affected_files,
passed,
stages,
errors,
tests_run,
tests_passed,
test_failures,
files,
} => match task.task_type {
crate::ai::AITaskType::AnalyzeContribution => {
serde_json::json!({
"risk_score": risk_score.unwrap_or(0.5),
"requires_review": requires_review.unwrap_or(false),
"affected_files": affected_files.clone().unwrap_or_default()
})
}
crate::ai::AITaskType::GenerateCode => {
serde_json::json!({
"summary": summary.clone().unwrap_or_default(),
"files": files.clone().unwrap_or_default()
})
}
crate::ai::AITaskType::ValidateCode => {
serde_json::json!({
"passed": passed.unwrap_or(false),
"stages": stages.clone().unwrap_or_default(),
"errors": errors.clone().unwrap_or_default()
})
}
crate::ai::AITaskType::RunTests => {
serde_json::json!({
"passed": passed.unwrap_or(false),
"tests_run": tests_run.unwrap_or(0),
"tests_passed": tests_passed.unwrap_or(0),
"failures": test_failures.clone()
})
}
crate::ai::AITaskType::PrepareReview | crate::ai::AITaskType::MergeChanges => {
serde_json::json!({})
}
},
};
task.status = AITaskStatus::Completed;
task.completed_at = Some(chrono::Utc::now());
task.output = Some(output.clone());
let doc_value = serde_json::to_value(&task)
.map_err(|e| DbError::InternalError(format!("Serialization error: {}", e)))?;
tasks_coll.update(&task_id, doc_value)?;
let contrib_doc = contribs_coll.get(&task.contribution_id)?;
let mut contribution: crate::ai::Contribution = serde_json::from_value(contrib_doc.to_value())
.map_err(|e| DbError::InternalError(format!("Corrupted contribution data: {}", e)))?;
if matches!(task.task_type, crate::ai::AITaskType::AnalyzeContribution) {
if let Some(risk_score) = output.get("risk_score").and_then(|v| v.as_f64()) {
contribution.risk_score = Some(risk_score);
}
if let Some(affected) = output.get("affected_files").and_then(|v| v.as_array()) {
contribution.affected_files = affected
.iter()
.filter_map(|v| v.as_str().map(|s| s.to_string()))
.collect();
}
}
let orchestration_result =
TaskOrchestrator::on_task_complete(&task, &contribution, Some(&output));
let mut next_stage = None;
for next_task in &orchestration_result.next_tasks {
let task_value = serde_json::to_value(next_task)
.map_err(|e| DbError::InternalError(format!("Serialization error: {}", e)))?;
tasks_coll.insert(task_value)?;
next_stage = Some(next_task.task_type.to_string());
}
if let Some(new_status) = orchestration_result.contribution_status {
contribution.status = new_status;
contribution.updated_at = chrono::Utc::now();
let contrib_value = serde_json::to_value(&contribution)
.map_err(|e| DbError::InternalError(format!("Serialization error: {}", e)))?;
contribs_coll.update(&contribution.id, contrib_value)?;
}
Ok(Json(CompleteTaskResponse {
task,
next_stage,
message: orchestration_result.message,
}))
}
#[derive(Debug, Deserialize)]
pub struct FailTaskRequest {
pub error: String,
}
pub async fn fail_task_handler(
State(state): State<AppState>,
Path((db_name, task_id)): Path<(String, String)>,
Json(request): Json<FailTaskRequest>,
) -> Result<Json<AITask>, DbError> {
let db = state.storage.get_database(&db_name)?;
let coll = db.get_collection("_ai_tasks")?;
let doc = coll.get(&task_id)?;
let mut task: AITask = serde_json::from_value(doc.to_value())
.map_err(|e| DbError::InternalError(format!("Corrupted task data: {}", e)))?;
if task.status != AITaskStatus::Running {
return Err(DbError::BadRequest(format!(
"Task {} is not in progress (current status: {})",
task_id, task.status
)));
}
task.status = AITaskStatus::Failed;
task.completed_at = Some(chrono::Utc::now());
task.fail(request.error);
let doc_value = serde_json::to_value(&task)
.map_err(|e| DbError::InternalError(format!("Serialization error: {}", e)))?;
coll.update(&task_id, doc_value)?;
Ok(Json(task))
}