feat(backend): replace auth stub with real session, create users and sessions tables

This commit is contained in:
2026-05-26 02:34:23 +03:00
parent a54ec4ed34
commit e7eaa42064
11 changed files with 636 additions and 135 deletions

40
backend/Cargo.lock generated
View File

@@ -32,6 +32,18 @@ version = "1.0.102"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c"
[[package]]
name = "argon2"
version = "0.5.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "3c3610892ee6e0cbce8ae2700349fcf8f98adb0dbfbee85aec3c9179d29cc072"
dependencies = [
"base64ct",
"blake2",
"cpufeatures 0.2.17",
"password-hash",
]
[[package]]
name = "atoi"
version = "2.0.0"
@@ -173,12 +185,14 @@ name = "backend"
version = "0.1.0"
dependencies = [
"anyhow",
"argon2",
"axum",
"axum-test",
"chrono",
"dotenvy",
"lazy-regex",
"opendal",
"rand 0.10.1",
"serde",
"serde_json",
"serial_test",
@@ -222,6 +236,15 @@ dependencies = [
"serde_core",
]
[[package]]
name = "blake2"
version = "0.10.6"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "46502ad458c9a52b69d4d4d32775c788b7a1b85e8bc9d482d92250fc0e3f8efe"
dependencies = [
"digest",
]
[[package]]
name = "block-buffer"
version = "0.10.4"
@@ -1735,6 +1758,17 @@ dependencies = [
"windows-link",
]
[[package]]
name = "password-hash"
version = "0.5.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "346f04948ba92c43e8469c1ee6736c7563d71012b17d40745260fe106aac2166"
dependencies = [
"base64ct",
"rand_core 0.6.4",
"subtle",
]
[[package]]
name = "pem-rfc7468"
version = "0.7.0"
@@ -1967,9 +2001,9 @@ dependencies = [
[[package]]
name = "rand"
version = "0.10.0"
version = "0.10.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "bc266eb313df6c5c09c1c7b1fbe2510961e5bcd3add930c1e31f7ed9da0feff8"
checksum = "d2e8e8bcc7961af1fdac401278c6a831614941f6164ee3bf4ce61b7edb162207"
dependencies = [
"chacha20",
"getrandom 0.4.2",
@@ -2181,7 +2215,7 @@ dependencies = [
"futures-util",
"http",
"mime",
"rand 0.10.0",
"rand 0.10.1",
"thiserror",
]

View File

@@ -18,6 +18,8 @@ sqlx = { version = "0.8", features = ["runtime-tokio-rustls", "postgres", "uuid"
chrono = { version = "0.4.44", features = ["serde"] }
dotenvy = "0.15.7"
opendal = { version = "0.56.0", features = ["tests", "services-fs"] }
argon2 = "0.5.3"
rand = "0.10.1"
[dev-dependencies]
axum-test = "20.0.0"

View File

@@ -0,0 +1,6 @@
CREATE TABLE users (
id BIGSERIAL PRIMARY KEY,
username TEXT NOT NULL,
password_hash TEXT NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
);

View File

@@ -0,0 +1,6 @@
CREATE TABLE sessions (
id TEXT PRIMARY KEY,
user_id BIGINT NOT NULL REFERENCES users(id),
expires_at TIMESTAMPTZ NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
);

View File

@@ -1,14 +1,14 @@
#[derive(Clone)]
pub struct Ctx {
user_id: u64,
user_id: i64,
}
impl Ctx {
pub fn new(user_id: u64) -> Self {
pub fn new(user_id: i64) -> Self {
Self { user_id }
}
pub fn user_id(&self) -> u64 {
pub fn user_id(&self) -> i64 {
self.user_id
}
}

View File

@@ -6,8 +6,8 @@ use tracing::info;
#[derive(Debug, Clone)]
pub enum LoftError {
LoginFail,
RegisterFail,
AuthFailNoAuthTokenCookie,
AuthFailTokenWrongFormat,
AuthFailCtxNotInRequestExt,
FileIdNotFound,
UndefinedErrorType,
@@ -25,8 +25,8 @@ impl IntoResponse for LoftError {
fn into_response(self) -> axum::response::Response {
match self {
Self::LoginFail
| Self::RegisterFail
| Self::AuthFailNoAuthTokenCookie
| Self::AuthFailTokenWrongFormat
| Self::AuthFailCtxNotInRequestExt => {
info!("UNAUTHORIZED");
StatusCode::UNAUTHORIZED.into_response()

View File

@@ -11,18 +11,19 @@ use axum::{
middleware,
response::Response,
};
use sqlx::PgPool;
use tower_cookies::CookieManagerLayer;
use tower_http::{cors::CorsLayer, services::ServeDir, trace::TraceLayer};
use tracing::info;
use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt};
use crate::{
model::FileRepository,
model::{FileRepository, UserRepository},
web::{
mw_auth::{mw_ctx_resolver, mw_require_auth},
routes_file::routes_file,
routes_health::routes_health,
routes_login::routes_login,
routes_login::routes_auth,
},
};
@@ -37,20 +38,25 @@ async fn main() -> Result<()> {
.with(tracing_subscriber::fmt::layer())
.init();
let file_repository = FileRepository::new().await?;
dotenvy::dotenv().ok();
let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set");
let pool = PgPool::connect(&database_url).await.unwrap();
let file_repository = FileRepository::new(pool.clone())?;
let routes_file = routes_file(file_repository.clone())
.route_layer(middleware::from_fn(mw_require_auth))
.layer(DefaultBodyLimit::disable());
let user_repository = UserRepository::new(pool)?;
let routes_auth = routes_auth(user_repository.clone());
let app = Router::new()
.nest("/api", routes_file)
.nest("/api/auth", routes_auth)
.merge(routes_health())
.merge(routes_login())
.layer(TraceLayer::new_for_http())
.layer(middleware::map_response(main_response_mapper))
.layer(middleware::from_fn_with_state(
file_repository,
user_repository,
mw_ctx_resolver,
))
.layer(CookieManagerLayer::new())

View File

@@ -43,11 +43,7 @@ pub struct FileRepository {
}
impl FileRepository {
pub async fn new() -> Result<Self, LoftError> {
dotenvy::dotenv().ok();
let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set");
let pool = PgPool::connect(&database_url).await.unwrap();
pub fn new(pool: PgPool) -> Result<Self, LoftError> {
let storage_path = std::env::var("STORAGE_PATH").expect("STORAGE_PATH must be set");
let op = Operator::new(services::Fs::default().root(&storage_path))
.unwrap()
@@ -129,10 +125,10 @@ impl FileRepository {
file_id
);
let record = self.get_file(file_id).await?;
info!("Downloading file bytes \"{}\"", file_id);
info!("Deleting file bytes \"{}\"", file_id);
self.op.delete(&record.storage_path).await.unwrap();
info!("Downloading file record \"{}\"", file_id);
info!("Deleting file record \"{}\"", file_id);
sqlx::query_as!(
FileRecord,
r#"
@@ -164,11 +160,143 @@ impl FileRepository {
}
}
#[derive(Clone, Debug, FromRow)]
pub struct User {
pub id: i64,
pub username: String,
pub password_hash: String,
pub created_at: chrono::DateTime<chrono::Utc>,
}
#[derive(Clone, Debug, FromRow)]
pub struct Session {
pub id: String,
pub user_id: i64,
pub expires_at: chrono::DateTime<chrono::Utc>,
pub created_at: chrono::DateTime<chrono::Utc>,
}
#[derive(Clone)]
pub struct UserRepository {
pub pool: PgPool,
}
impl UserRepository {
pub fn new(pool: PgPool) -> Result<Self, LoftError> {
Ok(Self { pool })
}
pub async fn create_user(
&self,
username: String,
password_hash: String,
) -> Result<User, LoftError> {
let user = sqlx::query_as!(
User,
r#"
INSERT INTO users (username, password_hash)
VALUES ($1, $2)
RETURNING *
"#,
username,
password_hash
)
.fetch_one(&self.pool)
.await
.unwrap();
info!("Persisted user: {}", username);
Ok(user)
}
pub async fn find_by_username(&self, username: String) -> Result<User, LoftError> {
info!("Fetching username \"{}\" from users", username);
let user = sqlx::query_as!(
User,
r#"
SELECT *
FROM users u
WHERE u.username = $1
"#,
username
)
.fetch_optional(&self.pool)
.await
.unwrap()
.ok_or(LoftError::LoginFail)?;
Ok(user)
}
pub async fn create_session(
&self,
user_id: i64,
token: String,
expires_at: chrono::DateTime<chrono::Utc>,
) -> Result<Session, LoftError> {
let session = sqlx::query_as!(
Session,
r#"
INSERT INTO sessions (id, user_id, expires_at)
VALUES ($1, $2, $3)
RETURNING *
"#,
token,
user_id,
expires_at
)
.fetch_one(&self.pool)
.await
.unwrap();
info!(
"Persisted session: {} with expiration date: {}",
token, expires_at
);
Ok(session)
}
pub async fn get_session(&self, token: String) -> Result<Session, LoftError> {
info!("Fetching session \"{}\" from sessions", token);
let session = sqlx::query_as!(
Session,
r#"
SELECT *
FROM sessions s
WHERE s.id = $1
AND s.expires_at > NOW()
"#,
token
)
.fetch_optional(&self.pool)
.await
.unwrap()
.ok_or(LoftError::LoginFail)?;
Ok(session)
}
pub async fn delete_session(&self, token: String) -> Result<(), LoftError> {
sqlx::query!(r#"DELETE FROM sessions WHERE id = $1"#, token)
.execute(&self.pool)
.await
.unwrap();
info!("Deleted session \"{}\"", token);
Ok(())
}
}
#[cfg(test)]
mod tests {
use rand::RngExt;
use super::*;
async fn truncate(pool: &PgPool) {
async fn truncate_file_records(pool: &PgPool) {
sqlx::query!("TRUNCATE TABLE file_records")
.execute(pool)
.await
@@ -176,14 +304,17 @@ mod tests {
}
async fn file_repository() -> Result<FileRepository, LoftError> {
Ok(FileRepository::new().await.unwrap())
dotenvy::dotenv().ok();
let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set");
let pool = PgPool::connect(&database_url).await.unwrap();
Ok(FileRepository::new(pool)?)
}
#[tokio::test]
#[serial_test::serial]
async fn test_upload_and_list() {
let file_repository = file_repository().await.unwrap();
truncate(&file_repository.pool).await;
truncate_file_records(&file_repository.pool).await;
file_repository
.upload_file(vec![0u8; 10], "a.jpg".to_string())
@@ -196,14 +327,14 @@ mod tests {
let files = file_repository.list_files().await.unwrap();
assert_eq!(files.len(), 2);
truncate(&file_repository.pool).await;
truncate_file_records(&file_repository.pool).await;
}
#[tokio::test]
#[serial_test::serial]
async fn test_download() {
let file_repository = file_repository().await.unwrap();
truncate(&file_repository.pool).await;
truncate_file_records(&file_repository.pool).await;
let uploaded = file_repository
.upload_file(vec![0u8; 10], "a.jpg".to_string())
@@ -212,14 +343,14 @@ mod tests {
let downloaded = file_repository.download_file(uploaded.id).await.unwrap();
assert!(!downloaded.is_empty());
truncate(&file_repository.pool).await;
truncate_file_records(&file_repository.pool).await;
}
#[tokio::test]
#[serial_test::serial]
async fn test_download_not_found() {
let file_repository = file_repository().await.unwrap();
truncate(&file_repository.pool).await;
truncate_file_records(&file_repository.pool).await;
assert!(matches!(
file_repository.download_file(i64::MAX).await,
@@ -231,7 +362,7 @@ mod tests {
#[serial_test::serial]
async fn test_delete() {
let file_repository = file_repository().await.unwrap();
truncate(&file_repository.pool).await;
truncate_file_records(&file_repository.pool).await;
let uploaded = file_repository
.upload_file(vec![0u8; 10], "a.jpg".to_string())
@@ -243,7 +374,7 @@ mod tests {
file_repository.download_file(uploaded.id).await,
Err(LoftError::FileIdNotFound)
));
truncate(&file_repository.pool).await;
truncate_file_records(&file_repository.pool).await;
}
#[tokio::test]
@@ -256,4 +387,161 @@ mod tests {
Err(LoftError::FileIdNotFound)
));
}
async fn truncate_users(pool: &PgPool) {
sqlx::query!("TRUNCATE TABLE users CASCADE")
.execute(pool)
.await
.unwrap();
}
async fn truncate_sessions(pool: &PgPool) {
sqlx::query!("TRUNCATE TABLE sessions")
.execute(pool)
.await
.unwrap();
}
async fn user_repository() -> Result<UserRepository, LoftError> {
dotenvy::dotenv().ok();
let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set");
let pool = PgPool::connect(&database_url).await.unwrap();
Ok(UserRepository::new(pool)?)
}
#[tokio::test]
#[serial_test::serial]
async fn test_create_user() {
let user_repository = user_repository().await.unwrap();
truncate_users(&user_repository.pool).await;
let username = "picolo".to_string();
let user = user_repository
.create_user(username.clone(), "pw".to_string())
.await
.unwrap();
assert_eq!(user.username, username);
truncate_users(&user_repository.pool).await;
}
#[tokio::test]
#[serial_test::serial]
async fn test_find_by_username() {
let user_repository = user_repository().await.unwrap();
truncate_users(&user_repository.pool).await;
let username = "picolo".to_string();
user_repository
.create_user(username.clone(), "pw".to_string())
.await
.unwrap();
let fetched_user = user_repository
.find_by_username(username.clone())
.await
.unwrap();
assert_eq!(fetched_user.username, username);
truncate_users(&user_repository.pool).await;
}
fn token() -> String {
rand::rng()
.sample_iter(&rand::distr::Alphanumeric)
.take(64)
.map(char::from)
.collect()
}
#[tokio::test]
#[serial_test::serial]
async fn test_create_session() {
let user_repository = user_repository().await.unwrap();
truncate_sessions(&user_repository.pool).await;
truncate_users(&user_repository.pool).await;
let user = user_repository
.create_user("picolo".to_string(), "pw".to_string())
.await
.unwrap();
let token = token();
let expires_at = chrono::Utc::now() + chrono::Duration::minutes(2);
let session = user_repository
.create_session(user.id, token.clone(), expires_at)
.await
.unwrap();
assert_eq!(session.id, token);
assert_eq!(session.user_id, user.id);
assert!(expires_at > chrono::Utc::now());
truncate_sessions(&user_repository.pool).await;
truncate_users(&user_repository.pool).await;
}
#[tokio::test]
#[serial_test::serial]
async fn test_get_session() {
let user_repository = user_repository().await.unwrap();
truncate_sessions(&user_repository.pool).await;
truncate_users(&user_repository.pool).await;
let user = user_repository
.create_user("picolo".to_string(), "pw".to_string())
.await
.unwrap();
let token = token();
let expires_at = chrono::Utc::now() + chrono::Duration::minutes(2);
user_repository
.create_session(user.id, token.clone(), expires_at)
.await
.unwrap();
let fetched_session = user_repository.get_session(token.clone()).await.unwrap();
assert_eq!(fetched_session.id, token);
assert_eq!(fetched_session.user_id, user.id);
assert!(fetched_session.expires_at > chrono::Utc::now());
truncate_sessions(&user_repository.pool).await;
truncate_users(&user_repository.pool).await;
}
#[tokio::test]
#[serial_test::serial]
async fn test_delete_session() {
let user_repository = user_repository().await.unwrap();
truncate_sessions(&user_repository.pool).await;
truncate_users(&user_repository.pool).await;
let user = user_repository
.create_user("picolo".to_string(), "pw".to_string())
.await
.unwrap();
let token = token();
let expires_at = chrono::Utc::now() + chrono::Duration::minutes(2);
let session = user_repository
.create_session(user.id, token.clone(), expires_at)
.await
.unwrap();
assert_eq!(session.id, token);
assert_eq!(session.user_id, user.id);
assert!(expires_at > chrono::Utc::now());
user_repository.delete_session(token.clone()).await.unwrap();
assert!(matches!(
user_repository.get_session(token.clone()).await,
Err(LoftError::LoginFail)
));
truncate_sessions(&user_repository.pool).await;
truncate_users(&user_repository.pool).await;
}
}

View File

@@ -1,14 +1,12 @@
use axum::{
extract::{FromRequestParts, Request},
extract::{FromRequestParts, Request, State},
middleware::Next,
response::Response,
};
use lazy_regex::regex_captures;
use tower_cookies::{Cookie, Cookies};
use tower_cookies::Cookies;
use crate::{ctx::Ctx, error::LoftError, web::AUTH_TOKEN};
use crate::{ctx::Ctx, error::LoftError, model::UserRepository, web::AUTH_TOKEN};
/// validates the cookie exists and is well-formed (3-part format)
pub async fn mw_require_auth(
ctx: Result<Ctx, LoftError>,
req: Request,
@@ -20,35 +18,27 @@ pub async fn mw_require_auth(
}
pub async fn mw_ctx_resolver(
State(user_repository): State<UserRepository>,
cookies: Cookies,
mut req: Request,
next: Next,
) -> Result<Response, LoftError> {
let auth_token = cookies.get(AUTH_TOKEN).map(|c| c.value().to_string());
let result_ctx = match auth_token
.ok_or(LoftError::AuthFailNoAuthTokenCookie)
.and_then(parse_auth_token)
{
Ok((user_id, _, _)) => {
//TODO: add validation
Ok(Ctx::new(user_id))
}
Err(e) => Err(e),
let result_ctx = match auth_token {
Some(token) => user_repository
.get_session(token)
.await
.map(|s| Ctx::new(s.user_id)),
None => Err(LoftError::AuthFailNoAuthTokenCookie),
};
if result_ctx.is_err() && !matches!(result_ctx, Err(LoftError::AuthFailNoAuthTokenCookie)) {
cookies.remove(Cookie::from(AUTH_TOKEN))
}
req.extensions_mut().insert(result_ctx);
Ok(next.run(req).await)
}
impl<S: Send + Sync> FromRequestParts<S> for Ctx {
type Rejection = LoftError;
// extracts user_id from the token and makes it available to handlers as an extractor
fn from_request_parts(
parts: &mut axum::http::request::Parts,
_: &S,
@@ -62,18 +52,3 @@ impl<S: Send + Sync> FromRequestParts<S> for Ctx {
}
}
}
fn parse_auth_token(auth_token: String) -> Result<(u64, u64, String), LoftError> {
let (_, user_id, expiration, signature) =
regex_captures!(r"^user-(\d+)\.(\d+)\.([a-f0-9]+)$", &auth_token)
.ok_or(LoftError::AuthFailTokenWrongFormat)?;
let user_id: u64 = user_id
.parse()
.map_err(|_| LoftError::AuthFailTokenWrongFormat)?;
let expiration: u64 = expiration
.parse()
.map_err(|_| LoftError::AuthFailTokenWrongFormat)?;
Ok((user_id, expiration, signature.to_string()))
}

View File

@@ -93,36 +93,71 @@ mod tests {
TestServer,
multipart::{MultipartForm, Part},
};
use rand::RngExt;
use serde_json::json;
use sqlx::PgPool;
use tower_cookies::CookieManagerLayer;
use crate::{
model::FileRepository,
model::{FileRepository, UserRepository},
web::{
mw_auth::{mw_ctx_resolver, mw_require_auth},
routes_file::routes_file,
},
};
// Cookie format: user-[user-id].[expiration].[signature]
const AUTH_COOKIE: &str = "auth-token=user-1.0123456789.a1b2c3d4e5f6";
const BAD_AUTH_COOKIE: &str = "auth-token=user-1.0123456789";
const BAD_AUTH_COOKIE: &str = "auth-token=user-0123456789";
async fn file_repository() -> FileRepository {
dotenvy::dotenv().ok();
let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set");
let pool = PgPool::connect(&database_url).await.unwrap();
let file_repository = FileRepository::new(pool).unwrap();
file_repository
}
async fn user_repository() -> UserRepository {
dotenvy::dotenv().ok();
let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set");
let pool = PgPool::connect(&database_url).await.unwrap();
let user_repository = UserRepository::new(pool).unwrap();
user_repository
}
async fn test_server() -> TestServer {
let file_repository = FileRepository::new().await.unwrap();
let user_repository = user_repository().await;
let file_repository = file_repository().await;
let routes_file =
routes_file(file_repository.clone()).route_layer(middleware::from_fn(mw_require_auth));
let app = Router::new()
.nest("/api", routes_file)
.layer(middleware::from_fn_with_state(
file_repository,
user_repository,
mw_ctx_resolver,
))
.layer(CookieManagerLayer::new());
TestServer::new(app)
}
async fn create_test_session(user_repository: &UserRepository) -> String {
let user = user_repository
.create_user("testuser".to_string(), "hash".to_string())
.await
.unwrap();
let token: String = rand::rng()
.sample_iter(&rand::distr::Alphanumeric)
.take(64)
.map(char::from)
.collect();
let expires_at = chrono::Utc::now() + chrono::Duration::days(1);
user_repository
.create_session(user.id, token.clone(), expires_at)
.await
.unwrap();
token
}
async fn truncate(pool: &PgPool) {
sqlx::query!("TRUNCATE TABLE file_records")
.execute(pool)
@@ -149,7 +184,7 @@ mod tests {
#[tokio::test]
#[serial_test::serial]
async fn test_requires_auth_post() {
let file_repository = FileRepository::new().await.unwrap();
let file_repository = file_repository().await;
truncate(&file_repository.pool).await;
let server = test_server().await;
@@ -167,13 +202,15 @@ mod tests {
#[tokio::test]
#[serial_test::serial]
async fn test_list_files_empty() {
let file_repository = FileRepository::new().await.unwrap();
let user_repository = user_repository().await;
let token = create_test_session(&user_repository).await;
let file_repository = file_repository().await;
truncate(&file_repository.pool).await;
let server = test_server().await;
server
.get("/api/files")
.add_header(axum::http::header::COOKIE, AUTH_COOKIE)
.add_header(axum::http::header::COOKIE, format!("auth-token={token}"))
.await
.assert_status_ok()
.assert_json(&json!([]));
@@ -182,13 +219,15 @@ mod tests {
#[tokio::test]
#[serial_test::serial]
async fn test_upload_and_list_files() {
let file_repository = FileRepository::new().await.unwrap();
let user_repository = user_repository().await;
let token = create_test_session(&user_repository).await;
let file_repository = file_repository().await;
truncate(&file_repository.pool).await;
let server = test_server().await;
let res = server
.post("/api/files")
.add_header(axum::http::header::COOKIE, AUTH_COOKIE)
.add_header(axum::http::header::COOKIE, format!("auth-token={token}"))
.multipart(MultipartForm::new().add_part(
"file",
Part::bytes(b"fake_bytes".to_vec()).file_name("a.jpg"),
@@ -201,7 +240,7 @@ mod tests {
let list = server
.get("/api/files")
.add_header(axum::http::header::COOKIE, AUTH_COOKIE)
.add_header(axum::http::header::COOKIE, format!("auth-token={token}"))
.await
.json::<serde_json::Value>();
@@ -212,13 +251,15 @@ mod tests {
#[tokio::test]
#[serial_test::serial]
async fn test_download_file() {
let file_repository = FileRepository::new().await.unwrap();
let user_repository = user_repository().await;
let token = create_test_session(&user_repository).await;
let file_repository = file_repository().await;
truncate(&file_repository.pool).await;
let server = test_server().await;
let post_res = server
.post("/api/files")
.add_header(axum::http::header::COOKIE, AUTH_COOKIE)
.add_header(axum::http::header::COOKIE, format!("auth-token={token}"))
.multipart(MultipartForm::new().add_part(
"file",
Part::bytes(b"fake_bytes".to_vec()).file_name("a.jpg"),
@@ -227,7 +268,7 @@ mod tests {
let id = post_res.json::<serde_json::Value>()["id"].as_i64().unwrap();
let res = server
.get(&format!("/api/files/{id}/download"))
.add_header(axum::http::header::COOKIE, AUTH_COOKIE)
.add_header(axum::http::header::COOKIE, format!("auth-token={token}"))
.await;
res.assert_status_ok();
@@ -238,10 +279,12 @@ mod tests {
#[tokio::test]
#[serial_test::serial]
async fn test_download_file_not_found() {
let user_repository = user_repository().await;
let token = create_test_session(&user_repository).await;
let server = test_server().await;
server
.get("/api/files/99/download")
.add_header(axum::http::header::COOKIE, AUTH_COOKIE)
.add_header(axum::http::header::COOKIE, format!("auth-token={token}"))
.await
.assert_status_not_found();
}
@@ -249,13 +292,15 @@ mod tests {
#[tokio::test]
#[serial_test::serial]
async fn test_delete_file() {
let file_repository = FileRepository::new().await.unwrap();
let user_repository = user_repository().await;
let token = create_test_session(&user_repository).await;
let file_repository = file_repository().await;
truncate(&file_repository.pool).await;
let server = test_server().await;
let post_res = server
.post("/api/files")
.add_header(axum::http::header::COOKIE, AUTH_COOKIE)
.add_header(axum::http::header::COOKIE, format!("auth-token={token}"))
.multipart(MultipartForm::new().add_part(
"file",
Part::bytes(b"fake_bytes".to_vec()).file_name("a.jpg"),
@@ -265,13 +310,13 @@ mod tests {
server
.delete(&format!("/api/files/{id}"))
.add_header(axum::http::header::COOKIE, AUTH_COOKIE)
.add_header(axum::http::header::COOKIE, format!("auth-token={token}"))
.await
.assert_status_ok();
server
.get(&format!("/api/files/{id}"))
.add_header(axum::http::header::COOKIE, AUTH_COOKIE)
.add_header(axum::http::header::COOKIE, format!("auth-token={token}"))
.await
.assert_status_not_found();
@@ -281,10 +326,12 @@ mod tests {
#[tokio::test]
#[serial_test::serial]
async fn test_delete_file_not_found() {
let user_repository = user_repository().await;
let token = create_test_session(&user_repository).await;
let server = test_server().await;
server
.delete("/api/files/99")
.add_header(axum::http::header::COOKIE, AUTH_COOKIE)
.add_header(axum::http::header::COOKIE, format!("auth-token={token}"))
.await
.assert_status_not_found();
}

View File

@@ -1,41 +1,103 @@
use axum::{
Json, Router,
routing::{get, post},
use argon2::{
Argon2, PasswordHash, PasswordHasher, PasswordVerifier,
password_hash::{SaltString, rand_core::OsRng},
};
use axum::{Json, Router, extract::State, http::StatusCode, routing::post};
use rand::RngExt;
use serde::Deserialize;
use serde_json::{Value, json};
use sqlx::PgPool;
use tower_cookies::{Cookie, Cookies};
use crate::{error::LoftError, web::AUTH_TOKEN};
pub fn routes_login() -> Router {
use crate::{error::LoftError, model::UserRepository, web::AUTH_TOKEN};
pub fn routes_auth(user_repository: UserRepository) -> Router {
Router::new()
.route("/login", post(login))
.route("/register", get(register))
.route("/logout", post(logout))
.route("/register", post(register))
.with_state(user_repository)
}
async fn login(
State(user_repository): State<UserRepository>,
cookies: Cookies,
Json(payload): Json<LoginPayload>,
) -> Result<Json<Value>, LoftError> {
//TODO: real db/auth logic
if payload.username != "x" || payload.password != "y" {
) -> Result<StatusCode, LoftError> {
let user = user_repository.find_by_username(payload.username).await?;
// TODO: replace unwrap with ?
let parsed_hash = PasswordHash::new(&user.password_hash).unwrap();
if Argon2::default()
.verify_password(payload.password.as_bytes(), &parsed_hash)
.is_ok()
{
let expires_at = chrono::Utc::now() + chrono::Duration::days(1);
let cookie = create_cookie();
let auth_token = cookie.value().to_string();
cookies.add(cookie);
user_repository
.create_session(user.id, auth_token, expires_at)
.await?;
} else {
return Err(LoftError::LoginFail);
}
// FIXME: real auth-token generation-signature
cookies.add(Cookie::new(AUTH_TOKEN, "user-1.exp.sign"));
let body = Json(json!({
"result": {
"success": true
}
}));
Ok(body)
Ok(StatusCode::OK)
}
async fn register() -> &'static str {
"register"
async fn logout(
State(user_repository): State<UserRepository>,
cookies: Cookies,
) -> Result<StatusCode, LoftError> {
let auth_token = cookies.get(AUTH_TOKEN).map(|c| c.value().to_string());
if let Some(auth_token) = auth_token {
user_repository.delete_session(auth_token.clone()).await?;
cookies.remove(Cookie::build(AUTH_TOKEN).path("/").build());
}
Ok(StatusCode::OK)
}
async fn register(
State(user_repository): State<UserRepository>,
Json(payload): Json<LoginPayload>,
) -> Result<StatusCode, LoftError> {
if user_repository
.find_by_username(payload.username.clone())
.await
.is_ok()
{
return Err(LoftError::RegisterFail);
}
let salt = SaltString::generate(&mut OsRng);
let argon2 = Argon2::default();
// TODO: replace unwrap
let password_hash = argon2
.hash_password(payload.password.as_bytes(), &salt)
.unwrap()
.to_string();
user_repository
.create_user(payload.username, password_hash)
.await?;
Ok(StatusCode::CREATED)
}
fn create_cookie() -> Cookie<'static> {
let auth_token: String = rand::rng()
.sample_iter(&rand::distr::Alphanumeric)
.take(64)
.map(char::from)
.collect();
Cookie::build((AUTH_TOKEN, auth_token))
.http_only(true)
.same_site(tower_cookies::cookie::SameSite::Lax)
.path("/")
.secure(std::env::var("ENVIRONMENT").unwrap_or_default() == "prod")
.build()
}
#[derive(Debug, Deserialize)]
@@ -46,23 +108,54 @@ struct LoginPayload {
#[cfg(test)]
mod tests {
use axum::{
Router,
routing::{get, post},
};
use axum::{Router, http::StatusCode};
use axum_test::TestServer;
use serde_json::json;
use sqlx::PgPool;
use crate::web::routes_login::{login, register};
use crate::{
model::UserRepository,
web::{AUTH_TOKEN, routes_login::routes_auth},
};
async fn truncate_users(pool: &PgPool) {
sqlx::query!("TRUNCATE TABLE users CASCADE")
.execute(pool)
.await
.unwrap();
}
async fn truncate_sessions(pool: &PgPool) {
sqlx::query!("TRUNCATE TABLE sessions")
.execute(pool)
.await
.unwrap();
}
async fn user_repository() -> UserRepository {
dotenvy::dotenv().ok();
let database_url = std::env::var("DATABASE_URL").expect("DATABASE_URL must be set");
let pool = PgPool::connect(&database_url).await.unwrap();
let user_repository = UserRepository::new(pool).unwrap();
user_repository
}
async fn test_server() -> TestServer {
let user_repository = user_repository().await;
let routes_auth = routes_auth(user_repository);
let app = Router::new()
.nest("/api/auth", routes_auth)
.layer(tower_cookies::CookieManagerLayer::new());
TestServer::new(app)
}
#[tokio::test]
#[serial_test::serial]
async fn test_routes_login_wrong_credentials() {
let app = Router::new()
.route(&"/login", post(login))
.layer(tower_cookies::CookieManagerLayer::new());
let server = TestServer::new(app);
let server = test_server().await;
let response = server
.post("/login")
.post("/api/auth/login")
.json(&json!({
"username": "wrong",
"password": "wrong",
@@ -72,30 +165,74 @@ mod tests {
}
#[tokio::test]
#[serial_test::serial]
async fn test_routes_login() {
let app = Router::new()
.route(&"/login", post(login))
.layer(tower_cookies::CookieManagerLayer::new());
let server = TestServer::new(app);
let user_repository = user_repository().await;
truncate_sessions(&user_repository.pool).await;
truncate_users(&user_repository.pool).await;
let server = test_server().await;
server
.post("/api/auth/register")
.json(&json!({ "username": "picolo", "password": "picolo" }))
.await;
let response = server
.post("/login")
.post("/api/auth/login")
.json(&json!({
"username": "x",
"password": "y",
"username": "picolo",
"password": "picolo",
}))
.await;
response.assert_status_ok().assert_json(&json!({
"result": {
"success": true
}
}));
response.assert_status_ok();
}
#[tokio::test]
#[serial_test::serial]
async fn test_routes_logout() {
let user_repository = user_repository().await;
truncate_sessions(&user_repository.pool).await;
truncate_users(&user_repository.pool).await;
let server = test_server().await;
let register_response = server
.post("/api/auth/register")
.json(&json!({ "username": "picolo", "password": "picolo" }))
.await;
register_response.assert_status(StatusCode::CREATED);
let login_response = server
.post("/api/auth/login")
.json(&json!({
"username": "picolo",
"password": "picolo",
}))
.await;
login_response.assert_status_ok();
let cookie = login_response.cookie(AUTH_TOKEN);
let logout_response = server.post("/api/auth/logout").add_cookie(cookie).await;
logout_response.assert_status_ok();
}
#[tokio::test]
#[serial_test::serial]
async fn test_routes_register() {
let app = Router::new().route(&"/register", get(register));
let server = TestServer::new(app);
let response = server.get("/register").await;
response.assert_status_ok().assert_text("register");
let user_repository = user_repository().await;
truncate_sessions(&user_repository.pool).await;
truncate_users(&user_repository.pool).await;
let server = test_server().await;
let register_response = server
.post("/api/auth/register")
.json(&json!({ "username": "picolo", "password": "picolo" }))
.await;
register_response.assert_status(StatusCode::CREATED);
}
}