diff --git a/api/src/auth.rs b/api/src/auth.rs index 317c53a..981bbc3 100644 --- a/api/src/auth.rs +++ b/api/src/auth.rs @@ -22,14 +22,18 @@ impl FromRequestParts for User { parts: &mut Parts, state: &AppState, ) -> Result { - if let Some(user) = &state.test_user { - return Ok(user.clone()); - } - let headers = HeaderMap::from_request_parts(parts, state) .await .map_err(|_| (StatusCode::BAD_REQUEST, "Failed to extract headers"))?; + if state.allow_test_user + && let Some(user) = headers.get("x-user-id") + { + return Ok(User { + id: user.to_str().unwrap_or_default().to_string(), + }); + } + if let Some(api_key) = headers.get("x-api-key") { return from_api_key(api_key, state).await; } diff --git a/api/src/lib.rs b/api/src/lib.rs index b0d3f09..d2e8ebf 100644 --- a/api/src/lib.rs +++ b/api/src/lib.rs @@ -11,7 +11,7 @@ use tower_http::{ }; use crate::{auth::JwksCache, jobs::web_push_notifications::run_web_push_job}; -use services::{User, events::EventService}; +use services::events::EventService; pub mod auth; mod jobs; @@ -21,7 +21,7 @@ mod routes; struct AppState { db_connection: OrmConnection, event_service: EventService, - test_user: Option, + allow_test_user: bool, constants: Constants, jwks_cache: JwksCache, } @@ -35,7 +35,7 @@ pub struct Constants { pub async fn app( db_url: &str, - test_user: Option, + allow_test_user: bool, web_push_key: Option, auth_api_base: String, self_api_base: String, @@ -43,7 +43,7 @@ pub async fn app( let state = AppState { db_connection: init_db(db_url).await?, event_service: EventService::default(), - test_user, + allow_test_user, constants: Constants { web_push_key, auth_api_base, diff --git a/api/src/main.rs b/api/src/main.rs index c3ea658..3c80238 100644 --- a/api/src/main.rs +++ b/api/src/main.rs @@ -20,7 +20,7 @@ async fn main() -> eyre::Result<()> { let port = env("PORT", "8080"); - let app = app(&db_url, None, web_push_key, auth_api_base, self_api_base).await?; + let app = app(&db_url, false, web_push_key, auth_api_base, self_api_base).await?; info!("Starting server"); let listener = tokio::net::TcpListener::bind(format!("0.0.0.0:{port}")) diff --git a/api/tests/common/mod.rs b/api/tests/common/mod.rs index 96f86d0..4d5881b 100644 --- a/api/tests/common/mod.rs +++ b/api/tests/common/mod.rs @@ -16,6 +16,7 @@ use types::{ use uuid::Uuid; pub struct Client { + pub uuid: Uuid, pub server: TestServer, } @@ -184,13 +185,9 @@ pub async fn get_client() -> Client { .init(); }); - let user = User { - id: Uuid::new_v4().to_string(), - }; - let app = app( "sqlite::memory:", - Some(user), + true, None, "http://localhost:3000".to_string(), "http://localhost:8080".to_string(), @@ -198,7 +195,13 @@ pub async fn get_client() -> Client { .await .unwrap(); + let uuid = Uuid::new_v4(); + + let mut test_server = TestServer::builder().http_transport().build(app); + test_server.add_header("x-user-id", uuid.to_string()); + Client { - server: TestServer::builder().http_transport().build(app), + uuid, + server: test_server, } } diff --git a/api/tests/sse.rs b/api/tests/sse.rs index 4ad568f..7f56288 100644 --- a/api/tests/sse.rs +++ b/api/tests/sse.rs @@ -4,8 +4,7 @@ use std::time::Duration; -use crate::common::get_client; -use axum_test::TestServer; +use crate::common::{Client, get_client}; use bytes::Bytes; use futures::Stream; use reqwest::Error; @@ -20,8 +19,14 @@ use uuid::Uuid; mod common; -pub async fn connect_sse(server: &TestServer) -> impl Stream> { - let response = server.reqwest_get("/api/updates/sse").send().await.unwrap(); +pub async fn connect_sse(client: &Client) -> impl Stream> { + let response = client + .server + .reqwest_get("/api/updates/sse") + .header("x-user-id", client.uuid.to_string()) + .send() + .await + .unwrap(); response.bytes_stream() } @@ -64,7 +69,7 @@ pub async fn read_event( async fn sends_events_on_add_todo() { let client = get_client().await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; let todo_uuid = Uuid::new_v4(); client @@ -104,7 +109,7 @@ async fn sends_events_on_update_todo() { })) .await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; client .update_todo_json(json!({ @@ -141,7 +146,7 @@ async fn sends_todo_events_on_remove_todo() { })) .await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; client.delete_todo_json(todo_uuid).await; @@ -168,7 +173,7 @@ async fn sends_deleted_events_on_remove_todo() { })) .await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; client.delete_todo_json(todo_uuid).await; @@ -195,7 +200,7 @@ async fn sends_todo_events_on_reactivate_todo() { })) .await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; client.delete_todo_json(todo_uuid).await; let todos: Vec = timeout(Duration::from_secs(2), read_event(stream, "todos")) @@ -203,7 +208,7 @@ async fn sends_todo_events_on_reactivate_todo() { .unwrap(); assert_eq!(todos.len(), 0); - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; client.reactivate_todo(todo_uuid).await; let todos: Vec = timeout(Duration::from_secs(2), read_event(stream, "todos")) @@ -229,7 +234,7 @@ async fn sends_deleted_events_on_reactivate_todo() { })) .await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; client.delete_todo_json(todo_uuid).await; let todos: Vec = timeout(Duration::from_secs(2), read_event(stream, "todos-deleted")) @@ -237,7 +242,7 @@ async fn sends_deleted_events_on_reactivate_todo() { .unwrap(); assert_eq!(todos.len(), 1); - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; client.reactivate_todo(todo_uuid).await; let todos: Vec = timeout(Duration::from_secs(2), read_event(stream, "todos-deleted")) @@ -263,7 +268,7 @@ async fn sends_events_on_check_todo() { })) .await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; client.check_todo(todo_uuid).await; @@ -292,7 +297,7 @@ async fn sends_events_on_check_remove_todo() { client.check_todo(todo_uuid).await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; client.remove_check_todo(todo_uuid).await; @@ -306,7 +311,7 @@ async fn sends_events_on_check_remove_todo() { async fn sends_events_on_add_tag() { let client = get_client().await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; let tag_uuid = Uuid::new_v4(); client @@ -342,7 +347,7 @@ async fn sends_events_on_update_tag() { )) .await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; client .update_tag_json(json!({ @@ -376,7 +381,7 @@ async fn sends_events_on_remove_tag() { )) .await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; client.delete_tag_json(tag_uuid).await; @@ -391,7 +396,7 @@ async fn sends_events_on_remove_tag() { async fn sends_events_on_add_category() { let client = get_client().await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; let category_uuid = Uuid::new_v4(); client @@ -430,7 +435,7 @@ async fn sends_events_on_update_category() { )) .await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; client .update_category_json(json!({ @@ -466,7 +471,7 @@ async fn sends_events_on_remove_category() { )) .await; - let stream = connect_sse(&client.server).await; + let stream = connect_sse(&client).await; client.delete_category_json(category_uuid).await; diff --git a/flake.nix b/flake.nix index 2d9c80f..88df415 100644 --- a/flake.nix +++ b/flake.nix @@ -53,6 +53,7 @@ pkgs.pkg-config pkgs.openssl pkgs.sea-orm-cli + pkgs.cargo-nextest ]; };