security upgrade

This commit is contained in:
Mateusz Gruszczyński
2026-07-30 09:57:43 +02:00
parent 2274cf57c9
commit f6b8e016a8
24 changed files with 1100 additions and 272 deletions
+33 -6
View File
@@ -18,16 +18,33 @@ pub struct AccessTokenRequest {
#[derive(Debug, Serialize)]
pub struct AccessTokenResponse {
access_token: String,
granted: bool,
expires_at: String,
}
pub async fn create_resource_access_token(
State(state): State<SharedState>,
headers: HeaderMap,
Json(payload): Json<AccessTokenRequest>,
) -> Result<Json<AccessTokenResponse>, ApiError> {
) -> Result<Response, ApiError> {
let kind = payload.kind.trim();
let slug = payload.slug.trim();
let client_key = crate::security::client_key(&headers);
let client_limit_key = format!("resource-password-client:{client_key}");
let limit_key = format!("resource-password:{client_key}:{kind}:{slug}");
let window = std::time::Duration::from_secs(15 * 60);
state
.check_rate_limit(client_limit_key, 50, window)
.await
.map_err(|seconds| ApiError::rate_limited(&format!(
"Too many password attempts. Try again in {seconds} seconds."
)))?;
state
.check_rate_limit(limit_key.clone(), 10, window)
.await
.map_err(|seconds| ApiError::rate_limited(&format!(
"Too many password attempts. Try again in {seconds} seconds."
)))?;
match kind {
"workspace" => {
let workspace = db::find_workspace(&state.db, slug)
@@ -63,10 +80,20 @@ pub async fn create_resource_access_token(
.bind(&expires_at)
.execute(state.db.pool())
.await?;
Ok(Json(AccessTokenResponse {
access_token: token,
state.clear_rate_limit(&limit_key).await;
let cookie = crate::security::resource_cookie(
kind,
slug,
&token,
state.anonymous_access_token_ttl_days,
);
let mut response = Json(AccessTokenResponse {
granted: true,
expires_at,
}))
})
.into_response();
response.headers_mut().insert(header::SET_COOKIE, cookie);
Ok(response)
}
pub async fn verify_resource_access_token(
@@ -78,7 +105,7 @@ pub async fn verify_resource_access_token(
let Some(token) = token.map(str::trim).filter(|value| !value.is_empty()) else {
return Ok(false);
};
if crate::auth::resource_permission(state, kind, slug, Some(token))
if crate::auth::share_link_permission(state, kind, slug, Some(token))
.await
.map_err(|error| ApiError::forbidden(&error.message))?
.is_some()
+6
View File
@@ -25,6 +25,12 @@ impl ApiError {
message: message.into(),
}
}
pub(crate) fn rate_limited(message: &str) -> Self {
Self {
status: StatusCode::TOO_MANY_REQUESTS,
message: message.into(),
}
}
pub(crate) fn payload_too_large(max_bytes: usize) -> Self {
let max_mb = max_bytes / (1024 * 1024);
Self {
+77 -12
View File
@@ -16,6 +16,7 @@ pub async fn upload_pad_file(
Path(slug): Path<String>,
mut multipart: Multipart,
) -> Result<Json<serde_json::Value>, ApiError> {
require_authenticated_upload(&state, &headers).await?;
let mut password: Option<String> = None;
let mut access_token: Option<String> = None;
let mut file: Option<(String, Vec<u8>)> = None;
@@ -55,8 +56,9 @@ pub async fn upload_pad_file(
&state,
&slug,
password.as_deref(),
access_token.as_deref(),
resource_request_token(&headers, "pad", &slug, access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
let level = if db::verify_pad_password(&pad, password.as_deref())
@@ -68,7 +70,7 @@ pub async fn upload_pad_file(
&state,
"pad",
&slug,
access_token.as_deref(),
resource_request_token(&headers, "pad", &slug, access_token.as_deref()),
bearer_token(&headers),
)
.await?
@@ -144,8 +146,9 @@ pub async fn pad_files(
&state,
&slug,
payload.password.as_deref(),
payload.access_token.as_deref(),
resource_request_token(&headers, "pad", &slug, payload.access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
let mut files = db::list_pad_files(&state.db, pad.id).await?;
@@ -175,8 +178,9 @@ pub async fn delete_pad_file(
&state,
&slug,
payload.password.as_deref(),
payload.access_token.as_deref(),
resource_request_token(&headers, "pad", &slug, payload.access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
if !crate::auth::is_resource_owner(&state, "pad", &pad.slug, bearer_token(&headers))
@@ -201,6 +205,7 @@ pub async fn upload_note_file(
Path((workspace_slug, note_slug)): Path<(String, String)>,
mut multipart: Multipart,
) -> Result<Json<serde_json::Value>, ApiError> {
require_authenticated_upload(&state, &headers).await?;
let mut password: Option<String> = None;
let mut access_token: Option<String> = None;
let mut file: Option<(String, Vec<u8>)> = None;
@@ -241,8 +246,9 @@ pub async fn upload_note_file(
&workspace_slug,
&note_slug,
password.as_deref(),
access_token.as_deref(),
resource_request_token(&headers, "workspace", &workspace_slug, access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
@@ -255,7 +261,7 @@ pub async fn upload_note_file(
&state,
"workspace",
&workspace_slug,
access_token.as_deref(),
resource_request_token(&headers, "workspace", &workspace_slug, access_token.as_deref()),
bearer_token(&headers),
)
.await?
@@ -309,8 +315,9 @@ pub async fn delete_note(
&workspace_slug,
&note_slug,
payload.password.as_deref(),
payload.access_token.as_deref(),
resource_request_token(&headers, "workspace", &workspace_slug, payload.access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
@@ -323,7 +330,7 @@ pub async fn delete_note(
&state,
"workspace",
&workspace_slug,
payload.access_token.as_deref(),
resource_request_token(&headers, "workspace", &workspace_slug, payload.access_token.as_deref()),
bearer_token(&headers),
)
.await?
@@ -360,8 +367,9 @@ pub async fn note_files(
&workspace_slug,
&note_slug,
payload.password.as_deref(),
payload.access_token.as_deref(),
resource_request_token(&headers, "workspace", &workspace_slug, payload.access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
let mut files = db::list_note_files(&state.db, note.id).await?;
@@ -392,8 +400,9 @@ pub async fn delete_note_file(
&workspace_slug,
&note_slug,
payload.password.as_deref(),
payload.access_token.as_deref(),
resource_request_token(&headers, "workspace", &workspace_slug, payload.access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
let workspace_owner = crate::auth::is_resource_owner(
@@ -429,6 +438,19 @@ pub async fn delete_note_file(
Ok(Json(serde_json::json!({"ok": true})))
}
async fn require_authenticated_upload(
state: &SharedState,
headers: &HeaderMap,
) -> Result<(), ApiError> {
let user = crate::auth::optional_user(state, headers)
.await
.map_err(|error| ApiError::forbidden(&error.message))?;
if user.is_none() {
return Err(ApiError::forbidden("Log in to upload files."));
}
Ok(())
}
pub async fn download_file(
State(state): State<SharedState>,
Path((token, filename)): Path<(String, String)>,
@@ -476,13 +498,29 @@ async fn serve_token_file(
.get_local_with_legacy(&key, &legacy_key)
.await
.map_err(|_| ApiError::not_found_file())?;
let mime = mime_guess::from_path(&safe).first_or_octet_stream();
let guessed_mime = mime_guess::from_path(&safe).first_or_octet_stream();
let inline_image = is_safe_inline_image_mime(guessed_mime.essence_str());
let served_mime = if inline_image {
guessed_mime.as_ref()
} else {
"application/octet-stream"
};
let disposition = if inline_image { "inline" } else { "attachment" };
let mut response = bytes.into_response();
response.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_str(mime.as_ref())
HeaderValue::from_str(served_mime)
.unwrap_or_else(|_| HeaderValue::from_static("application/octet-stream")),
);
response.headers_mut().insert(
header::CONTENT_DISPOSITION,
HeaderValue::from_str(&format!("{disposition}; filename=\"{safe}\""))
.expect("sanitized attachment filename"),
);
response.headers_mut().insert(
HeaderName::from_static("content-security-policy"),
HeaderValue::from_static("default-src 'none'; sandbox"),
);
response.headers_mut().insert(
header::X_CONTENT_TYPE_OPTIONS,
HeaderValue::from_static("nosniff"),
@@ -501,6 +539,19 @@ async fn serve_token_file(
Ok(response)
}
fn is_safe_inline_image_mime(value: &str) -> bool {
matches!(
value,
"image/png"
| "image/jpeg"
| "image/gif"
| "image/webp"
| "image/avif"
| "image/bmp"
| "image/x-icon"
)
}
fn sanitize_filename(value: &str) -> String {
let name = std::path::Path::new(value)
.file_name()
@@ -522,3 +573,17 @@ fn sanitize_filename(value: &str) -> String {
clean.chars().take(160).collect()
}
}
#[cfg(test)]
mod tests {
use super::is_safe_inline_image_mime;
#[test]
fn only_raster_images_are_inline() {
assert!(is_safe_inline_image_mime("image/png"));
assert!(is_safe_inline_image_mime("image/jpeg"));
assert!(!is_safe_inline_image_mime("image/svg+xml"));
assert!(!is_safe_inline_image_mime("text/html"));
assert!(!is_safe_inline_image_mime("application/xml"));
}
}
+138 -39
View File
@@ -40,21 +40,28 @@ const MAX_PASSWORD_LENGTH: usize = 128;
const MIN_WORKSPACE_SLUG_LENGTH: usize = 6;
fn bearer_token(headers: &HeaderMap) -> Option<&str> {
headers
.get(header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.strip_prefix("Bearer "))
.map(str::trim)
.filter(|value| !value.is_empty())
crate::security::session_token(headers)
}
fn authorization_token(headers: &HeaderMap) -> Option<&str> {
crate::security::bearer_token(headers)
}
fn user_session_token(headers: &HeaderMap) -> Option<&str> {
headers
.get("x-rustpad-user-token")
.and_then(|value| value.to_str().ok())
crate::security::session_token(headers)
}
fn resource_request_token<'a>(
headers: &'a HeaderMap,
kind: &str,
slug: &str,
supplied: Option<&'a str>,
) -> Option<&'a str> {
supplied
.map(str::trim)
.filter(|value| !value.is_empty())
.or_else(|| bearer_token(headers))
.filter(|value| !value.is_empty() && *value != "cookie")
.or_else(|| crate::security::resource_token(headers, kind, slug))
.or_else(|| authorization_token(headers))
}
async fn session_user(
@@ -75,14 +82,19 @@ async fn has_write_permission(
kind: &str,
slug: &str,
) -> Result<bool, ApiError> {
let bearer = bearer_token(headers);
if token_access_level(state, kind, slug, bearer).await? >= AccessLevel::Write {
let resource = crate::security::resource_token(headers, kind, slug);
if external_token_access_level(state, kind, slug, resource).await? >= AccessLevel::Write {
return Ok(true);
}
let session = user_session_token(headers);
if session.is_some() && session != bearer {
return Ok(token_access_level(state, kind, slug, session).await? >= AccessLevel::Write);
let authorization = authorization_token(headers);
if authorization != resource
&& external_token_access_level(state, kind, slug, authorization).await? >= AccessLevel::Write
{
return Ok(true);
}
let session = crate::security::session_cookie_token(headers);
if session != resource && session != authorization {
return Ok(account_token_access_level(state, kind, slug, session).await? >= AccessLevel::Write);
}
Ok(false)
}
@@ -476,10 +488,10 @@ pub async fn workspace_info(
.ok_or_else(ApiError::not_found_workspace)?;
ensure_private_resource_access(
&state,
&headers,
"workspace",
&workspace.slug,
workspace.is_private,
bearer_token(&headers),
)
.await?;
Ok(Json(workspace_info_from(&workspace)))
@@ -495,8 +507,9 @@ pub async fn open_workspace(
&state,
&workspace_slug,
payload.password.as_deref(),
payload.access_token.as_deref(),
resource_request_token(&headers, "workspace", &workspace_slug, payload.access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
let stats = db::list_note_stats(&state.db, workspace.id)
@@ -541,8 +554,9 @@ pub async fn create_note(
&state,
&workspace_slug,
payload.password.as_deref(),
payload.access_token.as_deref(),
resource_request_token(&headers, "workspace", &workspace_slug, payload.access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
let level = if db::verify_workspace_password(&workspace, payload.password.as_deref())
@@ -554,7 +568,7 @@ pub async fn create_note(
&state,
"workspace",
&workspace_slug,
payload.access_token.as_deref(),
resource_request_token(&headers, "workspace", &workspace_slug, payload.access_token.as_deref()),
bearer_token(&headers),
)
.await?
@@ -692,10 +706,10 @@ pub async fn note_info(
.ok_or_else(ApiError::not_found_workspace)?;
ensure_private_resource_access(
&state,
&headers,
"workspace",
&workspace.slug,
workspace.is_private,
bearer_token(&headers),
)
.await?;
let note = db::find_note(&state.db, workspace.id, &note_slug)
@@ -807,8 +821,9 @@ pub async fn history(
&workspace_slug,
&note_slug,
payload.password.as_deref(),
payload.access_token.as_deref(),
resource_request_token(&headers, "workspace", &workspace_slug, payload.access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
let _ = workspace;
@@ -834,8 +849,9 @@ pub async fn restore(
&workspace_slug,
&note_slug,
payload.password.as_deref(),
payload.access_token.as_deref(),
resource_request_token(&headers, "workspace", &workspace_slug, payload.access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
let level = if db::verify_workspace_password(&workspace, payload.password.as_deref())
@@ -847,7 +863,7 @@ pub async fn restore(
&state,
"workspace",
&workspace_slug,
payload.access_token.as_deref(),
resource_request_token(&headers, "workspace", &workspace_slug, payload.access_token.as_deref()),
bearer_token(&headers),
)
.await?
@@ -919,13 +935,13 @@ async fn anonymous_access_token_valid(
Ok(count > 0)
}
async fn token_access_level(
async fn external_token_access_level(
state: &SharedState,
kind: &str,
slug: &str,
token: Option<&str>,
) -> Result<AccessLevel, ApiError> {
let permission = crate::auth::resource_permission(state, kind, slug, token)
let permission = crate::auth::share_link_permission(state, kind, slug, token)
.await
.map_err(|error| ApiError::forbidden(&error.message))?;
let level = permission_level(permission.as_deref());
@@ -940,16 +956,28 @@ async fn token_access_level(
Ok(AccessLevel::None)
}
async fn account_token_access_level(
state: &SharedState,
kind: &str,
slug: &str,
token: Option<&str>,
) -> Result<AccessLevel, ApiError> {
let permission = crate::auth::account_resource_permission(state, kind, slug, token)
.await
.map_err(|error| ApiError::forbidden(&error.message))?;
Ok(permission_level(permission.as_deref()))
}
async fn combined_token_access_level(
state: &SharedState,
kind: &str,
slug: &str,
access_token: Option<&str>,
bearer: Option<&str>,
account_token: Option<&str>,
) -> Result<AccessLevel, ApiError> {
Ok(std::cmp::max(
token_access_level(state, kind, slug, access_token).await?,
token_access_level(state, kind, slug, bearer).await?,
external_token_access_level(state, kind, slug, access_token).await?,
account_token_access_level(state, kind, slug, account_token).await?,
))
}
@@ -961,28 +989,89 @@ fn require_write(level: AccessLevel) -> Result<(), ApiError> {
}
}
async fn has_header_resource_access(
state: &SharedState,
headers: &HeaderMap,
kind: &str,
slug: &str,
) -> Result<bool, ApiError> {
for token in [
crate::security::resource_token(headers, kind, slug),
authorization_token(headers),
] {
if external_token_access_level(state, kind, slug, token).await? != AccessLevel::None {
return Ok(true);
}
}
Ok(account_token_access_level(
state,
kind,
slug,
crate::security::session_cookie_token(headers),
)
.await?
!= AccessLevel::None)
}
async fn ensure_private_resource_access(
state: &SharedState,
headers: &HeaderMap,
kind: &str,
slug: &str,
is_private: i64,
token: Option<&str>,
) -> Result<(), ApiError> {
if is_private == 0 {
return Ok(());
}
if verify_resource_access_token(state, kind, slug, token).await? {
if has_header_resource_access(state, headers, kind, slug).await? {
return Ok(());
}
Err(ApiError::not_found_workspace())
}
async fn check_resource_password_attempt(
state: &SharedState,
headers: &HeaderMap,
kind: &str,
slug: &str,
password: Option<&str>,
password_ok: bool,
) -> Result<(), ApiError> {
let Some(_) = password.map(str::trim).filter(|value| !value.is_empty()) else {
return Ok(());
};
let client_key = crate::security::client_key(headers);
let window = std::time::Duration::from_secs(15 * 60);
state
.check_rate_limit(format!("resource-password-client:{client_key}"), 50, window)
.await
.map_err(|seconds| {
ApiError::rate_limited(&format!(
"Too many password attempts. Try again in {seconds} seconds."
))
})?;
let limit_key = format!("resource-password:{client_key}:{kind}:{slug}");
state
.check_rate_limit(limit_key.clone(), 10, window)
.await
.map_err(|seconds| {
ApiError::rate_limited(&format!(
"Too many password attempts. Try again in {seconds} seconds."
))
})?;
if password_ok {
state.clear_rate_limit(&limit_key).await;
}
Ok(())
}
pub async fn authorized_workspace(
state: &SharedState,
slug: &str,
password: Option<&str>,
access_token: Option<&str>,
bearer: Option<&str>,
headers: &HeaderMap,
) -> Result<db::Workspace, ApiError> {
let workspace = db::find_workspace(&state.db, slug)
.await?
@@ -992,11 +1081,13 @@ pub async fn authorized_workspace(
if workspace.is_private != 0 && token_level == AccessLevel::None {
return Err(ApiError::not_found_workspace());
}
if workspace.password_hash.is_some()
&& !db::verify_workspace_password(&workspace, password)
&& token_level == AccessLevel::None
{
return Err(ApiError::unauthorized());
if workspace.password_hash.is_some() && token_level < AccessLevel::Write {
let password_ok = db::verify_workspace_password(&workspace, password);
check_resource_password_attempt(state, headers, "workspace", slug, password, password_ok)
.await?;
if token_level == AccessLevel::None && !password_ok {
return Err(ApiError::unauthorized());
}
}
Ok(workspace)
}
@@ -1008,9 +1099,17 @@ async fn authorized_note(
password: Option<&str>,
access_token: Option<&str>,
bearer: Option<&str>,
headers: &HeaderMap,
) -> Result<(db::Workspace, db::Note), ApiError> {
let workspace =
authorized_workspace(state, workspace_slug, password, access_token, bearer).await?;
let workspace = authorized_workspace(
state,
workspace_slug,
password,
access_token,
bearer,
headers,
)
.await?;
let note = db::find_note(&state.db, workspace.id, note_slug)
.await?
.ok_or_else(ApiError::not_found_note)?;
+47 -18
View File
@@ -94,10 +94,10 @@ pub async fn pad_info(
.ok_or_else(ApiError::not_found_note)?;
ensure_private_resource_access(
&state,
&headers,
"pad",
&pad.slug,
pad.is_private,
bearer_token(&headers),
)
.await?;
let (global_color, note_color) = editor_colors(&state, &headers, "pad", &slug).await?;
@@ -238,8 +238,9 @@ pub async fn publish_pad_page(
&state,
&slug,
payload.password.as_deref(),
payload.access_token.as_deref(),
resource_request_token(&headers, "pad", &slug, payload.access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
let level = if db::verify_pad_password(&pad, payload.password.as_deref())
@@ -251,7 +252,7 @@ pub async fn publish_pad_page(
&state,
"pad",
&slug,
payload.access_token.as_deref(),
resource_request_token(&headers, "pad", &slug, payload.access_token.as_deref()),
bearer_token(&headers),
)
.await?
@@ -287,8 +288,14 @@ pub async fn publish_note_page(
&workspace_slug,
&note_slug,
payload.password.as_deref(),
payload.access_token.as_deref(),
resource_request_token(
&headers,
"workspace",
&workspace_slug,
payload.access_token.as_deref(),
),
bearer_token(&headers),
&headers,
)
.await?;
let level = if db::verify_workspace_password(&workspace, payload.password.as_deref())
@@ -300,7 +307,12 @@ pub async fn publish_note_page(
&state,
"workspace",
&workspace_slug,
payload.access_token.as_deref(),
resource_request_token(
&headers,
"workspace",
&workspace_slug,
payload.access_token.as_deref(),
),
bearer_token(&headers),
)
.await?
@@ -339,7 +351,6 @@ async fn ensure_public_page_access(
page: &db::PublishedPage,
) -> Result<(), ApiError> {
let password = page_password(headers);
let bearer = bearer_token(headers);
if let Some(pad_id) = page.pad_id {
if db::pad_public_page_unprotected(&state.db, pad_id).await? {
return Ok(());
@@ -358,10 +369,13 @@ async fn ensure_public_page_access(
let pad = db::find_pad(&state.db, &slug)
.await?
.ok_or_else(ApiError::not_found_note)?;
if db::verify_pad_password(&pad, password) {
if has_header_resource_access(state, headers, "pad", &slug).await? {
return Ok(());
}
if verify_resource_access_token(state, "pad", &slug, bearer).await? {
let password_ok = db::verify_pad_password(&pad, password);
check_resource_password_attempt(state, headers, "pad", &slug, password, password_ok)
.await?;
if password_ok {
return Ok(());
}
return if pad.password_hash.is_some() {
@@ -392,10 +406,20 @@ async fn ensure_public_page_access(
let workspace = db::find_workspace(&state.db, &slug)
.await?
.ok_or_else(ApiError::not_found_workspace)?;
if db::verify_workspace_password(&workspace, password) {
if has_header_resource_access(state, headers, "workspace", &slug).await? {
return Ok(());
}
if verify_resource_access_token(state, "workspace", &slug, bearer).await? {
let password_ok = db::verify_workspace_password(&workspace, password);
check_resource_password_attempt(
state,
headers,
"workspace",
&slug,
password,
password_ok,
)
.await?;
if password_ok {
return Ok(());
}
return if workspace.password_hash.is_some() {
@@ -476,8 +500,9 @@ pub async fn pad_history(
&state,
&slug,
payload.password.as_deref(),
payload.access_token.as_deref(),
resource_request_token(&headers, "pad", &slug, payload.access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
let revisions = db::list_pad_revisions(&state.db, pad.id)
@@ -501,8 +526,9 @@ pub async fn pad_restore(
&state,
&slug,
payload.password.as_deref(),
payload.access_token.as_deref(),
resource_request_token(&headers, "pad", &slug, payload.access_token.as_deref()),
bearer_token(&headers),
&headers,
)
.await?;
let level = if db::verify_pad_password(&pad, payload.password.as_deref())
@@ -514,7 +540,7 @@ pub async fn pad_restore(
&state,
"pad",
&slug,
payload.access_token.as_deref(),
resource_request_token(&headers, "pad", &slug, payload.access_token.as_deref()),
bearer_token(&headers),
)
.await?
@@ -555,6 +581,7 @@ pub(super) async fn authorized_pad(
password: Option<&str>,
access_token: Option<&str>,
bearer: Option<&str>,
headers: &HeaderMap,
) -> Result<db::Pad, ApiError> {
let pad = db::find_pad(&state.db, slug)
.await?
@@ -563,11 +590,13 @@ pub(super) async fn authorized_pad(
if pad.is_private != 0 && token_level == AccessLevel::None {
return Err(ApiError::not_found_note());
}
if pad.password_hash.is_some()
&& !db::verify_pad_password(&pad, password)
&& token_level == AccessLevel::None
{
return Err(ApiError::forbidden("Password required or incorrect."));
if pad.password_hash.is_some() && token_level < AccessLevel::Write {
let password_ok = db::verify_pad_password(&pad, password);
check_resource_password_attempt(state, headers, "pad", slug, password, password_ok)
.await?;
if token_level == AccessLevel::None && !password_ok {
return Err(ApiError::forbidden("Password required or incorrect."));
}
}
Ok(pad)
}
+161 -65
View File
@@ -17,7 +17,8 @@ use argon2::{
use axum::{
Json,
extract::{Path as AxumPath, State},
http::{HeaderMap, StatusCode},
http::{HeaderMap, StatusCode, header},
response::{IntoResponse, Response},
};
use chrono::{Duration, Utc};
use lettre::{
@@ -53,8 +54,6 @@ pub struct User {
#[derive(Deserialize)]
pub struct IdentityRequest {
nickname: String,
#[serde(default)]
session_token: Option<String>,
}
#[derive(Deserialize)]
pub struct RegisterRequest {
@@ -259,6 +258,7 @@ pub struct ResourceList {
}
#[derive(Serialize)]
pub struct SessionResponse {
#[serde(skip_serializing)]
token: String,
nickname: String,
email: String,
@@ -276,6 +276,7 @@ pub struct IdentityResponse {
}
#[derive(Serialize)]
pub struct RegisterResponse {
#[serde(skip_serializing)]
token: Option<String>,
nickname: String,
email: String,
@@ -286,10 +287,12 @@ pub struct RegisterResponse {
pub async fn identity(
State(state): State<SharedState>,
headers: HeaderMap,
Json(req): Json<IdentityRequest>,
) -> Result<Json<IdentityResponse>, AuthError> {
let nickname = validate_nickname(&req.nickname)?;
debug!(nickname = %nickname, has_session = req.session_token.is_some(), "identity check requested");
let request_token = crate::security::session_token(&headers);
debug!(nickname = %nickname, has_session = request_token.is_some(), "identity check requested");
match find_user_by_nickname(&state, &nickname).await? {
None => {
debug!(nickname = %nickname, "nickname is available for guest use");
@@ -299,7 +302,7 @@ pub async fn identity(
}))
}
Some(user) => {
let token = req.session_token.as_deref().ok_or_else(|| {
let token = request_token.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(|| {
@@ -322,7 +325,7 @@ pub async fn identity(
pub async fn register(
State(state): State<SharedState>,
Json(req): Json<RegisterRequest>,
) -> Result<(StatusCode, Json<RegisterResponse>), AuthError> {
) -> Result<Response, AuthError> {
if state.ldap.is_some() {
return Err(AuthError::forbidden(
"Local registration is disabled while LDAP authentication is enabled.",
@@ -434,12 +437,13 @@ pub async fn register(
"Account created. Check your e-mail and confirm the account before logging in."
.into(),
}),
));
).into_response());
}
let session = create_session(&state, &user).await?;
info!(user_id = user.id, nickname = %user.nickname, "account registered and session created");
Ok((
let cookie = crate::security::session_cookie(&session.token, state.user_session_ttl_days);
let mut response = (
StatusCode::CREATED,
Json(RegisterResponse {
token: Some(session.token),
@@ -449,22 +453,44 @@ pub async fn register(
confirmation_required: false,
message: "Account created.".into(),
}),
))
).into_response();
response.headers_mut().insert(header::SET_COOKIE, cookie);
Ok(response)
}
pub async fn login(
State(state): State<SharedState>,
headers: HeaderMap,
Json(req): Json<LoginRequest>,
) -> Result<Json<SessionResponse>, AuthError> {
if state.ldap.is_some() {
ldap::login(&state, &req.email, &req.password)
.await
.map(Json)
) -> Result<Response, AuthError> {
let client_key = crate::security::client_key(&headers);
let login_key: String = normalize(&req.email).chars().take(320).collect();
let client_limit_key = format!("login-client:{client_key}");
let limit_key = format!("login:{client_key}:{login_key}");
let window = std::time::Duration::from_secs(15 * 60);
state
.check_rate_limit(client_limit_key.clone(), 30, window)
.await
.map_err(|seconds| AuthError::rate_limited(&format!(
"Too many login attempts. Try again in {seconds} seconds."
)))?;
state
.check_rate_limit(limit_key.clone(), 5, window)
.await
.map_err(|seconds| AuthError::rate_limited(&format!(
"Too many login attempts. Try again in {seconds} seconds."
)))?;
let session = if state.ldap.is_some() {
ldap::login(&state, &req.email, &req.password).await?
} else {
local::login(&state, &req.email, &req.password)
.await
.map(Json)
}
local::login(&state, &req.email, &req.password).await?
};
state.clear_rate_limit(&limit_key).await;
Ok(session_json_response(
StatusCode::OK,
session,
state.user_session_ttl_days,
))
}
pub async fn resend_confirmation(
@@ -673,11 +699,16 @@ fn suggested_directory_nickname(display_name: Option<&str>, email: &str) -> Opti
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."))?;
) -> Result<Response, AuthError> {
let token = crate::security::session_token(&headers)
.ok_or_else(|| AuthError::unauthorized("Not logged in."))?;
let Some(user) = user_from_token(&state, token).await? else {
let mut response = AuthError::unauthorized("Your session has expired.").into_response();
response
.headers_mut()
.insert(header::SET_COOKIE, crate::security::clear_session_cookie());
return Ok(response);
};
debug!(user_id = user.id, "session validation successful");
let expires_at: String = sqlx::query_scalar(queries::get(
state.db.kind(),
@@ -697,17 +728,21 @@ pub async fn me(
.fetch_one(state.db.pool())
.await
.map_err(AuthError::database)?;
Ok(Json(SessionResponse {
token: token.into(),
nickname: user.nickname,
email: user.email,
expires_at,
directory_managed,
directory_display_name,
directory_organization,
suggested_nickname,
editor_color,
}))
Ok(session_json_response(
StatusCode::OK,
SessionResponse {
token: token.into(),
nickname: user.nickname,
email: user.email,
expires_at,
directory_managed,
directory_display_name,
directory_organization,
suggested_nickname,
editor_color,
},
state.user_session_ttl_days,
))
}
pub async fn update_profile(
@@ -1225,7 +1260,7 @@ pub async fn optional_user(
state: &SharedState,
headers: &HeaderMap,
) -> Result<Option<User>, AuthError> {
match bearer(headers) {
match crate::security::session_token(headers) {
Some(token) => user_from_token(state, token).await,
None => Ok(None),
}
@@ -1672,32 +1707,43 @@ pub async fn is_resource_owner(
Ok(ensure_owner(state, user.id, kind, slug).await.is_ok())
}
pub async fn resource_permission(
pub async fn account_resource_permission(
state: &SharedState,
kind: &str,
slug: &str,
token: Option<&str>,
) -> Result<Option<String>, AuthError> {
let Some(token) = token.filter(|v| !v.is_empty()) else {
let Some(token) = token.filter(|value| !value.is_empty()) else {
return Ok(None);
};
if let Some(user) = user_from_token(state, token).await? {
let owns = ensure_owner(state, user.id, kind, slug).await.is_ok();
if owns {
return Ok(Some("rw".into()));
}
let permission: Option<String> = sqlx::query_scalar(queries::get(
state.db.kind(),
queries::RESOURCE_PERMISSION_BY_USER,
))
.bind(kind)
.bind(slug)
.bind(user.id)
.fetch_optional(state.db.pool())
.await
.map_err(AuthError::database)?;
return Ok(permission);
let Some(user) = user_from_token(state, token).await? else {
return Ok(None);
};
if ensure_owner(state, user.id, kind, slug).await.is_ok() {
return Ok(Some("rw".into()));
}
let permission: Option<String> = sqlx::query_scalar(queries::get(
state.db.kind(),
queries::RESOURCE_PERMISSION_BY_USER,
))
.bind(kind)
.bind(slug)
.bind(user.id)
.fetch_optional(state.db.pool())
.await
.map_err(AuthError::database)?;
Ok(permission)
}
pub async fn share_link_permission(
state: &SharedState,
kind: &str,
slug: &str,
token: Option<&str>,
) -> Result<Option<String>, AuthError> {
let Some(token) = token.filter(|value| !value.is_empty()) else {
return Ok(None);
};
let now = Utc::now().to_rfc3339();
let permission: Option<String> = sqlx::query_scalar(queries::get(
state.db.kind(),
@@ -1716,8 +1762,8 @@ pub async fn resource_permission(
pub async fn logout(
State(state): State<SharedState>,
headers: HeaderMap,
) -> Result<Json<serde_json::Value>, AuthError> {
if let Some(token) = bearer(&headers) {
) -> Result<Response, AuthError> {
if let Some(token) = crate::security::session_token(&headers) {
let result = sqlx::query(queries::get(
state.db.kind(),
queries::AUTH_DELETE_SESSION_BY_TOKEN,
@@ -1730,14 +1776,37 @@ pub async fn logout(
} else {
debug!("logout requested without an active session");
}
Ok(Json(serde_json::json!({"ok": true})))
let mut response = Json(serde_json::json!({"ok": true})).into_response();
response
.headers_mut()
.insert(header::SET_COOKIE, crate::security::clear_session_cookie());
Ok(response)
}
pub async fn request_reset(
State(state): State<SharedState>,
headers: HeaderMap,
Json(req): Json<ResetRequest>,
) -> Result<Json<serde_json::Value>, AuthError> {
let email = validate_email(&req.email)?;
let client_key = crate::security::client_key(&headers);
let window = std::time::Duration::from_secs(60 * 60);
state
.check_rate_limit(format!("password-reset-client:{client_key}"), 10, window)
.await
.map_err(|seconds| AuthError::rate_limited(&format!(
"Too many password reset requests. Try again in {seconds} seconds."
)))?;
state
.check_rate_limit(
format!("password-reset:{client_key}:{}", normalize(&email)),
3,
window,
)
.await
.map_err(|seconds| AuthError::rate_limited(&format!(
"Too many password reset requests. Try again in {seconds} seconds."
)))?;
info!(email_domain = %email_domain(&email), "password reset requested");
let smtp = state.smtp.as_ref().ok_or_else(|| {
AuthError::service_unavailable("Password reset is not configured on this server.")
@@ -1777,13 +1846,34 @@ pub async fn request_reset(
pub async fn confirm_reset(
State(state): State<SharedState>,
headers: HeaderMap,
Json(req): Json<ResetConfirmRequest>,
) -> Result<Json<serde_json::Value>, AuthError> {
validate_password(&req.password)?;
let token_fingerprint = hash_token(req.token.trim());
let client_key = crate::security::client_key(&headers);
let client_limit_key = format!("password-reset-confirm-client:{client_key}");
let limit_key = format!(
"password-reset-confirm:{client_key}:{}",
&token_fingerprint[..16.min(token_fingerprint.len())]
);
let window = std::time::Duration::from_secs(15 * 60);
state
.check_rate_limit(client_limit_key.clone(), 20, window)
.await
.map_err(|seconds| AuthError::rate_limited(&format!(
"Too many reset attempts. Try again in {seconds} seconds."
)))?;
state
.check_rate_limit(limit_key.clone(), 10, window)
.await
.map_err(|seconds| AuthError::rate_limited(&format!(
"Too many reset attempts. Try again in {seconds} seconds."
)))?;
info!("password reset confirmation requested");
let now_time = Utc::now();
let now = now_time.to_rfc3339();
let token_hash = hash_token(req.token.trim());
let token_hash = token_fingerprint;
let token_row = sqlx::query(queries::get(
state.db.kind(),
queries::AUTH_FIND_RESET_TOKEN,
@@ -1855,6 +1945,7 @@ pub async fn confirm_reset(
.await
.map_err(AuthError::database)?;
tx.commit().await.map_err(AuthError::database)?;
state.clear_rate_limit(&limit_key).await;
info!(
user_id,
"password reset completed and existing sessions revoked"
@@ -2103,14 +2194,6 @@ fn hash_token(value: &str) -> String {
format!("{:x}", Sha256::digest(value.as_bytes()))
}
fn bearer(headers: &HeaderMap) -> Option<&str> {
headers
.get("authorization")?
.to_str()
.ok()?
.strip_prefix("Bearer ")
}
async fn send_registration_email(
smtp: &SmtpConfig,
user: &User,
@@ -2407,6 +2490,12 @@ impl AuthError {
message: m.into(),
}
}
fn rate_limited(m: &str) -> Self {
Self {
status: StatusCode::TOO_MANY_REQUESTS,
message: m.into(),
}
}
fn service_unavailable(m: &str) -> Self {
Self {
status: StatusCode::SERVICE_UNAVAILABLE,
@@ -2419,6 +2508,13 @@ impl AuthError {
}
}
fn session_json_response(status: StatusCode, session: SessionResponse, ttl_days: i64) -> Response {
let cookie = crate::security::session_cookie(&session.token, ttl_days);
let mut response = (status, Json(session)).into_response();
response.headers_mut().insert(header::SET_COOKIE, cookie);
response
}
fn email_domain(email: &str) -> &str {
email
.rsplit_once('@')
+1
View File
@@ -17,6 +17,7 @@ mod database;
mod db;
mod queries;
mod row_decode;
mod security;
mod state;
mod storage;
mod websocket;
+226
View File
@@ -0,0 +1,226 @@
use axum::http::{HeaderMap, HeaderValue, Uri, header};
use sha2::{Digest, Sha256};
pub const SESSION_COOKIE: &str = "__Host-rustpad_session";
pub fn session_token(headers: &HeaderMap) -> Option<&str> {
session_cookie_token(headers)
}
pub fn session_cookie_token(headers: &HeaderMap) -> Option<&str> {
cookie_value(headers, SESSION_COOKIE)
}
pub fn bearer_token(headers: &HeaderMap) -> Option<&str> {
headers
.get(header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.strip_prefix("Bearer "))
.map(str::trim)
.filter(|value| !value.is_empty())
}
pub fn resource_token<'a>(headers: &'a HeaderMap, kind: &str, slug: &str) -> Option<&'a str> {
let name = resource_cookie_name(kind, slug);
cookie_value(headers, &name)
}
pub fn session_cookie(token: &str, ttl_days: i64) -> HeaderValue {
secure_cookie(SESSION_COOKIE, token, ttl_days.saturating_mul(86_400))
}
pub fn clear_session_cookie() -> HeaderValue {
clear_cookie(SESSION_COOKIE)
}
pub fn resource_cookie(kind: &str, slug: &str, token: &str, ttl_days: i64) -> HeaderValue {
secure_cookie(
&resource_cookie_name(kind, slug),
token,
ttl_days.saturating_mul(86_400),
)
}
pub fn client_key(headers: &HeaderMap) -> String {
let forwarded_ip = header_ip(headers, "cf-connecting-ip")
.or_else(|| header_ip(headers, "x-real-ip"))
.or_else(|| header_ip(headers, "x-forwarded-for"));
if let Some(value) = forwarded_ip {
return format!("ip:{value}");
}
let user_agent = first_header_value(headers, header::USER_AGENT.as_str()).unwrap_or("");
let language = first_header_value(headers, header::ACCEPT_LANGUAGE.as_str()).unwrap_or("");
if user_agent.is_empty() && language.is_empty() {
return "unknown".into();
}
let digest = Sha256::digest(format!("{user_agent}|{language}").as_bytes());
format!("browser:{}", hex::encode(&digest[..12]))
}
pub fn websocket_origin_allowed(headers: &HeaderMap) -> bool {
let Some(origin) = headers
.get(header::ORIGIN)
.and_then(|value| value.to_str().ok())
.map(str::trim)
.filter(|value| !value.is_empty() && *value != "null")
else {
return false;
};
let Ok(uri) = origin.parse::<Uri>() else {
return false;
};
let Some(origin_authority) = uri.authority().map(|value| value.as_str()) else {
return false;
};
if uri.path() != "/" || uri.query().is_some() {
return false;
}
let Some(expected_authority) = first_header_value(headers, header::HOST.as_str()) else {
return false;
};
if !origin_authority.eq_ignore_ascii_case(expected_authority) {
return false;
}
if let Some(expected_scheme) = first_header_value(headers, "x-forwarded-proto") {
let Some(origin_scheme) = uri.scheme_str() else {
return false;
};
if !origin_scheme.eq_ignore_ascii_case(expected_scheme) {
return false;
}
}
matches!(uri.scheme_str(), Some("http" | "https"))
}
fn resource_cookie_name(kind: &str, slug: &str) -> String {
let digest = Sha256::digest(format!("{kind}:{slug}").as_bytes());
format!("__Host-rustpad_access_{}", hex::encode(&digest[..12]))
}
fn cookie_value<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> {
headers
.get(header::COOKIE)
.and_then(|value| value.to_str().ok())
.and_then(|cookies| {
cookies.split(';').find_map(|part| {
let (cookie_name, cookie_value) = part.trim().split_once('=')?;
(cookie_name == name && !cookie_value.is_empty()).then_some(cookie_value)
})
})
}
fn secure_cookie(name: &str, value: &str, max_age: i64) -> HeaderValue {
HeaderValue::from_str(&format!(
"{name}={value}; Path=/; Max-Age={}; HttpOnly; Secure; SameSite=Lax",
max_age.max(1)
))
.expect("valid secure cookie")
}
fn clear_cookie(name: &str) -> HeaderValue {
HeaderValue::from_str(&format!(
"{name}=; Path=/; Max-Age=0; HttpOnly; Secure; SameSite=Lax"
))
.expect("valid clear cookie")
}
fn header_ip(headers: &HeaderMap, name: &str) -> Option<std::net::IpAddr> {
first_header_value(headers, name)?.parse().ok()
}
fn first_header_value<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> {
headers
.get(name)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.split(',').next())
.map(str::trim)
.filter(|value| !value.is_empty())
}
#[cfg(test)]
mod tests {
use super::*;
fn websocket_headers(origin: &'static str, host: &'static str) -> HeaderMap {
let mut headers = HeaderMap::new();
headers.insert(header::ORIGIN, HeaderValue::from_static(origin));
headers.insert(header::HOST, HeaderValue::from_static(host));
headers
}
#[test]
fn account_sessions_are_cookie_only() {
let mut headers = HeaderMap::new();
headers.insert(
header::AUTHORIZATION,
HeaderValue::from_static("Bearer legacy-account-token"),
);
assert_eq!(session_token(&headers), None);
assert_eq!(bearer_token(&headers), Some("legacy-account-token"));
headers.insert(
header::COOKIE,
HeaderValue::from_static("__Host-rustpad_session=cookie-token"),
);
assert_eq!(session_token(&headers), Some("cookie-token"));
}
#[test]
fn prefers_proxy_controlled_real_ip() {
let mut headers = HeaderMap::new();
headers.insert(
axum::http::HeaderName::from_static("x-forwarded-for"),
HeaderValue::from_static("203.0.113.10"),
);
headers.insert(
axum::http::HeaderName::from_static("x-real-ip"),
HeaderValue::from_static("198.51.100.20"),
);
assert_eq!(client_key(&headers), "ip:198.51.100.20");
}
#[test]
fn accepts_same_origin_websocket() {
let headers = websocket_headers("https://pad.example.com", "pad.example.com");
assert!(websocket_origin_allowed(&headers));
}
#[test]
fn rejects_cross_origin_websocket() {
let headers = websocket_headers("https://evil.example", "pad.example.com");
assert!(!websocket_origin_allowed(&headers));
}
#[test]
fn does_not_trust_forwarded_host_for_websocket_origin() {
let mut headers = websocket_headers("https://evil.example", "pad.example.com");
headers.insert(
axum::http::HeaderName::from_static("x-forwarded-host"),
HeaderValue::from_static("evil.example"),
);
assert!(!websocket_origin_allowed(&headers));
}
#[test]
fn rejects_origin_with_path() {
let headers = websocket_headers("https://pad.example.com/other", "pad.example.com");
assert!(!websocket_origin_allowed(&headers));
}
#[test]
fn rejects_missing_websocket_origin() {
let mut headers = HeaderMap::new();
headers.insert(header::HOST, HeaderValue::from_static("pad.example.com"));
assert!(!websocket_origin_allowed(&headers));
}
#[test]
fn secure_cookies_are_not_script_readable() {
let value = session_cookie("abc123", 7).to_str().unwrap();
assert!(value.contains("HttpOnly"));
assert!(value.contains("Secure"));
assert!(value.contains("SameSite=Lax"));
assert!(value.starts_with("__Host-rustpad_session=abc123;"));
}
}
+56 -1
View File
@@ -15,10 +15,19 @@ use std::{
Arc,
atomic::{AtomicU64, Ordering},
},
time::{Duration, Instant},
};
use tokio::sync::{RwLock, broadcast};
use tokio::sync::{Mutex, RwLock, broadcast};
const CHANNEL_CAPACITY: usize = 256;
const MAX_RATE_LIMIT_ENTRIES: usize = 50_000;
#[derive(Debug)]
struct RateLimitEntry {
started_at: Instant,
attempts: u32,
window: Duration,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SmtpSecurity {
@@ -100,6 +109,7 @@ pub struct AppState {
channels: RwLock<HashMap<String, broadcast::Sender<RoomEvent>>>,
presence: RwLock<HashMap<String, HashMap<u64, PresenceConnection>>>,
next_connection_id: AtomicU64,
rate_limits: Mutex<HashMap<String, RateLimitEntry>>,
}
impl AppState {
@@ -137,8 +147,53 @@ impl AppState {
channels: RwLock::new(HashMap::new()),
presence: RwLock::new(HashMap::new()),
next_connection_id: AtomicU64::new(1),
rate_limits: Mutex::new(HashMap::new()),
}
}
pub async fn check_rate_limit(
&self,
key: String,
max_attempts: u32,
window: Duration,
) -> Result<(), u64> {
let now = Instant::now();
let mut limits = self.rate_limits.lock().await;
if !limits.contains_key(&key) && limits.len() >= MAX_RATE_LIMIT_ENTRIES {
limits.retain(|_, entry| now.duration_since(entry.started_at) < entry.window);
if limits.len() >= MAX_RATE_LIMIT_ENTRIES {
if let Some(oldest_key) = limits
.iter()
.min_by_key(|(_, entry)| entry.started_at)
.map(|(key, _)| key.clone())
{
limits.remove(&oldest_key);
}
}
}
let entry = limits.entry(key).or_insert(RateLimitEntry {
started_at: now,
attempts: 0,
window,
});
if now.duration_since(entry.started_at) >= entry.window {
entry.started_at = now;
entry.attempts = 0;
}
entry.window = window;
if entry.attempts >= max_attempts {
let remaining = entry
.window
.saturating_sub(now.duration_since(entry.started_at));
return Err(remaining.as_secs().max(1));
}
entry.attempts += 1;
Ok(())
}
pub async fn clear_rate_limit(&self, key: &str) {
self.rate_limits.lock().await.remove(key);
}
async fn channel_for_key(&self, key: String) -> broadcast::Sender<RoomEvent> {
if let Some(sender) = self.channels.read().await.get(&key) {
return sender.clone();
+114 -16
View File
@@ -7,7 +7,8 @@ use axum::{
Path, State, WebSocketUpgrade,
ws::{Message, WebSocket},
},
response::Response,
http::{HeaderMap, StatusCode},
response::{IntoResponse, Response},
};
use futures_util::{SinkExt, StreamExt};
use serde::{Deserialize, Serialize};
@@ -25,7 +26,6 @@ enum ClientMessage {
password: Option<String>,
access_token: Option<String>,
nickname: Option<String>,
session_token: Option<String>,
guest_id: Option<String>,
color: Option<String>,
},
@@ -76,13 +76,69 @@ enum ServerMessage {
},
}
async fn resource_permission_from_tokens(
state: &SharedState,
kind: &str,
slug: &str,
access_token: Option<&str>,
session_token: Option<&str>,
) -> Option<String> {
let access_permission = auth::share_link_permission(state, kind, slug, access_token)
.await
.ok()
.flatten();
let session_permission = auth::account_resource_permission(state, kind, slug, session_token)
.await
.ok()
.flatten();
if access_permission.as_deref() == Some("rw") || session_permission.as_deref() == Some("rw") {
Some("rw".into())
} else if access_permission.as_deref() == Some("ro")
|| session_permission.as_deref() == Some("ro")
{
Some("ro".into())
} else {
None
}
}
async fn anonymous_access_from_tokens(
state: &SharedState,
kind: &str,
slug: &str,
access_token: Option<&str>,
) -> bool {
crate::api::verify_resource_access_token(state, kind, slug, access_token)
.await
.unwrap_or(false)
}
// Merged from note.rs
pub async fn upgrade(
ws: WebSocketUpgrade,
headers: HeaderMap,
Path((workspace_slug, note_slug)): Path<(String, String)>,
State(state): State<SharedState>,
) -> Response {
ws.on_upgrade(move |socket| handle_socket(socket, state, workspace_slug, note_slug))
if !crate::security::websocket_origin_allowed(&headers) {
warn!(%workspace_slug, %note_slug, "note websocket rejected: invalid origin");
return (StatusCode::FORBIDDEN, "Invalid WebSocket origin").into_response();
}
let account_token = crate::security::session_token(&headers).map(str::to_owned);
let resource_token = crate::security::resource_token(&headers, "workspace", &workspace_slug)
.map(str::to_owned);
let client_key = crate::security::client_key(&headers);
ws.on_upgrade(move |socket| {
handle_socket(
socket,
state,
workspace_slug,
note_slug,
account_token,
resource_token,
client_key,
)
})
}
async fn handle_socket(
@@ -90,6 +146,9 @@ async fn handle_socket(
state: SharedState,
workspace_slug: String,
note_slug: String,
cookie_session_token: Option<String>,
cookie_access_token: Option<String>,
client_key: String,
) {
info!(%workspace_slug, %note_slug, "note websocket connected");
let Some(workspace) = db::find_workspace(&state.db, &workspace_slug)
@@ -110,21 +169,19 @@ async fn handle_socket(
let _ = send_error(&mut socket, "Note not found").await;
return;
};
let (password, access_token, nickname, session_token, guest_id, color) =
let (password, access_token, nickname, guest_id, color) =
match socket.recv().await {
Some(Ok(Message::Text(text))) => match serde_json::from_str::<ClientMessage>(&text) {
Ok(ClientMessage::Authenticate {
password,
access_token,
nickname,
session_token,
guest_id,
color,
}) => (
password,
access_token,
clean_nickname(nickname),
session_token,
clean_guest_id(guest_id),
clean_color(color),
),
@@ -135,6 +192,13 @@ async fn handle_socket(
},
_ => return,
};
let session_token = cookie_session_token;
let explicit_access_token = access_token
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty() && *value != "cookie")
.map(str::to_owned);
let access_token = explicit_access_token.or(cookie_access_token);
let nickname = match auth::authorize_nickname(&state, nickname, session_token.clone()).await {
Ok(value) => value,
Err(message) => {
@@ -154,22 +218,56 @@ async fn handle_socket(
.map(|name| format!("guest:{id}:{}", name.to_lowercase()))
}),
};
let supplied_token = session_token.as_deref().or(access_token.as_deref());
let permission =
auth::resource_permission(&state, "workspace", &workspace_slug, supplied_token)
.await
.ok()
.flatten();
let permission = resource_permission_from_tokens(
&state,
"workspace",
&workspace_slug,
access_token.as_deref(),
session_token.as_deref(),
)
.await;
let anonymous_token_ok = permission.is_none()
&& crate::api::verify_resource_access_token(
&& anonymous_access_from_tokens(
&state,
"workspace",
&workspace_slug,
supplied_token,
access_token.as_deref(),
)
.await
.unwrap_or(false);
.await;
let password_limit_key = format!("resource-password:{client_key}:workspace:{workspace_slug}");
let password_attempted = password
.as_deref()
.map(str::trim)
.is_some_and(|value| !value.is_empty());
if workspace.password_hash.is_some()
&& permission.as_deref() != Some("rw")
&& !anonymous_token_ok
&& password_attempted
{
let window = std::time::Duration::from_secs(15 * 60);
if let Err(seconds) = state
.check_rate_limit(format!("resource-password-client:{client_key}"), 50, window)
.await
{
let _ = send_error(&mut socket, &format!("Too many password attempts. Try again in {seconds} seconds.")).await;
return;
}
if let Err(seconds) = state
.check_rate_limit(password_limit_key.clone(), 10, window)
.await
{
let _ = send_error(
&mut socket,
&format!("Too many password attempts. Try again in {seconds} seconds."),
)
.await;
return;
}
}
let password_ok = db::verify_workspace_password(&workspace, password.as_deref());
if password_ok {
state.clear_rate_limit(&password_limit_key).await;
}
if workspace.is_private != 0 && permission.is_none() && !anonymous_token_ok {
let _ = send_error(&mut socket, "Workspace not found").await;
return;
+87 -12
View File
@@ -32,12 +32,30 @@ enum PadServerMessage {
}
pub async fn upgrade_pad(
ws: WebSocketUpgrade,
headers: HeaderMap,
Path(slug): Path<String>,
State(state): State<SharedState>,
) -> Response {
ws.on_upgrade(move |socket| handle_pad_socket(socket, state, slug))
if !crate::security::websocket_origin_allowed(&headers) {
warn!(%slug, "pad websocket rejected: invalid origin");
return (StatusCode::FORBIDDEN, "Invalid WebSocket origin").into_response();
}
let account_token = crate::security::session_token(&headers).map(str::to_owned);
let resource_token = crate::security::resource_token(&headers, "pad", &slug)
.map(str::to_owned);
let client_key = crate::security::client_key(&headers);
ws.on_upgrade(move |socket| {
handle_pad_socket(socket, state, slug, account_token, resource_token, client_key)
})
}
async fn handle_pad_socket(mut socket: WebSocket, state: SharedState, slug: String) {
async fn handle_pad_socket(
mut socket: WebSocket,
state: SharedState,
slug: String,
cookie_session_token: Option<String>,
cookie_access_token: Option<String>,
client_key: String,
) {
info!(%slug, "pad websocket connected");
let Some(pad) = db::find_pad(&state.db, &slug).await.ok().flatten() else {
warn!(%slug, "pad websocket rejected: pad not found");
@@ -50,21 +68,19 @@ async fn handle_pad_socket(mut socket: WebSocket, state: SharedState, slug: Stri
.await;
return;
};
let (password, access_token, nickname, session_token, guest_id, color) =
let (password, access_token, nickname, guest_id, color) =
match socket.recv().await {
Some(Ok(Message::Text(text))) => match serde_json::from_str::<ClientMessage>(&text) {
Ok(ClientMessage::Authenticate {
password,
access_token,
nickname,
session_token,
guest_id,
color,
}) => (
password,
access_token,
clean_nickname(nickname),
session_token,
clean_guest_id(guest_id),
clean_color(color),
),
@@ -81,6 +97,13 @@ async fn handle_pad_socket(mut socket: WebSocket, state: SharedState, slug: Stri
},
_ => return,
};
let session_token = cookie_session_token;
let explicit_access_token = access_token
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty() && *value != "cookie")
.map(str::to_owned);
let access_token = explicit_access_token.or(cookie_access_token);
let nickname = match auth::authorize_nickname(&state, nickname, session_token.clone()).await {
Ok(value) => value,
Err(message) => {
@@ -100,16 +123,68 @@ async fn handle_pad_socket(mut socket: WebSocket, state: SharedState, slug: Stri
.map(|name| format!("guest:{id}:{}", name.to_lowercase()))
}),
};
let supplied_token = session_token.as_deref().or(access_token.as_deref());
let permission = auth::resource_permission(&state, "pad", &slug, supplied_token)
.await
.ok()
.flatten();
let permission = resource_permission_from_tokens(
&state,
"pad",
&slug,
access_token.as_deref(),
session_token.as_deref(),
)
.await;
let anonymous_token_ok = permission.is_none()
&& crate::api::verify_resource_access_token(&state, "pad", &slug, supplied_token)
&& anonymous_access_from_tokens(
&state,
"pad",
&slug,
access_token.as_deref(),
)
.await;
let password_limit_key = format!("resource-password:{client_key}:pad:{slug}");
let password_attempted = password
.as_deref()
.map(str::trim)
.is_some_and(|value| !value.is_empty());
if pad.password_hash.is_some()
&& permission.as_deref() != Some("rw")
&& !anonymous_token_ok
&& password_attempted
{
let window = std::time::Duration::from_secs(15 * 60);
if let Err(seconds) = state
.check_rate_limit(format!("resource-password-client:{client_key}"), 50, window)
.await
.unwrap_or(false);
{
let _ = send_pad(
&mut socket,
&PadServerMessage::Error {
message: format!(
"Too many password attempts. Try again in {seconds} seconds."
),
},
)
.await;
return;
}
if let Err(seconds) = state
.check_rate_limit(password_limit_key.clone(), 10, window)
.await
{
let _ = send_pad(
&mut socket,
&PadServerMessage::Error {
message: format!(
"Too many password attempts. Try again in {seconds} seconds."
),
},
)
.await;
return;
}
}
let password_ok = db::verify_pad_password(&pad, password.as_deref());
if password_ok {
state.clear_rate_limit(&password_limit_key).await;
}
if pad.is_private != 0 && permission.is_none() && !anonymous_token_ok {
let _ = send_pad(
&mut socket,