use std::future::Future; use std::sync::Arc; use std::time::Duration; use axum::{ Router, extract::{ Form, FromRequest, FromRequestParts, Json, Path, Query, Request, State, ws::{Message, WebSocket, WebSocketUpgrade}, }, http::{HeaderMap, HeaderValue, StatusCode, header, request::Parts}, middleware::{self, Next}, response::{Html, IntoResponse, Redirect, Response}, routing::{get, post}, }; use futures_util::{SinkExt, StreamExt}; use maud::PreEscaped; use serde::{Deserialize, de::DeserializeOwned}; use thiserror::Error; use tower_http::trace::{DefaultMakeSpan, DefaultOnResponse, TraceLayer}; use tracing::{Level, error, warn}; use crate::assets; use crate::domain::{DomainError, SessionUser}; use crate::ports::{HubEvent, RealtimeNotifier}; use crate::services::{AuthService, InvitationService, ListService, MealService}; use crate::views; use crate::webauthn::WebAuthnService; #[derive(Clone)] pub struct AppState { pub auth: Arc, pub lists: Arc, pub meals: Arc, pub invitations: Arc, pub webauthn: Arc, pub realtime: Arc, pub cookie_secure: bool, pub public_base_url: String, } #[derive(Debug, Error)] pub enum AppError { #[error("database error")] Database(#[from] DomainError), #[error("bad request: {0}")] BadRequest(String), #[error("not found")] NotFound, #[error("list is archived")] Archived, } impl IntoResponse for AppError { fn into_response(self) -> Response { match self { AppError::Database(error) => { error!(%error, "request failed"); status_html_response( StatusCode::INTERNAL_SERVER_ERROR, views::error_page("500", "Something went wrong."), ) } AppError::BadRequest(message) => { status_html_response(StatusCode::BAD_REQUEST, views::error_page("400", &message)) } AppError::NotFound => status_html_response( StatusCode::NOT_FOUND, views::error_page("404", "That page could not be found."), ), AppError::Archived => status_html_response( StatusCode::CONFLICT, views::error_page("409", "This list is archived and cannot be modified."), ), } } } pub fn build_router(state: AppState) -> Router { Router::new() .route("/", get(home)) .route("/login", get(login_page).post(login)) .route("/register", get(register_page).post(register)) .route("/logout", post(logout)) .route("/account", get(account_page)) .route("/auth/passkey/register/start", post(passkey_register_start)) .route( "/auth/passkey/register/finish", post(passkey_register_finish), ) .route("/auth/passkey/login/start", post(passkey_login_start)) .route("/auth/passkey/login/finish", post(passkey_login_finish)) .route( "/account/passkeys/{passkey_id}/delete", post(delete_passkey), ) .route("/account/password", post(change_password)) .route("/lists", get(lists_page).post(create_list)) .route("/archive", get(archive_page)) .route("/lists/{list_id}", get(list_page)) .route("/lists/{list_id}/archive", post(archive_list)) .route("/lists/{list_id}/unarchive", post(unarchive_list)) .route("/lists/{list_id}/items", post(add_item)) .route("/lists/{list_id}/items/{item_id}/check", post(check_item)) .route("/lists/{list_id}/items/{item_id}/edit", post(edit_item)) .route("/lists/{list_id}/items/{item_id}/delete", post(delete_item)) .route("/categories", post(create_category)) .route("/invitations", post(create_invitation)) .route("/meals", get(meals_page).post(create_meal)) .route("/meals/new", get(new_meal_page)) .route("/meals/categories", post(create_meal_category)) .route( "/meals/categories/{category_id}/delete", post(delete_meal_category), ) .route("/meals/{meal_id}", get(meal_page)) .route("/meals/{meal_id}/edit", post(edit_meal)) .route("/meals/{meal_id}/delete", post(delete_meal)) .route("/meals/{meal_id}/ingredients", post(add_ingredient)) .route( "/meals/{meal_id}/ingredients/{ingredient_id}/edit", post(edit_ingredient), ) .route( "/meals/{meal_id}/ingredients/{ingredient_id}/delete", post(delete_ingredient), ) .route("/lists/{list_id}/add-meal", post(add_meal_to_list)) .route( "/lists/{list_id}/meals/{list_meal_id}/remove", post(remove_meal_from_list), ) .route("/lists/{list_id}/stream", get(list_stream)) .route("/invite/{token}", get(invitation_page)) .route("/invite/{token}/accept", post(accept_invitation)) .route("/static/{path}", get(static_asset)) .layer( TraceLayer::new_for_http() .make_span_with(DefaultMakeSpan::new().level(Level::INFO)) .on_response(DefaultOnResponse::new().level(Level::INFO)), ) .layer(middleware::from_fn(log_response_status)) .with_state(state) } #[derive(Clone)] struct CurrentUser { session_token: String, session: SessionUser, } impl FromRequestParts for CurrentUser { type Rejection = Response; fn from_request_parts( parts: &mut Parts, state: &AppState, ) -> impl Future> + Send { let session_token = cookie_value(&parts.headers, "session"); let auth = Arc::clone(&state.auth); async move { let Some(session_token) = session_token else { return Err(Redirect::to("/login").into_response()); }; match auth.session_user(session_token.clone()).await { Ok(Some(session)) => Ok(Self { session_token, session, }), Ok(None) => Err(Redirect::to("/login").into_response()), Err(error) => Err(AppError::Database(error).into_response()), } } } } struct LoggedForm(T); impl FromRequest for LoggedForm where S: Send + Sync, T: DeserializeOwned + Send, { type Rejection = Response; fn from_request( request: Request, state: &S, ) -> impl Future> + Send { let method = request.method().clone(); let uri = request.uri().clone(); async move { match Form::::from_request(request, state).await { Ok(Form(value)) => Ok(Self(value)), Err(rejection) => { warn!( %method, %uri, rejection = ?rejection, "request form deserialization failed" ); Err(rejection.into_response()) } } } } } #[derive(Debug, Deserialize)] struct InviteQuery { invite: Option, } #[derive(Debug, Deserialize)] struct RegisterForm { display_name: String, email: String, password: String, invite: Option, } #[derive(Debug, Deserialize)] struct LoginForm { email: String, password: String, invite: Option, } #[derive(Debug, Deserialize)] struct CreateListForm { name: String, csrf: String, } #[derive(Debug, Deserialize)] struct ItemForm { name: String, quantity: String, #[serde(default)] note: String, #[serde(default)] category_id: Option, csrf: String, } #[derive(Debug, Deserialize)] struct CheckForm { checked: String, csrf: String, } #[derive(Debug, Deserialize)] struct CsrfForm { csrf: String, } #[derive(Debug, Deserialize)] struct CategoryForm { name: String, csrf: String, } #[derive(Debug, Deserialize)] struct PasskeyRegisterStartForm { csrf: String, } #[derive(Debug, Deserialize)] struct PasskeyRegisterFinishForm { csrf: String, response: webauthn_rs::proto::RegisterPublicKeyCredential, } #[derive(Debug, Deserialize)] struct PasskeyLoginStartForm { #[serde(default)] email: String, } #[derive(Debug, Deserialize)] struct PasskeyLoginFinishForm { token: String, response: webauthn_rs::proto::PublicKeyCredential, } #[derive(Debug, Deserialize)] struct DeletePasskeyForm { csrf: String, } #[derive(Debug, Deserialize)] struct ChangePasswordForm { csrf: String, new_password: String, confirm_password: String, } #[derive(Debug, Deserialize)] struct MealForm { name: String, #[serde(default)] description: String, #[serde(default)] category_id: Option, csrf: String, } #[derive(Debug, Deserialize)] struct IngredientForm { name: String, #[serde(default)] quantity: String, #[serde(default)] note: String, #[serde(default)] category_id: Option, csrf: String, } #[derive(Debug, Deserialize)] struct AddMealForm { meal_id: i64, csrf: String, } #[derive(Debug, Deserialize)] struct MealPickerQuery { picker: Option, } async fn home() -> Redirect { Redirect::to("/lists") } /// Serves a file embedded in the binary from the `static/` folder. /// /// Every asset is served with an immutable, long-lived cache header. Because /// the HTML references each asset under a URL that is versioned by its content /// hash, a changed file gets a new URL and the cache is never stale. async fn static_asset(Path(path): Path) -> Response { let Some(asset) = assets::STORE.get(&path) else { return StatusCode::NOT_FOUND.into_response(); }; let mime = match path.as_str() { "style.css" => "text/css", _ => "application/javascript", }; let mut response = ([(header::CONTENT_TYPE, mime)], asset.data).into_response(); response.headers_mut().insert( header::CACHE_CONTROL, HeaderValue::from_static("public, max-age=31536000, immutable"), ); response } async fn log_response_status(request: Request, next: Next) -> Response { let method = request.method().clone(); let uri = request.uri().clone(); let response = next.run(request).await; let status = response.status(); if status.is_server_error() { error!(%method, %uri, %status, "request returned server error"); } else if status.is_client_error() { warn!(%method, %uri, %status, "request returned client error"); } response } async fn login_page(Query(query): Query) -> Result { Ok(html_response(views::login_page( None, query.invite.as_deref(), ))) } async fn register_page( State(state): State, Query(query): Query, ) -> Result { if state.auth.can_register(query.invite.as_deref()).await? { Ok(html_response(views::register_page( None, query.invite.as_deref(), ))) } else { Ok(status_html_response( StatusCode::FORBIDDEN, views::registration_closed_page(), )) } } async fn register( State(state): State, LoggedForm(form): LoggedForm, ) -> Result { if !state.auth.can_register(form.invite.as_deref()).await? { return Ok(status_html_response( StatusCode::FORBIDDEN, views::registration_closed_page(), )); } let display_name = form.display_name.trim().to_owned(); let email = form.email.trim().to_lowercase(); if display_name.is_empty() || display_name.chars().count() > 50 { return Ok(html_response(views::register_page( Some("Enter a name between 1 and 50 characters."), form.invite.as_deref(), ))); } if !email.contains('@') || email.len() > 200 { return Ok(html_response(views::register_page( Some("Enter a valid email address."), form.invite.as_deref(), ))); } let (_user, session_token) = match state .auth .register(display_name, email, form.password, form.invite.as_deref()) .await { Ok(result) => result, Err(DomainError::Conflict) => { return Ok(html_response(views::register_page( Some("An account with that email already exists."), form.invite.as_deref(), ))); } Err(error) => return Err(AppError::Database(error)), }; let destination = form .invite .filter(|invite| !invite.is_empty()) .map(|invite| format!("/invite/{invite}")) .unwrap_or_else(|| "/lists".into()); let mut response = Redirect::to(&destination).into_response(); set_session_cookie(&mut response, &session_token, state.cookie_secure); Ok(response) } async fn login( State(state): State, LoggedForm(form): LoggedForm, ) -> Result { let email = form.email.trim().to_lowercase(); let result = state.auth.login(email, form.password).await?; let Some((_user, session_token)) = result else { return Ok(html_response(views::login_page( Some("Email or password is incorrect."), form.invite.as_deref(), ))); }; let destination = form .invite .filter(|invite| !invite.is_empty()) .map(|invite| format!("/invite/{invite}")) .unwrap_or_else(|| "/lists".into()); let mut response = Redirect::to(&destination).into_response(); set_session_cookie(&mut response, &session_token, state.cookie_secure); Ok(response) } async fn logout(State(state): State, user: CurrentUser) -> Result { state.auth.logout(user.session_token).await?; let mut response = Redirect::to("/login").into_response(); clear_session_cookie(&mut response, state.cookie_secure); Ok(response) } async fn account_page( State(state): State, user: CurrentUser, ) -> Result { let passkeys = state.webauthn.list_passkeys(user.session.user.id).await?; Ok(html_response(views::account_page( &user.session.user, &passkeys, &user.session.csrf_token, None, false, ))) } async fn change_password( State(state): State, user: CurrentUser, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; let passkeys = state.webauthn.list_passkeys(user.session.user.id).await?; let render = |error: Option<&str>, success: bool| { html_response(views::account_page( &user.session.user, &passkeys, &user.session.csrf_token, error, success, )) }; if form.new_password != form.confirm_password { return Ok(render( Some("New password and confirmation do not match."), false, )); } state .auth .change_password(user.session.user.id, form.new_password) .await?; Ok(render(None, true)) } async fn passkey_register_start( State(state): State, user: CurrentUser, Json(form): Json, ) -> Result { verify_csrf(&user, &form.csrf)?; let challenge = state .webauthn .start_registration(&user.session.user) .map_err(AppError::Database)?; Ok(Json(challenge).into_response()) } async fn passkey_register_finish( State(state): State, user: CurrentUser, Json(form): Json, ) -> Result { verify_csrf(&user, &form.csrf)?; state .webauthn .finish_registration(&user.session.user, form.response) .await?; Ok(Redirect::to("/account").into_response()) } async fn passkey_login_start( State(state): State, Json(form): Json, ) -> Result { let email = form.email.trim().to_lowercase(); let (challenge, token) = if email.is_empty() { // Userless sign-in: no email needed, the authenticator selects a // discoverable credential and returns a user handle. state .webauthn .start_userless_authentication() .await .map_err(AppError::Database)? } else { let Some((user, _)) = state.auth.find_user_by_email(email).await? else { return Err(AppError::NotFound); }; state .webauthn .start_authentication(user.id) .await .map_err(AppError::Database)? }; Ok( Json(serde_json::json!({ "token": token, "publicKey": challenge.public_key })) .into_response(), ) } async fn passkey_login_finish( State(state): State, Json(form): Json, ) -> Result { let user_id = state .webauthn .finish_authentication(form.token, form.response) .await?; let (session_token, _) = state.auth.create_session_for_user(user_id).await?; let mut response = Redirect::to("/lists").into_response(); set_session_cookie(&mut response, &session_token, state.cookie_secure); Ok(response) } async fn delete_passkey( State(state): State, user: CurrentUser, Path(passkey_id): Path, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; state .webauthn .delete_passkey(user.session.user.id, passkey_id) .await?; Ok(Redirect::to("/account").into_response()) } async fn lists_page( State(state): State, user: CurrentUser, ) -> Result { let lists = state.lists.list_summaries().await?; let archived = state.lists.list_archived_summaries().await?; let categories = state.lists.categories().await?; Ok(html_response(views::lists_page( &user.session.user, &lists, archived.len(), &categories, &user.session.csrf_token, ))) } async fn archive_page( State(state): State, user: CurrentUser, ) -> Result { let archived_lists = state.lists.list_archived_summaries().await?; Ok(html_response(views::archive_page( &user.session.user, &archived_lists, ))) } async fn create_list( State(state): State, user: CurrentUser, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; let name = form.name.trim().to_owned(); if name.is_empty() || name.chars().count() > 80 { return Err(AppError::BadRequest( "List names must be between 1 and 80 characters.".into(), )); } let list = state.lists.create_list(name).await?; Ok(Redirect::to(&format!("/lists/{}", list.id)).into_response()) } async fn list_page( State(state): State, user: CurrentUser, Path(list_id): Path, ) -> Result { let access = require_list(&state, list_id).await?; let items = state.lists.items(list_id).await?; let categories = state.lists.categories().await?; let list_meals = state.meals.list_meals_on_list(list_id).await?; let presence = state.realtime.presence(list_id).await; Ok(html_response(views::list_page( &user.session.user, &access, &items, &categories, &list_meals, &presence, &user.session.csrf_token, ))) } async fn archive_list( State(state): State, user: CurrentUser, Path(list_id): Path, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; require_list(&state, list_id).await?; state.lists.archive_list(list_id).await?; Ok(Redirect::to("/lists").into_response()) } async fn unarchive_list( State(state): State, user: CurrentUser, Path(list_id): Path, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; require_list(&state, list_id).await?; state.lists.unarchive_list(list_id).await?; Ok(Redirect::to("/lists").into_response()) } async fn add_item( State(state): State, user: CurrentUser, Path(list_id): Path, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; require_mutable_list(&state, list_id).await?; let name = form.name.trim().to_owned(); let quantity = form.quantity.trim().to_owned(); let note = form.note.trim().to_owned(); let category_id = parse_category_id(form.category_id); if name.is_empty() || name.chars().count() > 120 { return Err(AppError::BadRequest( "Item names must be between 1 and 120 characters.".into(), )); } state .lists .add_item(list_id, name, quantity, note, category_id) .await?; list_fragment_response(&state, &user, list_id).await } async fn check_item( State(state): State, user: CurrentUser, Path((list_id, item_id)): Path<(i64, i64)>, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; require_mutable_list(&state, list_id).await?; let checked = match form.checked.as_str() { "1" | "true" => true, "0" | "false" => false, _ => return Err(AppError::BadRequest("Invalid checked value.".into())), }; state .lists .set_item_checked(list_id, item_id, checked) .await?; list_fragment_response(&state, &user, list_id).await } async fn edit_item( State(state): State, user: CurrentUser, Path((list_id, item_id)): Path<(i64, i64)>, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; require_mutable_list(&state, list_id).await?; let name = form.name.trim().to_owned(); if name.is_empty() || name.chars().count() > 120 { return Err(AppError::BadRequest( "Item names must be between 1 and 120 characters.".into(), )); } state .lists .update_item( list_id, item_id, name, form.quantity.trim().to_owned(), form.note.trim().to_owned(), parse_category_id(form.category_id), ) .await?; list_fragment_response(&state, &user, list_id).await } async fn delete_item( State(state): State, user: CurrentUser, Path((list_id, item_id)): Path<(i64, i64)>, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; require_mutable_list(&state, list_id).await?; state.lists.delete_item(list_id, item_id).await?; list_fragment_response(&state, &user, list_id).await } async fn create_category( State(state): State, user: CurrentUser, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; let name = form.name.trim().to_owned(); if name.is_empty() || name.chars().count() > 60 { return Err(AppError::BadRequest( "Category names must be between 1 and 60 characters.".into(), )); } state.lists.create_category(name).await?; Ok(Redirect::to("/lists").into_response()) } async fn create_meal_category( State(state): State, user: CurrentUser, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; let name = form.name.trim().to_owned(); if name.is_empty() || name.chars().count() > 60 { return Err(AppError::BadRequest( "Category names must be between 1 and 60 characters.".into(), )); } state.meals.create_meal_category(name).await?; Ok(Redirect::to("/meals").into_response()) } async fn delete_meal_category( State(state): State, user: CurrentUser, Path(category_id): Path, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; state.meals.delete_meal_category(category_id).await?; Ok(Redirect::to("/meals").into_response()) } async fn meals_page( State(state): State, user: CurrentUser, Query(query): Query, ) -> Result { let meals = state.meals.list_meals().await?; let meal_categories = state.meals.list_meal_categories().await?; if let Some(list_id) = query.picker { return Ok(html_response(views::meal_picker( &meals, &meal_categories, list_id, &user.session.csrf_token, ))); } Ok(html_response(views::meals_page( &user.session.user, &meals, &meal_categories, &user.session.csrf_token, ))) } async fn new_meal_page( State(state): State, user: CurrentUser, ) -> Result { let meal_categories = state.meals.list_meal_categories().await?; Ok(html_response(views::meal_form_page( &user.session.user, None, &meal_categories, &user.session.csrf_token, ))) } async fn create_meal( State(state): State, user: CurrentUser, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; let name = form.name.trim().to_owned(); if name.is_empty() || name.chars().count() > 120 { return Err(AppError::BadRequest( "Meal names must be between 1 and 120 characters.".into(), )); } let meal = state .meals .create_meal( name, form.description.trim().to_owned(), parse_category_id(form.category_id), ) .await?; Ok(Redirect::to(&format!("/meals/{}", meal.id)).into_response()) } async fn meal_page( State(state): State, user: CurrentUser, Path(meal_id): Path, ) -> Result { let meal = state .meals .get_meal(meal_id) .await? .ok_or(AppError::NotFound)?; let categories = state.lists.categories().await?; let meal_categories = state.meals.list_meal_categories().await?; Ok(html_response(views::meal_page( &user.session.user, &meal, &categories, &meal_categories, &user.session.csrf_token, ))) } async fn edit_meal( State(state): State, user: CurrentUser, Path(meal_id): Path, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; let name = form.name.trim().to_owned(); if name.is_empty() || name.chars().count() > 120 { return Err(AppError::BadRequest( "Meal names must be between 1 and 120 characters.".into(), )); } state .meals .update_meal( meal_id, name, form.description.trim().to_owned(), parse_category_id(form.category_id), ) .await?; Ok(Redirect::to(&format!("/meals/{meal_id}")).into_response()) } async fn delete_meal( State(state): State, user: CurrentUser, Path(meal_id): Path, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; state.meals.delete_meal(meal_id).await?; Ok(Redirect::to("/meals").into_response()) } async fn add_ingredient( State(state): State, user: CurrentUser, Path(meal_id): Path, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; let name = form.name.trim().to_owned(); if name.is_empty() || name.chars().count() > 120 { return Err(AppError::BadRequest( "Ingredient names must be between 1 and 120 characters.".into(), )); } state .meals .add_ingredient( meal_id, name, form.quantity.trim().to_owned(), form.note.trim().to_owned(), parse_category_id(form.category_id), ) .await?; Ok(Redirect::to(&format!("/meals/{meal_id}")).into_response()) } async fn edit_ingredient( State(state): State, user: CurrentUser, Path((meal_id, ingredient_id)): Path<(i64, i64)>, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; let name = form.name.trim().to_owned(); if name.is_empty() || name.chars().count() > 120 { return Err(AppError::BadRequest( "Ingredient names must be between 1 and 120 characters.".into(), )); } state .meals .update_ingredient( meal_id, ingredient_id, name, form.quantity.trim().to_owned(), form.note.trim().to_owned(), parse_category_id(form.category_id), ) .await?; Ok(Redirect::to(&format!("/meals/{meal_id}")).into_response()) } async fn delete_ingredient( State(state): State, user: CurrentUser, Path((meal_id, ingredient_id)): Path<(i64, i64)>, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; state .meals .delete_ingredient(meal_id, ingredient_id) .await?; Ok(Redirect::to(&format!("/meals/{meal_id}")).into_response()) } async fn add_meal_to_list( State(state): State, user: CurrentUser, Path(list_id): Path, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; require_mutable_list(&state, list_id).await?; state.meals.add_meal_to_list(form.meal_id, list_id).await?; list_fragment_response(&state, &user, list_id).await } async fn remove_meal_from_list( State(state): State, user: CurrentUser, Path((list_id, list_meal_id)): Path<(i64, i64)>, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; require_mutable_list(&state, list_id).await?; state .meals .remove_meal_from_list(list_id, list_meal_id) .await?; list_fragment_response(&state, &user, list_id).await } async fn create_invitation( State(state): State, user: CurrentUser, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; let token = state .invitations .create_invitation(user.session.user.id) .await?; let url = format!( "{}/invite/{token}", state.public_base_url.trim_end_matches('/') ); Ok(html_response(views::invite_result(&url))) } async fn invitation_page( State(state): State, Path(token): Path, headers: HeaderMap, ) -> Result { let info = state.invitations.invitation(token.clone()).await?; if !info { return Err(AppError::NotFound); } let user = optional_user(&state, &headers).await?; Ok(html_response(views::invite_page( user.as_ref().map(|current| ¤t.session.user), &token, None, user.as_ref() .map(|current| current.session.csrf_token.as_str()), ))) } async fn accept_invitation( State(state): State, user: CurrentUser, Path(token): Path, LoggedForm(form): LoggedForm, ) -> Result { verify_csrf(&user, &form.csrf)?; state.invitations.accept_invitation(token).await?; Ok(Redirect::to("/lists").into_response()) } async fn list_stream( State(state): State, user: CurrentUser, Path(list_id): Path, websocket: WebSocketUpgrade, ) -> Result { require_list(&state, list_id).await?; let state_for_socket = state.clone(); let user_for_socket = user.clone(); Ok(websocket .on_upgrade(move |socket| handle_socket(state_for_socket, user_for_socket, list_id, socket)) .into_response()) } async fn handle_socket(state: AppState, user: CurrentUser, list_id: i64, socket: WebSocket) { let subscription = state .realtime .join( list_id, user.session.user.id, user.session.user.display_name.clone(), ) .await; let connection_id = subscription.connection_id.clone(); let (mut sender, mut receiver) = socket.split(); let mut heartbeat = tokio::time::interval(Duration::from_secs(30)); heartbeat.tick().await; match websocket_snapshot(&state, &user, list_id, &subscription.presence).await { Ok(snapshot) => { if sender.send(Message::Text(snapshot.into())).await.is_err() { state.realtime.leave(list_id, &connection_id).await; return; } } Err(error) => { error!(%error, "could not render websocket snapshot"); state.realtime.leave(list_id, &connection_id).await; return; } } let mut events = subscription.receiver; loop { tokio::select! { event = events.recv() => { match event { Ok(HubEvent::ListChanged { list_id: event_list_id, revision }) if event_list_id == list_id => { tracing::debug!(%list_id, revision, "list changed on websocket"); match websocket_list_update(&state, &user, list_id).await { Ok(update) => { if sender.send(Message::Text(update.into())).await.is_err() { break; } } Err(error) => { error!(%error, "could not render websocket list update"); break; } } } Ok(HubEvent::PresenceChanged { list_id: event_list_id }) if event_list_id == list_id => { let presence = state.realtime.presence(list_id).await; let update = views::presence_panel(&presence, true).into_string(); if sender.send(Message::Text(update.into())).await.is_err() { break; } } Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => { match websocket_snapshot(&state, &user, list_id, &state.realtime.presence(list_id).await).await { Ok(snapshot) => { if sender.send(Message::Text(snapshot.into())).await.is_err() { break; } } Err(error) => { error!(%error, "could not resync websocket"); break; } } } Err(tokio::sync::broadcast::error::RecvError::Closed) => break, Ok(_) => {} } } _ = heartbeat.tick() => { if sender.send(Message::Ping(Vec::new().into())).await.is_err() { break; } } incoming = receiver.next() => { match incoming { Some(Ok(Message::Ping(payload))) => { if sender.send(Message::Pong(payload)).await.is_err() { break; } } Some(Ok(Message::Close(_))) | None => break, Some(Ok(_)) => {} Some(Err(_)) => break, } } } } state.realtime.leave(list_id, &connection_id).await; } async fn websocket_snapshot( state: &AppState, user: &CurrentUser, list_id: i64, presence: &[crate::domain::PresenceUser], ) -> Result { let access = require_list(state, list_id).await?; let items = state.lists.items(list_id).await?; let categories = state.lists.categories().await?; let list_meals = state.meals.list_meals_on_list(list_id).await?; Ok(views::live_list_fragments( &access, &items, &categories, &list_meals, &user.session.csrf_token, ) .into_string() + &views::presence_panel(presence, true).into_string()) } async fn websocket_list_update( state: &AppState, user: &CurrentUser, list_id: i64, ) -> Result { let access = require_list(state, list_id).await?; let items = state.lists.items(list_id).await?; let categories = state.lists.categories().await?; let list_meals = state.meals.list_meals_on_list(list_id).await?; Ok(views::live_list_fragments( &access, &items, &categories, &list_meals, &user.session.csrf_token, ) .into_string()) } async fn list_fragment_response( state: &AppState, user: &CurrentUser, list_id: i64, ) -> Result { let access = require_list(state, list_id).await?; let items = state.lists.items(list_id).await?; let categories = state.lists.categories().await?; let list_meals = state.meals.list_meals_on_list(list_id).await?; let editable = access.archived_at.is_none(); Ok(html_response(PreEscaped( views::list_items_fragment( &access, &items, &categories, &user.session.csrf_token, false, editable, ) .into_string() + &views::list_meals_panel( &list_meals, list_id, &user.session.csrf_token, true, editable, ) .into_string(), ))) } async fn require_list( state: &AppState, list_id: i64, ) -> Result { state .lists .get_list(list_id) .await? .ok_or(AppError::NotFound) } /// Like [`require_list`], but also rejects archived lists so they stay /// immutable until restored. async fn require_mutable_list( state: &AppState, list_id: i64, ) -> Result { let list = require_list(state, list_id).await?; if list.archived_at.is_some() { return Err(AppError::Archived); } Ok(list) } async fn optional_user( state: &AppState, headers: &HeaderMap, ) -> Result, AppError> { let Some(session_token) = cookie_value(headers, "session") else { return Ok(None); }; Ok(state .auth .session_user(session_token.clone()) .await? .map(|session| CurrentUser { session_token, session, })) } fn verify_csrf(user: &CurrentUser, token: &str) -> Result<(), AppError> { if token.is_empty() || token != user.session.csrf_token { return Err(AppError::BadRequest( "Your form has expired. Refresh and try again.".into(), )); } Ok(()) } fn parse_category_id(category_id: Option) -> Option { category_id .filter(|category_id| !category_id.trim().is_empty()) .and_then(|category_id| category_id.trim().parse().ok()) } fn html_response(markup: maud::Markup) -> Response { Html(markup.into_string()).into_response() } fn status_html_response(status: StatusCode, markup: maud::Markup) -> Response { (status, Html(markup.into_string())).into_response() } fn cookie_value(headers: &HeaderMap, name: &str) -> Option { headers .get(header::COOKIE)? .to_str() .ok()? .split(';') .map(str::trim) .find_map(|cookie| { let (key, value) = cookie.split_once('=')?; (key == name).then(|| value.to_owned()) }) } fn set_session_cookie(response: &mut Response, token: &str, secure: bool) { let secure_attribute = if secure { "; Secure" } else { "" }; let cookie = format!( "session={token}; Path=/; HttpOnly; SameSite=Lax; Max-Age=2592000{secure_attribute}" ); response.headers_mut().append( header::SET_COOKIE, HeaderValue::from_str(&cookie).expect("session cookie is valid"), ); } fn clear_session_cookie(response: &mut Response, secure: bool) { let secure_attribute = if secure { "; Secure" } else { "" }; let cookie = format!("session=; Path=/; HttpOnly; SameSite=Lax; Max-Age=0{secure_attribute}"); response.headers_mut().append( header::SET_COOKIE, HeaderValue::from_str(&cookie).expect("session cookie is valid"), ); }