big features
This commit is contained in:
+51
-17
@@ -5,9 +5,10 @@ use axum::{
|
||||
routing::{get, post},
|
||||
Router,
|
||||
};
|
||||
use tower_http::{services::ServeDir, trace::TraceLayer};
|
||||
use tower::ServiceBuilder;
|
||||
use tower_http::{services::ServeDir, set_header::SetResponseHeaderLayer, trace::TraceLayer};
|
||||
|
||||
use crate::{api, db, state::SharedState, websocket};
|
||||
use crate::{api, auth, db, state::SharedState, websocket};
|
||||
|
||||
pub fn router(state: SharedState, static_dir: &str, upload_max_size_bytes: usize) -> Router {
|
||||
Router::new()
|
||||
@@ -19,6 +20,13 @@ pub fn router(state: SharedState, static_dir: &str, upload_max_size_bytes: usize
|
||||
.route("/health", get(health))
|
||||
.route("/f/{token}/{filename}", get(api::download_file))
|
||||
.route("/files/{directory}/{filename}", get(api::download_legacy_file))
|
||||
.route("/api/auth/identity", post(auth::identity))
|
||||
.route("/api/auth/register", post(auth::register))
|
||||
.route("/api/auth/login", post(auth::login))
|
||||
.route("/api/auth/me", get(auth::me))
|
||||
.route("/api/auth/logout", post(auth::logout))
|
||||
.route("/api/auth/password-reset", post(auth::request_reset))
|
||||
.route("/api/auth/password-reset/confirm", post(auth::confirm_reset))
|
||||
.route("/api/public/{token}", get(api::public_page))
|
||||
.route("/api/public/{token}/tasks", post(api::update_public_task))
|
||||
.route("/api/pads", post(api::create_pad))
|
||||
@@ -62,8 +70,17 @@ pub fn router(state: SharedState, static_dir: &str, upload_max_size_bytes: usize
|
||||
)
|
||||
.route("/static", get(static_not_found))
|
||||
.route("/static/{*path}", get(static_not_found))
|
||||
.nest_service("/assets", ServeDir::new(static_dir))
|
||||
.nest_service(
|
||||
"/assets",
|
||||
ServiceBuilder::new()
|
||||
.layer(SetResponseHeaderLayer::overriding(
|
||||
header::CACHE_CONTROL,
|
||||
HeaderValue::from_static("private, must-revalidate"),
|
||||
))
|
||||
.service(ServeDir::new(static_dir)),
|
||||
)
|
||||
.fallback(not_found)
|
||||
.method_not_allowed_fallback(method_not_allowed)
|
||||
.layer(DefaultBodyLimit::max(upload_max_size_bytes.saturating_add(1024 * 1024)))
|
||||
.layer(TraceLayer::new_for_http())
|
||||
.with_state(state)
|
||||
@@ -74,7 +91,7 @@ async fn health() -> &'static str {
|
||||
}
|
||||
|
||||
async fn home(State(state): State<SharedState>) -> Response {
|
||||
versioned_html(include_str!("../static/home.html"), &state.asset_version)
|
||||
versioned_html(include_str!("../static/home.html"), &state.asset_version, state.registration_enabled)
|
||||
}
|
||||
|
||||
async fn pad(
|
||||
@@ -85,7 +102,7 @@ async fn pad(
|
||||
Ok(Some(pad)) => {
|
||||
let html = include_str!("../static/pad.html")
|
||||
.replace("__PAD_TITLE__", &escape_html(&pad.title));
|
||||
versioned_html(&html, &state.asset_version)
|
||||
versioned_html(&html, &state.asset_version, state.registration_enabled)
|
||||
},
|
||||
Ok(None) => error_response(
|
||||
StatusCode::NOT_FOUND,
|
||||
@@ -108,7 +125,7 @@ async fn public_page(
|
||||
Path(token): Path<String>,
|
||||
) -> Response {
|
||||
match db::find_published_page(&state.db, &token).await {
|
||||
Ok(Some(_)) => versioned_html(include_str!("../static/public.html"), &state.asset_version),
|
||||
Ok(Some(_)) => versioned_html(include_str!("../static/public.html"), &state.asset_version, state.registration_enabled),
|
||||
Ok(None) => error_response(
|
||||
StatusCode::NOT_FOUND,
|
||||
"404",
|
||||
@@ -133,7 +150,7 @@ async fn workspace(
|
||||
Ok(Some(workspace)) => {
|
||||
let html = include_str!("../static/workspace.html")
|
||||
.replace("__WORKSPACE_TITLE__", &escape_html(&workspace.title));
|
||||
versioned_html(&html, &state.asset_version)
|
||||
versioned_html(&html, &state.asset_version, state.registration_enabled)
|
||||
},
|
||||
Ok(None) => error_response(
|
||||
StatusCode::NOT_FOUND,
|
||||
@@ -180,7 +197,7 @@ async fn note(
|
||||
.replace("__NOTE_TITLE__", &escape_html(¬e.title))
|
||||
.replace("__WORKSPACE_TITLE__", &escape_html(&workspace.title))
|
||||
.replace("__WORKSPACE_SLUG__", &escape_html(&workspace_slug));
|
||||
versioned_html(&html, &state.asset_version)
|
||||
versioned_html(&html, &state.asset_version, state.registration_enabled)
|
||||
},
|
||||
Ok(None) => error_response(
|
||||
StatusCode::NOT_FOUND,
|
||||
@@ -198,13 +215,28 @@ async fn note(
|
||||
}
|
||||
}
|
||||
|
||||
async fn static_not_found() -> Response {
|
||||
let mut response = (StatusCode::NOT_FOUND, "404").into_response();
|
||||
response.headers_mut().insert(
|
||||
header::CONTENT_TYPE,
|
||||
HeaderValue::from_static("text/plain; charset=utf-8"),
|
||||
);
|
||||
response
|
||||
async fn static_not_found(State(state): State<SharedState>) -> Response {
|
||||
error_response(
|
||||
StatusCode::NOT_FOUND,
|
||||
"404",
|
||||
"File not found",
|
||||
"The requested static file does not exist.",
|
||||
"/",
|
||||
"Home page",
|
||||
&state.asset_version,
|
||||
)
|
||||
}
|
||||
|
||||
async fn method_not_allowed(State(state): State<SharedState>) -> Response {
|
||||
error_response(
|
||||
StatusCode::METHOD_NOT_ALLOWED,
|
||||
"405",
|
||||
"Method not allowed",
|
||||
"This address does not support the requested operation.",
|
||||
"/",
|
||||
"Home page",
|
||||
&state.asset_version,
|
||||
)
|
||||
}
|
||||
|
||||
async fn not_found(State(state): State<SharedState>) -> Response {
|
||||
@@ -253,8 +285,10 @@ fn error_response(
|
||||
response
|
||||
}
|
||||
|
||||
fn versioned_html(template: &str, asset_version: &str) -> Response {
|
||||
let html = template.replace("__ASSET_VERSION__", asset_version);
|
||||
fn versioned_html(template: &str, asset_version: &str, registration_enabled: bool) -> Response {
|
||||
let html = template
|
||||
.replace("__ASSET_VERSION__", asset_version)
|
||||
.replace("__REGISTRATION_ENABLED__", if registration_enabled { "true" } else { "false" });
|
||||
let mut response = Html(html).into_response();
|
||||
no_store(&mut response);
|
||||
response
|
||||
|
||||
+188
@@ -0,0 +1,188 @@
|
||||
use argon2::{password_hash::{PasswordHash, PasswordHasher, PasswordVerifier, SaltString}, Argon2};
|
||||
use axum::{extract::State, http::{HeaderMap, StatusCode}, Json};
|
||||
use chrono::{Duration, Utc};
|
||||
use lettre::{message::Mailbox, AsyncSmtpTransport, AsyncTransport, Message, Tokio1Executor, transport::smtp::authentication::Credentials};
|
||||
use rand_core::{OsRng, RngCore};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sha2::{Digest, Sha256};
|
||||
use sqlx::FromRow;
|
||||
|
||||
use crate::{queries, state::{SharedState, SmtpConfig}};
|
||||
|
||||
const MIN_PASSWORD: usize = 8;
|
||||
const MAX_PASSWORD: usize = 128;
|
||||
const MAX_NICKNAME: usize = 40;
|
||||
|
||||
#[derive(Debug, Clone, FromRow)]
|
||||
pub struct User { pub id: i64, pub nickname: String, pub email: String, pub password_hash: String }
|
||||
|
||||
#[derive(Deserialize)] pub struct IdentityRequest { nickname: String, #[serde(default)] session_token: Option<String> }
|
||||
#[derive(Deserialize)] pub struct RegisterRequest { nickname: String, email: String, password: String }
|
||||
#[derive(Deserialize)] pub struct LoginRequest { email: String, password: String }
|
||||
#[derive(Deserialize)] pub struct ResetRequest { email: String }
|
||||
#[derive(Deserialize)] pub struct ResetConfirmRequest { token: String, password: String }
|
||||
#[derive(Serialize)] pub struct SessionResponse { token: String, nickname: String, email: String, expires_at: String }
|
||||
#[derive(Serialize)] pub struct IdentityResponse { nickname: String, registered: bool }
|
||||
|
||||
pub async fn identity(State(state): State<SharedState>, Json(req): Json<IdentityRequest>) -> Result<Json<IdentityResponse>, AuthError> {
|
||||
let nickname = validate_nickname(&req.nickname)?;
|
||||
match find_user_by_nickname(&state, &nickname).await? {
|
||||
None => Ok(Json(IdentityResponse { nickname, registered: false })),
|
||||
Some(user) => {
|
||||
let token = req.session_token.as_deref().ok_or_else(|| AuthError::unauthorized("This nickname is registered. Log in to use it."))?;
|
||||
let current = user_from_token(&state, token).await?.ok_or_else(|| AuthError::unauthorized("Your session has expired. Log in again."))?;
|
||||
if current.id != user.id { return Err(AuthError::unauthorized("This nickname belongs to another account.")); }
|
||||
Ok(Json(IdentityResponse { nickname: user.nickname, registered: true }))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn register(State(state): State<SharedState>, Json(req): Json<RegisterRequest>) -> Result<(StatusCode, Json<SessionResponse>), AuthError> {
|
||||
if !state.registration_enabled { return Err(AuthError::forbidden("Registration is disabled.")); }
|
||||
let nickname = validate_nickname(&req.nickname)?;
|
||||
let email = validate_email(&req.email)?;
|
||||
validate_password(&req.password)?;
|
||||
let nickname_key = normalize(&nickname);
|
||||
let email_key = normalize(&email);
|
||||
if find_user_by_nickname(&state, &nickname).await?.is_some() { return Err(AuthError::conflict("This nickname is already registered.")); }
|
||||
if find_user_by_email(&state, &email).await?.is_some() { return Err(AuthError::conflict("This e-mail address is already registered.")); }
|
||||
let hash = hash_password(&req.password)?;
|
||||
sqlx::query(queries::get(state.db.kind(), queries::AUTH_INSERT_USER))
|
||||
.bind(&nickname).bind(nickname_key).bind(&email).bind(email_key).bind(hash).execute(state.db.pool()).await
|
||||
.map_err(AuthError::database)?;
|
||||
let user = find_user_by_nickname(&state, &nickname).await?.ok_or_else(|| AuthError::internal("Failed to create the account."))?;
|
||||
Ok((StatusCode::CREATED, Json(create_session(&state, &user).await?)))
|
||||
}
|
||||
|
||||
pub async fn login(State(state): State<SharedState>, Json(req): Json<LoginRequest>) -> Result<Json<SessionResponse>, AuthError> {
|
||||
let email = validate_email(&req.email)?;
|
||||
let user = find_user_by_email(&state, &email).await?.ok_or_else(|| AuthError::unauthorized("Invalid e-mail address or password."))?;
|
||||
if !verify_password(&user.password_hash, &req.password) { return Err(AuthError::unauthorized("Invalid e-mail address or password.")); }
|
||||
Ok(Json(create_session(&state, &user).await?))
|
||||
}
|
||||
|
||||
pub async fn me(State(state): State<SharedState>, headers: HeaderMap) -> Result<Json<SessionResponse>, AuthError> {
|
||||
let token = bearer(&headers).ok_or_else(|| AuthError::unauthorized("Not logged in."))?;
|
||||
let user = user_from_token(&state, token).await?.ok_or_else(|| AuthError::unauthorized("Your session has expired."))?;
|
||||
let expires_at: String = sqlx::query_scalar(queries::get(state.db.kind(), queries::AUTH_SESSION_EXPIRES_AT))
|
||||
.bind(token).fetch_one(state.db.pool()).await.map_err(AuthError::database)?;
|
||||
Ok(Json(SessionResponse { token: token.into(), nickname: user.nickname, email: user.email, expires_at }))
|
||||
}
|
||||
|
||||
pub async fn logout(State(state): State<SharedState>, headers: HeaderMap) -> Result<Json<serde_json::Value>, AuthError> {
|
||||
if let Some(token) = bearer(&headers) { sqlx::query(queries::get(state.db.kind(), queries::AUTH_DELETE_SESSION_BY_TOKEN)).bind(token).execute(state.db.pool()).await.map_err(AuthError::database)?; }
|
||||
Ok(Json(serde_json::json!({"ok": true})))
|
||||
}
|
||||
|
||||
pub async fn request_reset(State(state): State<SharedState>, Json(req): Json<ResetRequest>) -> Result<Json<serde_json::Value>, AuthError> {
|
||||
let email = validate_email(&req.email)?;
|
||||
let smtp = state.smtp.as_ref().ok_or_else(|| AuthError::service_unavailable("Password reset is not configured on this server."))?;
|
||||
if let Some(user) = find_user_by_email(&state, &email).await? {
|
||||
let token = random_token();
|
||||
let expires = (Utc::now() + Duration::minutes(30)).to_rfc3339();
|
||||
sqlx::query(queries::get(state.db.kind(), queries::AUTH_DELETE_RESET_TOKENS_BY_USER)).bind(user.id).execute(state.db.pool()).await.map_err(AuthError::database)?;
|
||||
sqlx::query(queries::get(state.db.kind(), queries::AUTH_INSERT_RESET_TOKEN))
|
||||
.bind(hash_token(&token)).bind(user.id).bind(expires).execute(state.db.pool()).await.map_err(AuthError::database)?;
|
||||
send_reset(smtp, &user, &token).await?;
|
||||
}
|
||||
Ok(Json(serde_json::json!({"ok": true, "message": "If the account exists, a reset link has been sent."})))
|
||||
}
|
||||
|
||||
pub async fn confirm_reset(State(state): State<SharedState>, Json(req): Json<ResetConfirmRequest>) -> Result<Json<serde_json::Value>, AuthError> {
|
||||
validate_password(&req.password)?;
|
||||
let now_time = Utc::now();
|
||||
let now = now_time.to_rfc3339();
|
||||
let token_hash = hash_token(req.token.trim());
|
||||
let token_row: Option<(i64, String, Option<String>)> = sqlx::query_as(queries::get(state.db.kind(), queries::AUTH_FIND_RESET_TOKEN))
|
||||
.bind(&token_hash).fetch_optional(state.db.pool()).await.map_err(AuthError::database)?;
|
||||
let (user_id, expires_at, used_at) = token_row.ok_or_else(|| AuthError::bad_request("The reset link is invalid or has expired."))?;
|
||||
let expires_at = chrono::DateTime::parse_from_rfc3339(&expires_at)
|
||||
.map_err(|_| AuthError::bad_request("The reset link is invalid or has expired."))?
|
||||
.with_timezone(&Utc);
|
||||
if used_at.is_some() || expires_at <= now_time {
|
||||
return Err(AuthError::bad_request("The reset link is invalid or has expired."));
|
||||
}
|
||||
let password_hash = hash_password(&req.password)?;
|
||||
let mut tx = state.db.pool().begin().await.map_err(AuthError::database)?;
|
||||
let updated = sqlx::query(queries::get(state.db.kind(), queries::AUTH_UPDATE_PASSWORD))
|
||||
.bind(password_hash).bind(&now).bind(user_id).execute(&mut *tx).await.map_err(AuthError::database)?;
|
||||
if updated.rows_affected() != 1 { return Err(AuthError::internal("The account could not be updated.")); }
|
||||
sqlx::query(queries::get(state.db.kind(), queries::AUTH_MARK_RESET_TOKEN_USED))
|
||||
.bind(&now).bind(&token_hash).execute(&mut *tx).await.map_err(AuthError::database)?;
|
||||
sqlx::query(queries::get(state.db.kind(), queries::AUTH_DELETE_SESSIONS_BY_USER))
|
||||
.bind(user_id).execute(&mut *tx).await.map_err(AuthError::database)?;
|
||||
tx.commit().await.map_err(AuthError::database)?;
|
||||
Ok(Json(serde_json::json!({"ok": true})))
|
||||
}
|
||||
|
||||
pub async fn user_from_token(state: &SharedState, token: &str) -> Result<Option<User>, AuthError> {
|
||||
let now = Utc::now().to_rfc3339();
|
||||
sqlx::query_as::<_, User>(queries::get(state.db.kind(), queries::AUTH_USER_BY_SESSION))
|
||||
.bind(token).bind(now).fetch_optional(state.db.pool()).await.map_err(AuthError::database)
|
||||
}
|
||||
|
||||
pub async fn authorize_nickname(state: &SharedState, nickname: Option<String>, token: Option<String>) -> Result<Option<String>, String> {
|
||||
let Some(nickname) = nickname else { return Ok(None); };
|
||||
let nickname = validate_nickname(&nickname).map_err(|e| e.message)?;
|
||||
let registered = find_user_by_nickname(state, &nickname).await.map_err(|_| "Database error".to_string())?;
|
||||
match registered {
|
||||
None => Ok(Some(nickname)),
|
||||
Some(owner) => {
|
||||
let Some(token) = token else { return Err("This nickname is registered. Log in to use it.".into()); };
|
||||
let current = user_from_token(state, &token).await.map_err(|_| "Database error".to_string())?;
|
||||
match current { Some(user) if user.id == owner.id => Ok(Some(owner.nickname)), _ => Err("This nickname belongs to another account or the session expired.".into()) }
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn create_session(state: &SharedState, user: &User) -> Result<SessionResponse, AuthError> {
|
||||
let token = random_token();
|
||||
let expires_at = (Utc::now() + Duration::days(30)).to_rfc3339();
|
||||
sqlx::query(queries::get(state.db.kind(), queries::AUTH_INSERT_SESSION))
|
||||
.bind(&token).bind(user.id).bind(&expires_at).execute(state.db.pool()).await.map_err(AuthError::database)?;
|
||||
Ok(SessionResponse { token, nickname: user.nickname.clone(), email: user.email.clone(), expires_at })
|
||||
}
|
||||
async fn find_user_by_nickname(state: &SharedState, nickname: &str) -> Result<Option<User>, AuthError> {
|
||||
sqlx::query_as::<_, User>(queries::get(state.db.kind(), queries::AUTH_USER_BY_NICKNAME))
|
||||
.bind(normalize(nickname)).fetch_optional(state.db.pool()).await.map_err(AuthError::database)
|
||||
}
|
||||
async fn find_user_by_email(state: &SharedState, email: &str) -> Result<Option<User>, AuthError> {
|
||||
sqlx::query_as::<_, User>(queries::get(state.db.kind(), queries::AUTH_USER_BY_EMAIL))
|
||||
.bind(normalize(email)).fetch_optional(state.db.pool()).await.map_err(AuthError::database)
|
||||
}
|
||||
fn validate_nickname(v: &str) -> Result<String, AuthError> { let v=v.trim(); if v.is_empty() || v.chars().count()>MAX_NICKNAME { return Err(AuthError::bad_request("Nickname must contain 1 to 40 characters.")); } if v.chars().any(|c| c.is_control()) { return Err(AuthError::bad_request("Nickname contains invalid characters.")); } Ok(v.into()) }
|
||||
fn validate_email(v: &str) -> Result<String, AuthError> { let v=v.trim(); if v.len()>320 || !v.contains('@') || v.starts_with('@') || v.ends_with('@') { return Err(AuthError::bad_request("Enter a valid e-mail address.")); } Ok(v.into()) }
|
||||
fn validate_password(v: &str) -> Result<(), AuthError> { if v.len()<MIN_PASSWORD || v.len()>MAX_PASSWORD { Err(AuthError::bad_request("Password must contain 8 to 128 characters.")) } else { Ok(()) } }
|
||||
fn normalize(v: &str)->String { v.trim().to_lowercase() }
|
||||
fn hash_password(v:&str)->Result<String,AuthError>{let salt=SaltString::generate(&mut OsRng);Argon2::default().hash_password(v.as_bytes(),&salt).map(|h|h.to_string()).map_err(|_|AuthError::internal("Failed to secure the password."))}
|
||||
fn verify_password(hash:&str,v:&str)->bool{PasswordHash::new(hash).ok().and_then(|h|Argon2::default().verify_password(v.as_bytes(),&h).ok()).is_some()}
|
||||
fn random_token()->String{let mut bytes=[0u8;32];let mut rng=OsRng;rng.fill_bytes(&mut bytes);bytes.iter().map(|b|format!("{b:02x}")).collect()}
|
||||
fn hash_token(v:&str)->String{format!("{:x}",Sha256::digest(v.as_bytes()))}
|
||||
fn bearer(headers:&HeaderMap)->Option<&str>{headers.get("authorization")?.to_str().ok()?.strip_prefix("Bearer ")}
|
||||
async fn send_reset(smtp:&SmtpConfig,user:&User,token:&str)->Result<(),AuthError>{
|
||||
let url=format!("{}/?reset_token={}",smtp.public_url.trim_end_matches('/'),token);
|
||||
let message=Message::builder().from(smtp.from.parse::<Mailbox>().map_err(|_|AuthError::internal("Invalid SMTP_FROM."))?).to(user.email.parse::<Mailbox>().map_err(|_|AuthError::internal("Invalid recipient address."))?).subject("RustPad password reset").body(format!("Hello {},\n\nUse this link within 30 minutes to set a new password:\n{}\n\nIf you did not request this, ignore this message.",user.nickname,url)).map_err(|_|AuthError::internal("Failed to build reset e-mail."))?;
|
||||
// Port 465 uses implicit TLS. Standard submission ports (usually 587)
|
||||
// require STARTTLS; using implicit TLS there causes an immediate SMTP failure.
|
||||
let mut builder = if smtp.port == 465 {
|
||||
AsyncSmtpTransport::<Tokio1Executor>::relay(&smtp.host)
|
||||
} else {
|
||||
AsyncSmtpTransport::<Tokio1Executor>::starttls_relay(&smtp.host)
|
||||
}
|
||||
.map_err(|error| {
|
||||
tracing::error!(error=%error, host=%smtp.host, port=smtp.port, "invalid SMTP configuration");
|
||||
AuthError::internal("Invalid SMTP configuration.")
|
||||
})?
|
||||
.port(smtp.port);
|
||||
if !smtp.username.is_empty() { builder=builder.credentials(Credentials::new(smtp.username.clone(),smtp.password.clone())); }
|
||||
let mailer=builder.build();
|
||||
mailer.send(message).await.map_err(|error| {
|
||||
tracing::error!(error=%error, host=%smtp.host, port=smtp.port, "password reset e-mail failed");
|
||||
AuthError::service_unavailable("The reset e-mail could not be sent. Check the SMTP configuration.")
|
||||
})?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub struct AuthError { status: StatusCode, pub message: String }
|
||||
impl AuthError { fn bad_request(m:&str)->Self{Self{status:StatusCode::BAD_REQUEST,message:m.into()}} fn unauthorized(m:&str)->Self{Self{status:StatusCode::UNAUTHORIZED,message:m.into()}} fn forbidden(m:&str)->Self{Self{status:StatusCode::FORBIDDEN,message:m.into()}} fn conflict(m:&str)->Self{Self{status:StatusCode::CONFLICT,message:m.into()}} fn internal(m:&str)->Self{Self{status:StatusCode::INTERNAL_SERVER_ERROR,message:m.into()}} fn service_unavailable(m:&str)->Self{Self{status:StatusCode::SERVICE_UNAVAILABLE,message:m.into()}} fn database(e:sqlx::Error)->Self{tracing::error!(error=%e,"authentication database error");Self::internal("Database error.")} }
|
||||
impl axum::response::IntoResponse for AuthError { fn into_response(self)->axum::response::Response{(self.status,Json(serde_json::json!({"error":self.message}))).into_response()} }
|
||||
+28
-1
@@ -10,6 +10,8 @@ pub struct Config {
|
||||
pub files_dir: String,
|
||||
pub upload_max_size_bytes: usize,
|
||||
pub asset_version: String,
|
||||
pub smtp: Option<crate::state::SmtpConfig>,
|
||||
pub registration_enabled: bool,
|
||||
}
|
||||
|
||||
impl Config {
|
||||
@@ -26,6 +28,18 @@ impl Config {
|
||||
return Err("UPLOAD_MAX_SIZE_MB must be greater than 0".into());
|
||||
}
|
||||
|
||||
let smtp_host = std::env::var("SMTP_HOST").ok().filter(|v| !v.trim().is_empty());
|
||||
let smtp = if let Some(host) = smtp_host {
|
||||
Some(crate::state::SmtpConfig {
|
||||
host,
|
||||
port: env_var("SMTP_PORT", "587").parse()?,
|
||||
username: std::env::var("SMTP_USERNAME").unwrap_or_default(),
|
||||
password: std::env::var("SMTP_PASSWORD").unwrap_or_default(),
|
||||
from: std::env::var("SMTP_FROM").map_err(|_| "SMTP_FROM is required when SMTP_HOST is set")?,
|
||||
public_url: std::env::var("PUBLIC_URL").map_err(|_| "PUBLIC_URL is required when SMTP_HOST is set")?,
|
||||
})
|
||||
} else { None };
|
||||
|
||||
Ok(Self {
|
||||
host,
|
||||
port,
|
||||
@@ -39,7 +53,9 @@ impl Config {
|
||||
upload_max_size_bytes: upload_max_size_mb
|
||||
.checked_mul(1024 * 1024)
|
||||
.ok_or("UPLOAD_MAX_SIZE_MB is too large")?,
|
||||
asset_version: env!("CARGO_PKG_VERSION").to_owned(),
|
||||
asset_version: env_var("ASSET_VERSION", env!("CARGO_PKG_VERSION")),
|
||||
smtp,
|
||||
registration_enabled: env_bool("REGISTRATION_ENABLED", false)?,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -47,3 +63,14 @@ impl Config {
|
||||
fn env_var(name: &str, default: &str) -> String {
|
||||
env::var(name).unwrap_or_else(|_| default.to_owned())
|
||||
}
|
||||
|
||||
fn env_bool(name: &str, default: bool) -> Result<bool, Box<dyn std::error::Error>> {
|
||||
match env::var(name) {
|
||||
Ok(value) => match value.trim().to_ascii_lowercase().as_str() {
|
||||
"1" | "true" | "yes" | "on" => Ok(true),
|
||||
"0" | "false" | "no" | "off" => Ok(false),
|
||||
_ => Err(format!("{name} must be true or false").into()),
|
||||
},
|
||||
Err(_) => Ok(default),
|
||||
}
|
||||
}
|
||||
|
||||
+4
-3
@@ -1,3 +1,4 @@
|
||||
use crate::queries;
|
||||
use sqlx::{any::AnyPoolOptions, AnyPool};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
@@ -22,9 +23,9 @@ impl Database {
|
||||
.connect(url)
|
||||
.await?;
|
||||
if kind == DatabaseKind::Sqlite {
|
||||
sqlx::query("PRAGMA foreign_keys = ON").execute(&pool).await?;
|
||||
sqlx::query("PRAGMA journal_mode = WAL").execute(&pool).await?;
|
||||
sqlx::query("PRAGMA busy_timeout = 5000").execute(&pool).await?;
|
||||
sqlx::query(queries::SQLITE_FOREIGN_KEYS_ON).execute(&pool).await?;
|
||||
sqlx::query(queries::SQLITE_JOURNAL_WAL).execute(&pool).await?;
|
||||
sqlx::query(queries::SQLITE_BUSY_TIMEOUT).execute(&pool).await?;
|
||||
}
|
||||
Ok(Self { pool, kind })
|
||||
}
|
||||
|
||||
@@ -11,11 +11,11 @@ use sqlx::{Any, Transaction};
|
||||
|
||||
async fn inserted_id(kind: DatabaseKind, tx: &mut Transaction<'_, Any>, table: &str) -> Result<i64, sqlx::Error> {
|
||||
let query = match kind {
|
||||
DatabaseKind::Sqlite => "SELECT last_insert_rowid()",
|
||||
DatabaseKind::MySql => "SELECT LAST_INSERT_ID()",
|
||||
DatabaseKind::Sqlite => queries::SQLITE_LAST_INSERT_ID,
|
||||
DatabaseKind::MySql => queries::MYSQL_LAST_INSERT_ID,
|
||||
DatabaseKind::Postgres => match table {
|
||||
"note_revisions" => "SELECT currval(pg_get_serial_sequence('note_revisions', 'id'))",
|
||||
"revisions" => "SELECT currval(pg_get_serial_sequence('revisions', 'id'))",
|
||||
"note_revisions" => queries::POSTGRES_NOTE_REVISION_LAST_INSERT_ID,
|
||||
"revisions" => queries::POSTGRES_PAD_REVISION_LAST_INSERT_ID,
|
||||
_ => unreachable!("unsupported identity table"),
|
||||
},
|
||||
};
|
||||
@@ -36,7 +36,8 @@ pub struct Workspace {
|
||||
pub struct Note {
|
||||
pub id: i64,
|
||||
#[serde(skip_serializing)]
|
||||
pub workspace_id: i64,
|
||||
#[sqlx(rename = "workspace_id")]
|
||||
pub _workspace_id: i64,
|
||||
pub slug: String,
|
||||
pub title: String,
|
||||
#[serde(skip_serializing)]
|
||||
@@ -66,7 +67,7 @@ impl From<SqliteNote> for Note {
|
||||
fn from(value: SqliteNote) -> Self {
|
||||
Self {
|
||||
id: value.id,
|
||||
workspace_id: value.workspace_id,
|
||||
_workspace_id: value.workspace_id,
|
||||
slug: value.slug,
|
||||
title: value.title,
|
||||
content: value.content,
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
mod api;
|
||||
mod auth;
|
||||
mod app;
|
||||
mod config;
|
||||
mod database;
|
||||
@@ -34,6 +35,8 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
|
||||
config.asset_version.clone(),
|
||||
config.files_dir.clone(),
|
||||
config.upload_max_size_bytes,
|
||||
config.smtp.clone(),
|
||||
config.registration_enabled,
|
||||
));
|
||||
let app = app::router(
|
||||
state,
|
||||
|
||||
@@ -1,6 +1,31 @@
|
||||
use std::{collections::HashMap, sync::{Mutex, OnceLock}};
|
||||
use crate::database::DatabaseKind;
|
||||
|
||||
|
||||
// Database bootstrap and identity helpers.
|
||||
pub const SQLITE_FOREIGN_KEYS_ON: &str = "PRAGMA foreign_keys = ON";
|
||||
pub const SQLITE_JOURNAL_WAL: &str = "PRAGMA journal_mode = WAL";
|
||||
pub const SQLITE_BUSY_TIMEOUT: &str = "PRAGMA busy_timeout = 5000";
|
||||
pub const SQLITE_LAST_INSERT_ID: &str = "SELECT last_insert_rowid()";
|
||||
pub const MYSQL_LAST_INSERT_ID: &str = "SELECT LAST_INSERT_ID()";
|
||||
pub const POSTGRES_NOTE_REVISION_LAST_INSERT_ID: &str = "SELECT currval(pg_get_serial_sequence('note_revisions', 'id'))";
|
||||
pub const POSTGRES_PAD_REVISION_LAST_INSERT_ID: &str = "SELECT currval(pg_get_serial_sequence('revisions', 'id'))";
|
||||
|
||||
// Authentication queries.
|
||||
pub const AUTH_INSERT_USER: &str = "INSERT INTO users (nickname, nickname_key, email, email_key, password_hash) VALUES (?, ?, ?, ?, ?)";
|
||||
pub const AUTH_SESSION_EXPIRES_AT: &str = "SELECT expires_at FROM user_sessions WHERE token = ?";
|
||||
pub const AUTH_DELETE_SESSION_BY_TOKEN: &str = "DELETE FROM user_sessions WHERE token = ?";
|
||||
pub const AUTH_DELETE_RESET_TOKENS_BY_USER: &str = "DELETE FROM password_reset_tokens WHERE user_id = ?";
|
||||
pub const AUTH_INSERT_RESET_TOKEN: &str = "INSERT INTO password_reset_tokens (token, user_id, expires_at) VALUES (?, ?, ?)";
|
||||
pub const AUTH_FIND_RESET_TOKEN: &str = "SELECT user_id, expires_at, used_at FROM password_reset_tokens WHERE token = ?";
|
||||
pub const AUTH_UPDATE_PASSWORD: &str = "UPDATE users SET password_hash = ?, updated_at = ? WHERE id = ?";
|
||||
pub const AUTH_MARK_RESET_TOKEN_USED: &str = "UPDATE password_reset_tokens SET used_at = ? WHERE token = ?";
|
||||
pub const AUTH_DELETE_SESSIONS_BY_USER: &str = "DELETE FROM user_sessions WHERE user_id = ?";
|
||||
pub const AUTH_USER_BY_SESSION: &str = "SELECT u.id, u.nickname, u.email, u.password_hash FROM user_sessions s JOIN users u ON u.id = s.user_id WHERE s.token = ? AND s.expires_at > ?";
|
||||
pub const AUTH_INSERT_SESSION: &str = "INSERT INTO user_sessions (token, user_id, expires_at) VALUES (?, ?, ?)";
|
||||
pub const AUTH_USER_BY_NICKNAME: &str = "SELECT id, nickname, email, password_hash FROM users WHERE nickname_key = ?";
|
||||
pub const AUTH_USER_BY_EMAIL: &str = "SELECT id, nickname, email, password_hash FROM users WHERE email_key = ?";
|
||||
|
||||
pub const Q001: &str = "SELECT id, slug, title, password_hash, created_at, updated_at FROM workspaces WHERE slug = ?";
|
||||
pub const Q002: &str = "INSERT INTO workspaces (slug, title, password_hash) VALUES (?, ?, ?)";
|
||||
pub const Q003: &str = "SELECT id, workspace_id, slug, title, content, created_at, updated_at, owner_map, protected, created_by FROM notes WHERE workspace_id = ? ORDER BY updated_at DESC, id DESC";
|
||||
|
||||
+7
-2
@@ -4,6 +4,9 @@ use tokio::sync::{broadcast, RwLock};
|
||||
|
||||
const CHANNEL_CAPACITY: usize = 256;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SmtpConfig { pub host: String, pub port: u16, pub username: String, pub password: String, pub from: String, pub public_url: String }
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct NoteUpdate {
|
||||
pub content: String,
|
||||
@@ -19,12 +22,14 @@ pub struct AppState {
|
||||
pub asset_version: String,
|
||||
pub files_dir: String,
|
||||
pub upload_max_size_bytes: usize,
|
||||
pub smtp: Option<SmtpConfig>,
|
||||
pub registration_enabled: bool,
|
||||
channels: RwLock<HashMap<String, broadcast::Sender<NoteUpdate>>>,
|
||||
}
|
||||
|
||||
impl AppState {
|
||||
pub fn new(db: Database, asset_version: String, files_dir: String, upload_max_size_bytes: usize) -> Self {
|
||||
Self { db, asset_version, files_dir, upload_max_size_bytes, channels: RwLock::new(HashMap::new()) }
|
||||
pub fn new(db: Database, asset_version: String, files_dir: String, upload_max_size_bytes: usize, smtp: Option<SmtpConfig>, registration_enabled: bool) -> Self {
|
||||
Self { db, asset_version, files_dir, upload_max_size_bytes, smtp, registration_enabled, channels: RwLock::new(HashMap::new()) }
|
||||
}
|
||||
async fn channel_for_key(&self, key: String) -> broadcast::Sender<NoteUpdate> {
|
||||
if let Some(sender) = self.channels.read().await.get(&key) { return sender.clone(); }
|
||||
|
||||
+8
-6
@@ -2,12 +2,12 @@ use axum::{extract::{ws::{Message, WebSocket}, Path, State, WebSocketUpgrade}, r
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use tracing::{debug, warn};
|
||||
use crate::{db, state::{NoteUpdate, SharedState}};
|
||||
use crate::{auth, db, state::{NoteUpdate, SharedState}};
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
enum ClientMessage {
|
||||
Authenticate { password: Option<String>, nickname: Option<String> },
|
||||
Authenticate { password: Option<String>, nickname: Option<String>, session_token: Option<String> },
|
||||
Update { content: String, owner_map: Option<String> },
|
||||
}
|
||||
|
||||
@@ -26,12 +26,13 @@ pub async fn upgrade(ws: WebSocketUpgrade, Path((workspace_slug, note_slug)): Pa
|
||||
async fn handle_socket(mut socket: WebSocket, state: SharedState, workspace_slug: String, note_slug: String) {
|
||||
let Some(workspace) = db::find_workspace(&state.db, &workspace_slug).await.ok().flatten() else { let _=send_error(&mut socket,"Workspace not found").await; return; };
|
||||
let Some(note) = db::find_note(&state.db, workspace.id, ¬e_slug).await.ok().flatten() else { let _=send_error(&mut socket,"Note not found").await; return; };
|
||||
let (password, nickname) = match socket.recv().await {
|
||||
let (password, nickname, session_token) = match socket.recv().await {
|
||||
Some(Ok(Message::Text(text))) => match serde_json::from_str::<ClientMessage>(&text) {
|
||||
Ok(ClientMessage::Authenticate { password, nickname }) => (password, clean_nickname(nickname)),
|
||||
Ok(ClientMessage::Authenticate { password, nickname, session_token }) => (password, clean_nickname(nickname), session_token),
|
||||
_ => { let _=send_error(&mut socket,"Wymagane uwierzytelnienie").await; return; }
|
||||
}, _ => return
|
||||
};
|
||||
let nickname = match auth::authorize_nickname(&state, nickname, session_token).await { Ok(value) => value, Err(message) => { let _=send_error(&mut socket,&message).await; return; } };
|
||||
if !db::verify_workspace_password(&workspace, password.as_deref()) { let _=send_error(&mut socket,"Invalid password").await; return; }
|
||||
if send(&mut socket,&ServerMessage::Authenticated { workspace_title:workspace.title.clone(), note_title:note.title.clone(), content:note.content.clone(), owner_map:note.owner_map.clone() }).await.is_err(){return;}
|
||||
let channel=state.note_channel(&workspace_slug,¬e_slug).await;
|
||||
@@ -76,12 +77,13 @@ pub async fn upgrade_pad(ws:WebSocketUpgrade,Path(slug):Path<String>,State(state
|
||||
}
|
||||
async fn handle_pad_socket(mut socket:WebSocket,state:SharedState,slug:String){
|
||||
let Some(pad)=db::find_pad(&state.db,&slug).await.ok().flatten() else {let _=send_pad(&mut socket,&PadServerMessage::Error{message:"Note not found".into()}).await;return;};
|
||||
let (password,nickname)=match socket.recv().await{
|
||||
let (password,nickname,session_token)=match socket.recv().await{
|
||||
Some(Ok(Message::Text(text)))=>match serde_json::from_str::<ClientMessage>(&text){
|
||||
Ok(ClientMessage::Authenticate{password,nickname})=>(password,clean_nickname(nickname)),
|
||||
Ok(ClientMessage::Authenticate{password,nickname,session_token})=>(password,clean_nickname(nickname),session_token),
|
||||
_=>{let _=send_pad(&mut socket,&PadServerMessage::Error{message:"Wymagane uwierzytelnienie".into()}).await;return;}
|
||||
},_=>return
|
||||
};
|
||||
let nickname=match auth::authorize_nickname(&state,nickname,session_token).await{Ok(value)=>value,Err(message)=>{let _=send_pad(&mut socket,&PadServerMessage::Error{message}).await;return;}};
|
||||
if !db::verify_pad_password(&pad,password.as_deref()){let _=send_pad(&mut socket,&PadServerMessage::Error{message:"Invalid password".into()}).await;return;}
|
||||
if send_pad(&mut socket,&PadServerMessage::Authenticated{title:pad.title.clone(),content:pad.content.clone(),owner_map:pad.owner_map.clone()}).await.is_err(){return;}
|
||||
let channel=state.pad_channel(&slug).await;
|
||||
|
||||
Reference in New Issue
Block a user