diff --git a/.env.example b/.env.example
index c062b55..1390efc 100644
--- a/.env.example
+++ b/.env.example
@@ -1,5 +1,6 @@
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
@@ -16,13 +17,9 @@ 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
-R2_PUBLIC_BASE_URL=https://media.example.com
+MEDIA_BASE_URL=https://media.example.com
diff --git a/Cargo.lock b/Cargo.lock
index b496651..c44764d 100644
--- a/Cargo.lock
+++ b/Cargo.lock
@@ -675,13 +675,14 @@ dependencies = [
"chrono",
"dotenvy",
"futures-util",
- "jsonwebtoken",
"reqwest",
"sea-orm",
"sea-orm-migration",
"serde",
"serde_json",
+ "sha2",
"tokio",
+ "tower",
"tower-http",
"tracing",
"tracing-subscriber",
@@ -1648,21 +1649,6 @@ 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"
@@ -1825,16 +1811,6 @@ 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"
@@ -1967,16 +1943,6 @@ 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"
@@ -2765,18 +2731,6 @@ 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"
@@ -3131,7 +3085,6 @@ 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 b6a6c27..6dbd768 100644
--- a/Cargo.toml
+++ b/Cargo.toml
@@ -10,14 +10,15 @@ 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"] }
+sea-orm = { version = "1.1", default-features = false, features = ["sqlx-postgres", "runtime-tokio-rustls", "macros", "with-chrono", "with-uuid", "mock"] }
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"] }
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 af5e8f0..2708a60 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`)
-- Signed session cookie auth
+- HttpOnly 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,6 +41,11 @@ 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
@@ -55,7 +60,7 @@ Web UI is available at `http://localhost:${PORT}/`.
## Authentik setup notes
Create an Authentik OAuth2/OIDC provider + application and set:
-- Redirect URI: `http://localhost:3000/auth/callback`
+- Redirect URI: `${APP_BASE_URL}/auth/callback`
- Scopes including at least: `openid profile email`
If you change `PORT`, update `OIDC_REDIRECT_URL` and this redirect URI to match.
@@ -91,6 +96,7 @@ 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
@@ -98,6 +104,7 @@ 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 e04011c..1684077 100644
--- a/desktop/index.html
+++ b/desktop/index.html
@@ -6,12 +6,9 @@
Chattz
-
-
-
-
+
diff --git a/desktop/preload.js b/desktop/preload.js
index dfcc665..f43670e 100644
--- a/desktop/preload.js
+++ b/desktop/preload.js
@@ -1,18 +1,20 @@
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) => 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))
+ onUpdateState: (callback) => onIpc('update-state', callback),
+ onUpdateAvailable: (callback) => onIpc('update-available', callback),
+ onUpdateDownloaded: (callback) => onIpc('update-downloaded', callback),
+ onUpdateError: (callback) => onIpc('update-error', callback)
});
diff --git a/main.js b/main.js
index 3a32c94..1baa6e2 100644
--- a/main.js
+++ b/main.js
@@ -1,8 +1,6 @@
const { app, BrowserWindow, session, desktopCapturer, ipcMain, clipboard } = require('electron');
const { autoUpdater } = require('electron-updater');
const path = require('path');
-const fs = require('fs');
-const url = require('url');
const updateState = {
status: 'idle', // idle | checking | available | downloading | downloaded | not-available | error
@@ -19,33 +17,6 @@ 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();
@@ -84,8 +55,9 @@ async function runUpdateCheck(reason = 'manual') {
}
function createWindow() {
- // Create a persistent session for chattz to keep the user logged in
const sess = session.fromPartition('persist:chattz');
+ const backendUrl = (process.env.CHATTZ_URL || process.env.APP_BASE_URL || 'http://localhost:3000').replace(/\/$/, '');
+ const backendOrigin = new URL(backendUrl).origin;
const win = new BrowserWindow({
width: 1200,
@@ -94,7 +66,6 @@ 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')
}
@@ -103,32 +74,33 @@ function createWindow() {
win.setAutoHideMenuBar(true);
win.setMenuBarVisibility(true);
- // Auto-approve media permissions (camera, microphone)
sess.setPermissionCheckHandler((webContents, permission) => {
- if (permission === 'media' || permission === 'clipboard-read' || permission === 'clipboard-write') {
- return true;
- }
- return false;
+ const origin = safeOrigin(webContents.getURL());
+ if (origin !== backendOrigin) return false;
+ return permission === 'media' || permission === 'clipboard-write';
});
sess.setPermissionRequestHandler((webContents, permission, callback) => {
- if (permission === 'media' || permission === 'clipboard-read' || permission === 'clipboard-write') {
+ const origin = safeOrigin(webContents.getURL());
+ if (origin === backendOrigin && (permission === 'media' || 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); // Reject safely
+ callback(null);
}
}).catch(err => {
console.error("Failed to get desktop sources for screen share", err);
@@ -136,55 +108,9 @@ function createWindow() {
});
});
- 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();
+ win.loadURL(backendUrl).catch((err) => {
+ console.error(`Failed to load desktop app: ${err}`);
+ });
// Ensure renderer receives latest updater state after any (re)load.
// Delay broadcast by 300ms to give the renderer time to register its
@@ -203,40 +129,11 @@ 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');
@@ -314,3 +211,11 @@ app.on('window-all-closed', () => {
app.quit();
}
});
+
+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 1b3497b..a773454 100644
--- a/scripts/generate-html.js
+++ b/scripts/generate-html.js
@@ -5,26 +5,30 @@ 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: '',
+ desktop_updater: updaterMarkup,
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: `
-
-
-
`,
+ desktop_updater: updaterMarkup,
outPath: path.join(root, 'desktop', 'index.html'),
},
};
diff --git a/shared-html/index.template.html b/shared-html/index.template.html
index bbf97e8..3425611 100644
--- a/shared-html/index.template.html
+++ b/shared-html/index.template.html
@@ -5,12 +5,9 @@
Chattz
-
-
-
-
+
diff --git a/src/auth.rs b/src/auth.rs
index 667c64d..535953e 100644
--- a/src/auth.rs
+++ b/src/auth.rs
@@ -1,33 +1,25 @@
-use std::time::{Duration, SystemTime, UNIX_EPOCH};
-
-use anyhow::{Context, Result, anyhow};
+use anyhow::{Result, anyhow};
use axum::{
Json,
extract::{FromRef, FromRequestParts},
- http::{StatusCode, request::Parts},
+ http::{HeaderMap, StatusCode, request::Parts},
response::{IntoResponse, Response},
};
-use jsonwebtoken::{DecodingKey, EncodingKey, Header, Validation, decode, encode};
-use serde::{Deserialize, Serialize};
+use chrono::{Duration, Utc};
+use serde::Serialize;
+use sha2::{Digest, Sha256};
use uuid::Uuid;
use crate::{AppState, db};
-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,
-}
+pub const OAUTH_STATE_COOKIE: &str = "chattz_oauth_state";
+pub const SESSION_COOKIE: &str = "chattz_session";
+const SESSION_TTL_DAYS: i64 = 30;
#[derive(Debug, Clone)]
pub struct AuthUser {
pub id: Uuid,
+ pub session_id: String,
}
#[derive(Debug)]
@@ -44,6 +36,13 @@ 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,
@@ -57,6 +56,13 @@ impl ApiError {
message: msg.to_string(),
}
}
+
+ pub fn service_unavailable(msg: &str) -> Self {
+ Self {
+ status: StatusCode::SERVICE_UNAVAILABLE,
+ message: msg.to_string(),
+ }
+ }
}
#[derive(Serialize)]
@@ -84,29 +90,25 @@ impl From for ApiError {
impl FromRequestParts for AuthUser
where
- AppState: axum::extract::FromRef,
+ AppState: 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 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)
+ let user_id = db::touch_active_session(&app.db, &session_id)
.await
- .map_err(|_| ApiError::unauthorized("session user not found"))?;
- if !exists {
- return Err(ApiError::unauthorized("session user not found"));
- }
+ .map_err(|_| ApiError::unauthorized("invalid or expired session"))?
+ .ok_or_else(|| ApiError::unauthorized("invalid or expired session"))?;
- Ok(Self { id: user_id })
+ Ok(Self {
+ id: user_id,
+ session_id,
+ })
}
}
@@ -114,6 +116,14 @@ 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}",
@@ -123,76 +133,32 @@ 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}=; Path=/; HttpOnly; SameSite=Lax; Max-Age=0{secure_flag}",
- name = OAUTH_STATE_COOKIE,
+ "{name}={value}; Path=/; HttpOnly; SameSite=Strict; Max-Age={ttl}{secure_flag}",
+ name = SESSION_COOKIE,
+ ttl = Duration::days(SESSION_TTL_DAYS).num_seconds(),
secure_flag = if secure { "; Secure" } else { "" }
)
}
-pub fn read_oauth_state_from_headers(headers: &axum::http::HeaderMap) -> Option {
+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 {
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 == OAUTH_STATE_COOKIE).then(|| value.to_string())
+ (key == cookie_name).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 {
@@ -201,37 +167,113 @@ pub fn validate_oauth_state(expected_cookie: Option, query_state: &str)
Ok(())
}
-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
+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_query_token(parts: &Parts) -> Option {
- let query = parts.uri.query()?;
+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 { "" }
+ )
+}
- // 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());
- }
+#[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"
+ );
}
- None
-}
-fn now_ts() -> u64 {
- SystemTime::now()
- .duration_since(UNIX_EPOCH)
- .unwrap_or_else(|_| Duration::from_secs(0))
- .as_secs()
+ #[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"));
+ }
}
diff --git a/src/chat.rs b/src/chat.rs
index 8b59bf6..565f08b 100644
--- a/src/chat.rs
+++ b/src/chat.rs
@@ -1,27 +1,30 @@
-use crate::AppState;
+use std::collections::HashMap;
+
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)]
+#[derive(Serialize, Clone, Debug, PartialEq, Eq)]
pub struct OnlineUser {
pub user_id: Uuid,
+ pub online: bool,
pub idle: bool,
}
#[derive(Default)]
pub struct ChatHub {
- // user_id -> client
- clients: RwLock>,
+ // user_id -> connection_id -> client
+ clients: RwLock>>,
}
#[derive(Serialize, Clone)]
@@ -45,85 +48,113 @@ 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, tx: mpsc::UnboundedSender) {
+ pub async fn add_client(
+ &self,
+ user_id: Uuid,
+ connection_id: Uuid,
+ tx: mpsc::UnboundedSender,
+ ) -> Option {
let mut clients = self.clients.write().await;
- clients.insert(user_id, ChatClient { tx, is_idle: false });
- }
+ let previous = aggregate_presence(clients.get(&user_id), user_id);
- pub async fn remove_client(&self, user_id: Uuid) {
- let mut clients = self.clients.write().await;
- clients.remove(&user_id);
- }
-
- pub async fn get_online_users(&self) -> Vec {
- let clients = self.clients.read().await;
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)
+ }
+
+ pub async fn remove_client(&self, user_id: Uuid, connection_id: Uuid) -> 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) {
+ 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)
+ }
+
+ pub async fn get_online_users_for(&self, visible_user_ids: &[Uuid]) -> Vec {
+ let clients = self.clients.read().await;
+ visible_user_ids
.iter()
- .map(|(id, client)| OnlineUser {
- user_id: *id,
- idle: client.is_idle,
- })
+ .filter_map(|user_id| aggregate_presence(clients.get(user_id), *user_id))
+ .filter(|presence| presence.online)
.collect()
}
- pub async fn broadcast_all(&self, event: ServerEvent) {
- let clients = self.clients.read().await;
- 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(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) {
+ if let Some(connections) = clients.get(&user_id) {
+ for client in connections.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() {
let _ = client.tx.send(event.clone());
}
}
}
- pub async fn set_idle_status(&self, user_id: Uuid, is_idle: bool) {
+ 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)
{
- let mut clients = self.clients.write().await;
- if let Some(client) = clients.get_mut(&user_id) {
- client.is_idle = is_idle;
- }
+ client.is_idle = is_idle;
}
- self.broadcast_all(ServerEvent::UserPresence {
- user_id,
- online: true,
- idle: is_idle,
- })
- .await;
+
+ let current = aggregate_presence(clients.get(&user_id), user_id);
+ presence_delta(previous, current)
}
}
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();
- state.chat.add_client(user_id, tx).await;
- state
- .chat
- .broadcast_all(ServerEvent::UserPresence {
- user_id,
- online: true,
- idle: false,
- })
- .await;
+ 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;
+ }
+ }
let send_task = tokio::spawn(async move {
while let Some(event) = rx.recv().await {
@@ -142,8 +173,24 @@ 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.set_idle_status(user_id, is_idle).await;
+ state
+ .chat
+ .broadcast_to_many(
+ visible_user_ids,
+ ServerEvent::UserPresence {
+ user_id: presence.user_id,
+ online: presence.online,
+ idle: presence.idle,
+ },
+ )
+ .await;
}
}
_ => {}
@@ -151,13 +198,157 @@ pub async fn handle_socket(state: AppState, socket: WebSocket, user_id: Uuid) {
}
send_task.abort();
- state.chat.remove_client(user_id).await;
- state
- .chat
- .broadcast_all(ServerEvent::UserPresence {
- user_id,
+ 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,
online: false,
idle: false,
- })
- .await;
+ }),
+ (_, 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 8f55aa0..fbe7ecb 100644
--- a/src/config.rs
+++ b/src/config.rs
@@ -1,8 +1,11 @@
-use anyhow::{Context, Result};
+use anyhow::{Context, Result, anyhow};
+use reqwest::Url;
#[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,
@@ -11,8 +14,7 @@ pub struct Settings {
pub oidc_userinfo_url: String,
pub oidc_redirect_url: String,
pub oidc_scopes: String,
- pub session_secret: String,
- pub cookie_secure: bool,
+ pub session_cookie_secure: bool,
pub stun_urls: Vec,
pub turn_urls: Vec,
pub turn_username: Option,
@@ -31,11 +33,17 @@ 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")?,
@@ -45,11 +53,7 @@ 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_secret: required("SESSION_SECRET")?,
- cookie_secure: std::env::var("COOKIE_SECURE")
- .unwrap_or_else(|_| "false".into())
- .parse()
- .context("COOKIE_SECURE must be true/false")?,
+ session_cookie_secure: requires_secure_cookie(&app_url)?,
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"),
@@ -65,7 +69,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("R2_PUBLIC_BASE_URL");
+ let public_base_url = optional("MEDIA_BASE_URL").or_else(|| optional("R2_PUBLIC_BASE_URL"));
if account_id.is_none()
&& access_key_id.is_none()
@@ -81,7 +85,9 @@ 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: public_base_url.context("missing env var R2_PUBLIC_BASE_URL")?,
+ public_base_url: trim_trailing_slash(
+ &public_base_url.context("missing env var MEDIA_BASE_URL")?,
+ ),
}))
}
@@ -109,3 +115,105 @@ 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 42d980b..c86dd8c 100644
--- a/src/db.rs
+++ b/src/db.rs
@@ -1,15 +1,17 @@
+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,
- TransactionTrait, sea_query::OnConflict,
+ Statement, TransactionTrait, sea_query::OnConflict,
};
use uuid::Uuid;
use crate::{
entity::{
- attachments, channels, direct_messages, guild_members, guilds, invites, messages,
+ attachments, channels, direct_messages, guild_members, guilds, invites, messages, sessions,
soundboard_sounds, users,
},
models::{
@@ -26,6 +28,88 @@ 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