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 @@ - - 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 downloads: Windows Linux (RPM)
`, - 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 @@ - - 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> { - let session = sessions::Entity::find_by_id(session_id.to_string()) - .one(db) - .await?; - - let Some(session) = session else { - return Ok(None); - }; - - if session.revoked_at.is_some() || session.expires_at <= Utc::now().fixed_offset() { - return Ok(None); - } - - let user_id = session.user_id; - sessions::Entity::update_many() - .col_expr( - sessions::Column::LastSeenAt, - sea_orm::sea_query::Expr::value(Utc::now().fixed_offset()), - ) - .filter(sessions::Column::Id.eq(session_id.to_string())) - .exec(db) - .await?; - Ok(Some(user_id)) -} - -pub async fn revoke_session(db: &DatabaseConnection, session_id: &str) -> Result<()> { - if let Some(session) = sessions::Entity::find_by_id(session_id.to_string()) - .one(db) - .await? - { - sessions::Entity::update_many() - .col_expr( - sessions::Column::RevokedAt, - sea_orm::sea_query::Expr::value(Some(Utc::now().fixed_offset())), - ) - .filter(sessions::Column::Id.eq(session.id)) - .exec(db) - .await?; - } - - Ok(()) -} - -pub async fn cleanup_sessions(db: &DatabaseConnection) -> Result<()> { - sessions::Entity::delete_many() - .filter( - Condition::any() - .add(sessions::Column::ExpiresAt.lte(Utc::now().fixed_offset())) - .add(sessions::Column::RevokedAt.is_not_null()), - ) - .exec(db) - .await?; - Ok(()) -} - pub async fn upsert_user_from_oidc( db: &DatabaseConnection, oidc_sub: &str, @@ -190,58 +106,18 @@ pub async fn list_guilds_for_user(db: &DatabaseConnection, user_id: Uuid) -> Res .collect()) } -pub async fn list_visible_user_ids(db: &DatabaseConnection, user_id: Uuid) -> Result> { - let guild_ids: Vec = guild_members::Entity::find() - .filter(guild_members::Column::UserId.eq(user_id)) - .select_only() - .column(guild_members::Column::GuildId) - .into_tuple() - .all(db) - .await?; - - let mut visible = HashSet::from([user_id]); - - if !guild_ids.is_empty() { - let guild_users = guild_members::Entity::find() - .filter(guild_members::Column::GuildId.is_in(guild_ids)) - .all(db) - .await?; - visible.extend(guild_users.into_iter().map(|membership| membership.user_id)); - } - - let dm_rows = direct_messages::Entity::find() - .filter( - Condition::any() - .add(direct_messages::Column::SenderUserId.eq(user_id)) - .add(direct_messages::Column::RecipientUserId.eq(user_id)), - ) - .all(db) - .await?; - - for row in dm_rows { - if row.sender_user_id == user_id { - visible.insert(row.recipient_user_id); - } else { - visible.insert(row.sender_user_id); - } - } - - Ok(visible.into_iter().collect()) -} - pub async fn create_guild( db: &DatabaseConnection, owner_user_id: Uuid, name: &str, ) -> Result { - let txn = db.begin().await?; let guild = guilds::Entity::insert(guilds::ActiveModel { id: Set(Uuid::new_v4()), name: Set(name.to_string()), owner_user_id: Set(owner_user_id), ..Default::default() }) - .exec_with_returning(&txn) + .exec_with_returning(db) .await?; guild_members::Entity::insert(guild_members::ActiveModel { @@ -257,13 +133,17 @@ pub async fn create_guild( .do_nothing() .to_owned(), ) - .exec(&txn) + .exec(db) .await?; - txn.commit().await?; Ok(map_guild(guild)) } +pub async fn get_guild_by_id(db: &DatabaseConnection, guild_id: Uuid) -> Result> { + let row = guilds::Entity::find_by_id(guild_id).one(db).await?; + Ok(row.map(map_guild)) +} + pub async fn is_guild_owner( db: &DatabaseConnection, guild_id: Uuid, @@ -317,12 +197,7 @@ pub async fn create_invite( pub async fn join_invite(db: &DatabaseConnection, code: &str, user_id: Uuid) -> Result { let txn = db.begin().await?; - let invite = invites::Entity::find() - .from_raw_sql(Statement::from_sql_and_values( - sea_orm::DatabaseBackend::Postgres, - r#"SELECT * FROM invites WHERE code = $1 FOR UPDATE"#, - [code.into()], - )) + let invite = invites::Entity::find_by_id(code.to_string()) .one(&txn) .await? .ok_or_else(|| anyhow!("invite not found"))?; @@ -922,10 +797,7 @@ pub async fn create_sound( created_by_user_id: Uuid, name: &str, icon: &str, - object_key: &str, - media_url: &str, - mime_type: &str, - size_bytes: i64, + file_path: &str, ) -> Result { let model = soundboard_sounds::Entity::insert(soundboard_sounds::ActiveModel { id: Set(Uuid::new_v4()), @@ -933,13 +805,7 @@ pub async fn create_sound( created_by_user_id: Set(created_by_user_id), name: Set(name.to_string()), icon: Set(icon.to_string()), - object_key: Set(Some(object_key.to_string())), - media_url: Set(media_url.to_string()), - mime_type: Set(Some(mime_type.to_string())), - size_bytes: Set(Some(size_bytes)), - // Keep the legacy column populated until every deployment has applied - // the nullable migration and old fallback paths are fully removed. - file_path: Set(Some(media_url.to_string())), + file_path: Set(file_path.to_string()), ..Default::default() }) .exec_with_returning(db) @@ -953,11 +819,6 @@ pub async fn get_sound_by_id(db: &DatabaseConnection, id: Uuid) -> Result