diff --git a/.env.example b/.env.example
index 1390efc..c062b55 100644
--- a/.env.example
+++ b/.env.example
@@ -1,6 +1,5 @@
DATABASE_URL=postgres://postgres:postgres@localhost:5432/chattz
PORT=3000
-APP_BASE_URL=http://localhost:3000
# Authentik OIDC app values
OIDC_CLIENT_ID=replace-me
@@ -17,9 +16,13 @@ TURN_URLS=turn:turn.example.com:3478?transport=udp,turn:turn.example.com:3478?tr
TURN_USERNAME=replace-me
TURN_PASSWORD=replace-me
+# 32+ random chars; used to sign session cookies
+SESSION_SECRET=replace-with-long-random-secret
+COOKIE_SECURE=false
+
# Cloudflare R2 media uploads
R2_ACCOUNT_ID=replace-me
R2_ACCESS_KEY_ID=replace-me
R2_SECRET_ACCESS_KEY=replace-me
R2_BUCKET=chattz-media
-MEDIA_BASE_URL=https://media.example.com
+R2_PUBLIC_BASE_URL=https://media.example.com
diff --git a/Cargo.lock b/Cargo.lock
index c44764d..b496651 100644
--- a/Cargo.lock
+++ b/Cargo.lock
@@ -675,14 +675,13 @@ dependencies = [
"chrono",
"dotenvy",
"futures-util",
+ "jsonwebtoken",
"reqwest",
"sea-orm",
"sea-orm-migration",
"serde",
"serde_json",
- "sha2",
"tokio",
- "tower",
"tower-http",
"tracing",
"tracing-subscriber",
@@ -1649,6 +1648,21 @@ dependencies = [
"wasm-bindgen",
]
+[[package]]
+name = "jsonwebtoken"
+version = "9.3.1"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "5a87cc7a48537badeae96744432de36f4be2b4a34a05a5ef32e9dd8a1c169dde"
+dependencies = [
+ "base64",
+ "js-sys",
+ "pem",
+ "ring",
+ "serde",
+ "serde_json",
+ "simple_asn1",
+]
+
[[package]]
name = "lazy_static"
version = "1.5.0"
@@ -1811,6 +1825,16 @@ dependencies = [
"windows-sys 0.61.2",
]
+[[package]]
+name = "num-bigint"
+version = "0.4.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "a5e44f723f1133c9deac646763579fdb3ac745e418f2a7af9cd0c431da1f20b9"
+dependencies = [
+ "num-integer",
+ "num-traits",
+]
+
[[package]]
name = "num-bigint-dig"
version = "0.8.6"
@@ -1943,6 +1967,16 @@ dependencies = [
"windows-link",
]
+[[package]]
+name = "pem"
+version = "3.0.6"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "1d30c53c26bc5b31a98cd02d20f25a7c8567146caf63ed593a9d87b2775291be"
+dependencies = [
+ "base64",
+ "serde_core",
+]
+
[[package]]
name = "pem-rfc7468"
version = "0.7.0"
@@ -2731,6 +2765,18 @@ dependencies = [
"rand_core 0.6.4",
]
+[[package]]
+name = "simple_asn1"
+version = "0.6.4"
+source = "registry+https://github.com/rust-lang/crates.io-index"
+checksum = "0d585997b0ac10be3c5ee635f1bab02d512760d14b7c468801ac8a01d9ae5f1d"
+dependencies = [
+ "num-bigint",
+ "num-traits",
+ "thiserror",
+ "time",
+]
+
[[package]]
name = "slab"
version = "0.4.12"
@@ -3085,6 +3131,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "743bd48c283afc0388f9b8827b976905fb217ad9e647fae3a379a9283c4def2c"
dependencies = [
"deranged",
+ "itoa",
"num-conv",
"powerfmt",
"serde_core",
diff --git a/Cargo.toml b/Cargo.toml
index 6dbd768..d7ece96 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -10,15 +10,14 @@ aws-sdk-s3 = { version = "1", default-features = false, features = ["rt-tokio",
axum = { version = "0.8", features = ["macros", "ws", "multipart"] }
chrono = { version = "0.4", features = ["serde"] }
dotenvy = "0.15"
+jsonwebtoken = "9"
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls"] }
-sea-orm = { version = "1.1", default-features = false, features = ["sqlx-postgres", "runtime-tokio-rustls", "macros", "with-chrono", "with-uuid", "mock"] }
+sea-orm = { version = "1.1", default-features = false, features = ["sqlx-postgres", "runtime-tokio-rustls", "macros", "with-chrono", "with-uuid"] }
sea-orm-migration = { version = "1.1", default-features = false, features = ["sqlx-postgres", "runtime-tokio-rustls"] }
serde = { version = "1", features = ["derive"] }
serde_json = "1"
-sha2 = "0.10"
futures-util = "0.3"
-tokio = { version = "1", features = ["fs", "io-util", "macros", "rt-multi-thread"] }
-tower = { version = "0.5", features = ["util"] }
+tokio = { version = "1", features = ["macros", "rt-multi-thread"] }
tower-http = { version = "0.6", features = ["trace", "fs"] }
tracing = "0.1"
tracing-subscriber = { version = "0.3", features = ["env-filter", "fmt"] }
diff --git a/README.md b/README.md
index 2708a60..af5e8f0 100644
--- a/README.md
+++ b/README.md
@@ -8,7 +8,7 @@ A simple single-instance Discord-style monolith in Rust using:
## What this includes
- OIDC login flow (`/auth/login`, `/auth/callback`, `/auth/logout`)
-- HttpOnly session cookie auth
+- Signed session cookie auth
- Channel voice chat over WebRTC (P2P mesh) with server WebSocket signaling
- Guild invite codes (create + join)
- Direct messages (DM) between users
@@ -41,11 +41,6 @@ For voice reliability on restrictive networks, configure TURN in `.env`:
- `TURN_USERNAME`
- `TURN_PASSWORD`
-For production deployments:
-- `APP_BASE_URL` must be your public app origin and should use `https`
-- `MEDIA_BASE_URL` should be a separate media origin for user uploads
-- uploads and soundboard require R2/object storage to be configured
-
3. Run app:
```bash
@@ -60,7 +55,7 @@ Web UI is available at `http://localhost:${PORT}/`.
## Authentik setup notes
Create an Authentik OAuth2/OIDC provider + application and set:
-- Redirect URI: `${APP_BASE_URL}/auth/callback`
+- Redirect URI: `http://localhost:3000/auth/callback`
- Scopes including at least: `openid profile email`
If you change `PORT`, update `OIDC_REDIRECT_URL` and this redirect URI to match.
@@ -96,7 +91,6 @@ For Authentik these are commonly under `/application/o/...` for the app slug.
- `GET /channels/:channel_id/voice/ws` (WebSocket signaling)
All endpoints except health and auth flow require the session cookie from successful login.
-Authenticated WebSocket connections (`/ws`, `/channels/:channel_id/voice/ws`) also use the same cookie session.
## Notes
@@ -104,7 +98,6 @@ This is intentionally minimal and monolithic (single process, single Postgres in
Voice is implemented as browser-to-browser WebRTC audio with signaling in this server.
For two users behind strict NAT/firewall, you may need TURN for reliable connectivity.
The web UI remembers the last selected guild in browser local storage and auto-selects it on reload.
-User uploads are served from the configured media origin, not from `/static`.
Mic filter modes in the UI:
- `NSNet2 (Compat)`: always-on denoising mode (implemented using DeepFilterNet3 with lighter suppression preset)
diff --git a/desktop/index.html b/desktop/index.html
index 1684077..cb1de11 100644
--- a/desktop/index.html
+++ b/desktop/index.html
@@ -6,9 +6,12 @@
Chattz
+
+
+
-
+
@@ -266,16 +269,6 @@
-
-
-
Upload Too Large
-
Uploads are limited to 50 MB per file.
-
-
-
-
diff --git a/desktop/preload.js b/desktop/preload.js
index f43670e..dfcc665 100644
--- a/desktop/preload.js
+++ b/desktop/preload.js
@@ -1,20 +1,18 @@
const { contextBridge, ipcRenderer } = require('electron');
-function onIpc(channel, callback) {
- const listener = (_event, payload) => callback(payload);
- ipcRenderer.on(channel, listener);
- return () => ipcRenderer.removeListener(channel, listener);
-}
-
contextBridge.exposeInMainWorld('electronAPI', {
copyToClipboard: (text) => ipcRenderer.invoke('clipboard-write', text),
+ getConfig: () => ipcRenderer.invoke('get-config'),
+ storageGet: (key) => ipcRenderer.invoke('storage-get', key),
+ storageSet: (key, value) => ipcRenderer.invoke('storage-set', key, value),
+ storageRemove: (key) => ipcRenderer.invoke('storage-remove', key),
getUpdateState: () => ipcRenderer.invoke('get-update-state'),
checkForUpdatesNow: () => ipcRenderer.invoke('check-for-updates-now'),
checkForUpdates: () => ipcRenderer.send('check-for-updates'),
downloadUpdate: () => ipcRenderer.send('download-update'),
quitAndInstall: () => ipcRenderer.send('quit-and-install'),
- onUpdateState: (callback) => onIpc('update-state', callback),
- onUpdateAvailable: (callback) => onIpc('update-available', callback),
- onUpdateDownloaded: (callback) => onIpc('update-downloaded', callback),
- onUpdateError: (callback) => onIpc('update-error', callback)
+ onUpdateState: (callback) => ipcRenderer.on('update-state', (event, state) => callback(state)),
+ onUpdateAvailable: (callback) => ipcRenderer.on('update-available', (event, info) => callback(info)),
+ onUpdateDownloaded: (callback) => ipcRenderer.on('update-downloaded', (event, info) => callback(info)),
+ onUpdateError: (callback) => ipcRenderer.on('update-error', (event, error) => callback(error))
});
diff --git a/main.js b/main.js
index 030c413..3a32c94 100644
--- a/main.js
+++ b/main.js
@@ -1,11 +1,11 @@
const { app, BrowserWindow, session, desktopCapturer, ipcMain, clipboard } = require('electron');
const { autoUpdater } = require('electron-updater');
const path = require('path');
-const packageJson = require('./package.json');
-const PERIODIC_UPDATE_CHECK_INTERVAL_MS = 1 * 60 * 1000;
+const fs = require('fs');
+const url = require('url');
const updateState = {
- status: 'idle', // idle | checking | available | downloading | downloaded | installing | not-available | error
+ status: 'idle', // idle | checking | available | downloading | downloaded | not-available | error
info: null,
error: null,
};
@@ -19,43 +19,41 @@ function broadcastUpdateState() {
}
}
+function getStorageFilePath() {
+ return path.join(app.getPath('userData'), 'renderer-storage.json');
+}
+
+function readPersistentStore() {
+ const filePath = getStorageFilePath();
+ try {
+ if (!fs.existsSync(filePath)) return {};
+ const raw = fs.readFileSync(filePath, 'utf8');
+ const parsed = JSON.parse(raw);
+ return parsed && typeof parsed === 'object' ? parsed : {};
+ } catch (err) {
+ console.error('Failed to read persistent store', err);
+ return {};
+ }
+}
+
+function writePersistentStore(store) {
+ const filePath = getStorageFilePath();
+ try {
+ fs.mkdirSync(path.dirname(filePath), { recursive: true });
+ fs.writeFileSync(filePath, JSON.stringify(store), 'utf8');
+ } catch (err) {
+ console.error('Failed to write persistent store', err);
+ }
+}
+
function isNewerVersionAvailable(info) {
const next = info && typeof info.version === 'string' ? info.version : '';
return Boolean(next) && next !== app.getVersion();
}
let updateCheckInProgress = false;
-let installInProgress = false;
-let periodicUpdateTimer = null;
-
-function resolveBackendUrl() {
- const configuredUrl = (process.env.CHATTZ_URL || process.env.APP_BASE_URL || '').trim();
- if (configuredUrl) {
- return configuredUrl.replace(/\/$/, '');
- }
-
- const packagedFallback = typeof packageJson.homepage === 'string'
- ? packageJson.homepage.trim()
- : '';
- if (app.isPackaged && packagedFallback) {
- return packagedFallback.replace(/\/$/, '');
- }
-
- return 'http://localhost:3000';
-}
-
-function webContentsOrigin(webContents) {
- if (!webContents || typeof webContents.getURL !== 'function') {
- return null;
- }
- return safeOrigin(webContents.getURL());
-}
async function runUpdateCheck(reason = 'manual') {
- if (installInProgress) {
- console.log(`Update check skipped (${reason}): install already in progress`);
- return;
- }
if (updateCheckInProgress) {
console.log(`Update check skipped (${reason}): another check is already in progress`);
return;
@@ -85,62 +83,9 @@ async function runUpdateCheck(reason = 'manual') {
}
}
-function installDownloadedUpdate() {
- if (installInProgress) {
- console.log('Update install already in progress');
- return;
- }
-
- installInProgress = true;
- updateState.status = 'installing';
- updateState.error = null;
- broadcastUpdateState();
-
- const wins = BrowserWindow.getAllWindows();
- for (const win of wins) {
- if (win.isDestroyed()) continue;
- try {
- win.removeAllListeners('close');
- win.destroy();
- } catch (err) {
- console.warn('Failed to destroy window before update install', err);
- }
- }
-
- setImmediate(() => {
- try {
- autoUpdater.quitAndInstall(false, true);
- } catch (err) {
- installInProgress = false;
- updateState.status = 'error';
- updateState.error = err && err.message ? err.message : String(err);
- broadcastUpdateState();
- console.error('quitAndInstall failed', err);
- }
- });
-
- setTimeout(() => {
- if (installInProgress) {
- console.warn('Update install is still waiting for app shutdown');
- }
- }, 15000);
-}
-
-function startPeriodicUpdateChecks() {
- if (periodicUpdateTimer || !app.isPackaged) {
- return;
- }
-
- periodicUpdateTimer = setInterval(() => {
- void runUpdateCheck('periodic');
- }, PERIODIC_UPDATE_CHECK_INTERVAL_MS);
-}
-
function createWindow() {
+ // Create a persistent session for chattz to keep the user logged in
const sess = session.fromPartition('persist:chattz');
- const backendUrl = resolveBackendUrl();
- const backendOrigin = new URL(backendUrl).origin;
- console.log(`Desktop backend URL: ${backendUrl}`);
const win = new BrowserWindow({
width: 1200,
@@ -149,6 +94,7 @@ function createWindow() {
webPreferences: {
nodeIntegration: false,
contextIsolation: true,
+ webSecurity: false, // Required for cross-origin fetch/ws from file:// with cookies
session: sess,
preload: path.join(__dirname, 'desktop', 'preload.js')
}
@@ -157,33 +103,32 @@ function createWindow() {
win.setAutoHideMenuBar(true);
win.setMenuBarVisibility(true);
+ // Auto-approve media permissions (camera, microphone)
sess.setPermissionCheckHandler((webContents, permission) => {
- const origin = webContentsOrigin(webContents);
- if (origin !== backendOrigin) return false;
- return permission === 'media' || permission === 'clipboard-write';
+ if (permission === 'media' || permission === 'clipboard-read' || permission === 'clipboard-write') {
+ return true;
+ }
+ return false;
});
sess.setPermissionRequestHandler((webContents, permission, callback) => {
- const origin = webContentsOrigin(webContents);
- if (origin === backendOrigin && (permission === 'media' || permission === 'clipboard-write')) {
+ if (permission === 'media' || permission === 'clipboard-read' || permission === 'clipboard-write') {
callback(true);
} else {
callback(false);
}
});
+ // Handle screen share requests
sess.setDisplayMediaRequestHandler((request, callback) => {
- const origin = safeOrigin(request.frame?.url || win.webContents.getURL());
- if (origin !== backendOrigin) {
- callback(null);
- return;
- }
desktopCapturer.getSources({ types: ['screen', 'window'] }).then((sources) => {
+ // Provide the first screen source by default, or implement a picker window here
if (sources && sources.length > 0) {
+ // We prefer a screen over a window if available, simple heuristic
const screenSource = sources.find(s => s.id.startsWith('screen')) || sources[0];
callback({ video: screenSource, audio: 'loopback' });
} else {
- callback(null);
+ callback(null); // Reject safely
}
}).catch(err => {
console.error("Failed to get desktop sources for screen share", err);
@@ -191,9 +136,55 @@ function createWindow() {
});
});
- win.loadURL(backendUrl).catch((err) => {
- console.error(`Failed to load desktop app: ${err}`);
- });
+ const backendUrl = (process.env.CHATTZ_URL || 'https://discord.flegr.me').replace(/\/$/, '');
+ const indexPath = path.join(__dirname, 'desktop', 'index.html');
+
+ const loadDesktopApp = (queryParams = '') => {
+ const options = {};
+ if (queryParams) {
+ try {
+ options.query = Object.fromEntries(new URLSearchParams(queryParams));
+ } catch (e) {
+ console.error("Failed to parse query params", e);
+ }
+ }
+
+ // Use setImmediate to ensure the current navigation tick is cleared,
+ // which prevents ERR_ABORTED (-3) on some platforms when interrupting a redirect.
+ setImmediate(() => {
+ if (win.isDestroyed()) return;
+ win.loadFile(indexPath, options).catch((err) => {
+ // Ignore aborted errors as they often happen during fast redirects
+ if (err.toString().includes('-3') || err.code === -3) return;
+ console.error(`Failed to load desktop file: ${err}`);
+ });
+ });
+ };
+
+ // Intercept navigations/redirects that return to the backend with a token
+ // after the OAuth login flow. Two handlers are needed because:
+ // - will-navigate fires for client-side navigations (location.href, link clicks)
+ // - will-redirect fires for server-side 302 redirects (OAuth callback chain)
+ const interceptLoginRedirect = (event, navigatedUrl) => {
+ try {
+ const urlObj = new URL(navigatedUrl);
+ const backendObj = new URL(backendUrl);
+ if (urlObj.origin === backendObj.origin && urlObj.pathname === '/') {
+ if (urlObj.searchParams.has('token')) {
+ console.log("Detected login success redirect, returning to desktop UI...");
+ event.preventDefault();
+ loadDesktopApp(urlObj.search.slice(1));
+ }
+ }
+ } catch (e) {
+ console.error(e);
+ }
+ };
+
+ win.webContents.on('will-navigate', interceptLoginRedirect);
+ win.webContents.on('will-redirect', interceptLoginRedirect);
+
+ loadDesktopApp();
// Ensure renderer receives latest updater state after any (re)load.
// Delay broadcast by 300ms to give the renderer time to register its
@@ -212,11 +203,40 @@ app.commandLine.appendSwitch('disable-webrtc-hw-encoding'); // Sometime helps re
app.commandLine.appendSwitch('log-level', '3'); // Silences standard Chromium warning logs (like SRTP transport unprotect logs which are benign timing issues)
app.whenReady().then(() => {
+ // Register IPC handlers once (before creating any windows)
+ ipcMain.handle('get-config', () => {
+ return {
+ backendUrl: (process.env.CHATTZ_URL || 'https://discord.flegr.me').replace(/\/$/, '')
+ };
+ });
+
ipcMain.handle('clipboard-write', (event, text) => {
clipboard.writeText(text);
return true;
});
+ ipcMain.handle('storage-get', (event, key) => {
+ if (typeof key !== 'string' || key.length === 0) return null;
+ const store = readPersistentStore();
+ return Object.prototype.hasOwnProperty.call(store, key) ? store[key] : null;
+ });
+
+ ipcMain.handle('storage-set', (event, key, value) => {
+ if (typeof key !== 'string' || key.length === 0) return false;
+ const store = readPersistentStore();
+ store[key] = value;
+ writePersistentStore(store);
+ return true;
+ });
+
+ ipcMain.handle('storage-remove', (event, key) => {
+ if (typeof key !== 'string' || key.length === 0) return false;
+ const store = readPersistentStore();
+ delete store[key];
+ writePersistentStore(store);
+ return true;
+ });
+
ipcMain.handle('get-update-state', () => ({ ...updateState }));
ipcMain.handle('check-for-updates-now', async () => {
await runUpdateCheck('renderer-direct');
@@ -267,7 +287,6 @@ app.whenReady().then(() => {
});
ipcMain.on('download-update', () => {
- if (installInProgress) return;
updateState.status = 'downloading';
updateState.error = null;
broadcastUpdateState();
@@ -275,14 +294,13 @@ app.whenReady().then(() => {
});
ipcMain.on('quit-and-install', () => {
- installDownloadedUpdate();
+ autoUpdater.quitAndInstall();
});
// Check once on startup
setTimeout(() => {
void runUpdateCheck('startup');
}, 5000);
- startPeriodicUpdateChecks();
app.on('activate', () => {
if (BrowserWindow.getAllWindows().length === 0) {
@@ -296,18 +314,3 @@ app.on('window-all-closed', () => {
app.quit();
}
});
-
-app.on('before-quit', () => {
- if (periodicUpdateTimer) {
- clearInterval(periodicUpdateTimer);
- periodicUpdateTimer = null;
- }
-});
-
-function safeOrigin(value) {
- try {
- return new URL(value).origin;
- } catch (_) {
- return null;
- }
-}
diff --git a/scripts/generate-html.js b/scripts/generate-html.js
index a773454..1b3497b 100644
--- a/scripts/generate-html.js
+++ b/scripts/generate-html.js
@@ -5,30 +5,26 @@ const root = path.resolve(__dirname, '..');
const templatePath = path.join(root, 'shared-html', 'index.template.html');
const template = fs.readFileSync(templatePath, 'utf8');
-const updaterMarkup = `
-
-
-
`;
-
const variants = {
web: {
styles_href: '/static/styles.css',
- lucide_src: '/static/vendor/lucide.min.js',
app_src: '/static/app.js?v=20260227-shared-core-1',
web_downloads: ``,
- desktop_updater: updaterMarkup,
+ desktop_updater: '',
outPath: path.join(root, 'static', 'index.html'),
},
desktop: {
styles_href: '../static/styles.css',
- lucide_src: '../static/vendor/lucide.min.js',
app_src: 'app.js?v=20260227-shared-core-1',
web_downloads: '',
- desktop_updater: updaterMarkup,
+ desktop_updater: `
+
+
+
`,
outPath: path.join(root, 'desktop', 'index.html'),
},
};
diff --git a/shared-html/index.template.html b/shared-html/index.template.html
index 3425611..9780068 100644
--- a/shared-html/index.template.html
+++ b/shared-html/index.template.html
@@ -5,9 +5,12 @@
Chattz
+
+
+
-
+
@@ -263,16 +266,6 @@
-
-
-
Upload Too Large
-
Uploads are limited to 50 MB per file.
-
-
-
-
diff --git a/src/auth.rs b/src/auth.rs
index 535953e..667c64d 100644
--- a/src/auth.rs
+++ b/src/auth.rs
@@ -1,25 +1,33 @@
-use anyhow::{Result, anyhow};
+use std::time::{Duration, SystemTime, UNIX_EPOCH};
+
+use anyhow::{Context, Result, anyhow};
use axum::{
Json,
extract::{FromRef, FromRequestParts},
- http::{HeaderMap, StatusCode, request::Parts},
+ http::{StatusCode, request::Parts},
response::{IntoResponse, Response},
};
-use chrono::{Duration, Utc};
-use serde::Serialize;
-use sha2::{Digest, Sha256};
+use jsonwebtoken::{DecodingKey, EncodingKey, Header, Validation, decode, encode};
+use serde::{Deserialize, Serialize};
use uuid::Uuid;
use crate::{AppState, db};
-pub const OAUTH_STATE_COOKIE: &str = "chattz_oauth_state";
-pub const SESSION_COOKIE: &str = "chattz_session";
-const SESSION_TTL_DAYS: i64 = 30;
+const OAUTH_STATE_COOKIE: &str = "chattz_oauth_state";
+const SESSION_TTL_SECS: u64 = 60 * 15; // 15 minutes
+const REFRESH_TTL_SECS: u64 = 60 * 60 * 24 * 30; // 30 days
+
+#[derive(Debug, Serialize, Deserialize, Clone)]
+pub struct SessionClaims {
+ pub sub: String,
+ pub kind: String, // "access" or "refresh"
+ pub exp: usize,
+ pub iat: usize,
+}
#[derive(Debug, Clone)]
pub struct AuthUser {
pub id: Uuid,
- pub session_id: String,
}
#[derive(Debug)]
@@ -36,13 +44,6 @@ impl ApiError {
}
}
- pub fn forbidden(msg: &str) -> Self {
- Self {
- status: StatusCode::FORBIDDEN,
- message: msg.to_string(),
- }
- }
-
pub fn bad_request(msg: &str) -> Self {
Self {
status: StatusCode::BAD_REQUEST,
@@ -56,13 +57,6 @@ impl ApiError {
message: msg.to_string(),
}
}
-
- pub fn service_unavailable(msg: &str) -> Self {
- Self {
- status: StatusCode::SERVICE_UNAVAILABLE,
- message: msg.to_string(),
- }
- }
}
#[derive(Serialize)]
@@ -90,25 +84,29 @@ impl From for ApiError {
impl FromRequestParts for AuthUser
where
- AppState: FromRef,
+ AppState: axum::extract::FromRef,
S: Send + Sync,
{
type Rejection = ApiError;
async fn from_request_parts(parts: &mut Parts, state: &S) -> Result {
let app = AppState::from_ref(state);
- let session_id = read_cookie_from_headers(&parts.headers, SESSION_COOKIE)
- .ok_or_else(|| ApiError::unauthorized("missing session cookie"))?;
- let user_id = db::touch_active_session(&app.db, &session_id)
+ let token = read_bearer_token(parts)
+ .or_else(|| read_query_token(parts))
+ .ok_or_else(|| ApiError::unauthorized("missing jwt token"))?;
+
+ let user_id = verify_session(&token, &app.settings.session_secret, "access")
+ .map_err(|_| ApiError::unauthorized("invalid or expired token"))?;
+
+ let exists = db::user_exists(&app.db, user_id)
.await
- .map_err(|_| ApiError::unauthorized("invalid or expired session"))?
- .ok_or_else(|| ApiError::unauthorized("invalid or expired session"))?;
+ .map_err(|_| ApiError::unauthorized("session user not found"))?;
+ if !exists {
+ return Err(ApiError::unauthorized("session user not found"));
+ }
- Ok(Self {
- id: user_id,
- session_id,
- })
+ Ok(Self { id: user_id })
}
}
@@ -116,14 +114,6 @@ pub fn new_oauth_state() -> String {
Uuid::new_v4().to_string()
}
-pub fn new_session_id() -> String {
- Uuid::new_v4().simple().to_string()
-}
-
-pub fn session_expiry() -> chrono::DateTime {
- (Utc::now() + Duration::days(SESSION_TTL_DAYS)).fixed_offset()
-}
-
pub fn make_oauth_state_cookie(value: &str, secure: bool) -> String {
format!(
"{name}={value}; Path=/; HttpOnly; SameSite=Lax; Max-Age=600{secure_flag}",
@@ -133,32 +123,76 @@ pub fn make_oauth_state_cookie(value: &str, secure: bool) -> String {
}
pub fn clear_oauth_state_cookie(secure: bool) -> String {
- clear_cookie(OAUTH_STATE_COOKIE, secure, "Lax")
-}
-
-pub fn make_session_cookie(value: &str, secure: bool) -> String {
format!(
- "{name}={value}; Path=/; HttpOnly; SameSite=Strict; Max-Age={ttl}{secure_flag}",
- name = SESSION_COOKIE,
- ttl = Duration::days(SESSION_TTL_DAYS).num_seconds(),
+ "{name}=; Path=/; HttpOnly; SameSite=Lax; Max-Age=0{secure_flag}",
+ name = OAUTH_STATE_COOKIE,
secure_flag = if secure { "; Secure" } else { "" }
)
}
-pub fn clear_session_cookie(secure: bool) -> String {
- clear_cookie(SESSION_COOKIE, secure, "Strict")
-}
-
-pub fn read_cookie_from_headers(headers: &HeaderMap, cookie_name: &str) -> Option {
+pub fn read_oauth_state_from_headers(headers: &axum::http::HeaderMap) -> Option {
let raw = headers.get(axum::http::header::COOKIE)?.to_str().ok()?;
raw.split(';').find_map(|pair| {
let mut kv = pair.trim().splitn(2, '=');
let key = kv.next()?;
let value = kv.next()?;
- (key == cookie_name).then(|| value.to_string())
+ (key == OAUTH_STATE_COOKIE).then(|| value.to_string())
})
}
+pub fn new_jwt_tokens(user_id: Uuid, secret: &str) -> Result<(String, String)> {
+ let now = now_ts();
+
+ let access_claims = SessionClaims {
+ sub: user_id.to_string(),
+ kind: "access".to_string(),
+ iat: now as usize,
+ exp: (now + SESSION_TTL_SECS) as usize,
+ };
+
+ let refresh_claims = SessionClaims {
+ sub: user_id.to_string(),
+ kind: "refresh".to_string(),
+ iat: now as usize,
+ exp: (now + REFRESH_TTL_SECS) as usize,
+ };
+
+ let access_token = encode(
+ &Header::default(),
+ &access_claims,
+ &EncodingKey::from_secret(secret.as_bytes()),
+ )
+ .context("failed to encode access token")?;
+
+ let refresh_token = encode(
+ &Header::default(),
+ &refresh_claims,
+ &EncodingKey::from_secret(secret.as_bytes()),
+ )
+ .context("failed to encode refresh token")?;
+
+ Ok((access_token, refresh_token))
+}
+
+pub fn verify_session(token: &str, secret: &str, expected_kind: &str) -> Result {
+ let mut validation = Validation::default();
+ validation.validate_exp = true;
+
+ let data = decode::(
+ token,
+ &DecodingKey::from_secret(secret.as_bytes()),
+ &validation,
+ )
+ .context("failed to decode session token")?;
+
+ if data.claims.kind != expected_kind {
+ return Err(anyhow!("invalid token kind"));
+ }
+
+ let user_id = Uuid::parse_str(&data.claims.sub).context("invalid sub in session token")?;
+ Ok(user_id)
+}
+
pub fn validate_oauth_state(expected_cookie: Option, query_state: &str) -> Result<()> {
let expected = expected_cookie.ok_or_else(|| anyhow!("missing oauth state cookie"))?;
if expected != query_state {
@@ -167,113 +201,37 @@ pub fn validate_oauth_state(expected_cookie: Option, query_state: &str)
Ok(())
}
-pub fn user_agent_hash(headers: &HeaderMap) -> Option {
- header_hash(headers, axum::http::header::USER_AGENT.as_str())
-}
-
-pub fn ip_hash(headers: &HeaderMap) -> Option {
- headers
- .get("x-forwarded-for")
- .and_then(|v| v.to_str().ok())
- .and_then(|raw| raw.split(',').next().map(str::trim))
- .filter(|value| !value.is_empty())
- .map(hash_string)
- .or_else(|| header_hash(headers, "x-real-ip"))
-}
-
-pub fn origin_matches(headers: &HeaderMap, expected_origin: &str) -> bool {
- headers
- .get(axum::http::header::ORIGIN)
- .and_then(|value| value.to_str().ok())
- .map(|origin| origin == expected_origin)
- .unwrap_or(false)
-}
-
-fn header_hash(headers: &HeaderMap, header_name: &str) -> Option {
- headers
- .get(header_name)
- .and_then(|v| v.to_str().ok())
- .filter(|value| !value.trim().is_empty())
- .map(hash_string)
-}
-
-fn hash_string(value: &str) -> String {
- let mut hasher = Sha256::new();
- hasher.update(value.as_bytes());
- let digest = hasher.finalize();
- let mut out = String::with_capacity(digest.len() * 2);
- for byte in digest {
- out.push(nibble_to_hex(byte >> 4));
- out.push(nibble_to_hex(byte & 0x0f));
- }
- out
-}
-
-fn nibble_to_hex(value: u8) -> char {
- match value {
- 0..=9 => (b'0' + value) as char,
- 10..=15 => (b'a' + (value - 10)) as char,
- _ => unreachable!(),
+fn read_bearer_token(parts: &Parts) -> Option {
+ let raw = parts
+ .headers
+ .get(axum::http::header::AUTHORIZATION)?
+ .to_str()
+ .ok()?;
+ if raw.starts_with("Bearer ") {
+ Some(raw["Bearer ".len()..].trim().to_string())
+ } else {
+ None
}
}
-fn clear_cookie(name: &str, secure: bool, same_site: &str) -> String {
- format!(
- "{name}=; Path=/; HttpOnly; SameSite={same_site}; Max-Age=0{secure_flag}",
- secure_flag = if secure { "; Secure" } else { "" }
- )
+fn read_query_token(parts: &Parts) -> Option {
+ let query = parts.uri.query()?;
+
+ // Simple query param parsing without pulling in url::Url overhead
+ for pair in query.split('&') {
+ let mut kv = pair.splitn(2, '=');
+ let key = kv.next()?;
+ let value = kv.next()?;
+ if key == "token" {
+ return Some(value.to_string());
+ }
+ }
+ None
}
-#[cfg(test)]
-mod tests {
- use super::{
- SESSION_COOKIE, clear_session_cookie, hash_string, make_session_cookie, origin_matches,
- read_cookie_from_headers,
- };
- use axum::http::{HeaderMap, header};
-
- #[test]
- fn hash_string_is_stable() {
- assert_eq!(
- hash_string("example"),
- "50d858e0985ecc7f60418aaf0cc5ab587f42c2570a884095a9e8ccacd0f6545c"
- );
- }
-
- #[test]
- fn reads_named_cookie_from_header() {
- let mut headers = HeaderMap::new();
- headers.insert(
- header::COOKIE,
- "other=value; chattz_session=session-123; another=ok"
- .parse()
- .unwrap(),
- );
-
- assert_eq!(
- read_cookie_from_headers(&headers, SESSION_COOKIE),
- Some("session-123".to_string())
- );
- }
-
- #[test]
- fn origin_match_requires_exact_origin() {
- let mut headers = HeaderMap::new();
- headers.insert(header::ORIGIN, "http://localhost:3000".parse().unwrap());
-
- assert!(origin_matches(&headers, "http://localhost:3000"));
- assert!(!origin_matches(&headers, "https://localhost:3000"));
- }
-
- #[test]
- fn session_cookie_has_strict_policy_and_clear_cookie_expires() {
- let session_cookie = make_session_cookie("session-123", false);
- let cleared_cookie = clear_session_cookie(false);
-
- assert!(session_cookie.contains("HttpOnly"));
- assert!(session_cookie.contains("SameSite=Strict"));
- assert!(session_cookie.contains("Max-Age="));
- assert!(cleared_cookie.contains("SameSite=Strict"));
- assert!(cleared_cookie.contains("Max-Age=0"));
- }
+fn now_ts() -> u64 {
+ SystemTime::now()
+ .duration_since(UNIX_EPOCH)
+ .unwrap_or_else(|_| Duration::from_secs(0))
+ .as_secs()
}
diff --git a/src/chat.rs b/src/chat.rs
index 565f08b..8b59bf6 100644
--- a/src/chat.rs
+++ b/src/chat.rs
@@ -1,30 +1,27 @@
-use std::collections::HashMap;
-
+use crate::AppState;
use axum::extract::ws::{Message, WebSocket};
use futures_util::{SinkExt, StreamExt};
use serde::{Deserialize, Serialize};
+use std::collections::HashMap;
use tokio::sync::{RwLock, mpsc};
use uuid::Uuid;
-use crate::{AppState, db};
-
#[derive(Clone)]
pub struct ChatClient {
tx: mpsc::UnboundedSender,
is_idle: bool,
}
-#[derive(Serialize, Clone, Debug, PartialEq, Eq)]
+#[derive(Serialize, Clone)]
pub struct OnlineUser {
pub user_id: Uuid,
- pub online: bool,
pub idle: bool,
}
#[derive(Default)]
pub struct ChatHub {
- // user_id -> connection_id -> client
- clients: RwLock>>,
+ // user_id -> client
+ clients: RwLock>,
}
#[derive(Serialize, Clone)]
@@ -48,113 +45,85 @@ pub enum ServerEvent {
#[derive(Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum ClientEvent {
+ // Currently no interactive client events for the general chat WS
Ping,
SetIdleStatus { is_idle: bool },
}
impl ChatHub {
- pub async fn add_client(
- &self,
- user_id: Uuid,
- connection_id: Uuid,
- tx: mpsc::UnboundedSender,
- ) -> Option {
+ pub async fn add_client(&self, user_id: Uuid, tx: mpsc::UnboundedSender) {
let mut clients = self.clients.write().await;
- let previous = aggregate_presence(clients.get(&user_id), user_id);
-
- clients
- .entry(user_id)
- .or_default()
- .insert(connection_id, ChatClient { tx, is_idle: false });
-
- let current = aggregate_presence(clients.get(&user_id), user_id);
- presence_delta(previous, current)
+ clients.insert(user_id, ChatClient { tx, is_idle: false });
}
- pub async fn remove_client(&self, user_id: Uuid, connection_id: Uuid) -> Option {
+ pub async fn remove_client(&self, user_id: Uuid) {
let mut clients = self.clients.write().await;
- let previous = aggregate_presence(clients.get(&user_id), user_id);
-
- if let Some(connections) = clients.get_mut(&user_id) {
- connections.remove(&connection_id);
- if connections.is_empty() {
- clients.remove(&user_id);
- }
- }
-
- let current = aggregate_presence(clients.get(&user_id), user_id);
- presence_delta(previous, current)
+ clients.remove(&user_id);
}
- pub async fn get_online_users_for(&self, visible_user_ids: &[Uuid]) -> Vec {
+ pub async fn get_online_users(&self) -> Vec {
let clients = self.clients.read().await;
- visible_user_ids
+ clients
.iter()
- .filter_map(|user_id| aggregate_presence(clients.get(user_id), *user_id))
- .filter(|presence| presence.online)
+ .map(|(id, client)| OnlineUser {
+ user_id: *id,
+ idle: client.is_idle,
+ })
.collect()
}
- pub async fn broadcast_to_many(&self, user_ids: Vec, event: ServerEvent) {
+ pub async fn broadcast_all(&self, event: ServerEvent) {
let clients = self.clients.read().await;
- for user_id in user_ids {
- if let Some(connections) = clients.get(&user_id) {
- for client in connections.values() {
- let _ = client.tx.send(event.clone());
- }
- }
+ for client in clients.values() {
+ let _ = client.tx.send(event.clone());
}
}
pub async fn broadcast_to_user(&self, user_id: Uuid, event: ServerEvent) {
let clients = self.clients.read().await;
- if let Some(connections) = clients.get(&user_id) {
- for client in connections.values() {
+ if let Some(client) = clients.get(&user_id) {
+ let _ = client.tx.send(event);
+ }
+ }
+
+ pub async fn broadcast_to_many(&self, user_ids: Vec, event: ServerEvent) {
+ let clients = self.clients.read().await;
+ for user_id in user_ids {
+ if let Some(client) = clients.get(&user_id) {
let _ = client.tx.send(event.clone());
}
}
}
- pub async fn set_idle_status(
- &self,
- user_id: Uuid,
- connection_id: Uuid,
- is_idle: bool,
- ) -> Option {
- let mut clients = self.clients.write().await;
- let previous = aggregate_presence(clients.get(&user_id), user_id);
-
- if let Some(connections) = clients.get_mut(&user_id)
- && let Some(client) = connections.get_mut(&connection_id)
+ pub async fn set_idle_status(&self, user_id: Uuid, is_idle: bool) {
{
- client.is_idle = is_idle;
+ let mut clients = self.clients.write().await;
+ if let Some(client) = clients.get_mut(&user_id) {
+ client.is_idle = is_idle;
+ }
}
-
- let current = aggregate_presence(clients.get(&user_id), user_id);
- presence_delta(previous, current)
+ self.broadcast_all(ServerEvent::UserPresence {
+ user_id,
+ online: true,
+ idle: is_idle,
+ })
+ .await;
}
}
pub async fn handle_socket(state: AppState, socket: WebSocket, user_id: Uuid) {
let (mut ws_sender, mut ws_receiver) = socket.split();
let (tx, mut rx) = mpsc::unbounded_channel::();
- let connection_id = Uuid::new_v4();
- if let Some(presence) = state.chat.add_client(user_id, connection_id, tx).await {
- if let Ok(visible_user_ids) = db::list_visible_user_ids(&state.db, user_id).await {
- state
- .chat
- .broadcast_to_many(
- visible_user_ids,
- ServerEvent::UserPresence {
- user_id: presence.user_id,
- online: presence.online,
- idle: presence.idle,
- },
- )
- .await;
- }
- }
+ state.chat.add_client(user_id, tx).await;
+ state
+ .chat
+ .broadcast_all(ServerEvent::UserPresence {
+ user_id,
+ online: true,
+ idle: false,
+ })
+ .await;
let send_task = tokio::spawn(async move {
while let Some(event) = rx.recv().await {
@@ -173,24 +142,8 @@ pub async fn handle_socket(state: AppState, socket: WebSocket, user_id: Uuid) {
Message::Text(text) => {
if let Ok(ClientEvent::SetIdleStatus { is_idle }) =
serde_json::from_str::(&text)
- && let Some(presence) = state
- .chat
- .set_idle_status(user_id, connection_id, is_idle)
- .await
- && let Ok(visible_user_ids) =
- db::list_visible_user_ids(&state.db, user_id).await
{
- state
- .chat
- .broadcast_to_many(
- visible_user_ids,
- ServerEvent::UserPresence {
- user_id: presence.user_id,
- online: presence.online,
- idle: presence.idle,
- },
- )
- .await;
+ state.chat.set_idle_status(user_id, is_idle).await;
}
}
_ => {}
@@ -198,157 +151,13 @@ pub async fn handle_socket(state: AppState, socket: WebSocket, user_id: Uuid) {
}
send_task.abort();
- if let Some(presence) = state.chat.remove_client(user_id, connection_id).await
- && let Ok(visible_user_ids) = db::list_visible_user_ids(&state.db, user_id).await
- {
- state
- .chat
- .broadcast_to_many(
- visible_user_ids,
- ServerEvent::UserPresence {
- user_id: presence.user_id,
- online: presence.online,
- idle: presence.idle,
- },
- )
- .await;
- }
-}
-
-fn aggregate_presence(
- connections: Option<&HashMap>,
- user_id: Uuid,
-) -> Option {
- let connections = connections?;
- if connections.is_empty() {
- return None;
- }
-
- let idle = connections.values().all(|client| client.is_idle);
- Some(OnlineUser {
- user_id,
- online: true,
- idle,
- })
-}
-
-fn presence_delta(previous: Option, current: Option) -> Option {
- match (previous, current) {
- (None, None) => None,
- (Some(prev), Some(curr)) if prev == curr => None,
- (Some(prev), None) => Some(OnlineUser {
- user_id: prev.user_id,
+ state.chat.remove_client(user_id).await;
+ state
+ .chat
+ .broadcast_all(ServerEvent::UserPresence {
+ user_id,
online: false,
idle: false,
- }),
- (_, Some(curr)) => Some(curr),
- }
-}
-
-#[cfg(test)]
-mod tests {
- use super::{ChatHub, OnlineUser, ServerEvent};
- use serde_json::json;
- use tokio::sync::mpsc;
- use uuid::Uuid;
-
- #[tokio::test]
- async fn multiple_connections_keep_user_online_until_last_disconnect() {
- let hub = ChatHub::default();
- let user_id = Uuid::new_v4();
- let first_connection = Uuid::new_v4();
- let second_connection = Uuid::new_v4();
- let (tx1, _rx1) = mpsc::unbounded_channel();
- let (tx2, _rx2) = mpsc::unbounded_channel();
-
- let first_presence = hub.add_client(user_id, first_connection, tx1).await;
- let second_presence = hub.add_client(user_id, second_connection, tx2).await;
- let after_first_disconnect = hub.remove_client(user_id, first_connection).await;
- let after_last_disconnect = hub.remove_client(user_id, second_connection).await;
-
- assert_eq!(
- first_presence,
- Some(OnlineUser {
- user_id,
- online: true,
- idle: false,
- })
- );
- assert_eq!(second_presence, None);
- assert_eq!(after_first_disconnect, None);
- assert_eq!(
- after_last_disconnect,
- Some(OnlineUser {
- user_id,
- online: false,
- idle: false,
- })
- );
- }
-
- #[tokio::test]
- async fn idle_status_only_flips_when_all_connections_are_idle() {
- let hub = ChatHub::default();
- let user_id = Uuid::new_v4();
- let first_connection = Uuid::new_v4();
- let second_connection = Uuid::new_v4();
- let (tx1, _rx1) = mpsc::unbounded_channel();
- let (tx2, _rx2) = mpsc::unbounded_channel();
-
- hub.add_client(user_id, first_connection, tx1).await;
- hub.add_client(user_id, second_connection, tx2).await;
-
- let first_idle = hub.set_idle_status(user_id, first_connection, true).await;
- let second_idle = hub.set_idle_status(user_id, second_connection, true).await;
- let active_again = hub.set_idle_status(user_id, first_connection, false).await;
-
- assert_eq!(first_idle, None);
- assert_eq!(
- second_idle,
- Some(OnlineUser {
- user_id,
- online: true,
- idle: true,
- })
- );
- assert_eq!(
- active_again,
- Some(OnlineUser {
- user_id,
- online: true,
- idle: false,
- })
- );
- }
-
- #[tokio::test]
- async fn broadcast_to_user_reaches_all_active_connections() {
- let hub = ChatHub::default();
- let user_id = Uuid::new_v4();
- let first_connection = Uuid::new_v4();
- let second_connection = Uuid::new_v4();
- let (tx1, mut rx1) = mpsc::unbounded_channel();
- let (tx2, mut rx2) = mpsc::unbounded_channel();
-
- hub.add_client(user_id, first_connection, tx1).await;
- hub.add_client(user_id, second_connection, tx2).await;
-
- hub.broadcast_to_user(
- user_id,
- ServerEvent::DmCreated {
- other_user_id: Uuid::new_v4(),
- message: json!({ "body": "hello" }),
- },
- )
+ })
.await;
-
- assert!(matches!(
- rx1.try_recv(),
- Ok(ServerEvent::DmCreated { message, .. }) if message["body"] == "hello"
- ));
- assert!(matches!(
- rx2.try_recv(),
- Ok(ServerEvent::DmCreated { message, .. }) if message["body"] == "hello"
- ));
- }
}
diff --git a/src/config.rs b/src/config.rs
index fbe7ecb..8f55aa0 100644
--- a/src/config.rs
+++ b/src/config.rs
@@ -1,11 +1,8 @@
-use anyhow::{Context, Result, anyhow};
-use reqwest::Url;
+use anyhow::{Context, Result};
#[derive(Clone, Debug)]
pub struct Settings {
pub port: u16,
- pub app_base_url: String,
- pub app_origin: String,
pub database_url: String,
pub oidc_client_id: String,
pub oidc_client_secret: String,
@@ -14,7 +11,8 @@ pub struct Settings {
pub oidc_userinfo_url: String,
pub oidc_redirect_url: String,
pub oidc_scopes: String,
- pub session_cookie_secure: bool,
+ pub session_secret: String,
+ pub cookie_secure: bool,
pub stun_urls: Vec,
pub turn_urls: Vec,
pub turn_username: Option,
@@ -33,17 +31,11 @@ pub struct MediaSettings {
impl Settings {
pub fn from_env() -> Result {
- let app_base_url = required("APP_BASE_URL")?;
- let app_url = Url::parse(&app_base_url).context("APP_BASE_URL must be a valid URL")?;
- let app_origin = origin_from_url(&app_url)?;
-
Ok(Self {
port: std::env::var("PORT")
.unwrap_or_else(|_| "3000".into())
.parse()
.context("PORT must be a valid u16")?,
- app_base_url: trim_trailing_slash(&app_base_url),
- app_origin,
database_url: required("DATABASE_URL")?,
oidc_client_id: required("OIDC_CLIENT_ID")?,
oidc_client_secret: required("OIDC_CLIENT_SECRET")?,
@@ -53,7 +45,11 @@ impl Settings {
oidc_redirect_url: required("OIDC_REDIRECT_URL")?,
oidc_scopes: std::env::var("OIDC_SCOPES")
.unwrap_or_else(|_| "openid profile email".to_string()),
- session_cookie_secure: requires_secure_cookie(&app_url)?,
+ session_secret: required("SESSION_SECRET")?,
+ cookie_secure: std::env::var("COOKIE_SECURE")
+ .unwrap_or_else(|_| "false".into())
+ .parse()
+ .context("COOKIE_SECURE must be true/false")?,
stun_urls: parse_csv_env("STUN_URLS", "stun:stun.l.google.com:19302"),
turn_urls: parse_csv_env("TURN_URLS", ""),
turn_username: optional("TURN_USERNAME"),
@@ -69,7 +65,7 @@ impl MediaSettings {
let access_key_id = optional("R2_ACCESS_KEY_ID");
let secret_access_key = optional("R2_SECRET_ACCESS_KEY");
let bucket = optional("R2_BUCKET");
- let public_base_url = optional("MEDIA_BASE_URL").or_else(|| optional("R2_PUBLIC_BASE_URL"));
+ let public_base_url = optional("R2_PUBLIC_BASE_URL");
if account_id.is_none()
&& access_key_id.is_none()
@@ -85,9 +81,7 @@ impl MediaSettings {
access_key_id: access_key_id.context("missing env var R2_ACCESS_KEY_ID")?,
secret_access_key: secret_access_key.context("missing env var R2_SECRET_ACCESS_KEY")?,
bucket: bucket.context("missing env var R2_BUCKET")?,
- public_base_url: trim_trailing_slash(
- &public_base_url.context("missing env var MEDIA_BASE_URL")?,
- ),
+ public_base_url: public_base_url.context("missing env var R2_PUBLIC_BASE_URL")?,
}))
}
@@ -115,105 +109,3 @@ fn parse_csv_env(name: &str, default_value: &str) -> Vec {
.map(ToString::to_string)
.collect()
}
-
-fn origin_from_url(url: &Url) -> Result {
- let host = url
- .host_str()
- .ok_or_else(|| anyhow!("APP_BASE_URL must include a host"))?;
- let mut origin = format!("{}://{}", url.scheme(), host);
- if let Some(port) = url.port() {
- origin.push(':');
- origin.push_str(&port.to_string());
- }
- Ok(origin)
-}
-
-fn requires_secure_cookie(url: &Url) -> Result {
- match url.scheme() {
- "https" => Ok(true),
- "http" => {
- let host = url
- .host_str()
- .ok_or_else(|| anyhow!("APP_BASE_URL must include a host"))?;
- if matches!(host, "localhost" | "127.0.0.1" | "::1") {
- Ok(false)
- } else {
- Err(anyhow!(
- "APP_BASE_URL must use https outside localhost when session cookies are enabled"
- ))
- }
- }
- other => Err(anyhow!("APP_BASE_URL scheme {other} is not supported")),
- }
-}
-
-fn trim_trailing_slash(value: &str) -> String {
- value.trim_end_matches('/').to_string()
-}
-
-#[cfg(test)]
-mod tests {
- use super::{MediaSettings, origin_from_url, requires_secure_cookie, trim_trailing_slash};
- use reqwest::Url;
-
- #[test]
- fn origin_from_url_preserves_explicit_port() {
- let url = Url::parse("https://chat.example.com:8443/app").unwrap();
- let origin = origin_from_url(&url).unwrap();
-
- assert_eq!(origin, "https://chat.example.com:8443");
- }
-
- #[test]
- fn secure_cookie_is_required_for_https_origins() {
- let url = Url::parse("https://chat.example.com").unwrap();
-
- assert_eq!(requires_secure_cookie(&url).unwrap(), true);
- }
-
- #[test]
- fn localhost_http_is_allowed_without_secure_cookie() {
- let url = Url::parse("http://localhost:3000").unwrap();
-
- assert_eq!(requires_secure_cookie(&url).unwrap(), false);
- }
-
- #[test]
- fn non_localhost_http_is_rejected() {
- let url = Url::parse("http://chat.example.com").unwrap();
- let err = requires_secure_cookie(&url).unwrap_err();
-
- assert!(
- err.to_string()
- .contains("APP_BASE_URL must use https outside localhost")
- );
- }
-
- #[test]
- fn trim_trailing_slash_removes_only_suffix_slashes() {
- assert_eq!(
- trim_trailing_slash("https://chat.example.com///"),
- "https://chat.example.com"
- );
- assert_eq!(
- trim_trailing_slash("https://chat.example.com/app"),
- "https://chat.example.com/app"
- );
- }
-
- #[test]
- fn media_endpoint_url_uses_account_id() {
- let media = MediaSettings {
- account_id: "acct123".to_string(),
- access_key_id: "key".to_string(),
- secret_access_key: "secret".to_string(),
- bucket: "bucket".to_string(),
- public_base_url: "https://media.example.com".to_string(),
- };
-
- assert_eq!(
- media.endpoint_url(),
- "https://acct123.r2.cloudflarestorage.com"
- );
- }
-}
diff --git a/src/db.rs b/src/db.rs
index c86dd8c..42d980b 100644
--- a/src/db.rs
+++ b/src/db.rs
@@ -1,17 +1,15 @@
-use std::collections::HashSet;
-
use anyhow::{Result, anyhow};
use chrono::{Duration, Utc};
use sea_orm::{
ActiveModelTrait, ActiveValue::Set, ColumnTrait, Condition, DatabaseConnection,
DatabaseTransaction, EntityTrait, PaginatorTrait, QueryFilter, QueryOrder, QuerySelect,
- Statement, TransactionTrait, sea_query::OnConflict,
+ TransactionTrait, sea_query::OnConflict,
};
use uuid::Uuid;
use crate::{
entity::{
- attachments, channels, direct_messages, guild_members, guilds, invites, messages, sessions,
+ attachments, channels, direct_messages, guild_members, guilds, invites, messages,
soundboard_sounds, users,
},
models::{
@@ -28,88 +26,6 @@ pub async fn user_exists(db: &DatabaseConnection, user_id: Uuid) -> Result
Ok(count > 0)
}
-pub async fn create_session(
- db: &DatabaseConnection,
- session_id: &str,
- user_id: Uuid,
- expires_at: chrono::DateTime,
- user_agent_hash: Option,
- ip_hash: Option,
-) -> Result<()> {
- sessions::Entity::insert(sessions::ActiveModel {
- id: Set(session_id.to_string()),
- user_id: Set(user_id),
- expires_at: Set(expires_at),
- created_at: Set(Utc::now().fixed_offset()),
- last_seen_at: Set(Utc::now().fixed_offset()),
- revoked_at: Set(None),
- user_agent_hash: Set(user_agent_hash),
- ip_hash: Set(ip_hash),
- })
- .exec(db)
- .await?;
-
- Ok(())
-}
-
-pub async fn touch_active_session(
- db: &DatabaseConnection,
- session_id: &str,
-) -> Result