use axum::{ Json, Router, extract::{Path, State}, http::{Method, StatusCode}, routing::{get, post}, }; use chrono::Utc; use sea_orm::{ActiveModelTrait, Database, DatabaseConnection, EntityTrait, QueryOrder, Set}; use serde::{Deserialize, Serialize}; use std::sync::Arc; use tower_http::cors::{Any, CorsLayer}; use uuid::Uuid; use crate::agent::Agent; use crate::entities::task::{self, Entity as Task}; use migration::{Migrator, MigratorTrait}; #[derive(Clone)] pub struct AppState { pub db: DatabaseConnection, pub zen_api_key: Option, pub tavily_api_key: Option, } #[derive(Deserialize)] pub struct CreateTaskRequest { pub goal: String, } #[derive(Serialize)] pub struct TaskResponse { pub id: Uuid, pub goal: String, pub status: String, pub logs: String, pub answer: Option, pub created_at: chrono::DateTime, } pub async fn start(db_url: &str) -> Result<(), Box> { let db = Database::connect(db_url).await?; Migrator::up(&db, None).await?; let zen_api_key = std::env::var("ZEN_API_KEY").ok(); let tavily_api_key = std::env::var("TAVILY_API_KEY").ok(); let state = Arc::new(AppState { db, zen_api_key, tavily_api_key, }); let cors = CorsLayer::new() .allow_origin(Any) .allow_methods([Method::GET, Method::POST]) .allow_headers(Any); let app = Router::new() .route("/tasks", post(create_task).get(list_tasks)) .route("/tasks/:id", get(get_task)) .layer(cors) .with_state(state); let listener = tokio::net::TcpListener::bind("0.0.0.0:3000").await?; println!("Server running on http://localhost:3000"); axum::serve(listener, app).await?; Ok(()) } async fn list_tasks( State(state): State>, ) -> Result>, (StatusCode, String)> { let tasks = Task::find() .order_by_desc(task::Column::CreatedAt) .all(&state.db) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; let response = tasks .into_iter() .map(|t| TaskResponse { id: t.id, goal: t.goal, status: t.status, logs: t.logs, answer: t.answer, created_at: t.created_at, }) .collect(); Ok(Json(response)) } async fn create_task( State(state): State>, Json(payload): Json, ) -> Result, (StatusCode, String)> { let task_id = Uuid::new_v4(); // Initial save let new_task = task::ActiveModel { id: Set(task_id), goal: Set(payload.goal.clone()), status: Set("running".to_string()), logs: Set(String::new()), answer: Set(None), created_at: Set(Utc::now().into()), }; new_task .insert(&state.db) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; // Run agent in the background or synchronously for now as requested // "it can be created and executed on the agent with an api. the resulting logs after execution will get returned by the endpoint" // This implies we wait for it to finish. let mut agent = Agent::new( state.zen_api_key.clone(), state.tavily_api_key.clone(), payload.goal.clone(), ); let (logs, answer) = match agent.run().await { Ok((logs, answer)) => (logs, answer), Err(e) => (format!("Execution failed: {}", e), None), }; // Update with final logs and status let mut task: task::ActiveModel = Task::find_by_id(task_id) .one(&state.db) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .ok_or(( StatusCode::NOT_FOUND, "Task not found after insert".to_string(), ))? .into(); task.logs = Set(logs.clone()); task.answer = Set(answer.clone()); task.status = Set("completed".to_string()); let updated_task = task .update(&state.db) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; Ok(Json(TaskResponse { id: updated_task.id, goal: updated_task.goal, status: updated_task.status, logs: updated_task.logs, answer: updated_task.answer, created_at: updated_task.created_at, })) } async fn get_task( Path(id): Path, State(state): State>, ) -> Result, (StatusCode, String)> { let task = Task::find_by_id(id) .one(&state.db) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .ok_or((StatusCode::NOT_FOUND, "Task not found".to_string()))?; Ok(Json(TaskResponse { id: task.id, goal: task.goal, status: task.status, logs: task.logs, answer: task.answer, created_at: task.created_at, })) }