use axum::{ Json, RequestPartsExt, Router, extract::{FromRef, FromRequestParts, Path, Query, State}, http::{HeaderValue, StatusCode, request::Parts}, response::IntoResponse, routing::{get, post}, }; use axum_extra::{ TypedHeader, extract::cookie::{Cookie, CookieJar, SameSite}, headers::{Authorization, authorization::Bearer}, }; use chrono::Utc; use sea_orm::{ ActiveModelTrait, Database, DatabaseConnection, EntityTrait, QueryOrder, QuerySelect, Set, }; use serde::{Deserialize, Serialize}; use std::sync::Arc; use tower_http::cors::{AllowOrigin, CorsLayer}; use uuid::Uuid; use crate::agent::Agent; use crate::entities::task::{self, Entity as Task}; use crate::entities::task_run::{self, Entity as TaskRun}; use crate::scheduler::Scheduler; use migration::{Migrator, MigratorTrait}; #[derive(Clone)] pub struct AppState { pub db: DatabaseConnection, pub scheduler: Arc, pub zen_api_key: Option, pub tavily_api_key: Option, pub verifier: Arc, pub authenticator: Arc, } #[derive(Deserialize)] pub struct CreateTaskRequest { pub goal: String, pub cron: Option, } #[derive(Deserialize)] pub struct UpdateTaskRequest { pub goal: String, pub cron: Option, } #[derive(Serialize)] pub struct TaskResponse { pub id: Uuid, pub goal: String, pub cron: Option, pub created_at: chrono::DateTime, pub runs: Vec, } #[derive(Serialize)] pub struct TaskRunResponse { pub id: Uuid, pub status: String, pub logs: String, pub answer: Option, pub created_at: chrono::DateTime, } #[derive(Serialize)] pub struct RecentRunResponse { pub id: Uuid, pub task_id: Uuid, pub goal: String, pub status: String, 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 scheduler = Arc::new(Scheduler::new(db.clone(), zen_api_key.clone(), tavily_api_key.clone()).await?); // Load existing scheduled tasks let existing_tasks = Task::find().all(&db).await?; for task in existing_tasks { if let Some(cron) = task.cron { let _ = scheduler.add_task_job(task.id, &cron).await; } } let authentik_issuer = std::env::var("AUTHENTIK_ISSUER").map_err(|_| "AUTHENTIK_ISSUER not set")?; let authentik_client_id = std::env::var("AUTHENTIK_CLIENT_ID").map_err(|_| "AUTHENTIK_CLIENT_ID not set")?; let authentik_client_secret = std::env::var("AUTHENTIK_CLIENT_SECRET").map_err(|_| "AUTHENTIK_CLIENT_SECRET not set")?; let verifier = Arc::new( crate::auth::JwksVerifier::new(authentik_issuer.clone(), authentik_client_id.clone()) .await?, ); let authenticator = Arc::new( crate::auth::Authenticator::new( authentik_issuer, authentik_client_id, authentik_client_secret, ) .await?, ); let state = Arc::new(AppState { db, scheduler, zen_api_key, tavily_api_key, verifier, authenticator, }); let cors = build_cors_layer(); let app = Router::new() .route("/api/tasks", post(create_task).get(list_tasks)) .route("/api/tasks/:id", get(get_task).put(update_task)) .route("/api/tasks/:id/runs", post(rerun_task)) .route("/api/runs/recent", get(get_recent_runs)) .route("/api/auth/callback", get(auth_callback)) .route("/api/auth/refresh", post(auth_refresh)) .route("/api/auth/logout", post(auth_logout)) .layer(cors) .with_state(state); let port = std::env::var("PORT").unwrap_or_else(|_| "3000".to_string()); let addr = format!("0.0.0.0:{}", port); let listener = tokio::net::TcpListener::bind(&addr).await?; println!("Server running on http://localhost:{}", port); axum::serve(listener, app).await?; Ok(()) } fn build_cors_layer() -> CorsLayer { let origins = std::env::var("CORS_ALLOWED_ORIGINS").ok(); let allow_origin = if let Some(origins) = origins { let values: Vec = origins .split(',') .map(|origin| origin.trim()) .filter(|origin| !origin.is_empty()) .filter_map(|origin| HeaderValue::from_str(origin).ok()) .collect(); if values.is_empty() { AllowOrigin::mirror_request() } else { AllowOrigin::list(values) } } else { AllowOrigin::mirror_request() }; CorsLayer::new() .allow_origin(allow_origin) .allow_methods([ axum::http::Method::GET, axum::http::Method::POST, axum::http::Method::PUT, axum::http::Method::PATCH, axum::http::Method::DELETE, axum::http::Method::OPTIONS, ]) .allow_headers([ axum::http::header::CONTENT_TYPE, axum::http::header::AUTHORIZATION, axum::http::header::ACCEPT, ]) .allow_credentials(true) } async fn list_tasks( _user: AuthenticatedUser, State(state): State>, ) -> Result>, (StatusCode, String)> { let tasks = Task::find() .find_with_related(TaskRun) .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, mut runs)| { runs.sort_by(|a, b| b.created_at.cmp(&a.created_at)); TaskResponse { id: t.id, goal: t.goal, cron: t.cron, created_at: t.created_at, runs: runs .into_iter() .map(|r| TaskRunResponse { id: r.id, status: r.status, logs: r.logs, answer: r.answer, created_at: r.created_at, }) .collect(), } }) .collect(); Ok(Json(response)) } async fn create_task( _user: AuthenticatedUser, State(state): State>, Json(payload): Json, ) -> Result, (StatusCode, String)> { let task_id = Uuid::new_v4(); // Initial task save let new_task = task::ActiveModel { id: Set(task_id), goal: Set(payload.goal.clone()), cron: Set(payload.cron.clone()), created_at: Set(Utc::now().into()), }; new_task .insert(&state.db) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; if let Some(cron) = &payload.cron { let _ = state.scheduler.add_task_job(task_id, cron).await; } else { let _ = state.scheduler.remove_task_job(task_id).await; } get_task_inner(task_id, &state).await.map(Json) } async fn rerun_task( _user: AuthenticatedUser, State(state): State>, Path(id): Path, ) -> 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()))?; execute_agent_run(state, task.id, task.goal).await } async fn execute_agent_run( state: Arc, task_id: Uuid, goal: String, ) -> Result, (StatusCode, String)> { let run_id = Uuid::new_v4(); // Initial run save let new_run = task_run::ActiveModel { id: Set(run_id), task_id: Set(task_id), status: Set("running".to_string()), logs: Set(String::new()), answer: Set(None), created_at: Set(Utc::now().into()), }; new_run .insert(&state.db) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; let mut agent = Agent::new( state.zen_api_key.clone(), state.tavily_api_key.clone(), goal.clone(), ) .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; let (logs, answer, status) = match agent.run().await { Ok((logs, answer)) => (logs, answer, "completed".to_string()), Err(e) => ( format!("Execution failed: {}", e), None, "failed".to_string(), ), }; // Update with final logs and status let mut run: task_run::ActiveModel = TaskRun::find_by_id(run_id) .one(&state.db) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))? .ok_or(( StatusCode::NOT_FOUND, "Run not found after insert".to_string(), ))? .into(); run.logs = Set(logs.clone()); run.answer = Set(answer.clone()); run.status = Set(status); run.update(&state.db) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; get_task_inner(task_id, &state).await.map(Json) } async fn update_task( _user: AuthenticatedUser, State(state): State>, Path(id): Path, Json(payload): Json, ) -> Result, (StatusCode, String)> { let mut task: task::ActiveModel = 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()))? .into(); task.goal = Set(payload.goal.clone()); task.cron = Set(payload.cron.clone()); task.update(&state.db) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; if let Some(cron) = &payload.cron { let _ = state.scheduler.add_task_job(id, cron).await; } else { let _ = state.scheduler.remove_task_job(id).await; } get_task_inner(id, &state).await.map(Json) } #[allow(dead_code)] pub struct AuthenticatedUser(pub crate::auth::Claims); #[axum::async_trait] impl FromRequestParts for AuthenticatedUser where Arc: axum::extract::FromRef, S: Send + Sync, { type Rejection = (StatusCode, String); async fn from_request_parts(parts: &mut Parts, state: &S) -> Result { let app_state = Arc::::from_ref(state); let token = if let Ok(TypedHeader(Authorization(bearer))) = parts.extract::>>().await { Some(bearer.token().to_string()) } else { let jar = parts.extract::().await.unwrap(); jar.get("access_token") .map(|cookie| cookie.value().to_string()) }; let token = token.ok_or(( StatusCode::UNAUTHORIZED, "Missing or invalid access token".to_string(), ))?; let claims = app_state.verifier.verify(&token).await.map_err(|e| { ( StatusCode::UNAUTHORIZED, format!("Token verification failed: {}", e), ) })?; Ok(AuthenticatedUser(claims)) } } #[derive(Deserialize)] pub struct AuthCallbackQuery { pub code: String, pub redirect_uri: String, } #[derive(Deserialize)] struct RefreshRequest { refresh_token: Option, } async fn auth_refresh( State(state): State>, jar: CookieJar, Json(payload): Json, ) -> Result { let refresh_token = payload .refresh_token .filter(|token| !token.is_empty()) .or_else(|| { jar.get("refresh_token") .map(|cookie| cookie.value().to_string()) }) .ok_or(( StatusCode::UNAUTHORIZED, "Missing refresh token".to_string(), ))?; let data = state .authenticator .refresh_token(refresh_token) .await .map_err(|e| (StatusCode::UNAUTHORIZED, e.to_string()))?; let jar = update_auth_cookies(jar, &data); Ok((jar, Json(data))) } async fn auth_callback( State(state): State>, jar: CookieJar, Query(query): Query, ) -> Result { let data = state .authenticator .exchange_code(query.code, query.redirect_uri) .await .map_err(|e| { ( StatusCode::INTERNAL_SERVER_ERROR, format!("Token exchange failed: {}", e), ) })?; let jar = update_auth_cookies(jar, &data); Ok((jar, Json(data))) } async fn auth_logout(jar: CookieJar) -> impl IntoResponse { let jar = clear_auth_cookies(jar); (jar, StatusCode::NO_CONTENT) } fn secure() -> bool { std::env::var("COOKIE_SECURE") .map(|value| value == "true") .unwrap_or(false) } fn update_auth_cookies(jar: CookieJar, data: &serde_json::Value) -> CookieJar { let access_token = data.get("access_token"); let refresh_token = data.get("refresh_token"); let mut jar = jar; if let Some(token) = access_token.and_then(|t| t.as_str()) { let cookie = Cookie::build(("access_token", token.to_owned())) .path("/") .http_only(true) .same_site(SameSite::Lax) .secure(secure()) .build(); jar = jar.add(cookie); } if let Some(token) = refresh_token.and_then(|t| t.as_str()) { let cookie = Cookie::build(("refresh_token", token.to_owned())) .path("/") .http_only(true) .same_site(SameSite::Lax) .secure(secure()) .build(); jar = jar.add(cookie); } jar } fn clear_auth_cookies(jar: CookieJar) -> CookieJar { let mut jar = jar; for name in ["access_token", "refresh_token"] { let cookie = Cookie::build((name, "")) .path("/") .http_only(true) .same_site(SameSite::Lax) .secure(secure()) .max_age(cookie::time::Duration::seconds(0)) .build(); jar = jar.add(cookie); } jar } async fn get_task( _user: AuthenticatedUser, Path(id): Path, State(state): State>, ) -> Result, (StatusCode, String)> { get_task_inner(id, &state).await.map(Json) } async fn get_task_inner(id: Uuid, state: &AppState) -> Result { let results = Task::find_by_id(id) .find_with_related(TaskRun) .all(&state.db) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; let (t, mut runs) = results .into_iter() .next() .ok_or((StatusCode::NOT_FOUND, "Task not found".to_string()))?; runs.sort_by(|a, b| b.created_at.cmp(&a.created_at)); Ok(TaskResponse { id: t.id, goal: t.goal, cron: t.cron, created_at: t.created_at, runs: runs .into_iter() .map(|r| TaskRunResponse { id: r.id, status: r.status, logs: r.logs, answer: r.answer, created_at: r.created_at, }) .collect(), }) } async fn get_recent_runs( _user: AuthenticatedUser, State(state): State>, ) -> Result>, (StatusCode, String)> { let results = TaskRun::find() .find_also_related(Task) .order_by_desc(task_run::Column::CreatedAt) .limit(50) .all(&state.db) .await .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; let response = results .into_iter() .filter_map(|(run, task_opt)| { task_opt.map(|task| RecentRunResponse { id: run.id, task_id: run.task_id, goal: task.goal, status: run.status, created_at: run.created_at, }) }) .collect(); Ok(Json(response)) }