use std::sync::atomic::{AtomicBool, Ordering};
use actix_web::HttpRequest;
use argon2::password_hash::{SaltString, rand_core::RngCore};
use argon2::{Argon2, PasswordHash, PasswordHasher, PasswordVerifier};
use rand::rngs::OsRng;
use uuid::Uuid;
use crate::db::Database;
pub mod handlers;
#[derive(Debug, Clone)]
pub struct User {
pub id: String,
pub username: String,
pub email: Option<String>,
pub password_hash: String,
pub created_at: String,
pub is_admin: bool,
}
#[derive(Debug, Clone)]
pub struct Exercise {
pub id: String,
pub name: String,
pub description: String,
pub unit: String,
pub points_per_unit: f64,
pub created_at: String,
}
#[derive(Debug, Clone)]
pub struct Goal {
pub id: String,
pub name: String,
pub target_points: f64,
pub created_at: String,
}
#[derive(Debug, Clone)]
pub struct ExerciseLog {
pub id: String,
pub user_id: String,
pub exercise_id: String,
pub amount: f64,
pub points_earned: f64,
pub created_at: String,
}
pub fn hash_password(password: &str) -> Result<String, String> {
let argon2 = Argon2::default();
let salt = SaltString::generate(&mut OsRng);
let password_hash = argon2
.hash_password(password.as_bytes(), &salt)
.map_err(|e| e.to_string())?;
Ok(password_hash.to_string())
}
pub fn verify_password(password: &str, hash: &str) -> Result<bool, String> {
let argon2 = Argon2::default();
let parsed_hash = PasswordHash::new(hash).map_err(|e| e.to_string())?;
Ok(argon2
.verify_password(password.as_bytes(), &parsed_hash)
.is_ok())
}
pub fn generate_token() -> String {
let mut bytes = [0u8; 32];
OsRng.fill_bytes(&mut bytes);
hex::encode(bytes)
}
pub fn generate_short_code() -> String {
let mut bytes = [0u8; 4];
OsRng.fill_bytes(&mut bytes);
hex::encode(bytes)
}
pub fn create_user(username: String, email: String, password: String) -> Result<User, String> {
let password_hash = hash_password(&password)?;
let id = Uuid::new_v4().to_string();
let now = chrono::Utc::now().to_rfc3339();
let user = User {
id,
username,
email: Some(email),
password_hash,
created_at: now,
is_admin: false,
};
Ok(user)
}
pub fn create_exercise(
name: String,
description: String,
unit: String,
points_per_unit: f64,
) -> Exercise {
let id = Uuid::new_v4().to_string();
let now = chrono::Utc::now().to_rfc3339();
Exercise {
id,
name,
description,
unit,
points_per_unit,
created_at: now,
}
}
pub fn create_goal(name: String, target_points: f64) -> Goal {
let id = Uuid::new_v4().to_string();
let now = chrono::Utc::now().to_rfc3339();
Goal {
id,
name,
target_points,
created_at: now,
}
}
fn parse_basic_auth(encoded: &str) -> Option<(String, String)> {
let decoded = base64_decode(encoded)?;
let parts: Vec<&str> = decoded.splitn(2, ':').collect();
if parts.len() != 2 {
return None;
}
Some((parts[0].to_string(), parts[1].to_string()))
}
pub async fn is_admin_request(
req: &HttpRequest,
auth_state: &HodisContext,
config: &crate::config::Server,
) -> bool {
if let Some((username, password)) = extract_basic_auth(req)
&& username == config.admin_username()
&& password == config.admin_password()
{
return true;
}
if let Some(cookie) = req.cookie("session")
&& let Ok(Some(user_id)) = auth_state.db().get_token_user(cookie.value()).await
&& let Ok(Some(user)) = auth_state.db().get_user_by_id(&user_id).await
&& user.is_admin
{
return true;
}
false
}
pub fn extract_basic_auth(req: &HttpRequest) -> Option<(String, String)> {
if let Some(auth_header) = req.headers().get("Authorization") {
let auth_str = auth_header.to_str().ok()?;
if let Some(encoded) = auth_str.strip_prefix("Basic ") {
return parse_basic_auth(encoded);
}
}
if let Some(cookie) = req.cookie("admin_auth") {
return parse_basic_auth(cookie.value());
}
None
}
fn base64_decode(input: &str) -> Option<String> {
use base64::Engine;
let decoded = base64::engine::general_purpose::STANDARD
.decode(input)
.ok()?;
String::from_utf8(decoded).ok()
}
pub struct HodisContext {
db: Database,
initialized: AtomicBool,
}
impl HodisContext {
pub fn new(db: Database) -> Self {
Self {
db,
initialized: AtomicBool::new(false),
}
}
pub fn db(&self) -> &Database {
&self.db
}
pub fn set_initialized(&self) {
self.initialized.store(true, Ordering::SeqCst);
}
pub async fn create_session(&self, user_id: String) -> Result<String, String> {
let db = self.db();
let token = generate_token();
db.create_token(&token, &user_id).await?;
if let Err(e) = db.cleanup_expired_tokens().await {
log::warn!("Failed to cleanup expired tokens: {e}");
}
Ok(token)
}
pub async fn validate_token(&self, token: &str) -> Option<String> {
let db = self.db();
db.get_token_user(token).await.ok().flatten()
}
pub async fn invalidate_token(&self, token: &str) -> Result<(), String> {
let db = self.db();
db.delete_token(token).await
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_hash_and_verify_password() {
let hash = hash_password("testpassword").unwrap();
assert!(verify_password("testpassword", &hash).unwrap());
assert!(!verify_password("wrongpassword", &hash).unwrap());
}
#[test]
fn test_verify_password_invalid_hash() {
assert!(verify_password("test", "not-a-hash").is_err());
}
#[test]
fn test_generate_token_is_hex() {
let token = generate_token();
assert_eq!(token.len(), 64);
assert!(token.chars().all(|c| c.is_ascii_hexdigit()));
}
#[test]
fn test_generate_tokens_are_unique() {
let t1 = generate_token();
let t2 = generate_token();
assert_ne!(t1, t2);
}
#[test]
fn test_create_user_fields() {
let user = create_user(
"testuser".to_string(),
"test@example.com".to_string(),
"password123".to_string(),
)
.unwrap();
assert_eq!(user.username, "testuser");
assert_eq!(user.email, Some("test@example.com".to_string()));
assert!(!user.password_hash.is_empty());
assert!(!user.id.is_empty());
}
#[test]
fn test_create_user_password_is_hashed() {
let user = create_user(
"testuser".to_string(),
"test@example.com".to_string(),
"password123".to_string(),
)
.unwrap();
assert_ne!(user.password_hash, "password123");
assert!(verify_password("password123", &user.password_hash).unwrap());
}
#[test]
fn test_create_exercise_fields() {
let exercise = create_exercise(
"Running".to_string(),
"Go for a run".to_string(),
"km".to_string(),
2.0,
);
assert_eq!(exercise.name, "Running");
assert_eq!(exercise.unit, "km");
assert_eq!(exercise.points_per_unit, 2.0);
}
#[test]
fn test_create_goal_fields() {
let goal = create_goal("Marathon".to_string(), 500.0);
assert_eq!(goal.name, "Marathon");
assert_eq!(goal.target_points, 500.0);
}
#[test]
fn test_base64_decode_valid() {
let result = base64_decode("dXNlcjpwYXNz");
assert_eq!(result, Some("user:pass".to_string()));
}
#[test]
fn test_base64_decode_invalid() {
let result = base64_decode("not-valid-base64!!!");
assert!(result.is_none());
}
#[test]
fn test_base64_decode_empty() {
let result = base64_decode("");
assert_eq!(result, Some("".to_string()));
}
#[tokio::test]
async fn test_hodis_context_initialized_flag() {
let db = Database::new("/tmp/test_hodis_ctx_flag.db");
let ctx = HodisContext::new(db);
assert!(!ctx.initialized.load(Ordering::SeqCst));
ctx.set_initialized();
assert!(ctx.initialized.load(Ordering::SeqCst));
}
}