From 04ece6afb19df9de5884303ec68c478beb710bb7 Mon Sep 17 00:00:00 2001 From: pavel Date: Wed, 11 Feb 2026 01:15:56 +0100 Subject: [PATCH] refactor config --- src/agent.rs | 17 ++++------ src/config.rs | 67 ++++++++++++++++++++++++++++++++++++++ src/main.rs | 5 +-- src/scheduler.rs | 29 ++++++++--------- src/server.rs | 83 +++++++++++++++++++++--------------------------- 5 files changed, 127 insertions(+), 74 deletions(-) create mode 100644 src/config.rs diff --git a/src/agent.rs b/src/agent.rs index 09c93ce..7b433e5 100644 --- a/src/agent.rs +++ b/src/agent.rs @@ -66,19 +66,14 @@ impl Agent { self.logs.push('\n'); } - pub async fn run(&mut self) -> Result<(String, Option), Box> { + pub async fn run( + &mut self, + config: &crate::config::Config, + ) -> Result<(String, Option), Box> { let mut finished = false; let start_time = Instant::now(); - let max_duration_secs = std::env::var("AGENT_MAX_DURATION_SECS") - .ok() - .and_then(|s| s.parse().ok()) - .unwrap_or(120); - let max_duration = Duration::from_secs(max_duration_secs); - - let max_turns = std::env::var("AGENT_MAX_TURNS") - .ok() - .and_then(|s| s.parse().ok()) - .unwrap_or(20); + let max_duration = Duration::from_secs(config.agent_max_duration_secs); + let max_turns = config.agent_max_turns; let mut turns = 0; while !finished { diff --git a/src/config.rs b/src/config.rs new file mode 100644 index 0000000..8da6394 --- /dev/null +++ b/src/config.rs @@ -0,0 +1,67 @@ +use std::env; + +#[derive(Clone, Debug)] +pub struct Config { + pub database_url: String, + pub port: u16, + pub zen_api_key: Option, + pub tavily_api_key: Option, + pub authentik_issuer: String, + pub authentik_client_id: String, + pub authentik_client_secret: String, + pub cors_allowed_origins: Option, + pub cookie_secure: bool, + pub agent_max_turns: u32, + pub agent_max_duration_secs: u64, +} + +impl Config { + pub fn from_env() -> Result> { + let database_url = env::var("DATABASE_URL").map_err(|_| "DATABASE_URL must be set")?; + + let port = env::var("PORT") + .ok() + .and_then(|s| s.parse().ok()) + .unwrap_or(3000); + + let zen_api_key = env::var("ZEN_API_KEY").ok(); + let tavily_api_key = env::var("TAVILY_API_KEY").ok(); + + let authentik_issuer = + env::var("AUTHENTIK_ISSUER").map_err(|_| "AUTHENTIK_ISSUER must be set")?; + let authentik_client_id = + env::var("AUTHENTIK_CLIENT_ID").map_err(|_| "AUTHENTIK_CLIENT_ID must be set")?; + let authentik_client_secret = env::var("AUTHENTIK_CLIENT_SECRET") + .map_err(|_| "AUTHENTIK_CLIENT_SECRET must be set")?; + + let cors_allowed_origins = env::var("CORS_ALLOWED_ORIGINS").ok(); + + let cookie_secure = env::var("COOKIE_SECURE") + .map(|v| v == "true") + .unwrap_or(false); + + let agent_max_turns = env::var("AGENT_MAX_TURNS") + .ok() + .and_then(|s| s.parse().ok()) + .unwrap_or(20); + + let agent_max_duration_secs = env::var("AGENT_MAX_DURATION_SECS") + .ok() + .and_then(|s| s.parse().ok()) + .unwrap_or(120); + + Ok(Config { + database_url, + port, + zen_api_key, + tavily_api_key, + authentik_issuer, + authentik_client_id, + authentik_client_secret, + cors_allowed_origins, + cookie_secure, + agent_max_turns, + agent_max_duration_secs, + }) + } +} diff --git a/src/main.rs b/src/main.rs index 14ad58e..efc614d 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,6 +1,7 @@ mod agent; mod api; mod auth; +mod config; mod entities; mod scheduler; mod server; @@ -8,9 +9,9 @@ mod tools; #[tokio::main] async fn main() -> Result<(), Box> { - let db_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set"); + let config = config::Config::from_env()?; - server::start(&db_url).await?; + server::start(config).await?; Ok(()) } diff --git a/src/scheduler.rs b/src/scheduler.rs index d2c8d47..1c4aeaf 100644 --- a/src/scheduler.rs +++ b/src/scheduler.rs @@ -3,6 +3,7 @@ use crate::entities::task::Entity as Task; use crate::entities::task_run; use dashmap::DashMap; use sea_orm::{DatabaseConnection, EntityTrait, Set}; +use std::sync::Arc; use tokio_cron_scheduler::{Job, JobScheduler}; use uuid::Uuid; @@ -10,15 +11,13 @@ pub struct Scheduler { scheduler: JobScheduler, db: DatabaseConnection, tasks_to_jobs: DashMap, - zen_api_key: Option, - tavily_api_key: Option, + config: Arc, } impl Scheduler { pub async fn new( db: DatabaseConnection, - zen_api_key: Option, - tavily_api_key: Option, + config: Arc, ) -> Result> { let scheduler = JobScheduler::new().await?; scheduler.start().await?; @@ -26,8 +25,7 @@ impl Scheduler { scheduler, db, tasks_to_jobs: DashMap::new(), - zen_api_key, - tavily_api_key, + config, }) } @@ -42,15 +40,13 @@ impl Scheduler { } let db = self.db.clone(); - let zen_key = self.zen_api_key.clone(); - let tavily_key = self.tavily_api_key.clone(); + let config = self.config.clone(); let job = Job::new_async(cron_expr, move |_uuid, _l| { let db = db.clone(); - let zen_key = zen_key.clone(); - let tavily_key = tavily_key.clone(); + let config = config.clone(); Box::pin(async move { - if let Err(e) = Self::run_task(db, zen_key, tavily_key, task_id).await { + if let Err(e) = Self::run_task(db, config, task_id).await { eprintln!("Error in scheduled task {}: {}", task_id, e); } }) @@ -71,8 +67,7 @@ impl Scheduler { async fn run_task( db: DatabaseConnection, - zen_key: Option, - tavily_key: Option, + config: Arc, task_id: Uuid, ) -> Result<(), Box> { let task = Task::find_by_id(task_id) @@ -95,10 +90,14 @@ impl Scheduler { run.insert(&db).await?; // Start agent in background - let mut agent = Agent::new(zen_key, tavily_key, task.goal.clone())?; + let mut agent = Agent::new( + config.zen_api_key.clone(), + config.tavily_api_key.clone(), + task.goal.clone(), + )?; tokio::spawn(async move { - let (logs, answer, status) = match agent.run().await { + let (logs, answer, status) = match agent.run(&config).await { Ok((logs, answer)) => (logs, answer, "completed".to_string()), Err(e) => ( format!("Scheduled run failed: {}", e), diff --git a/src/server.rs b/src/server.rs index a7e97b1..443ee04 100644 --- a/src/server.rs +++ b/src/server.rs @@ -29,8 +29,7 @@ use migration::{Migrator, MigratorTrait}; pub struct AppState { pub db: DatabaseConnection, pub scheduler: Arc, - pub zen_api_key: Option, - pub tavily_api_key: Option, + pub config: Arc, pub verifier: Arc, pub authenticator: Arc, } @@ -74,15 +73,13 @@ pub struct RecentRunResponse { pub created_at: chrono::DateTime, } -pub async fn start(db_url: &str) -> Result<(), Box> { - let db = Database::connect(db_url).await?; +pub async fn start(config: crate::config::Config) -> Result<(), Box> { + let db = Database::connect(&config.database_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 config = Arc::new(config); - let scheduler = - Arc::new(Scheduler::new(db.clone(), zen_api_key.clone(), tavily_api_key.clone()).await?); + let scheduler = Arc::new(Scheduler::new(db.clone(), config.clone()).await?); // Load existing scheduled tasks let existing_tasks = Task::find().all(&db).await?; @@ -92,22 +89,18 @@ pub async fn start(db_url: &str) -> Result<(), Box> { } } - 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?, + crate::auth::JwksVerifier::new( + config.authentik_issuer.clone(), + config.authentik_client_id.clone(), + ) + .await?, ); let authenticator = Arc::new( crate::auth::Authenticator::new( - authentik_issuer, - authentik_client_id, - authentik_client_secret, + config.authentik_issuer.clone(), + config.authentik_client_id.clone(), + config.authentik_client_secret.clone(), ) .await?, ); @@ -115,13 +108,12 @@ pub async fn start(db_url: &str) -> Result<(), Box> { let state = Arc::new(AppState { db, scheduler, - zen_api_key, - tavily_api_key, + config: config.clone(), verifier, authenticator, }); - let cors = build_cors_layer(); + let cors = build_cors_layer(&config); let app = Router::new() .route("/api/tasks", post(create_task).get(list_tasks)) @@ -148,19 +140,16 @@ pub async fn start(db_url: &str) -> Result<(), Box> { .layer(tower_http::limit::RequestBodyLimitLayer::new(1024 * 1024)) // 1MB limit .with_state(state); - let port = std::env::var("PORT").unwrap_or_else(|_| "3000".to_string()); - let addr = format!("0.0.0.0:{}", port); + let addr = format!("0.0.0.0:{}", config.port); let listener = tokio::net::TcpListener::bind(&addr).await?; - println!("Server running on http://localhost:{}", port); + println!("Server running on http://localhost:{}", config.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 { +fn build_cors_layer(config: &crate::config::Config) -> CorsLayer { + let allow_origin = if let Some(origins) = &config.cors_allowed_origins { let values: Vec = origins .split(',') .map(|origin| origin.trim()) @@ -298,13 +287,13 @@ async fn execute_agent_run( .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(), + state.config.zen_api_key.clone(), + state.config.tavily_api_key.clone(), goal.clone(), ) .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?; - let (logs, answer, status) = match agent.run().await { + let (logs, answer, status) = match agent.run(&state.config).await { Ok((logs, answer)) => (logs, answer, "completed".to_string()), Err(e) => ( format!("Execution failed: {}", e), @@ -438,7 +427,7 @@ async fn auth_refresh( .await .map_err(|e| (StatusCode::UNAUTHORIZED, e.to_string()))?; - let jar = update_auth_cookies(jar, &data); + let jar = update_auth_cookies(jar, &data, &state.config); Ok((jar, Json(data))) } @@ -458,12 +447,12 @@ async fn auth_callback( ) })?; - let jar = update_auth_cookies(jar, &data); + let jar = update_auth_cookies(jar, &data, &state.config); Ok((jar, Json(data))) } -async fn auth_logout(jar: CookieJar) -> impl IntoResponse { - let jar = clear_auth_cookies(jar); +async fn auth_logout(State(state): State>, jar: CookieJar) -> impl IntoResponse { + let jar = clear_auth_cookies(jar, &state.config); (jar, StatusCode::NO_CONTENT) } @@ -474,13 +463,15 @@ async fn auth_session(user: AuthenticatedUser) -> Json { })) } -fn secure() -> bool { - std::env::var("COOKIE_SECURE") - .map(|value| value == "true") - .unwrap_or(false) +fn secure(config: &crate::config::Config) -> bool { + config.cookie_secure } -fn update_auth_cookies(jar: CookieJar, data: &serde_json::Value) -> CookieJar { +fn update_auth_cookies( + jar: CookieJar, + data: &serde_json::Value, + config: &crate::config::Config, +) -> CookieJar { let access_token = data.get("access_token"); let refresh_token = data.get("refresh_token"); @@ -491,7 +482,7 @@ fn update_auth_cookies(jar: CookieJar, data: &serde_json::Value) -> CookieJar { .path("/") .http_only(true) .same_site(SameSite::Lax) - .secure(secure()) + .secure(secure(config)) .build(); jar = jar.add(cookie); } @@ -501,7 +492,7 @@ fn update_auth_cookies(jar: CookieJar, data: &serde_json::Value) -> CookieJar { .path("/") .http_only(true) .same_site(SameSite::Lax) - .secure(secure()) + .secure(secure(config)) .build(); jar = jar.add(cookie); } @@ -509,14 +500,14 @@ fn update_auth_cookies(jar: CookieJar, data: &serde_json::Value) -> CookieJar { jar } -fn clear_auth_cookies(jar: CookieJar) -> CookieJar { +fn clear_auth_cookies(jar: CookieJar, config: &crate::config::Config) -> 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()) + .secure(secure(config)) .max_age(cookie::time::Duration::seconds(0)) .build(); jar = jar.add(cookie);