diff --git a/Cargo.lock b/Cargo.lock index 4a1dc3a..5538b67 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -316,6 +316,12 @@ dependencies = [ "zeroize", ] +[[package]] +name = "deranged" +version = "0.5.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7cd812cc2bc1d69d4764bd80df88b4317eaef9e773c75226407d9bc0876b211c" + [[package]] name = "digest" version = "0.10.7" @@ -969,6 +975,23 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "jsonwebtoken" +version = "10.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eba32bfb4ffdeaca3e34431072faf01745c9b26d25504aa7a6cf5684334fc4fc" +dependencies = [ + "base64", + "getrandom 0.2.17", + "js-sys", + "pem", + "serde", + "serde_json", + "signature", + "simple_asn1", + "zeroize", +] + [[package]] name = "lazy_static" version = "1.5.0" @@ -1125,6 +1148,16 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "num-bigint" +version = "0.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a5e44f723f1133c9deac646763579fdb3ac745e418f2a7af9cd0c431da1f20b9" +dependencies = [ + "num-integer", + "num-traits", +] + [[package]] name = "num-bigint-dig" version = "0.8.6" @@ -1141,6 +1174,12 @@ dependencies = [ "zeroize", ] +[[package]] +name = "num-conv" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "521739c6d2bac4aa25192232afe6841231376b2b26d4d9fae5ecf8ca5772e441" + [[package]] name = "num-integer" version = "0.1.46" @@ -1179,6 +1218,7 @@ dependencies = [ "axum", "chrono", "http-body-util", + "jsonwebtoken", "reqwest", "serde", "serde_json", @@ -1272,6 +1312,16 @@ dependencies = [ "windows-link", ] +[[package]] +name = "pem" +version = "3.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d30c53c26bc5b31a98cd02d20f25a7c8567146caf63ed593a9d87b2775291be" +dependencies = [ + "base64", + "serde_core", +] + [[package]] name = "pem-rfc7468" version = "0.7.0" @@ -1335,6 +1385,12 @@ dependencies = [ "zerovec", ] +[[package]] +name = "powerfmt" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "439ee305def115ba05938db6eb1644ff94165c5ab5e9420d1c1bcedbba909391" + [[package]] name = "ppv-lite86" version = "0.2.21" @@ -1834,6 +1890,18 @@ dependencies = [ "rand_core 0.6.4", ] +[[package]] +name = "simple_asn1" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0d585997b0ac10be3c5ee635f1bab02d512760d14b7c468801ac8a01d9ae5f1d" +dependencies = [ + "num-bigint", + "num-traits", + "thiserror 2.0.18", + "time", +] + [[package]] name = "slab" version = "0.4.12" @@ -2213,6 +2281,36 @@ dependencies = [ "cfg-if", ] +[[package]] +name = "time" +version = "0.3.53" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "18dfaaeddcb932337b5e7866ee7d0ce9b76d2fd092997146f187ec09b4558a50" +dependencies = [ + "deranged", + "num-conv", + "powerfmt", + "serde_core", + "time-core", + "time-macros", +] + +[[package]] +name = "time-core" +version = "0.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9e1c906769ad99c88eaa54e728060edef082f8e358ff32030cb7c7d315e81109" + +[[package]] +name = "time-macros" +version = "0.2.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c431b87111666e491a90baa837f914fb45cd5dc3c268591b0220ff5057f2085f" +dependencies = [ + "num-conv", + "time-core", +] + [[package]] name = "tinystr" version = "0.8.3" @@ -3169,6 +3267,20 @@ name = "zeroize" version = "1.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b97154e67e32c85465826e8bcc1c59429aaaf107c1e4a9e53c8d8ccd5eff88d0" +dependencies = [ + "zeroize_derive", +] + +[[package]] +name = "zeroize_derive" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3c50655cbb0fe3fc43170059e702f1ce5e19b84cec58dc87b037a09935c2f328" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] [[package]] name = "zerotrie" diff --git a/Cargo.toml b/Cargo.toml index 8aab50c..e8fa22d 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -19,6 +19,7 @@ chrono = { version = "0.4", features = ["serde"] } sqlx = { version = "0.8", default-features = false, features = ["runtime-tokio-rustls", "postgres", "chrono", "uuid", "migrate", "macros"] } url = "2" urlencoding = "2" +jsonwebtoken = "10" [dev-dependencies] http-body-util = "0.1" diff --git a/src/auth.rs b/src/auth.rs new file mode 100644 index 0000000..fedfdf4 --- /dev/null +++ b/src/auth.rs @@ -0,0 +1,78 @@ +use axum::{ + extract::Request, + http::{header::AUTHORIZATION, StatusCode}, + middleware::Next, + response::{IntoResponse, Response}, + Json, +}; +use jsonwebtoken::{decode, Algorithm, DecodingKey, Validation}; +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Serialize, Deserialize, Clone)] +pub struct Claims { + pub sub: String, + pub email: String, + pub roles: Option>, + pub active_role: Option, + pub exp: usize, + pub iat: usize, +} + +#[derive(Debug, Clone)] +pub struct AuthUser { + pub user_id: String, + pub email: String, + pub claims: Claims, +} + +pub async fn require_auth(mut request: Request, next: Next) -> Response { + let auth_header = request + .headers() + .get(AUTHORIZATION) + .and_then(|v| v.to_str().ok()) + .map(|s| s.to_owned()); + + let token = match auth_header.as_deref().and_then(|h| h.strip_prefix("Bearer ")) { + Some(t) => t.to_owned(), + None => { + return ( + StatusCode::UNAUTHORIZED, + Json(serde_json::json!({ + "error": "Authorization header required", + "code": "MISSING_TOKEN" + })), + ) + .into_response(); + } + }; + + let jwt_secret = std::env::var("JWT_SECRET").expect("JWT_SECRET must be set"); + + let token_data = match decode::( + &token, + &DecodingKey::from_secret(jwt_secret.as_bytes()), + &Validation::new(Algorithm::HS256), + ) { + Ok(data) => data, + Err(e) => { + tracing::debug!("JWT validation failed: {}", e); + return ( + StatusCode::UNAUTHORIZED, + Json(serde_json::json!({ + "error": "Token is invalid or expired", + "code": "INVALID_TOKEN" + })), + ) + .into_response(); + } + }; + + let auth_user = AuthUser { + user_id: token_data.claims.sub.clone(), + email: token_data.claims.email.clone(), + claims: token_data.claims, + }; + + request.extensions_mut().insert(auth_user); + next.run(request).await +} diff --git a/src/handlers/confirm_action.rs b/src/handlers/confirm_action.rs index 74a58b0..6aa0430 100644 --- a/src/handlers/confirm_action.rs +++ b/src/handlers/confirm_action.rs @@ -1,6 +1,10 @@ -use axum::{extract::State, Json}; +use axum::{ + extract::{Extension, State}, + Json, +}; use crate::{ + auth::AuthUser, error::AppError, handlers::actions::{ConfirmActionRequest, ConfirmActionResponse}, state::AppState, @@ -8,6 +12,7 @@ use crate::{ pub async fn confirm_action( State(state): State, + Extension(auth_user): Extension, Json(request): Json, ) -> Result, AppError> { if request.conversation_id.is_empty() { diff --git a/src/main.rs b/src/main.rs index b1a69dc..dc88577 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,4 +1,5 @@ mod actions; +mod auth; mod chat; mod config; mod cover_letter; diff --git a/src/providers/tickets/nxtgauge_ticket_provider.rs b/src/providers/tickets/nxtgauge_ticket_provider.rs index 0b1b94a..9450c3d 100644 --- a/src/providers/tickets/nxtgauge_ticket_provider.rs +++ b/src/providers/tickets/nxtgauge_ticket_provider.rs @@ -57,7 +57,7 @@ impl TicketProvider for NxtgaugeTicketProvider { }; let ai_service_key = std::env::var("AI_SERVICE_KEY") - .unwrap_or_else(|_| "nxtgauge-ai-assistant".to_string()); + .expect("AI_SERVICE_KEY must be set"); let response = self .client diff --git a/src/routes/mod.rs b/src/routes/mod.rs index b3a37f7..6e5be13 100644 --- a/src/routes/mod.rs +++ b/src/routes/mod.rs @@ -1,12 +1,31 @@ use axum::{ + http::HeaderValue, + middleware, routing::{get, post}, Router, }; -use tower_http::{cors::CorsLayer, trace::TraceLayer}; +use tower_http::{ + cors::{Any, CorsLayer}, + trace::TraceLayer, +}; -use crate::{handlers, state::AppState}; +use crate::{auth::require_auth, handlers, state::AppState}; pub fn build_router(state: AppState) -> Router { + let frontend_url: HeaderValue = std::env::var("FRONTEND_URL") + .unwrap_or_else(|_| "http://localhost:3000".to_string()) + .parse() + .expect("FRONTEND_URL is not a valid header value"); + let admin_url: HeaderValue = std::env::var("ADMIN_URL") + .unwrap_or_else(|_| "http://localhost:3001".to_string()) + .parse() + .expect("ADMIN_URL is not a valid header value"); + + let cors = CorsLayer::new() + .allow_origin([frontend_url, admin_url]) + .allow_methods(Any) + .allow_headers(Any); + Router::new() .route("/health", get(handlers::health::health)) .nest( @@ -24,9 +43,10 @@ pub fn build_router(state: AppState) -> Router { .route("/forms/extract", post(handlers::forms::extract)) .route("/tickets/create", post(handlers::tickets::create)) .route("/help/search", post(handlers::help::search)) - .route("/actions/confirm", post(handlers::confirm_action::confirm_action)), + .route("/actions/confirm", post(handlers::confirm_action::confirm_action)) + .layer(middleware::from_fn(require_auth)), ) - .layer(CorsLayer::permissive()) + .layer(cors) .layer(TraceLayer::new_for_http()) .with_state(state) }