254 lines
8.6 KiB
Rust
254 lines
8.6 KiB
Rust
mod row_decode;
|
|
mod api;
|
|
mod app;
|
|
mod assets;
|
|
mod auth;
|
|
mod config;
|
|
mod database;
|
|
mod db;
|
|
mod queries;
|
|
mod state;
|
|
mod storage;
|
|
mod websocket;
|
|
|
|
use std::{net::SocketAddr, path::PathBuf, sync::Arc};
|
|
|
|
use config::Config;
|
|
use database::{Database, DatabaseKind};
|
|
use state::AppState;
|
|
use tokio::net::TcpListener;
|
|
use tracing::{info, warn};
|
|
use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt};
|
|
|
|
|
|
#[tokio::main]
|
|
async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
|
rustls::crypto::ring::default_provider()
|
|
.install_default()
|
|
.map_err(|_| "failed to install rustls ring CryptoProvider")?;
|
|
|
|
dotenvy::dotenv().ok();
|
|
let cli = parse_command()?;
|
|
init_tracing();
|
|
|
|
let config = Config::load(cli.config.as_deref())?;
|
|
if matches!(cli.command, Command::CheckConfig) {
|
|
println!("configuration is valid{}", cli.config.as_ref().map(|path| format!(" ({})", path.display())).unwrap_or_default());
|
|
return Ok(());
|
|
}
|
|
info!(
|
|
host = %config.host,
|
|
port = config.port,
|
|
database_kind = %database_kind_label(&config.database_url),
|
|
database_max_connections = config.database_max_connections,
|
|
static_dir = %config.static_dir,
|
|
files_dir = %config.files_dir,
|
|
storage_driver = match &config.storage { storage::StorageConfig::Local { .. } => "local", storage::StorageConfig::S3 { .. } => "s3" },
|
|
upload_max_size_bytes = config.upload_max_size_bytes,
|
|
asset_cache_max_age_seconds = config.asset_cache_max_age_seconds,
|
|
file_cache_max_age_seconds = config.file_cache_max_age_seconds,
|
|
registration_enabled = config.registration_enabled,
|
|
account_confirmation_required = config.account_confirmation_required,
|
|
share_confirmation_required = config.share_confirmation_required,
|
|
frontend_log_level = %config.frontend_log_level,
|
|
anonymous_access_token_ttl_days = config.anonymous_access_token_ttl_days,
|
|
user_session_ttl_days = config.user_session_ttl_days,
|
|
smtp_configured = config.smtp.is_some(),
|
|
authorization_type = config.authorization_type.as_str(),
|
|
asset_version = %config.asset_version,
|
|
"configuration loaded"
|
|
);
|
|
if let Some(path) = config
|
|
.database_url
|
|
.strip_prefix("sqlite://")
|
|
.and_then(|v| v.split('?').next())
|
|
{
|
|
if let Some(parent) = std::path::Path::new(path).parent() {
|
|
std::fs::create_dir_all(parent)?;
|
|
}
|
|
}
|
|
info!("connecting to database");
|
|
let db = Database::connect(&config.database_url, config.database_max_connections).await?;
|
|
info!(database_kind = ?db.kind(), "database connection established");
|
|
run_migrations(&db).await?;
|
|
info!(database_kind = ?db.kind(), "database migrations completed");
|
|
if matches!(cli.command, Command::Migrate) {
|
|
println!("database migrations completed");
|
|
return Ok(());
|
|
}
|
|
|
|
let storage = storage::Storage::from_config(config.storage.clone()).await?;
|
|
info!(
|
|
storage_driver = storage.backend_name(),
|
|
"file storage ready"
|
|
);
|
|
let state = Arc::new(AppState::new(
|
|
db,
|
|
config.asset_version.clone(),
|
|
storage,
|
|
config.upload_max_size_bytes,
|
|
config.file_cache_max_age_seconds,
|
|
config.smtp.clone(),
|
|
config.registration_enabled && config.ldap.is_none(),
|
|
config.account_confirmation_required,
|
|
config.share_confirmation_required,
|
|
config.frontend_log_level.clone(),
|
|
config.anonymous_access_token_ttl_days,
|
|
config.user_session_ttl_days,
|
|
config.unconfirmed_account_ttl_days,
|
|
config.ldap.clone(),
|
|
));
|
|
let cleanup_state = state.clone();
|
|
tokio::spawn(async move {
|
|
let mut interval = tokio::time::interval(std::time::Duration::from_secs(24 * 60 * 60));
|
|
loop {
|
|
interval.tick().await;
|
|
let cutoff = (chrono::Utc::now() - chrono::Duration::days(cleanup_state.unconfirmed_account_ttl_days)).to_rfc3339();
|
|
match sqlx::query(crate::queries::get(cleanup_state.db.kind(), crate::queries::AUTH_DELETE_EXPIRED_UNCONFIRMED_USERS))
|
|
.bind(cutoff).execute(cleanup_state.db.pool()).await {
|
|
Ok(result) if result.rows_affected() > 0 => info!(deleted = result.rows_affected(), "removed expired unconfirmed accounts"),
|
|
Ok(_) => {},
|
|
Err(error) => tracing::error!(%error, "failed to remove expired unconfirmed accounts"),
|
|
}
|
|
}
|
|
});
|
|
let app = app::router(
|
|
state,
|
|
&config.static_dir,
|
|
config.upload_max_size_bytes,
|
|
config.asset_cache_max_age_seconds,
|
|
);
|
|
let address = SocketAddr::new(config.host, config.port);
|
|
let listener = TcpListener::bind(address).await?;
|
|
|
|
info!(%address, asset_version = %config.asset_version, "RustPad is running");
|
|
axum::serve(listener, app)
|
|
.with_graceful_shutdown(shutdown_signal())
|
|
.await?;
|
|
info!("RustPad stopped cleanly");
|
|
Ok(())
|
|
}
|
|
|
|
#[derive(Clone, Copy)]
|
|
enum Command { Run, CheckConfig, Migrate }
|
|
|
|
struct Cli {
|
|
command: Command,
|
|
config: Option<PathBuf>,
|
|
}
|
|
|
|
fn parse_command() -> Result<Cli, Box<dyn std::error::Error>> {
|
|
let mut command = Command::Run;
|
|
let mut config = None;
|
|
let mut args = std::env::args().skip(1);
|
|
|
|
while let Some(arg) = args.next() {
|
|
match arg.as_str() {
|
|
"-v" | "--version" => {
|
|
println!("rustpad {}", env!("CARGO_PKG_VERSION"));
|
|
std::process::exit(0);
|
|
}
|
|
"-h" | "--help" => {
|
|
print_help();
|
|
std::process::exit(0);
|
|
}
|
|
"-c" | "--config" => {
|
|
let path = args.next().ok_or("--config requires a file path")?;
|
|
if config.replace(PathBuf::from(path)).is_some() {
|
|
return Err("--config can only be specified once".into());
|
|
}
|
|
}
|
|
"check-config" => {
|
|
if !matches!(command, Command::Run) {
|
|
return Err("only one command may be specified".into());
|
|
}
|
|
command = Command::CheckConfig;
|
|
}
|
|
"migrate" => {
|
|
if !matches!(command, Command::Run) {
|
|
return Err("only one command may be specified".into());
|
|
}
|
|
command = Command::Migrate;
|
|
}
|
|
_ if arg.starts_with('-') => return Err(format!("unknown option: {arg}; use --help").into()),
|
|
_ => return Err(format!("unknown command: {arg}; use --help").into()),
|
|
}
|
|
}
|
|
|
|
Ok(Cli { command, config })
|
|
}
|
|
|
|
fn print_help() {
|
|
println!(
|
|
"rustpad {version}
|
|
|
|
USAGE:
|
|
rustpad [OPTIONS] [COMMAND]
|
|
|
|
OPTIONS:
|
|
-c, --config <FILE> Load YAML configuration file; environment variables override it
|
|
-h, --help Show help
|
|
-v, --version Show version
|
|
|
|
COMMANDS:
|
|
check-config Parse and validate configuration, then exit
|
|
migrate Validate configuration, apply database migrations, then exit",
|
|
version = env!("CARGO_PKG_VERSION")
|
|
);
|
|
}
|
|
|
|
fn init_tracing() {
|
|
tracing_subscriber::registry()
|
|
.with(
|
|
tracing_subscriber::EnvFilter::try_from_default_env()
|
|
.unwrap_or_else(|_| "rustpad=debug,tower_http=info".into()),
|
|
)
|
|
.with(tracing_subscriber::fmt::layer())
|
|
.init();
|
|
}
|
|
|
|
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 SIGTERM handler")
|
|
.recv()
|
|
.await;
|
|
};
|
|
#[cfg(not(unix))]
|
|
let terminate = std::future::pending::<()>();
|
|
tokio::select! {
|
|
() = ctrl_c => warn!("shutdown requested by Ctrl+C"),
|
|
() = terminate => warn!("shutdown requested by SIGTERM"),
|
|
}
|
|
}
|
|
|
|
async fn run_migrations(db: &Database) -> Result<(), sqlx::migrate::MigrateError> {
|
|
let path = match db.kind() {
|
|
DatabaseKind::Sqlite => std::path::Path::new("migrations/sqlite"),
|
|
DatabaseKind::Postgres => std::path::Path::new("migrations/postgres"),
|
|
DatabaseKind::MySql => std::path::Path::new("migrations/mysql"),
|
|
};
|
|
sqlx::migrate::Migrator::new(path)
|
|
.await?
|
|
.run(db.pool())
|
|
.await
|
|
}
|
|
|
|
fn database_kind_label(url: &str) -> &'static str {
|
|
if url.starts_with("sqlite:") {
|
|
"sqlite"
|
|
} else if url.starts_with("postgres:") || url.starts_with("postgresql:") {
|
|
"postgres"
|
|
} else if url.starts_with("mysql:") {
|
|
"mysql"
|
|
} else {
|
|
"unknown"
|
|
}
|
|
}
|