Files
sustenance/src/main.rs
T

215 lines
7.9 KiB
Rust

mod assets;
mod domain;
mod http;
mod hub;
mod ports;
mod security;
mod seed;
mod services;
mod sqlite;
mod views;
mod webauthn;
use std::env;
use std::path::Path as FilePath;
use std::sync::Arc;
use tokio::net::TcpListener;
use tracing::{info, warn};
use crate::http::{AppState, build_router};
use crate::hub::InMemoryHub;
use crate::ports::{
CategoryRepository, InvitationRepository, ItemRepository, ListMealRepository, ListRepository,
MealCategoryRepository, MealIngredientRepository, MealRepository, PasskeyRepository,
PasswordHasher, RealtimeNotifier, RewardsCardRepository, SessionRepository, TokenGenerator,
UserRepository,
};
use crate::security::{Argon2PasswordHasher, RandomTokenGenerator};
use crate::services::{
AuthService, InvitationService, ListService, MealService, RegistrationMode, RewardsCardService,
};
use crate::sqlite::{
SqliteCategoryRepository, SqliteDatabase, SqliteInvitationRepository, SqliteItemRepository,
SqliteListMealRepository, SqliteListRepository, SqliteMealCategoryRepository,
SqliteMealIngredientRepository, SqliteMealRepository, SqlitePasskeyRepository,
SqliteRewardsCardRepository, SqliteSessionRepository, SqliteUserRepository,
};
use crate::webauthn::{AppWebauthnConfig, WebAuthnService};
#[tokio::main]
async fn main() -> Result<(), Box<dyn std::error::Error>> {
tracing_subscriber::fmt()
.with_env_filter(
env::var("RUST_LOG").unwrap_or_else(|_| "sustenance=info,tower_http=info".into()),
)
.init();
let database_path = env::var("DATABASE_PATH").unwrap_or_else(|_| "sustenance.db".into());
let database_in_memory = env::var("DATABASE_IN_MEMORY")
.map(|value| value == "1" || value.eq_ignore_ascii_case("true"))
.unwrap_or(false);
let bind_address = env::var("BIND_ADDRESS").unwrap_or_else(|_| "127.0.0.1:3000".into());
// For loopback hosts, advertise `localhost` so WebAuthn works locally (browsers
// reject IP addresses as RP IDs). Access the app via http://localhost:PORT.
let bind_host = bind_address.split(':').next().unwrap_or("127.0.0.1");
let is_loopback = bind_host == "127.0.0.1" || bind_host == "::1" || bind_host == "localhost";
let public_host = if is_loopback { "localhost" } else { bind_host };
let public_base_url = env::var("PUBLIC_BASE_URL").unwrap_or_else(|_| {
let port = bind_address.rsplit(':').next().unwrap_or("3000");
format!("http://{}:{}", public_host, port)
});
let cookie_secure = env::var("COOKIE_SECURE")
.map(|value| value == "1" || value.eq_ignore_ascii_case("true"))
.unwrap_or(false);
let registration_mode = match env::var("REGISTRATION_MODE")
.unwrap_or_else(|_| "invite_only".into())
.to_ascii_lowercase()
.as_str()
{
"open" => RegistrationMode::Open,
"invite_only" | "invite-only" => RegistrationMode::InviteOnly,
value => {
warn!(value, "unknown REGISTRATION_MODE; using invite_only");
RegistrationMode::InviteOnly
}
};
// Build the adapters (ports) and wire them into application services.
let db = if database_in_memory {
SqliteDatabase::open_in_memory().await?
} else {
SqliteDatabase::open(&database_path).await?
};
let users: Arc<dyn UserRepository> = Arc::new(SqliteUserRepository);
let sessions: Arc<dyn SessionRepository> = Arc::new(SqliteSessionRepository);
let lists: Arc<dyn ListRepository> = Arc::new(SqliteListRepository);
let categories: Arc<dyn CategoryRepository> = Arc::new(SqliteCategoryRepository);
let items: Arc<dyn ItemRepository> = Arc::new(SqliteItemRepository);
let list_meals: Arc<dyn ListMealRepository> = Arc::new(SqliteListMealRepository);
let meals: Arc<dyn MealRepository> = Arc::new(SqliteMealRepository);
let meal_ingredients: Arc<dyn MealIngredientRepository> =
Arc::new(SqliteMealIngredientRepository);
let meal_categories: Arc<dyn MealCategoryRepository> = Arc::new(SqliteMealCategoryRepository);
let invitations: Arc<dyn InvitationRepository> = Arc::new(SqliteInvitationRepository);
let passkeys: Arc<dyn PasskeyRepository> = Arc::new(SqlitePasskeyRepository);
let rewards_cards: Arc<dyn RewardsCardRepository> = Arc::new(SqliteRewardsCardRepository);
let hasher: Arc<dyn PasswordHasher> = Arc::new(Argon2PasswordHasher);
let tokens: Arc<dyn TokenGenerator> = Arc::new(RandomTokenGenerator);
let realtime: Arc<dyn RealtimeNotifier> = Arc::new(InMemoryHub::default());
let auth = Arc::new(AuthService::new(
db.clone(),
Arc::clone(&users),
Arc::clone(&sessions),
Arc::clone(&invitations),
Arc::clone(&hasher),
registration_mode,
));
let lists_service = Arc::new(ListService::new(
db.clone(),
Arc::clone(&lists),
Arc::clone(&categories),
Arc::clone(&items),
Arc::clone(&realtime),
));
let invitations_service = Arc::new(InvitationService::new(
db.clone(),
Arc::clone(&invitations),
Arc::clone(&tokens),
));
let rewards_cards_service = Arc::new(RewardsCardService::new(
db.clone(),
Arc::clone(&rewards_cards),
));
let meals_service = Arc::new(MealService::new(
db.clone(),
Arc::clone(&meals),
Arc::clone(&meal_ingredients),
Arc::clone(&meal_categories),
Arc::clone(&lists),
Arc::clone(&items),
Arc::clone(&list_meals),
Arc::clone(&realtime),
));
// WebAuthn config from env vars. RP_ID must match the host users access the site from.
let rp_id = env::var("RP_ID").unwrap_or_else(|_| {
let host = public_base_url
.trim_start_matches("http://")
.trim_start_matches("https://")
.split('/')
.next()
.unwrap_or("localhost")
.split(':')
.next()
.unwrap_or("localhost")
.to_owned();
host
});
let rp_name = env::var("RP_NAME").unwrap_or_else(|_| "Sustenance".into());
let origin =
url::Url::parse(&public_base_url).map_err(|e| format!("invalid PUBLIC_BASE_URL: {e}"))?;
let webauthn_service = Arc::new(WebAuthnService::new(
db.clone(),
AppWebauthnConfig::new(rp_id, rp_name, origin),
Arc::clone(&users),
Arc::clone(&passkeys),
));
let seed_path = env::var("SEED_CONFIG").unwrap_or_else(|_| "seed.json".into());
seed::seed_if_needed(&db, &users, &hasher, FilePath::new(&seed_path)).await;
let state = AppState {
auth,
lists: lists_service,
meals: meals_service,
invitations: invitations_service,
rewards_cards: rewards_cards_service,
webauthn: webauthn_service,
realtime,
cookie_secure,
public_base_url,
};
let app = build_router(state);
let listener = TcpListener::bind(&bind_address).await?;
info!(address = %bind_address, "sustenance listening");
axum::serve(listener, app)
.with_graceful_shutdown(shutdown_signal())
.await?;
info!("shutdown complete; closing database");
// Explicitly close the pool so SQLite can checkpoint and remove the
// WAL/SHM sidecar files. Without this, the pool's background close task
// races with process exit and the sidecars can be left behind.
db.close().await;
Ok(())
}
async fn shutdown_signal() {
let ctrl_c = async {
tokio::signal::ctrl_c()
.await
.expect("failed to install Ctrl+C handler");
};
#[cfg(unix)]
let terminate = async {
tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())
.expect("failed to install signal handler")
.recv()
.await;
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! {
_ = ctrl_c => {},
_ = terminate => {},
}
info!("signal received; starting graceful shutdown");
}