Merge pull request #1229 from J3vb/fix/security-scan-2026-07-22

fix(security): 2026-07-22 scan remediation, D13 permission consolidation, dead-code sweep, lint zero
This commit is contained in:
J3vb
2026-07-23 17:08:53 +02:00
committed by GitHub
193 changed files with 4151 additions and 12675 deletions
File diff suppressed because it is too large Load Diff
-619
View File
@@ -1,619 +0,0 @@
{
"timestamp": "2026-03-31T09:23:48.747Z",
"sessions": [
{
"date": "2025-03-24",
"summary": "Comprehensive file reference documentation review — 28 bugs fixed, all docs updated to match code reality",
"tasksCompleted": 28,
"modulesTouched": [
"api",
"auth",
"db",
"ws",
"voice",
"permissions",
"config",
"stores",
"components",
"tauri-rust",
"protocol",
"tests",
"docs",
"security",
"ci"
],
"fileName": "2025-03-24-file-reference-review.md"
},
{
"date": "2026-03-17",
"summary": "Completed all 13 CEO review fixes and raised server test coverage to 80%+",
"tasksCompleted": 14,
"modulesTouched": [
"admin",
"api",
"auth",
"db",
"ws",
"config",
"stores",
"tauri-rust",
"tests",
"ci"
],
"fileName": "2026-03-17-ceo-fixes-and-coverage.md"
},
{
"date": "2026-03-17",
"summary": "CEO plan review sections 1-8 (HOLD SCOPE) for tauri-migration branch",
"tasksCompleted": 0,
"modulesTouched": [
"admin",
"api",
"auth",
"db",
"ws",
"storage",
"config",
"stores",
"pages",
"tauri-rust",
"tests",
"security",
"ci"
],
"fileName": "2026-03-17-ceo-review-sections-1-8.md"
},
{
"date": "2026-03-17",
"summary": "CEO plan review of tauri-migration branch — HOLD SCOPE",
"tasksCompleted": 0,
"modulesTouched": [
"admin",
"api",
"auth",
"db",
"ws",
"voice",
"permissions",
"storage",
"components",
"pages",
"tauri-rust",
"protocol",
"tests",
"security"
],
"fileName": "2026-03-17-ceo-review.md"
},
{
"date": "2026-03-17",
"summary": "Fixed all golangci-lint issues blocking CI, set up project brain vault",
"tasksCompleted": 2,
"modulesTouched": [
"api",
"ws",
"voice",
"tauri-rust",
"tests",
"docs",
"ci"
],
"fileName": "2026-03-17-lint-fixes-and-vault-setup.md"
},
{
"date": "2026-03-17",
"summary": "Pre-landing review of tauri-migration, commit, push, and PR #15 to dev",
"tasksCompleted": 0,
"modulesTouched": [
"auth",
"db",
"config",
"tauri-rust",
"tests",
"docs",
"security",
"ci"
],
"fileName": "2026-03-17-merge-review-and-pr.md"
},
{
"date": "2026-03-18",
"summary": "Channel management, file uploads, URL previews, voice fixes, UX polish",
"tasksCompleted": 15,
"modulesTouched": [
"admin",
"api",
"auth",
"db",
"voice",
"storage",
"components",
"tauri-rust",
"ci"
],
"fileName": "2026-03-18-features.md"
},
{
"date": "2026-03-18",
"summary": "Added native E2E testing via WebView2 CDP",
"tasksCompleted": 1,
"modulesTouched": [
"api",
"auth",
"voice",
"storage",
"config",
"lib",
"tauri-rust",
"e2e",
"tests",
"docs",
"ci"
],
"fileName": "2026-03-18-native-e2e.md"
},
{
"date": "2026-03-18",
"summary": "Voice chat NAT fix, audio pipeline overhaul, debugging infrastructure",
"tasksCompleted": 12,
"modulesTouched": [
"api",
"auth",
"voice",
"stores",
"components",
"tauri-rust",
"tests",
"ci"
],
"fileName": "2026-03-18-voice-audio-fixes.md"
},
{
"date": "2026-03-19",
"summary": "Go code review (4 passes, 19 fixes), admin panel redesign, live server logs, audit log filters, console output cleanup",
"tasksCompleted": 7,
"modulesTouched": [
"admin",
"api",
"auth",
"db",
"ws",
"voice",
"permissions",
"storage",
"config",
"lib",
"components",
"tests",
"docs",
"security",
"ci"
],
"fileName": "2026-03-19-admin-panel-and-server-review.md"
},
{
"date": "2026-03-20",
"summary": "Camera button delay fix, video feed flicker fix, security hardening, documentation sync",
"tasksCompleted": 4,
"modulesTouched": [
"auth",
"db",
"voice",
"permissions",
"config",
"stores",
"pages",
"tauri-rust",
"protocol",
"tests",
"docs",
"security",
"ci"
],
"fileName": "2026-03-20-livekit-camera-fixes.md"
},
{
"date": "2026-03-20",
"summary": "",
"tasksCompleted": 0,
"modulesTouched": [
"admin",
"auth",
"db",
"ws",
"stores",
"protocol",
"tests",
"docs",
"security",
"ci"
],
"fileName": "2026-03-20-research-report.md"
},
{
"date": "2026-03-21",
"summary": "Implemented 30 actionable fixes from code review audit across server and client",
"tasksCompleted": 30,
"modulesTouched": [
"api",
"auth",
"db",
"ws",
"voice",
"storage",
"stores",
"components",
"pages",
"tauri-rust",
"e2e",
"tests",
"security",
"ci"
],
"fileName": "2026-03-21-code-review-fixes.md"
},
{
"date": "2026-03-22",
"summary": "Competitive research, feature roadmap, GitHub issues, vault build-out",
"tasksCompleted": 10,
"modulesTouched": [
"admin",
"api",
"auth",
"db",
"ws",
"voice",
"permissions",
"config",
"lib",
"stores",
"components",
"tauri-rust",
"protocol",
"e2e",
"tests",
"docs",
"security",
"ci"
],
"fileName": "2026-03-22-competitive-research.md"
},
{
"date": "2026-03-26",
"summary": "Implemented client-side fixes from Codex/code review findings and added targeted regression coverage.",
"tasksCompleted": 1,
"modulesTouched": [
"api",
"auth",
"ws",
"voice",
"config",
"stores",
"components",
"tests",
"ci"
],
"fileName": "2026-03-26-client-review-fixes.md"
},
{
"date": "2026-03-27",
"summary": "Code review (27 fixes), connection quality indicator, remember-password fix, login redesign with OC branding, new app icon.",
"tasksCompleted": 38,
"modulesTouched": [
"admin",
"api",
"auth",
"db",
"ws",
"voice",
"storage",
"config",
"lib",
"stores",
"components",
"tauri-rust",
"protocol",
"tests",
"security",
"ci"
],
"fileName": "2026-03-27-dual-model-code-review.md"
},
{
"date": "2026-03-28",
"summary": "Spec audit (18 files, 50 fixes), 143 unit tests, E2E overhaul, CSS injection fix",
"tasksCompleted": 8,
"modulesTouched": [
"api",
"auth",
"db",
"ws",
"voice",
"storage",
"lib",
"stores",
"components",
"tauri-rust",
"e2e",
"tests",
"docs",
"security",
"ci"
],
"fileName": "2026-03-28-specs-tests-e2e.md"
},
{
"date": "2026-03-29",
"summary": "Client-side 2FA integration — TOTP enrollment/disable UI, api.ts fixes, documentation sync",
"tasksCompleted": 5,
"modulesTouched": [
"admin",
"api",
"auth",
"db",
"voice",
"permissions",
"config",
"stores",
"components",
"e2e",
"tests",
"docs",
"security",
"ci"
],
"fileName": "2026-03-29-2fa-client-integration.md"
},
{
"date": "2026-03-29",
"summary": "Full-project security + code quality audit, 30 issues fixed",
"tasksCompleted": 27,
"modulesTouched": [
"admin",
"api",
"auth",
"db",
"ws",
"voice",
"permissions",
"config",
"lib",
"tests",
"security",
"ci"
],
"fileName": "2026-03-29-security-audit-fixes.md"
},
{
"date": "2026-03-30",
"summary": "Fixed 10 test quality bugs (BUG-058 through BUG-067)",
"tasksCompleted": 10,
"modulesTouched": [
"auth",
"voice",
"permissions",
"config",
"tauri-rust",
"e2e",
"tests",
"docs",
"ci"
],
"fileName": "2026-03-30-bug-remediation.md"
},
{
"date": "2026-03-30",
"summary": "Full documentation update — README, Changelog, Dashboard, task tracking",
"tasksCompleted": 5,
"modulesTouched": [
"api",
"auth",
"db",
"voice",
"permissions",
"config",
"components",
"e2e",
"tests",
"docs",
"ci"
],
"fileName": "2026-03-30-documentation-update.md"
},
{
"date": "2026-03-30",
"summary": "Sidebar stream preview + screenshare focus fix",
"tasksCompleted": 0,
"modulesTouched": [
"api",
"auth",
"voice",
"components",
"ci"
],
"fileName": "2026-03-30-stream-preview.md"
}
],
"progressOverTime": [
{
"date": "2025-03-24",
"cumulativeDone": 28,
"sessionDone": 28
},
{
"date": "2026-03-17",
"cumulativeDone": 42,
"sessionDone": 14
},
{
"date": "2026-03-17",
"cumulativeDone": 42,
"sessionDone": 0
},
{
"date": "2026-03-17",
"cumulativeDone": 42,
"sessionDone": 0
},
{
"date": "2026-03-17",
"cumulativeDone": 44,
"sessionDone": 2
},
{
"date": "2026-03-17",
"cumulativeDone": 44,
"sessionDone": 0
},
{
"date": "2026-03-18",
"cumulativeDone": 59,
"sessionDone": 15
},
{
"date": "2026-03-18",
"cumulativeDone": 60,
"sessionDone": 1
},
{
"date": "2026-03-18",
"cumulativeDone": 72,
"sessionDone": 12
},
{
"date": "2026-03-19",
"cumulativeDone": 79,
"sessionDone": 7
},
{
"date": "2026-03-20",
"cumulativeDone": 83,
"sessionDone": 4
},
{
"date": "2026-03-20",
"cumulativeDone": 83,
"sessionDone": 0
},
{
"date": "2026-03-21",
"cumulativeDone": 113,
"sessionDone": 30
},
{
"date": "2026-03-22",
"cumulativeDone": 123,
"sessionDone": 10
},
{
"date": "2026-03-26",
"cumulativeDone": 124,
"sessionDone": 1
},
{
"date": "2026-03-27",
"cumulativeDone": 162,
"sessionDone": 38
},
{
"date": "2026-03-28",
"cumulativeDone": 170,
"sessionDone": 8
},
{
"date": "2026-03-29",
"cumulativeDone": 175,
"sessionDone": 5
},
{
"date": "2026-03-29",
"cumulativeDone": 202,
"sessionDone": 27
},
{
"date": "2026-03-30",
"cumulativeDone": 212,
"sessionDone": 10
},
{
"date": "2026-03-30",
"cumulativeDone": 217,
"sessionDone": 5
},
{
"date": "2026-03-30",
"cumulativeDone": 217,
"sessionDone": 0
}
],
"streaks": {
"current": 5,
"longest": 6,
"totalSessions": 22
},
"lastSession": {
"date": "2026-03-30",
"summary": "Sidebar stream preview + screenshare focus fix",
"tasksCompleted": 0,
"modulesTouched": [
"api",
"auth",
"voice",
"components",
"ci"
]
},
"inProgress": [],
"recentlyDone": [
{
"id": "T-197",
"description": "Discord-style video grid with fixed 16:9 aspect ratio",
"date": "2026-03-30"
},
{
"id": "T-198",
"description": "Sidebar stream preview + screenshare focus fix",
"date": "2026-03-30"
},
{
"id": "T-199",
"description": "Fix 10 test quality bugs (BUG-058067) and resolve 115 TS type errors",
"date": "2026-03-30"
},
{
"id": "T-200",
"description": "Regenerate codemaps from current codebase",
"date": "2026-03-30"
},
{
"id": "T-201",
"description": "Full documentation update (README, Changelog, Dashboard, all docs)",
"date": "2026-03-30"
},
{
"id": "T-192",
"description": "Client 2FA enrollment/disable settings UI",
"date": "2026-03-29"
},
{
"id": "T-193",
"description": "Client 2FA test coverage — 27 new tests",
"date": "2026-03-29"
},
{
"id": "T-194",
"description": "Full regression validation pass — all green",
"date": "2026-03-29"
},
{
"id": "T-023",
"description": "Add TOTP 2FA support — Login challenge: DONE; Server endpoints: DONE; Client enrollment UI: DONE; Client tests: DONE",
"date": "2026-03-29"
},
{
"id": "T-190",
"description": "Propagate `context.Context` from WS upgrade through all handlers — added `ctx context.Context` field to Client struct (set from `r.Context()` on WS upgrade), updated `MessageHandler` type signature, threaded ctx through all 17 WS handlers across 9 files (chat, presence, reaction, voice, ping). Added `ExecContext`/`QueryRowContext`/`QueryContext`/`BeginTx` context-accepting methods to DB wrapper. Go build + all ws/api/auth/db tests pass",
"date": "2026-03-29"
}
]
}
+2 -1
View File
@@ -4,7 +4,8 @@
"project": ["src/**/*.ts"],
"ignore": [
"public/**",
"src-tauri/**"
"src-tauri/**",
"src/lib/protocolTypes.ts"
],
"ignoreDependencies": [
"@tauri-apps/cli"
+3 -32
View File
@@ -12,19 +12,16 @@
"@tauri-apps/api": "^2.10.1",
"@tauri-apps/plugin-dialog": "^2.6.0",
"@tauri-apps/plugin-fs": "^2.4.5",
"@tauri-apps/plugin-global-shortcut": "^2",
"@tauri-apps/plugin-http": "^2.5.7",
"@tauri-apps/plugin-notification": "^2",
"@tauri-apps/plugin-opener": "^2.5.3",
"@tauri-apps/plugin-process": "^2.3.1",
"@tauri-apps/plugin-store": "^2",
"@tauri-apps/plugin-updater": "^2.10.0",
"livekit-client": "^2.18.0",
"zod": "^4.3.6"
"livekit-client": "^2.18.0"
},
"devDependencies": {
"@eslint/js": "^9.39.4",
"@playwright/test": "^1",
"@stryker-mutator/api": "^9.6.0",
"@stryker-mutator/core": "^9.6.0",
"@stryker-mutator/typescript-checker": "^9.6.0",
"@stryker-mutator/vitest-runner": "^9.6.0",
@@ -3873,15 +3870,6 @@
"@tauri-apps/api": "^2.8.0"
}
},
"node_modules/@tauri-apps/plugin-global-shortcut": {
"version": "2.3.1",
"resolved": "https://registry.npmjs.org/@tauri-apps/plugin-global-shortcut/-/plugin-global-shortcut-2.3.1.tgz",
"integrity": "sha512-vr40W2N6G63dmBPaha1TsBQLLURXG538RQbH5vAm0G/ovVZyXJrmZR1HF1W+WneNloQvwn4dm8xzwpEXRW560g==",
"license": "MIT OR Apache-2.0",
"dependencies": {
"@tauri-apps/api": "^2.8.0"
}
},
"node_modules/@tauri-apps/plugin-http": {
"version": "2.5.7",
"resolved": "https://registry.npmjs.org/@tauri-apps/plugin-http/-/plugin-http-2.5.7.tgz",
@@ -3918,24 +3906,6 @@
"@tauri-apps/api": "^2.8.0"
}
},
"node_modules/@tauri-apps/plugin-store": {
"version": "2.4.2",
"resolved": "https://registry.npmjs.org/@tauri-apps/plugin-store/-/plugin-store-2.4.2.tgz",
"integrity": "sha512-0ClHS50Oq9HEvLPhNzTNFxbWVOqoAp3dRvtewQBeqfIQ0z5m3JRnOISIn2ZVPCrQC0MyGyhTS9DWhHjpigQE7A==",
"license": "MIT OR Apache-2.0",
"dependencies": {
"@tauri-apps/api": "^2.8.0"
}
},
"node_modules/@tauri-apps/plugin-updater": {
"version": "2.10.0",
"resolved": "https://registry.npmjs.org/@tauri-apps/plugin-updater/-/plugin-updater-2.10.0.tgz",
"integrity": "sha512-ljN8jPlnT0aSn8ecYhuBib84alxfMx6Hc8vJSKMJyzGbTPFZAC44T2I1QNFZssgWKrAlofvJqCC6Rr472JWfkQ==",
"license": "MIT OR Apache-2.0",
"dependencies": {
"@tauri-apps/api": "^2.10.1"
}
},
"node_modules/@testing-library/dom": {
"version": "10.4.1",
"resolved": "https://registry.npmjs.org/@testing-library/dom/-/dom-10.4.1.tgz",
@@ -8571,6 +8541,7 @@
"version": "4.3.6",
"resolved": "https://registry.npmjs.org/zod/-/zod-4.3.6.tgz",
"integrity": "sha512-rftlrkhHZOcjDwkGlnUtZZkvaPHCsDATp4pGpuOOMDaTdDDXF91wuVDJoWoPsKX/3YPQ5fHuF3STjcYyKr+Qhg==",
"dev": true,
"license": "MIT",
"funding": {
"url": "https://github.com/sponsors/colinhacks"
+2 -5
View File
@@ -32,6 +32,7 @@
"devDependencies": {
"@eslint/js": "^9.39.4",
"@playwright/test": "^1",
"@stryker-mutator/api": "^9.6.0",
"@stryker-mutator/core": "^9.6.0",
"@stryker-mutator/typescript-checker": "^9.6.0",
"@stryker-mutator/vitest-runner": "^9.6.0",
@@ -62,14 +63,10 @@
"@tauri-apps/api": "^2.10.1",
"@tauri-apps/plugin-dialog": "^2.6.0",
"@tauri-apps/plugin-fs": "^2.4.5",
"@tauri-apps/plugin-global-shortcut": "^2",
"@tauri-apps/plugin-http": "^2.5.7",
"@tauri-apps/plugin-notification": "^2",
"@tauri-apps/plugin-opener": "^2.5.3",
"@tauri-apps/plugin-process": "^2.3.1",
"@tauri-apps/plugin-store": "^2",
"@tauri-apps/plugin-updater": "^2.10.0",
"livekit-client": "^2.18.0",
"zod": "^4.3.6"
"livekit-client": "^2.18.0"
}
}
-67
View File
@@ -1578,16 +1578,6 @@ dependencies = [
"version_check",
]
[[package]]
name = "gethostname"
version = "1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1bd49230192a3797a9a4d6abe9b3eed6f7fa4c8a8a4947977c6f80025f92cbd8"
dependencies = [
"rustix",
"windows-link 0.2.1",
]
[[package]]
name = "getrandom"
version = "0.1.16"
@@ -1725,24 +1715,6 @@ version = "0.3.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "0cc23270f6e1808e30a928bdc84dea0b9b4136a8bc82338574f23baf47bbd280"
[[package]]
name = "global-hotkey"
version = "0.7.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b9247516746aa8e53411a0db9b62b0e24efbcf6a76e0ba73e5a91b512ddabed7"
dependencies = [
"crossbeam-channel",
"keyboard-types",
"objc2",
"objc2-app-kit",
"once_cell",
"serde",
"thiserror 2.0.18",
"windows-sys 0.59.0",
"x11rb",
"xkeysym",
]
[[package]]
name = "globset"
version = "0.4.18"
@@ -2979,7 +2951,6 @@ dependencies = [
"tauri-build",
"tauri-plugin-dialog",
"tauri-plugin-fs",
"tauri-plugin-global-shortcut",
"tauri-plugin-http",
"tauri-plugin-notification",
"tauri-plugin-opener",
@@ -4931,21 +4902,6 @@ dependencies = [
"url",
]
[[package]]
name = "tauri-plugin-global-shortcut"
version = "2.3.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "424af23c7e88d05e4a1a6fc2c7be077912f8c76bd7900fd50aa2b7cbf5a2c405"
dependencies = [
"global-hotkey",
"log",
"serde",
"serde_json",
"tauri",
"tauri-plugin",
"thiserror 2.0.18",
]
[[package]]
name = "tauri-plugin-http"
version = "2.5.7"
@@ -6887,23 +6843,6 @@ dependencies = [
"pkg-config",
]
[[package]]
name = "x11rb"
version = "0.13.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9993aa5be5a26815fe2c3eacfc1fde061fc1a1f094bf1ad2a18bf9c495dd7414"
dependencies = [
"gethostname",
"rustix",
"x11rb-protocol",
]
[[package]]
name = "x11rb-protocol"
version = "0.13.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "ea6fc2961e4ef194dcbfe56bb845534d0dc8098940c7e5c012a258bfec6701bd"
[[package]]
name = "xattr"
version = "1.6.1"
@@ -6914,12 +6853,6 @@ dependencies = [
"rustix",
]
[[package]]
name = "xkeysym"
version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b9cc00251562a284751c9973bace760d86c0276c471b4be569fe6b068ee97a56"
[[package]]
name = "yoke"
version = "0.8.1"
-1
View File
@@ -19,7 +19,6 @@ devtools = ["tauri/devtools"]
[dependencies]
tauri = { version = "2", features = ["tray-icon"] }
tauri-plugin-store = "2"
tauri-plugin-global-shortcut = "2"
tauri-plugin-notification = "2"
serde = { version = "1", features = ["derive"] }
serde_json = "1"
@@ -20,12 +20,6 @@
"core:window:allow-outer-position",
"core:window:allow-outer-size",
"core:window:allow-available-monitors",
"store:default",
"global-shortcut:default",
"global-shortcut:allow-register",
"global-shortcut:allow-unregister",
"global-shortcut:allow-unregister-all",
"global-shortcut:allow-is-registered",
"notification:default",
"notification:allow-notify",
"notification:allow-request-permission",
@@ -63,7 +57,6 @@
"http:allow-fetch-cancel",
"opener:default",
"dialog:default",
"updater:default",
"process:allow-restart",
"fs:default",
{
+41 -142
View File
@@ -9,10 +9,12 @@
// host. The webview fetches http://127.0.0.1:{port}/api/v1/... and the proxy
// opens a TLS connection to the real server, enforcing the same TOFU
// (Trust On First Use) fingerprint pinning as ws_proxy:
// - Unknown host → accept, persist the fingerprint, emit `cert-tofu`
// (status "trusted_first_use") so the UI can show the banner. HTTP is the
// FIRST TLS contact with a server (login precedes the WS connect), so this
// proxy — not ws_proxy — usually establishes the pin.
// - Unknown host → REJECT (502) and emit `cert-tofu` (status "first_use") so
// the UI prompts the user to confirm the fingerprint. Nothing is pinned or
// forwarded until the user explicitly accepts (accept_cert_fingerprint), so
// no credential is ever sent to an unconfirmed host. HTTP is the FIRST TLS
// contact with a server (login precedes the WS connect), so this proxy
// usually surfaces the first-use prompt. (F4/F8)
// - Pinned host → fingerprint must match or the connection is refused and a
// `cert-tofu` mismatch event fires (CertMismatchModal flow).
//
@@ -27,21 +29,17 @@
// - The accept loop exits after 5 consecutive errors to prevent CPU spin.
use log::{debug, error, info, warn};
use ring::digest::{digest, SHA256};
use std::collections::HashMap;
use std::net::IpAddr;
use std::sync::Arc;
use rustls::pki_types::ServerName;
use serde_json::Value;
use tauri::{AppHandle, Emitter, Runtime};
use tauri_plugin_store::StoreExt;
use tokio::io::{self, AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::Mutex;
use tokio::time::{timeout, Duration};
use crate::constants::CERTS_STORE;
use crate::livekit_proxy::cert_store_key;
use crate::tofu::{self, TofuOutcome};
/// Tauri-managed state: one running tunnel per remote host.
pub struct HttpProxyState {
@@ -139,132 +137,10 @@ pub async fn stop_http_proxy(
}
// ---------------------------------------------------------------------------
// TOFU verification (mirrors ws_proxy semantics; shared cert store)
// TOFU verification lives in the shared `tofu` module (crate::tofu):
// CaptureVerifier, cert_store_key, evaluate/decide, and the mismatch message.
// ---------------------------------------------------------------------------
/// Fingerprint captured during the TLS handshake.
type CapturedFingerprint = Arc<std::sync::Mutex<Option<String>>>;
/// Accepts the handshake while recording the leaf certificate's SHA-256
/// fingerprint; the TOFU decision happens immediately after the handshake,
/// before any request bytes are forwarded.
#[derive(Debug)]
struct CaptureVerifier {
captured: CapturedFingerprint,
}
impl CaptureVerifier {
fn new() -> (Self, CapturedFingerprint) {
let fp = Arc::new(std::sync::Mutex::new(None));
(Self { captured: fp.clone() }, fp)
}
}
impl rustls::client::danger::ServerCertVerifier for CaptureVerifier {
fn verify_server_cert(
&self,
end_entity: &rustls::pki_types::CertificateDer<'_>,
_intermediates: &[rustls::pki_types::CertificateDer<'_>],
_server_name: &rustls::pki_types::ServerName<'_>,
_ocsp_response: &[u8],
_now: rustls::pki_types::UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
let hash = digest(&SHA256, end_entity.as_ref());
let hex = hash
.as_ref()
.iter()
.map(|b| format!("{b:02x}"))
.collect::<Vec<_>>()
.join(":");
if let Ok(mut guard) = self.captured.lock() {
*guard = Some(hex);
}
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls12_signature(
message,
cert,
dss,
&rustls::crypto::ring::default_provider().signature_verification_algorithms,
)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls13_signature(
message,
cert,
dss,
&rustls::crypto::ring::default_provider().signature_verification_algorithms,
)
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
rustls::crypto::ring::default_provider()
.signature_verification_algorithms
.supported_schemes()
}
}
/// TOFU decision for `host` (cert-store key, i.e. without a default :443):
/// first use stores the pin, match passes, mismatch fails. Same store, same
/// save-rollback behavior, and same event payloads as ws_proxy::tofu_check.
fn tofu_check<R: Runtime>(
app: &AppHandle<R>,
host: &str,
fingerprint: &str,
) -> Result<String, String> {
let store = app
.store(CERTS_STORE)
.map_err(|e| format!("failed to open certs store: {e}"))?;
let stored = store.get(host).and_then(|v| {
if let Value::String(s) = v {
Some(s)
} else {
None
}
});
match stored {
None => {
let old_value = store.get(host);
store.set(host, Value::String(fingerprint.to_string()));
if let Err(e) = store.save() {
match old_value {
Some(v) => {
store.set(host, v);
}
None => {
let _ = store.delete(host);
}
}
return Err(format!("failed to persist cert fingerprint: {e}"));
}
Ok("trusted_first_use".to_string())
}
Some(ref stored_fp) if stored_fp == fingerprint => Ok("trusted".to_string()),
Some(stored_fp) => Err(format!(
"Certificate fingerprint changed for {host}.\n\
Stored: {stored_fp}\n\
Current: {fingerprint}\n\
This may indicate a man-in-the-middle attack or a server certificate rotation.\n\
Use accept_cert_fingerprint to trust the new certificate."
)),
}
}
// ---------------------------------------------------------------------------
// Proxy internals
// ---------------------------------------------------------------------------
@@ -400,7 +276,7 @@ async fn handle_connection<R: Runtime>(
let modified = rewrite_request_headers(&buf, remote_host);
// ── 2. TLS connect + TOFU check ──────────────────────────────────────
let (verifier, captured_fp) = CaptureVerifier::new();
let (verifier, captured_fp) = tofu::CaptureVerifier::new();
let tls_config = rustls::ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(Arc::new(verifier))
@@ -437,22 +313,44 @@ async fn handle_connection<R: Runtime>(
return Err("TLS handshake completed but no certificate fingerprint was captured".into());
}
let store_key = cert_store_key(remote_host);
match tofu_check(&app, &store_key, &fingerprint) {
Ok(status) => {
if status == "trusted_first_use" {
info!("[http_proxy] TOFU first-use pin for {}", store_key);
}
let store_key = tofu::cert_store_key(remote_host);
match tofu::evaluate(&app, &store_key, &fingerprint)? {
TofuOutcome::Trusted => {
let _ = app.emit(
"cert-tofu",
serde_json::json!({
"host": store_key,
"fingerprint": fingerprint,
"status": status,
"status": "trusted",
}),
);
}
Err(mismatch_msg) => {
// F4/F8: a first-use cert is NOT silently pinned or forwarded to. Reject
// the request (502) and surface the fingerprint so the user can confirm
// it (accept_cert_fingerprint) before any credential-bearing request is
// sent. The connect page's health check triggers this before login.
TofuOutcome::FirstUse => {
info!("[http_proxy] first-use cert for {} — awaiting user confirmation", store_key);
let _ = app.emit(
"cert-tofu",
serde_json::json!({
"host": store_key,
"fingerprint": fingerprint,
"status": "first_use",
}),
);
let _ = local
.write_all(
b"HTTP/1.1 502 Bad Gateway\r\nConnection: close\r\nContent-Length: 0\r\n\r\n",
)
.await;
return Err(format!(
"certificate for {store_key} is not yet trusted; confirm the fingerprint to continue"
)
.into());
}
TofuOutcome::Mismatch { stored } => {
let mismatch_msg = tofu::mismatch_message(&store_key, &stored, &fingerprint);
warn!(
"[http_proxy] TOFU check FAILED for {} — certificate fingerprint mismatch",
store_key
@@ -464,6 +362,7 @@ async fn handle_connection<R: Runtime>(
"fingerprint": fingerprint,
"status": "mismatch",
"message": mismatch_msg,
"storedFingerprint": stored,
}),
);
// Give the local fetch a clean HTTP failure instead of a reset.
+1 -1
View File
@@ -4,6 +4,7 @@ mod credentials;
mod http_proxy;
mod livekit_proxy;
mod ptt;
mod tofu;
mod tray;
mod update_commands;
mod ws_proxy;
@@ -12,7 +13,6 @@ mod ws_proxy;
pub fn run() {
match tauri::Builder::default()
.plugin(tauri_plugin_store::Builder::new().build())
.plugin(tauri_plugin_global_shortcut::Builder::new().build())
.plugin(tauri_plugin_notification::init())
.plugin(tauri_plugin_http::init())
.plugin(tauri_plugin_opener::init())
@@ -26,13 +26,10 @@
// - The accept loop exits after 5 consecutive errors to prevent CPU spin.
use log::{debug, error, info, warn};
use ring::digest::{digest, SHA256};
use std::net::IpAddr;
use std::sync::Arc;
use rustls::pki_types::ServerName;
use serde_json::Value;
use tauri::Runtime;
use tauri_plugin_store::StoreExt;
use tokio::io::{self, AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use tokio::sync::Mutex;
@@ -65,118 +62,16 @@ impl LiveKitProxyState {
}
// ---------------------------------------------------------------------------
// TLS certificate verifier — pinned fingerprint check
// TLS verification & cert-store helpers live in the shared `tofu` module
// (crate::tofu): PinnedVerifier, cert_store_key, load_stored_fingerprint.
// ---------------------------------------------------------------------------
use crate::constants::CERTS_STORE;
/// Verifies the server certificate against a known SHA-256 fingerprint.
/// Reuses the fingerprint stored by ws_proxy's TOFU handshake for the same
/// host, so LiveKit connections are pinned to the same certificate the user
/// already trusted during WebSocket setup.
#[derive(Debug)]
pub(crate) struct PinnedVerifier {
/// Expected SHA-256 colon-hex fingerprint (e.g. "aa:bb:cc:...").
expected_fingerprint: String,
}
impl PinnedVerifier {
pub(crate) fn new(expected_fingerprint: String) -> Self {
Self { expected_fingerprint }
}
}
impl rustls::client::danger::ServerCertVerifier for PinnedVerifier {
fn verify_server_cert(
&self,
end_entity: &rustls::pki_types::CertificateDer<'_>,
_intermediates: &[rustls::pki_types::CertificateDer<'_>],
_server_name: &rustls::pki_types::ServerName<'_>,
_ocsp_response: &[u8],
_now: rustls::pki_types::UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
let hash = digest(&SHA256, end_entity.as_ref());
let hex = hash
.as_ref()
.iter()
.map(|b| format!("{b:02x}"))
.collect::<Vec<_>>()
.join(":");
if hex == self.expected_fingerprint {
Ok(rustls::client::danger::ServerCertVerified::assertion())
} else {
Err(rustls::Error::General(format!(
"certificate fingerprint mismatch: expected {}, got {}",
self.expected_fingerprint, hex
)))
}
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls12_signature(
message,
cert,
dss,
&rustls::crypto::ring::default_provider().signature_verification_algorithms,
)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls13_signature(
message,
cert,
dss,
&rustls::crypto::ring::default_provider().signature_verification_algorithms,
)
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
rustls::crypto::ring::default_provider()
.signature_verification_algorithms
.supported_schemes()
}
}
use crate::tofu;
// ---------------------------------------------------------------------------
// Tauri commands
// ---------------------------------------------------------------------------
/// Produce the cert store key matching ws_proxy's format.
/// ws_proxy extracts the host from "wss://host/path" which omits port 443.
/// We normalise by stripping the default ":443" suffix so the keys match.
pub(crate) fn cert_store_key(remote_host: &str) -> String {
remote_host.strip_suffix(":443").unwrap_or(remote_host).to_string()
}
/// Load the stored certificate fingerprint for a host from the Tauri cert store.
pub(crate) fn load_stored_fingerprint<R: Runtime>(
app: &tauri::AppHandle<R>,
host: &str,
) -> Result<Option<String>, String> {
let store = app
.store(CERTS_STORE)
.map_err(|e| format!("failed to open certs store: {e}"))?;
Ok(store.get(host).and_then(|v| {
if let Value::String(s) = v {
Some(s)
} else {
None
}
}))
}
/// Start a local TCP proxy that tunnels LiveKit signal connections to the
/// remote OwnCord server over TLS, pinning the certificate to the fingerprint
/// already trusted via ws_proxy's TOFU handshake.
@@ -221,8 +116,8 @@ pub async fn start_livekit_proxy<R: Runtime>(
// have connected first (establishing the TOFU trust), so the fingerprint
// should already be stored. If not, reject — we refuse to connect without
// a pinned cert.
let store_key = cert_store_key(&remote_host);
let fingerprint = load_stored_fingerprint(&app, &store_key)?
let store_key = tofu::cert_store_key(&remote_host);
let fingerprint = tofu::load_stored_fingerprint(&app, &store_key)?
.ok_or_else(|| format!(
"no trusted certificate fingerprint for {remote_host}. \
Connect via WebSocket first to establish TOFU trust."
@@ -384,7 +279,7 @@ async fn handle_connection(
let tls_config = rustls::ClientConfig::builder()
.dangerous()
.with_custom_certificate_verifier(Arc::new(
PinnedVerifier::new(pinned_fingerprint.to_string()),
tofu::PinnedVerifier::new(pinned_fingerprint.to_string()),
))
.with_no_client_auth();
@@ -426,36 +321,4 @@ async fn handle_connection(
Ok(())
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cert_store_key_strips_default_port() {
assert_eq!(cert_store_key("example.com:443"), "example.com");
}
#[test]
fn cert_store_key_keeps_non_default_port() {
assert_eq!(cert_store_key("example.com:8443"), "example.com:8443");
}
#[test]
fn cert_store_key_no_port() {
assert_eq!(cert_store_key("example.com"), "example.com");
}
#[test]
fn cert_store_key_ipv4_default_port() {
assert_eq!(cert_store_key("192.168.1.1:443"), "192.168.1.1");
}
#[test]
fn cert_store_key_ipv4_custom_port() {
assert_eq!(cert_store_key("192.168.1.1:7880"), "192.168.1.1:7880");
}
}
// cert_store_key is covered by unit tests in the shared `tofu` module.
+314
View File
@@ -0,0 +1,314 @@
// Shared TLS Trust-On-First-Use (TOFU) machinery for the http / ws / livekit
// proxies. Self-hosted servers use self-signed certs, so we pin the leaf cert's
// SHA-256 fingerprint on first use — like SSH's known_hosts.
//
// F4/F8: pinning is now EXPLICIT. A first-use certificate is never silently
// trusted or forwarded to. The proxies capture the fingerprint during the
// handshake, then reject the connection and surface the fingerprint so the user
// can confirm it (via `accept_cert_fingerprint`) before any credential-bearing
// request is sent. `decide` is a pure function with no persistence side effects;
// the only writer of a pin is the explicit `accept_cert_fingerprint` command.
use ring::digest::{digest, SHA256};
use serde_json::Value;
use std::sync::Arc;
use tauri::{AppHandle, Runtime};
use tauri_plugin_store::StoreExt;
use crate::constants::CERTS_STORE;
/// Shared fingerprint captured during the TLS handshake.
pub(crate) type CapturedFingerprint = Arc<std::sync::Mutex<Option<String>>>;
/// Format a DER-encoded certificate's SHA-256 as lowercase colon-hex
/// ("aa:bb:cc:..."), the canonical pin format used across the cert store.
pub(crate) fn fingerprint_hex(cert_der: &[u8]) -> String {
digest(&SHA256, cert_der)
.as_ref()
.iter()
.map(|b| format!("{b:02x}"))
.collect::<Vec<_>>()
.join(":")
}
// ── shared rustls signature-verification boilerplate ────────────────────────
// Identical across every verifier; single-homed here so the three proxies don't
// each re-implement it.
pub(crate) fn verify_tls12(
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls12_signature(
message,
cert,
dss,
&rustls::crypto::ring::default_provider().signature_verification_algorithms,
)
}
pub(crate) fn verify_tls13(
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls13_signature(
message,
cert,
dss,
&rustls::crypto::ring::default_provider().signature_verification_algorithms,
)
}
pub(crate) fn default_verify_schemes() -> Vec<rustls::SignatureScheme> {
rustls::crypto::ring::default_provider()
.signature_verification_algorithms
.supported_schemes()
}
// ── verifiers ───────────────────────────────────────────────────────────────
/// A rustls verifier that ACCEPTS any leaf cert but records its fingerprint for
/// the post-handshake TOFU decision. Used by the http and ws proxies. Accepting
/// here is safe only because `evaluate` + the caller gate on the pin afterward.
#[derive(Debug)]
pub(crate) struct CaptureVerifier {
captured: CapturedFingerprint,
}
impl CaptureVerifier {
pub(crate) fn new() -> (Self, CapturedFingerprint) {
let fp = Arc::new(std::sync::Mutex::new(None));
(Self { captured: fp.clone() }, fp)
}
}
impl rustls::client::danger::ServerCertVerifier for CaptureVerifier {
fn verify_server_cert(
&self,
end_entity: &rustls::pki_types::CertificateDer<'_>,
_intermediates: &[rustls::pki_types::CertificateDer<'_>],
_server_name: &rustls::pki_types::ServerName<'_>,
_ocsp_response: &[u8],
_now: rustls::pki_types::UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
if let Ok(mut guard) = self.captured.lock() {
*guard = Some(fingerprint_hex(end_entity.as_ref()));
}
// Accept — the TOFU decision happens after the handshake, before any
// request bytes are forwarded.
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
verify_tls12(message, cert, dss)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
verify_tls13(message, cert, dss)
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
default_verify_schemes()
}
}
/// A rustls verifier that requires the leaf cert to match a pinned fingerprint,
/// failing the handshake itself on mismatch. Used by the livekit proxy, which
/// refuses to start unless a pin already exists (no TOFU establishment).
#[derive(Debug)]
pub(crate) struct PinnedVerifier {
expected_fingerprint: String,
}
impl PinnedVerifier {
pub(crate) fn new(expected_fingerprint: String) -> Self {
Self { expected_fingerprint }
}
}
impl rustls::client::danger::ServerCertVerifier for PinnedVerifier {
fn verify_server_cert(
&self,
end_entity: &rustls::pki_types::CertificateDer<'_>,
_intermediates: &[rustls::pki_types::CertificateDer<'_>],
_server_name: &rustls::pki_types::ServerName<'_>,
_ocsp_response: &[u8],
_now: rustls::pki_types::UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
let hex = fingerprint_hex(end_entity.as_ref());
if hex == self.expected_fingerprint {
Ok(rustls::client::danger::ServerCertVerified::assertion())
} else {
Err(rustls::Error::General(format!(
"certificate fingerprint mismatch: expected {}, got {}",
self.expected_fingerprint, hex
)))
}
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
verify_tls12(message, cert, dss)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
verify_tls13(message, cert, dss)
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
default_verify_schemes()
}
}
// ── store keys ──────────────────────────────────────────────────────────────
/// Cert-store key for a host. Strips a default `:443` so the ws proxy (which
/// keys off `wss://host` with no explicit 443) and the http/livekit proxies
/// (which see `host:443`) resolve the SAME pin. Non-default ports are kept.
pub(crate) fn cert_store_key(host: &str) -> String {
host.strip_suffix(":443").unwrap_or(host).to_string()
}
/// Extract the host (with any non-default port) from a `wss://` URL.
pub(crate) fn extract_host(url: &str) -> String {
cert_store_key(
url.strip_prefix("wss://")
.unwrap_or(url)
.split('/')
.next()
.unwrap_or(url),
)
}
/// Load the stored pin for `host` from the Tauri cert store.
pub(crate) fn load_stored_fingerprint<R: Runtime>(
app: &AppHandle<R>,
host: &str,
) -> Result<Option<String>, String> {
let store = app
.store(CERTS_STORE)
.map_err(|e| format!("failed to open certs store: {e}"))?;
Ok(store.get(host).and_then(|v| match v {
Value::String(s) => Some(s),
_ => None,
}))
}
// ── the TOFU decision (pure) ────────────────────────────────────────────────
/// The trust decision for an observed fingerprint given the stored pin.
#[derive(Debug, PartialEq, Eq)]
pub(crate) enum TofuOutcome {
/// A pin exists and matches — proceed.
Trusted,
/// No pin exists — do NOT trust or forward; ask the user to confirm.
FirstUse,
/// A pin exists but differs — reject; possible MITM or cert rotation.
Mismatch { stored: String },
}
/// Pure trust decision. No I/O, no persistence — this is the whole point of the
/// F4/F8 fix: deciding never writes a pin.
pub(crate) fn decide(stored: Option<String>, current: &str) -> TofuOutcome {
match stored {
None => TofuOutcome::FirstUse,
Some(s) if s == current => TofuOutcome::Trusted,
Some(s) => TofuOutcome::Mismatch { stored: s },
}
}
/// Load the stored pin and decide. Never persists.
pub(crate) fn evaluate<R: Runtime>(
app: &AppHandle<R>,
host: &str,
fingerprint: &str,
) -> Result<TofuOutcome, String> {
let stored = load_stored_fingerprint(app, host)?;
Ok(decide(stored, fingerprint))
}
/// The human-readable mismatch message. The frontend parses `Stored:` out of it,
/// so keep this exact shape stable.
pub(crate) fn mismatch_message(host: &str, stored: &str, current: &str) -> String {
format!(
"Certificate fingerprint changed for {host}.\n\
Stored: {stored}\n\
Current: {current}\n\
This may indicate a man-in-the-middle attack or a server certificate rotation.\n\
Use accept_cert_fingerprint to trust the new certificate."
)
}
// ---------------------------------------------------------------------------
// Tests (pure logic only — no Tauri runtime required)
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn decide_first_use_when_no_pin() {
assert_eq!(decide(None, "aa:bb"), TofuOutcome::FirstUse);
}
#[test]
fn decide_trusted_when_pin_matches() {
assert_eq!(decide(Some("aa:bb".into()), "aa:bb"), TofuOutcome::Trusted);
}
#[test]
fn decide_mismatch_when_pin_differs() {
assert_eq!(
decide(Some("aa:bb".into()), "cc:dd"),
TofuOutcome::Mismatch { stored: "aa:bb".into() }
);
}
#[test]
fn cert_store_key_strips_default_443_only() {
assert_eq!(cert_store_key("example.com:443"), "example.com");
assert_eq!(cert_store_key("example.com"), "example.com");
assert_eq!(cert_store_key("example.com:8443"), "example.com:8443");
}
#[test]
fn extract_host_variants() {
assert_eq!(extract_host("wss://example.com/chat"), "example.com");
assert_eq!(extract_host("wss://example.com:8443/chat"), "example.com:8443");
assert_eq!(extract_host("wss://example.com:443/chat"), "example.com");
assert_eq!(extract_host("wss://example.com"), "example.com");
assert_eq!(extract_host("example.com/path"), "example.com");
assert_eq!(extract_host(""), "");
}
#[test]
fn fingerprint_hex_of_empty_is_known_sha256() {
// SHA-256("") = e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855
assert_eq!(
fingerprint_hex(b""),
"e3:b0:c4:42:98:fc:1c:14:9a:fb:f4:c8:99:6f:b9:24:27:ae:41:e4:64:9b:93:4c:a4:95:99:1b:78:52:b8:55"
);
}
}
+38 -208
View File
@@ -1,14 +1,17 @@
// WebSocket proxy — routes WSS through Rust to bypass self-signed cert rejection.
// JS sends/receives messages via Tauri events instead of native WebSocket.
//
// Implements TOFU (Trust On First Use) certificate pinning:
// - On first connect to a host, the cert SHA-256 fingerprint is stored.
// - On subsequent connects, the fingerprint is compared with the stored value.
// - If the fingerprint changes, the connection is rejected (potential MitM).
// Implements TOFU (Trust On First Use) certificate pinning via the shared
// `tofu` module:
// - The cert SHA-256 fingerprint is captured during the handshake.
// - On a known host it must match the stored pin, or the connection is rejected.
// - On first use (no pin yet) the connection is rejected and a `cert-tofu`
// "first_use" event is emitted so the user can confirm the fingerprint. F4/F8:
// the proxy never silently pins or forwards to an unconfirmed host — the only
// writer of a pin is the explicit `accept_cert_fingerprint` command.
use futures_util::{SinkExt, StreamExt};
use log::{debug, error, info, warn};
use ring::digest::{digest, SHA256};
use serde_json::Value;
use std::sync::Arc;
use std::time::Duration;
@@ -22,6 +25,7 @@ use tokio_tungstenite::tungstenite::Message;
const CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
use crate::constants::CERTS_STORE;
use crate::tofu::{self, TofuOutcome};
/// Sender half kept in Tauri state so `ws_send` can push messages.
/// `tx` is wrapped in `Arc` so the monitoring task can clone a reference
@@ -38,157 +42,6 @@ impl WsState {
}
}
/// Shared fingerprint captured during TLS handshake.
type CapturedFingerprint = Arc<std::sync::Mutex<Option<String>>>;
/// TOFU certificate verifier that captures the server cert fingerprint
/// during the TLS handshake. Still accepts self-signed certs (required
/// for self-hosted servers), but records the fingerprint for comparison
/// with the stored value after the connection is established.
#[derive(Debug)]
struct TofuVerifier {
captured: CapturedFingerprint,
}
impl TofuVerifier {
fn new() -> (Self, CapturedFingerprint) {
let fp = Arc::new(std::sync::Mutex::new(None));
(Self { captured: fp.clone() }, fp)
}
}
impl rustls::client::danger::ServerCertVerifier for TofuVerifier {
fn verify_server_cert(
&self,
end_entity: &rustls::pki_types::CertificateDer<'_>,
_intermediates: &[rustls::pki_types::CertificateDer<'_>],
_server_name: &rustls::pki_types::ServerName<'_>,
_ocsp_response: &[u8],
_now: rustls::pki_types::UnixTime,
) -> Result<rustls::client::danger::ServerCertVerified, rustls::Error> {
// Compute SHA-256 fingerprint of the DER-encoded leaf certificate.
let hash = digest(&SHA256, end_entity.as_ref());
let hex = hash
.as_ref()
.iter()
.map(|b| format!("{b:02x}"))
.collect::<Vec<_>>()
.join(":");
if let Ok(mut guard) = self.captured.lock() {
*guard = Some(hex);
}
// Accept the cert — TOFU check happens after the handshake completes.
Ok(rustls::client::danger::ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls12_signature(
message,
cert,
dss,
&rustls::crypto::ring::default_provider().signature_verification_algorithms,
)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &rustls::pki_types::CertificateDer<'_>,
dss: &rustls::DigitallySignedStruct,
) -> Result<rustls::client::danger::HandshakeSignatureValid, rustls::Error> {
rustls::crypto::verify_tls13_signature(
message,
cert,
dss,
&rustls::crypto::ring::default_provider().signature_verification_algorithms,
)
}
fn supported_verify_schemes(&self) -> Vec<rustls::SignatureScheme> {
vec![
rustls::SignatureScheme::RSA_PKCS1_SHA256,
rustls::SignatureScheme::RSA_PKCS1_SHA384,
rustls::SignatureScheme::RSA_PKCS1_SHA512,
rustls::SignatureScheme::ECDSA_NISTP256_SHA256,
rustls::SignatureScheme::ECDSA_NISTP384_SHA384,
rustls::SignatureScheme::ECDSA_NISTP521_SHA512,
rustls::SignatureScheme::RSA_PSS_SHA256,
rustls::SignatureScheme::RSA_PSS_SHA384,
rustls::SignatureScheme::RSA_PSS_SHA512,
rustls::SignatureScheme::ED25519,
rustls::SignatureScheme::ED448,
]
}
}
/// Extract the host (with port) from a wss:// URL.
fn extract_host(url: &str) -> String {
url.strip_prefix("wss://")
.unwrap_or(url)
.split('/')
.next()
.unwrap_or(url)
.to_string()
}
/// Perform TOFU fingerprint check against the Tauri cert store.
/// Returns Ok(()) if trusted, Err(message) if fingerprint mismatch.
fn tofu_check<R: Runtime>(
app: &AppHandle<R>,
host: &str,
fingerprint: &str,
) -> Result<String, String> {
let store = app
.store(CERTS_STORE)
.map_err(|e| format!("failed to open certs store: {e}"))?;
let stored = store.get(host).and_then(|v| {
if let Value::String(s) = v {
Some(s)
} else {
None
}
});
match stored {
None => {
// First use — store the fingerprint.
// Capture old value before mutating (None here, but consistent pattern).
let old_value = store.get(host);
store.set(host, Value::String(fingerprint.to_string()));
if let Err(e) = store.save() {
// Restore previous in-memory state: put back old value or delete
// if there was none, keeping in-memory consistent with on-disk.
match old_value {
Some(v) => { store.set(host, v); }
None => { let _ = store.delete(host); }
}
return Err(format!("failed to persist cert fingerprint: {e}"));
}
Ok("trusted_first_use".to_string())
}
Some(ref stored_fp) if stored_fp == fingerprint => {
Ok("trusted".to_string())
}
Some(stored_fp) => {
Err(format!(
"Certificate fingerprint changed for {host}.\n\
Stored: {stored_fp}\n\
Current: {fingerprint}\n\
This may indicate a man-in-the-middle attack or a server certificate rotation.\n\
Use accept_cert_fingerprint to trust the new certificate."
))
}
}
}
/// Single call site for ws-state events — keeps tauri-typegen from generating duplicates.
fn emit_ws_state<R: Runtime>(app: &AppHandle<R>, state: &str) {
let _ = app.emit("ws-state", state);
@@ -229,8 +82,9 @@ pub async fn ws_connect<R: Runtime>(
emit_ws_state(&app, "connecting");
// Create TOFU verifier that captures the cert fingerprint during handshake.
let (verifier, captured_fp) = TofuVerifier::new();
// Capture the cert fingerprint during the handshake; the TOFU decision runs
// afterward, before the socket is used.
let (verifier, captured_fp) = tofu::CaptureVerifier::new();
let tls_config = rustls::ClientConfig::builder()
.dangerous()
@@ -261,7 +115,7 @@ pub async fn ws_connect<R: Runtime>(
debug!("[ws_proxy] WebSocket handshake complete");
// ── TOFU check ───────────────────────────────────────────────────────
let host = extract_host(&url);
let host = tofu::extract_host(&url);
let fingerprint = captured_fp
.lock()
.map_err(|e| format!("failed to read captured fingerprint: {e}"))?
@@ -272,26 +126,41 @@ pub async fn ws_connect<R: Runtime>(
return Err("TLS handshake completed but no certificate fingerprint was captured".into());
}
match tofu_check(&app, &host, &fingerprint) {
Ok(status) => {
info!("[ws_proxy] TOFU check passed for {}: {}", host, status);
match tofu::evaluate(&app, &host, &fingerprint)? {
TofuOutcome::Trusted => {
info!("[ws_proxy] TOFU check passed for {}", host);
emit_cert_tofu(&app, serde_json::json!({
"host": host,
"fingerprint": fingerprint,
"status": status,
"status": "trusted",
}));
}
Err(mismatch_msg) => {
TofuOutcome::FirstUse => {
info!("[ws_proxy] first-use cert for {} — awaiting user confirmation", host);
emit_cert_tofu(&app, serde_json::json!({
"host": host,
"fingerprint": fingerprint,
"status": "first_use",
}));
// Do not open the socket: the user must confirm the fingerprint
// (accept_cert_fingerprint) before anything is sent over it.
return Err(format!(
"certificate for {host} is not yet trusted; confirm the fingerprint to continue"
));
}
TofuOutcome::Mismatch { stored } => {
let msg = tofu::mismatch_message(&host, &stored, &fingerprint);
warn!("[ws_proxy] TOFU check FAILED for {} — certificate fingerprint mismatch", host);
debug!("[ws_proxy] TOFU detail: {}", mismatch_msg);
debug!("[ws_proxy] TOFU detail: {}", msg);
emit_cert_tofu(&app, serde_json::json!({
"host": host,
"fingerprint": fingerprint,
"status": "mismatch",
"message": mismatch_msg,
"message": msg,
"storedFingerprint": stored,
}));
// Reject the connection — do not proceed.
return Err(mismatch_msg);
return Err(msg);
}
}
// ── End TOFU check ───────────────────────────────────────────────────
@@ -413,8 +282,8 @@ pub async fn ws_disconnect(state: tauri::State<'_, WsState>) -> Result<(), Strin
Ok(())
}
/// Accept a changed certificate fingerprint for a host.
/// Call this after the user acknowledges a cert-mismatch warning.
/// Accept a certificate fingerprint for a host — the ONLY path that writes a pin.
/// Called after the user acknowledges a first-use or cert-mismatch prompt.
#[tauri::command]
pub fn accept_cert_fingerprint<R: Runtime>(
app: AppHandle<R>,
@@ -458,42 +327,3 @@ pub fn accept_cert_fingerprint<R: Runtime>(
}
Ok(())
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn extract_host_basic_wss_url() {
assert_eq!(extract_host("wss://example.com/chat"), "example.com");
}
#[test]
fn extract_host_with_port() {
assert_eq!(extract_host("wss://example.com:8443/chat"), "example.com:8443");
}
#[test]
fn extract_host_no_path() {
assert_eq!(extract_host("wss://example.com"), "example.com");
}
#[test]
fn extract_host_no_scheme() {
assert_eq!(extract_host("example.com/path"), "example.com");
}
#[test]
fn extract_host_empty() {
assert_eq!(extract_host(""), "");
}
#[test]
fn extract_host_with_port_and_deep_path() {
assert_eq!(extract_host("wss://myhost:9443/api/v1/ws"), "myhost:9443");
}
}
@@ -108,6 +108,99 @@ export function createCertMismatchModal(options: CertMismatchModalOptions): Moun
return { mount, destroy };
}
export interface CertFirstUseModalOptions {
readonly host: string;
readonly fingerprint: string;
readonly onAccept: () => void;
readonly onReject: () => void;
}
/**
* createCertFirstUseModal — shown on the FIRST connection to a server, when no
* certificate is pinned yet (F4/F8). The proxy refuses to send anything until
* the user confirms this fingerprint, so an on-path attacker at first contact
* cannot silently capture credentials. Mirrors an SSH known-hosts prompt.
*/
export function createCertFirstUseModal(options: CertFirstUseModalOptions): MountableComponent {
const { host, fingerprint, onAccept, onReject } = options;
let overlay: HTMLDivElement | null = null;
const ac = new AbortController();
function mount(container: Element): void {
overlay = createElement("div", { class: "modal-overlay visible" });
const modal = createElement("div", { class: "modal" });
const header = createElement("div", { class: "modal-header" });
const title = createElement("h3", {}, "New Server Certificate");
const closeBtn = createElement("button", { class: "modal-close", type: "button" });
closeBtn.textContent = "";
closeBtn.appendChild(createIcon("x", 14));
closeBtn.addEventListener("click", onReject, { signal: ac.signal });
appendChildren(header, title, closeBtn);
const body = createElement("div", { class: "modal-body" });
const warning = createElement("div", { class: "cert-warning" });
warning.appendChild(createIcon("triangle-alert", 24));
const certTitle = createElement("div", { class: "cert-title" });
setText(certTitle, "Confirm the certificate fingerprint");
const desc = createElement("div", { class: "cert-desc" });
setText(
desc,
"This is the first connection to this server, so its certificate is not " +
"yet trusted. Verify the fingerprint below out-of-band (e.g. with the " +
"server operator) before trusting it — on an untrusted network an " +
"attacker could present a fake certificate.",
);
const details = createElement("div", { class: "cert-details" });
appendChildren(
details,
buildRow("Host", host, false),
buildRow("Fingerprint", fingerprint, true),
);
appendChildren(body, warning, certTitle, desc, details);
const footer = createElement("div", { class: "modal-footer" });
const rejectBtn = createElement("button", { class: "btn-ghost", type: "button" });
setText(rejectBtn, "Cancel");
rejectBtn.addEventListener("click", onReject, { signal: ac.signal });
const acceptBtn = createElement("button", { class: "btn-danger", type: "button" });
setText(acceptBtn, "Trust This Certificate");
acceptBtn.addEventListener("click", onAccept, { signal: ac.signal });
appendChildren(footer, rejectBtn, acceptBtn);
appendChildren(modal, header, body, footer);
overlay.appendChild(modal);
overlay.addEventListener(
"click",
(e) => {
if (e.target === overlay) onReject();
},
{ signal: ac.signal },
);
container.appendChild(overlay);
}
function destroy(): void {
ac.abort();
if (overlay !== null) {
overlay.remove();
overlay = null;
}
}
return { mount, destroy };
}
function buildRow(label: string, value: string, isFingerprint: boolean): HTMLDivElement {
const row = createElement("div", { class: "cert-row" });
const labelEl = createElement("span", { class: "cert-label" });
@@ -1,101 +0,0 @@
/**
* File upload validation and preview rendering for message input.
*/
import { createElement, appendChildren } from "@lib/dom";
import { createIcon } from "@lib/icons";
export const MAX_FILE_SIZE = 100 * 1024 * 1024; // 100MB matches server limit
export const ALLOWED_TYPES = [
"image/",
"video/",
"audio/",
"application/pdf",
"text/",
"application/zip",
"application/x-zip-compressed",
"application/json",
];
/** Read a File as a data: URL (more reliable than createObjectURL in WebView2). */
export function readFileAsDataUrl(file: File): Promise<string> {
return new Promise((resolve, reject) => {
const reader = new FileReader();
reader.addEventListener("load", () => resolve(reader.result as string));
reader.addEventListener("error", () => reject(new Error("Failed to read file")));
reader.readAsDataURL(file);
});
}
/** Validate file size and type. Returns an error message or null. */
export function validateFile(file: File): string | null {
if (file.size > MAX_FILE_SIZE) {
return `File too large: ${file.name} exceeds 100 MB limit`;
}
if (file.type === "" || !ALLOWED_TYPES.some((t) => file.type.startsWith(t))) {
return `${file.name} is not a supported file type`;
}
return null;
}
/** Build a preview item element for a file being uploaded. */
export function buildPreviewItem(
file: File,
signal: AbortSignal,
onRemove: () => void,
): HTMLDivElement {
const isImage = file.type.startsWith("image/");
const item = createElement("div", { class: "attachment-preview-item uploading" });
if (isImage) {
const img = createElement("img", {
class: "attachment-preview-img",
alt: file.name,
});
item.appendChild(img);
readFileAsDataUrl(file)
.then((dataUrl) => {
if (signal.aborted) return;
img.src = dataUrl;
})
.catch(() => {
if (signal.aborted) return;
const nameEl = createElement("span", { class: "attachment-preview-name" }, file.name);
img.replaceWith(nameEl);
});
} else {
const icon = createElement("div", { class: "attachment-preview-file" });
icon.appendChild(createIcon("file-text", 16));
const nameEl = createElement("span", { class: "attachment-preview-name" }, file.name);
appendChildren(item, icon, nameEl);
}
// Loading spinner overlay
const spinner = createElement("div", { class: "attachment-preview-spinner" });
spinner.appendChild(createIcon("loader", 16));
item.appendChild(spinner);
const removeBtn = createElement("button", {
class: "attachment-preview-remove",
"data-testid": "attachment-remove",
});
removeBtn.appendChild(createIcon("x", 14));
removeBtn.addEventListener(
"click",
(e) => {
e.stopPropagation();
onRemove();
},
{ signal },
);
item.appendChild(removeBtn);
return item;
}
/** Mark a preview item as uploaded (removes loading state). */
export function markPreviewUploaded(item: HTMLDivElement): void {
item.classList.remove("uploading");
const spinner = item.querySelector(".attachment-preview-spinner");
spinner?.remove();
}
@@ -1,77 +0,0 @@
/**
* Reusable picker toggle — manages open/close/click-outside lifecycle
* for floating panels (emoji picker, GIF picker, etc.).
*/
export interface PickerInstance {
readonly element: HTMLDivElement;
destroy(): void;
}
export interface PickerToggleOptions {
/** Creates and returns a new picker instance. */
readonly create: () => PickerInstance;
/** The trigger button element — clicks on it won't close the picker. */
readonly triggerEl: HTMLElement;
/** Parent element to append the picker to. */
readonly parentEl: HTMLElement | null;
/** Called before opening — use to close other pickers first. */
readonly onBeforeOpen?: () => void;
/** Timer set for deferred cleanup. */
readonly activeTimers: Set<ReturnType<typeof setTimeout>>;
}
export interface PickerToggleHandle {
toggle(): void;
close(): void;
}
export function createPickerToggle(opts: PickerToggleOptions): PickerToggleHandle {
let instance: PickerInstance | null = null;
let pendingTimer: ReturnType<typeof setTimeout> | null = null;
function handleClickOutside(e: MouseEvent): void {
if (instance === null) return;
const target = e.target as Node;
if (
!instance.element.contains(target) &&
target !== opts.triggerEl &&
!opts.triggerEl.contains(target)
) {
close();
}
}
function close(): void {
if (pendingTimer !== null) {
clearTimeout(pendingTimer);
opts.activeTimers.delete(pendingTimer);
pendingTimer = null;
}
if (instance !== null) {
instance.element.remove();
instance.destroy();
instance = null;
document.removeEventListener("mousedown", handleClickOutside);
}
}
function toggle(): void {
opts.onBeforeOpen?.();
if (instance !== null) {
close();
return;
}
instance = opts.create();
opts.parentEl?.appendChild(instance.element);
// Defer so this click doesn't immediately close it
pendingTimer = setTimeout(() => {
opts.activeTimers.delete(pendingTimer!);
pendingTimer = null;
document.addEventListener("mousedown", handleClickOutside);
}, 0);
opts.activeTimers.add(pendingTimer);
}
return { toggle, close };
}
@@ -37,48 +37,38 @@ export function clearEmbedCaches(): void {
// -- OG tag parsing -----------------------------------------------------------
/** Escape special regex characters in a string for safe use in `new RegExp()`. */
function escapeRegex(s: string): string {
return s.replace(/[.*+?^${}()|[\]\\]/g, "\\$&");
}
/** Extract Open Graph meta tags from raw HTML using regex (no DOM parser needed). */
/**
* Extract Open Graph meta tags from raw HTML.
*
* F7: parse with the platform HTML tokenizer (DOMParser) instead of backtracking
* regexes. Untrusted preview HTML previously ran through patterns with two
* `[^>]*` quantifiers around a required literal, which backtrack polynomially and
* froze the UI thread on crafted input (ReDoS). A real tokenizer is linear and
* additionally ignores meta-like strings inside comments/scripts.
*
* The 50 KB slice at the call site (fetchOgMeta) is kept as a plain memory bound.
*/
export function parseOgTags(html: string): OgMeta {
function getMetaContent(property: string): string | null {
// Match both property="og:X" and name="og:X" patterns
const escaped = escapeRegex(property);
const regex = new RegExp(
`<meta[^>]*(?:property|name)=["']${escaped}["'][^>]*content=["']([^"']*)["']` +
`|<meta[^>]*content=["']([^"']*)["'][^>]*(?:property|name)=["']${escaped}["']`,
"i",
);
const match = html.match(regex);
if (match !== null) {
return match[1] ?? match[2] ?? null;
const doc = new DOMParser().parseFromString(html, "text/html");
// First matching element wins (document order), mirroring the old first-match
// behaviour. Returns "" for an empty content attribute but null when the
// attribute (or element) is absent, so the title→host fallback still fires.
function metaContent(...ogNames: readonly string[]): string | null {
for (const name of ogNames) {
const el =
doc.querySelector(`meta[property="${name}"]`) ?? doc.querySelector(`meta[name="${name}"]`);
const content = el?.getAttribute("content");
if (content != null) return content;
}
return null;
}
// Fallback: extract <title> tag if no og:title
function getTitle(): string | null {
const og = getMetaContent("og:title");
if (og !== null) return og;
const titleMatch = html.match(/<title[^>]*>([^<]*)<\/title>/i);
return titleMatch?.[1]?.trim() ?? null;
}
// Fallback: extract meta description if no og:description
function getDescription(): string | null {
const og = getMetaContent("og:description");
if (og !== null) return og;
return getMetaContent("description");
}
return {
title: getTitle(),
description: getDescription(),
image: getMetaContent("og:image"),
siteName: getMetaContent("og:site_name"),
title: metaContent("og:title") ?? doc.querySelector("title")?.textContent?.trim() ?? null,
description: metaContent("og:description", "description"),
image: metaContent("og:image"),
siteName: metaContent("og:site_name"),
};
}
@@ -1,139 +0,0 @@
/**
* Virtual scroll manager — manages height estimation, Fenwick-tree-backed
* offset calculations, and spacer management for DOM windowing.
*/
import { FenwickTree } from "./fenwick";
export interface VirtualScrollItem {
readonly kind: string;
}
export interface VirtualScrollOptions {
/** Number of items to render beyond visible viewport in each direction. */
readonly overscan: number;
/** Estimate height for an item at given index. */
readonly estimateHeight: (index: number) => number;
/** Generate a stable cache key for an item at given index. */
readonly itemKey: (index: number) => string;
}
export interface VisibleRange {
readonly start: number;
readonly end: number;
}
export class VirtualScrollManager {
private readonly heightCache = new Map<string, number>();
private tree: FenwickTree | null = null;
private itemCount = 0;
private readonly opts: VirtualScrollOptions;
constructor(opts: VirtualScrollOptions) {
this.opts = opts;
}
/** Rebuild the Fenwick tree for a new item count, preserving cached heights. */
rebuild(count: number): void {
this.itemCount = count;
this.tree = new FenwickTree(count);
for (let i = 0; i < count; i++) {
const key = this.opts.itemKey(i);
const cached = this.heightCache.get(key);
const h = cached !== undefined ? cached : this.opts.estimateHeight(i);
this.tree.set(i, h);
}
}
/** Get height for item at index (cached or estimated). */
getHeight(index: number): number {
const cached = this.heightCache.get(this.opts.itemKey(index));
if (cached !== undefined) return cached;
return this.opts.estimateHeight(index);
}
/** Cache a measured height for an item. */
setMeasured(index: number, height: number): void {
if (height <= 0) return;
const key = this.opts.itemKey(index);
this.heightCache.set(key, height);
if (this.tree !== null && index < this.tree.size) {
this.tree.set(index, height);
}
}
/** Total estimated height of all items. */
totalHeight(): number {
if (this.tree !== null) return this.tree.total();
let h = 0;
for (let i = 0; i < this.itemCount; i++) {
h += this.getHeight(i);
}
return h;
}
/** Sum of heights for items [0, index). */
offsetBefore(index: number): number {
if (this.tree !== null && index > 0) return this.tree.prefixSum(index - 1);
if (this.tree !== null && index <= 0) return 0;
let offset = 0;
for (let i = 0; i < index && i < this.itemCount; i++) {
offset += this.getHeight(i);
}
return offset;
}
/** Find the item index at a given scroll offset. */
offsetToIndex(scrollTop: number): number {
if (this.tree !== null) return this.tree.findIndex(scrollTop);
let offset = 0;
for (let i = 0; i < this.itemCount; i++) {
const h = this.getHeight(i);
if (offset + h > scrollTop) return i;
offset += h;
}
return Math.max(0, this.itemCount - 1);
}
/** Compute the visible range with overscan. */
visibleRange(scrollTop: number, clientHeight: number): VisibleRange {
const firstVisible = this.offsetToIndex(scrollTop);
const lastVisible = this.offsetToIndex(scrollTop + clientHeight);
return {
start: Math.max(0, firstVisible - this.opts.overscan),
end: Math.min(this.itemCount, lastVisible + this.opts.overscan + 1),
};
}
/** Compute spacer heights for a rendered range. */
spacerHeights(start: number, end: number): { top: number; bottom: number } {
const top = this.offsetBefore(start);
let bottom: number;
if (this.tree !== null) {
const totalH = this.tree.total();
const endOffset = end > 0 ? this.tree.prefixSum(end - 1) : 0;
bottom = totalH - endOffset;
} else {
bottom = 0;
for (let i = end; i < this.itemCount; i++) {
bottom += this.getHeight(i);
}
}
return { top, bottom };
}
/** Clear all cached heights. */
clear(): void {
this.heightCache.clear();
this.tree = null;
this.itemCount = 0;
}
get size(): number {
return this.itemCount;
}
get treeSize(): number {
return this.tree?.size ?? 0;
}
}
@@ -1687,7 +1687,6 @@ owncordNs.lkDebug = session.getSessionDebugInfo.bind(session);
export const setWsClient = session.setWsClient.bind(session);
export const setServerHost = session.setServerHost.bind(session);
export const setOnError = session.setOnError.bind(session);
export const clearOnError = session.clearOnError.bind(session);
export const setOnRemoteVideo = session.setOnRemoteVideo.bind(session);
export const setOnRemoteVideoRemoved = session.setOnRemoteVideoRemoved.bind(session);
export const clearOnRemoteVideo = session.clearOnRemoteVideo.bind(session);
+4 -29
View File
@@ -5,7 +5,7 @@
// Rotation: keeps the most recent MAX_LOG_FILES days of logs.
import { appLogDir, join } from "@tauri-apps/api/path";
import { mkdir, writeTextFile, readDir, remove, exists, readTextFile } from "@tauri-apps/plugin-fs";
import { mkdir, writeTextFile, readDir, remove, exists } from "@tauri-apps/plugin-fs";
import { type LogEntry, addLogListener, createLogger } from "./logger";
const log = createLogger("logPersistence");
@@ -165,35 +165,10 @@ export async function flushLogs(): Promise<void> {
}
/**
* Get the log directory path (for use in debug bundle export).
* Returns null if persistence hasn't been initialized.
* Get the log directory path. Production-unused but exported as the test
* suite's observability point for persistence state.
* @public
*/
export function getLogDir(): string | null {
return logDir;
}
/**
* Read all persisted log files and return their combined content.
* Intended for on-demand export only (reads all files into memory).
*/
export async function readAllPersistedLogs(): Promise<string> {
if (!logDir) return "";
try {
const entries = await readDir(logDir);
const jsonlFiles = entries
.filter((e) => e.name?.endsWith(".jsonl") && !e.isDirectory)
.map((e) => e.name)
.toSorted((a, b) => a.localeCompare(b));
const parts: string[] = [];
for (const file of jsonlFiles) {
// oxlint-disable-next-line no-await-in-loop -- files must be read in sorted order for correct log concatenation
const content = await readTextFile(`${logDir}/${file}`);
parts.push(content);
}
return parts.join("");
} catch (err) {
log.warn("readAllPersistedLogs failed", err);
return "";
}
}
-172
View File
@@ -1,172 +0,0 @@
/**
* Phase C Step 9 — client-side plugin bridge.
*
* Mounts plugin UI tabs in sandboxed iframes and forwards postMessage traffic
* between the host client and each plugin. The host injects theme CSS
* variables on every load so plugin UIs match OwnCord's look and feel
* without each plugin re-implementing them.
*
* The bridge is intentionally tiny: it owns iframe lifecycles and message
* routing; everything else (rendering tabs, fetching the plugin list) lives
* in PluginContainer.tsx.
*/
export interface PluginTabBinding {
pluginId: number;
pluginName: string;
tabId: string;
label: string;
asset: string;
}
export interface PluginMessageEnvelope {
pluginId: number;
type: string;
payload?: unknown;
}
type Listener = (env: PluginMessageEnvelope) => void;
const HOST_ORIGIN_PREFIX = "owncord-plugin-host";
class PluginBridge {
private frames = new Map<number, HTMLIFrameElement>();
private listeners = new Set<Listener>();
private themeVars: Record<string, string> = {};
private hostOrigin: string;
constructor() {
// Plugin iframes are served from /api/v1/plugins/... on the same origin
// as the host page, so postMessage targets that origin explicitly. Using
// "*" as the target origin is unsafe — any frame the user navigates to
// would receive host messages. window.location.origin is undefined in
// some test runners (jsdom prior to 16); fall back to "/" which still
// restricts to same-origin under the strict postMessage matching rules.
this.hostOrigin =
typeof window !== "undefined" && window.location && window.location.origin
? window.location.origin
: "/";
if (typeof window !== "undefined") {
window.addEventListener("message", this.onMessage);
}
}
/**
* destroy unhooks the global message listener and clears all mounted
* frames. Intended for tests that create disposable bridge instances; the
* exported `pluginBridge` singleton lives for the lifetime of the page and
* does not need explicit teardown.
*/
destroy(): void {
if (typeof window !== "undefined") {
window.removeEventListener("message", this.onMessage);
}
for (const frame of this.frames.values()) {
frame.remove();
}
this.frames.clear();
this.listeners.clear();
}
/** Replace the theme variables broadcast to plugin iframes. */
setTheme(vars: Record<string, string>): void {
this.themeVars = { ...vars };
for (const [pid, frame] of this.frames) {
this.postToFrame(pid, frame, { type: "theme", payload: this.themeVars });
}
}
/** Mount an iframe for binding into parent. Returns a destroy function. */
mount(binding: PluginTabBinding, parent: HTMLElement): () => void {
const iframe = document.createElement("iframe");
iframe.className = "plugin-iframe";
iframe.sandbox.add("allow-scripts");
iframe.title = `${binding.pluginName}: ${binding.label}`;
iframe.src = `/api/v1/plugins/${encodeURIComponent(binding.pluginName)}/ui/${binding.asset}`;
iframe.dataset.pluginId = String(binding.pluginId);
iframe.addEventListener("load", () => {
this.postToFrame(binding.pluginId, iframe, { type: "theme", payload: this.themeVars });
this.postToFrame(binding.pluginId, iframe, { type: "ready", payload: null });
});
parent.appendChild(iframe);
this.frames.set(binding.pluginId, iframe);
return () => {
iframe.remove();
this.frames.delete(binding.pluginId);
};
}
/** Listen for messages emitted by any mounted plugin iframe. */
onMessageEnvelope(listener: Listener): () => void {
this.listeners.add(listener);
return () => this.listeners.delete(listener);
}
/** Send a host → plugin message. */
send(pluginId: number, type: string, payload?: unknown): void {
const frame = this.frames.get(pluginId);
if (!frame) return;
this.postToFrame(pluginId, frame, { type, payload });
}
private postToFrame(
pluginId: number,
frame: HTMLIFrameElement,
msg: { type: string; payload: unknown },
): void {
// Restrict the postMessage target origin to the host page origin so a
// navigated-away iframe (or one whose contentWindow has been swapped)
// cannot receive host messages intended for a sandboxed plugin. The
// plugin asset endpoint is same-origin with the host page, so this
// matches every legitimate plugin iframe.
frame.contentWindow?.postMessage(
{ source: HOST_ORIGIN_PREFIX, pluginId, ...msg },
this.hostOrigin,
);
}
/**
* Look up the pluginId of an iframe by its contentWindow. Returns null if
* the source is not one of our managed plugin frames. This is the key
* defense against postMessage spoofing: we never trust the pluginId field
* inside the message body, only the e.source pointer.
*/
private pluginIdForSource(source: MessageEventSource | null): number | null {
if (!source) return null;
for (const [pid, frame] of this.frames) {
if (frame.contentWindow === source) return pid;
}
return null;
}
private onMessage = (e: MessageEvent): void => {
const data = e.data;
if (!data || typeof data !== "object") return;
if ((data as { source?: unknown }).source === HOST_ORIGIN_PREFIX) return; // own echo
// SECURITY: validate the message originated from one of our managed
// plugin iframes by matching e.source against frame.contentWindow.
// Without this check, any arbitrary frame (including a malicious parent
// frame in an embedding scenario, or any same-origin script that
// obtained a window reference) could spoof messages from any plugin by
// claiming an arbitrary pluginId in the body. The pluginId from the
// message body is intentionally ignored — we use the trusted lookup.
const trustedPluginId = this.pluginIdForSource(e.source);
if (trustedPluginId === null) return;
const env = data as { type?: unknown; payload?: unknown };
if (typeof env.type !== "string") return;
const envelope: PluginMessageEnvelope = {
pluginId: trustedPluginId,
type: env.type,
payload: env.payload,
};
for (const l of this.listeners) {
try {
l(envelope);
} catch (err) {
console.error("plugin bridge listener threw", err);
}
}
};
}
export const pluginBridge = new PluginBridge();
+73 -41
View File
@@ -62,11 +62,15 @@ export type WsListener<T extends ServerMessage["type"]> = (
id?: string,
) => void;
/** TOFU certificate event emitted by the Rust WS proxy. */
/** TOFU certificate event emitted by the Rust proxies (http / ws).
* - "first_use": no pin yet — the proxy REJECTED the connection; the user must
* confirm this fingerprint (acceptCertFingerprint) before anything is sent.
* - "trusted": pin matches — proceed.
* - "mismatch": pin differs — reject (possible MITM or cert rotation). */
export interface CertTofuEvent {
readonly host: string;
readonly fingerprint: string;
readonly status: "trusted_first_use" | "trusted" | "mismatch";
readonly status: "first_use" | "trusted" | "mismatch";
readonly message?: string;
readonly storedFingerprint?: string;
}
@@ -79,7 +83,7 @@ export function parseStoredFingerprint(message?: string): string | undefined {
}
export type CertMismatchListener = (event: CertTofuEvent) => void;
export type CertFirstTrustListener = (event: CertTofuEvent) => void;
export type CertFirstUseListener = (event: CertTofuEvent) => void;
export interface WsClientConfig {
readonly host: string;
@@ -129,8 +133,13 @@ export function createWsClient() {
// TOFU cert mismatch listeners
const certMismatchListeners = new Set<CertMismatchListener>();
// TOFU first-trust listeners (BUG-133)
const certFirstTrustListeners = new Set<CertFirstTrustListener>();
// TOFU first-use confirmation listeners (F4/F8)
const certFirstUseListeners = new Set<CertFirstUseListener>();
// Global cert-tofu Tauri listener unsub (registered once via startCertListener,
// active for the whole app lifetime so first-use/mismatch events are received
// during the connect page's health checks — before any WS connect).
let certListenerUnsub: (() => void) | null = null;
function setState(newState: ConnectionState): void {
if (state !== newState) {
@@ -299,6 +308,38 @@ export function createWsClient() {
}
}
// Route a cert-tofu event (from the http or ws proxy) to the right listeners.
// Registered globally via startCertListener so first-use/mismatch events are
// received during the connect page's health checks, before any WS connect.
function handleCertTofu(raw: CertTofuEvent): void {
log.info("TOFU cert event", { host: raw.host, status: raw.status });
if (raw.status === "first_use") {
log.warn("TOFU: first-use certificate — awaiting user confirmation", {
host: raw.host,
fingerprint: raw.fingerprint,
});
for (const listener of certFirstUseListeners) {
listener(raw);
}
} else if (raw.status === "mismatch") {
const evt: CertTofuEvent = {
...raw,
storedFingerprint: raw.storedFingerprint ?? parseStoredFingerprint(raw.message),
};
log.error("Certificate fingerprint mismatch!", {
host: evt.host,
fingerprint: evt.fingerprint,
storedFingerprint: evt.storedFingerprint,
});
certMismatchBlock = true;
setState("disconnected");
for (const listener of certMismatchListeners) {
listener(evt);
}
}
// "trusted" → no action
}
async function setupEventListeners(): Promise<void> {
if (tauriListen === null) return;
@@ -356,38 +397,15 @@ export function createWsClient() {
});
eventUnsubs.push(unsubErr);
// TOFU certificate events
const unsubCert = await tauriListen("cert-tofu", (e) => {
if (gen !== wsGeneration) return;
const raw = e.payload as CertTofuEvent;
log.info("TOFU cert event", { host: raw.host, status: raw.status });
if (raw.status === "trusted_first_use") {
log.warn("TOFU: first-use certificate trust", {
host: raw.host,
fingerprint: raw.fingerprint,
});
for (const listener of certFirstTrustListeners) {
listener(raw);
}
} else if (raw.status === "mismatch") {
const evt: CertTofuEvent = {
...raw,
storedFingerprint: parseStoredFingerprint(raw.message),
};
log.error("Certificate fingerprint mismatch!", {
host: evt.host,
fingerprint: evt.fingerprint,
storedFingerprint: evt.storedFingerprint,
});
certMismatchBlock = true;
setState("disconnected");
for (const listener of certMismatchListeners) {
listener(evt);
}
}
});
eventUnsubs.push(unsubCert);
// Register the global cert-tofu listener on first connect (idempotent).
// startCertListener() registers the same listener at app bootstrap so
// first-use/mismatch events are also caught during the connect page's health
// checks, before any WS connection exists.
if (certListenerUnsub === null) {
certListenerUnsub = await tauriListen("cert-tofu", (e) => {
handleCertTofu(e.payload as CertTofuEvent);
});
}
}
function cleanupEventListeners(): void {
@@ -556,10 +574,24 @@ export function createWsClient() {
return () => sendFailureListeners.delete(listener);
},
/** Register a listener for TOFU first-trust events (BUG-133). */
onCertFirstTrust(listener: CertFirstTrustListener): () => void {
certFirstTrustListeners.add(listener);
return () => certFirstTrustListeners.delete(listener);
/**
* Register the global cert-tofu event listener. Idempotent. Call once at app
* bootstrap (before the connect page's health checks) so first-use and
* mismatch events are received even before a WS connection exists.
*/
async startCertListener(): Promise<void> {
if (certListenerUnsub !== null) return;
await ensureTauriApis();
if (tauriListen === null) return;
certListenerUnsub = await tauriListen("cert-tofu", (e) => {
handleCertTofu(e.payload as CertTofuEvent);
});
},
/** Register a listener for TOFU first-use confirmation events (F4/F8). */
onCertFirstUse(listener: CertFirstUseListener): () => void {
certFirstUseListeners.add(listener);
return () => certFirstUseListeners.delete(listener);
},
/** Register a listener for TOFU certificate mismatch events. */
+49 -27
View File
@@ -26,7 +26,7 @@ import { createLogger } from "@lib/logger";
import { initLogPersistence, flushLogs } from "@lib/logPersistence";
import { saveCredential, loadCredential, deleteCredential } from "@lib/credentials";
import { initWindowState } from "@lib/window-state";
import { createCertMismatchModal } from "@components/CertMismatchModal";
import { createCertMismatchModal, createCertFirstUseModal } from "@components/CertMismatchModal";
import { createProfileManager, createTauriBackend } from "@lib/profiles";
import type { CertTofuEvent } from "@lib/ws";
@@ -105,37 +105,49 @@ let dispatcherCleanup: (() => void) | null = null;
let connectedOverlay: ConnectedOverlayControl | null = null;
let lastConnectHost = "";
let lastConnectToken = "";
// Re-run the connect page's health checks (set while the connect page is
// mounted, cleared otherwise) — refreshes a server's status after its
// certificate is trusted for the first time.
let rerunConnectHealth: (() => void) | null = null;
// Certificate first-trust notification (BUG-133).
// Show a brief banner so the user is aware a new server cert was pinned.
ws.onCertFirstTrust((evt: CertTofuEvent) => {
log.warn("TOFU: first-use certificate pinned", {
// Shared guard so the first-use and mismatch cert modals never stack.
let certModalActive = false;
// First-use certificate confirmation (F4/F8). The Rust proxy REJECTS the first
// connection to a server until the user confirms its fingerprint, so no
// credential is ever sent to an unconfirmed host. This fires during the connect
// page's health check (the first TLS contact), before login.
ws.onCertFirstUse((evt: CertTofuEvent) => {
if (certModalActive) return;
certModalActive = true;
const modal = createCertFirstUseModal({
host: evt.host,
fingerprint: evt.fingerprint,
onAccept: () => {
modal.destroy?.();
certModalActive = false;
void (async () => {
try {
await ws.acceptCertFingerprint(evt.host, evt.fingerprint);
// Refresh server health so the now-trusted host becomes reachable,
// and resume a pending connect if one was in flight.
rerunConnectHealth?.();
if (lastConnectHost && lastConnectToken) {
ws.connect({ host: lastConnectHost, token: lastConnectToken });
}
} catch (err) {
log.error("Failed to trust first-use certificate", err);
}
})();
},
onReject: () => {
modal.destroy?.();
certModalActive = false;
},
});
const banner = document.createElement("div");
Object.assign(banner.style, {
position: "fixed",
top: "12px",
left: "50%",
transform: "translateX(-50%)",
background: "#2d5a27",
color: "#e0e0e0",
padding: "10px 20px",
borderRadius: "8px",
fontSize: "13px",
zIndex: "10000",
boxShadow: "0 4px 12px rgba(0,0,0,0.5)",
cursor: "default",
});
banner.textContent = `New server certificate trusted for ${evt.host}`;
banner.title = `SHA-256: ${evt.fingerprint}`;
document.body.appendChild(banner);
setTimeout(() => banner.remove(), 8000);
modal.mount(document.body);
});
// Certificate mismatch modal handler
let certModalActive = false;
ws.onCertMismatch((evt: CertTofuEvent) => {
if (certModalActive) return;
certModalActive = true;
@@ -169,6 +181,10 @@ ws.onCertMismatch((evt: CertTofuEvent) => {
modal.mount(document.body);
});
// Register the global cert-tofu listener now so first-use / mismatch prompts
// are received during the connect page's health checks, before any WS connect.
void ws.startCertListener();
// Current page component reference for cleanup
let currentPage: { destroy?(): void } | null = null;
@@ -224,6 +240,8 @@ function renderPage(pageId: "connect" | "main"): void {
currentPage?.destroy?.();
currentPage = null;
appEl!.textContent = "";
// Only valid while the connect page is mounted (re-set in its render branch).
rerunConnectHealth = null;
// Shared helper for post-auth WS connect + overlay flow
function wirePostAuth(
@@ -433,6 +451,10 @@ function renderPage(pageId: "connect" | "main"): void {
},
};
// Expose a health-refresh hook so trusting a first-use certificate can
// re-check the now-reachable server without a full page navigation.
rerunConnectHealth = () => runHealthChecks(connectPage, getProfileList());
// Load saved profiles and kick off health checks
void (async () => {
try {
@@ -286,5 +286,3 @@ export function createConnectPage(
},
};
}
export type ConnectPage = ReturnType<typeof createConnectPage>;
@@ -559,5 +559,3 @@ export function createMainPage(options: MainPageOptions): MountableComponent {
return { mount, destroy };
}
export type MainPage = ReturnType<typeof createMainPage>;
+5 -1
View File
@@ -78,7 +78,11 @@ export function createMockWsClient() {
return () => sendFailureListeners.delete(listener);
},
onCertFirstTrust(): () => void {
async startCertListener(): Promise<void> {
// no-op in mock
},
onCertFirstUse(): () => void {
return () => {};
},
@@ -67,7 +67,9 @@ function createMockWsClient(): MockWsClient {
return () => {};
},
onCertFirstTrust(): () => void {
async startCertListener(): Promise<void> {},
onCertFirstUse(): () => void {
return () => {};
},
@@ -0,0 +1,37 @@
import { describe, it, expect, vi } from "vitest";
import { createCertFirstUseModal } from "../../src/components/CertMismatchModal";
describe("createCertFirstUseModal (F4/F8)", () => {
it("renders the host + fingerprint and wires accept/reject", () => {
const onAccept = vi.fn();
const onReject = vi.fn();
const modal = createCertFirstUseModal({
host: "example.com:8443",
fingerprint: "aa:bb:cc:dd:ee:ff",
onAccept,
onReject,
});
const container = document.createElement("div");
modal.mount(container);
const text = container.textContent ?? "";
expect(text).toContain("example.com:8443");
expect(text).toContain("aa:bb:cc:dd:ee:ff");
const buttons = Array.from(container.querySelectorAll("button"));
const trustBtn = buttons.find((b) => b.textContent === "Trust This Certificate");
const cancelBtn = buttons.find((b) => b.textContent === "Cancel");
expect(trustBtn).toBeTruthy();
expect(cancelBtn).toBeTruthy();
trustBtn!.click();
expect(onAccept).toHaveBeenCalledTimes(1);
cancelBtn!.click();
expect(onReject).toHaveBeenCalledTimes(1);
modal.destroy?.();
expect(container.querySelector(".modal-overlay")).toBeNull();
});
});
@@ -65,7 +65,8 @@ function createMockWs() {
sendFailureListeners.add(listener);
return () => sendFailureListeners.delete(listener);
},
onCertFirstTrust: vi.fn(() => () => {}),
startCertListener: vi.fn(async () => {}),
onCertFirstUse: vi.fn(() => () => {}),
onCertMismatch: vi.fn(() => () => {}),
acceptCertFingerprint: vi.fn(async () => {}),
getState: vi.fn(() => "disconnected" as const),
@@ -509,6 +509,19 @@ describe("parseOgTags", () => {
const meta = parseOgTags(html);
expect(meta.title).toBe("Spaced Title");
});
it("ignores meta-like strings inside comments and scripts (F7: real HTML parsing, not backtracking regex)", () => {
// A real HTML tokenizer treats these as a comment node and script text, not
// <meta> elements — so attacker-controlled preview HTML can neither smuggle a
// fake OG tag nor drive a polynomial-backtracking regex. A string-matching
// regex would (wrongly) pick up "Commented Out".
const html = `<html><head>
<!-- <meta property="og:title" content="Commented Out"> -->
<script>var s = '<meta property="og:title" content="In Script">';</script>
</head></html>`;
const meta = parseOgTags(html);
expect(meta.title).toBeNull();
});
});
describe("applyOgMeta", () => {
@@ -576,99 +576,6 @@ describe("log persistence", () => {
});
});
// -----------------------------------------------------------------------
// readAllPersistedLogs
// -----------------------------------------------------------------------
describe("readAllPersistedLogs", () => {
it("returns empty string when not initialized (logDir is null)", async () => {
const { readAllPersistedLogs } = await freshImport();
const result = await readAllPersistedLogs();
expect(result).toBe("");
});
it("reads and concatenates all jsonl files in sorted order", async () => {
captureListener();
const { initLogPersistence, readAllPersistedLogs } = await freshImport();
await initLogPersistence();
mockReadDir.mockResolvedValueOnce([
{ name: "2025-06-14.jsonl", isDirectory: false },
{ name: "2025-06-15.jsonl", isDirectory: false },
{ name: "2025-06-13.jsonl", isDirectory: false },
]);
mockReadTextFile
.mockResolvedValueOnce('{"day":"13"}\n')
.mockResolvedValueOnce('{"day":"14"}\n')
.mockResolvedValueOnce('{"day":"15"}\n');
const result = await readAllPersistedLogs();
// Files should be read in sorted order: 13, 14, 15
expect(mockReadTextFile).toHaveBeenCalledTimes(3);
expect(mockReadTextFile.mock.calls[0]![0]).toContain("2025-06-13");
expect(mockReadTextFile.mock.calls[1]![0]).toContain("2025-06-14");
expect(mockReadTextFile.mock.calls[2]![0]).toContain("2025-06-15");
expect(result).toBe('{"day":"13"}\n{"day":"14"}\n{"day":"15"}\n');
});
it("filters out directories and non-jsonl entries", async () => {
captureListener();
const { initLogPersistence, readAllPersistedLogs } = await freshImport();
await initLogPersistence();
mockReadDir.mockResolvedValueOnce([
{ name: "2025-06-15.jsonl", isDirectory: false },
{ name: "subdir", isDirectory: true },
{ name: "readme.txt", isDirectory: false },
]);
mockReadTextFile.mockResolvedValueOnce('{"msg":"only"}\n');
const result = await readAllPersistedLogs();
expect(mockReadTextFile).toHaveBeenCalledTimes(1);
expect(result).toBe('{"msg":"only"}\n');
});
it("returns empty string when directory has no jsonl files", async () => {
captureListener();
const { initLogPersistence, readAllPersistedLogs } = await freshImport();
await initLogPersistence();
mockReadDir.mockResolvedValueOnce([{ name: "notes.txt", isDirectory: false }]);
const result = await readAllPersistedLogs();
expect(result).toBe("");
expect(mockReadTextFile).not.toHaveBeenCalled();
});
it("returns empty string on readDir failure", async () => {
captureListener();
const { initLogPersistence, readAllPersistedLogs } = await freshImport();
await initLogPersistence();
mockReadDir.mockRejectedValueOnce(new Error("no access"));
const result = await readAllPersistedLogs();
expect(result).toBe("");
});
it("returns empty string on readTextFile failure", async () => {
captureListener();
const { initLogPersistence, readAllPersistedLogs } = await freshImport();
await initLogPersistence();
mockReadDir.mockResolvedValueOnce([{ name: "2025-06-15.jsonl", isDirectory: false }]);
mockReadTextFile.mockRejectedValueOnce(new Error("corrupt file"));
const result = await readAllPersistedLogs();
// The entire function returns "" on any error
expect(result).toBe("");
});
});
// -----------------------------------------------------------------------
// JSONL format
// -----------------------------------------------------------------------
@@ -780,21 +687,6 @@ describe("log persistence", () => {
expect(mockRemove).not.toHaveBeenCalled();
});
it("handles entries with undefined name in readAllPersistedLogs", async () => {
captureListener();
const { initLogPersistence, readAllPersistedLogs } = await freshImport();
await initLogPersistence();
mockReadDir.mockResolvedValueOnce([
{ name: undefined, isDirectory: false },
{ name: "2025-06-15.jsonl", isDirectory: false },
]);
mockReadTextFile.mockResolvedValueOnce('{"msg":"ok"}\n');
const result = await readAllPersistedLogs();
expect(result).toBe('{"msg":"ok"}\n');
});
it("multiple rapid entries reuse the same debounce timer", async () => {
const { getListener } = captureListener();
const { initLogPersistence } = await freshImport();
@@ -27,7 +27,8 @@ function createMockWs(state: "connected" | "disconnected" = "connected"): WsClie
stateListeners.add(listener);
return () => stateListeners.delete(listener);
}),
onCertFirstTrust: vi.fn().mockReturnValue(() => {}),
startCertListener: vi.fn().mockResolvedValue(undefined),
onCertFirstUse: vi.fn().mockReturnValue(() => {}),
onCertMismatch: vi.fn().mockReturnValue(() => {}),
acceptCertFingerprint: vi.fn(),
getState: vi.fn(() => currentState),
+36
View File
@@ -629,6 +629,42 @@ describe("cert mismatch blocking", () => {
expect(mockInvoke).toHaveBeenCalledWith("ws_connect", expect.anything());
});
it("routes first_use cert events to onCertFirstUse, not onCertMismatch (F4/F8)", async () => {
const firstUse: unknown[] = [];
const mismatch: unknown[] = [];
client.onCertFirstUse((e) => firstUse.push(e));
client.onCertMismatch((e) => mismatch.push(e));
client.connect({ host: "localhost:8443", token: "t" });
await vi.advanceTimersByTimeAsync(10);
emitTauriEvent("cert-tofu", {
host: "localhost:8443",
fingerprint: "sha256:NEW",
status: "first_use",
});
expect(firstUse).toHaveLength(1);
expect(mismatch).toHaveLength(0);
});
it("startCertListener catches cert events before any WS connect (connect-page path)", async () => {
const firstUse: unknown[] = [];
client.onCertFirstUse((e) => firstUse.push(e));
// No connect() — main.ts registers the listener at bootstrap so first-use
// fires during the connect page's health check, before login.
await client.startCertListener();
emitTauriEvent("cert-tofu", {
host: "localhost:8443",
fingerprint: "sha256:NEW",
status: "first_use",
});
expect(firstUse).toHaveLength(1);
});
it("should not schedule reconnect when certMismatchBlock is true", async () => {
const mismatchEvents: unknown[] = [];
client.onCertMismatch((evt) => mismatchEvents.push(evt));
+1
View File
@@ -29,6 +29,7 @@ linters:
excludes:
- G104 # unhandled errors — errcheck covers this better
- G304 # file path from variable — expected in file storage code
- G306 # WriteFile perms ≤0600 — our only hits are generated source files (genprotocol), which must stay world-readable or multi-stage container builds break
- G706 # log injection — false positive with slog structured logging (values are typed key-value pairs, not interpolated)
exclusions:
-8
View File
@@ -56,11 +56,3 @@ func NewHandler(database *db.DB, version string, hub HubBroadcaster, u *updater.
return r
}
// Handler returns the admin panel http.Handler using a nil database.
//
// Deprecated: use NewHandler instead. Kept for backwards-compat with any
// caller that already imported this symbol before Phase 6.
func Handler() http.Handler {
return http.FileServer(http.FS(staticFiles))
}
+3 -29
View File
@@ -1,6 +1,7 @@
package admin_test
import (
"context"
"net/http"
"net/http/httptest"
"os"
@@ -115,33 +116,6 @@ func TestNewHandler_WithUpdater(t *testing.T) {
}
}
// ─── Handler (deprecated) ────────────────────────────────────────────────────
// TestHandler_ReturnsNonNil verifies the deprecated Handler() function returns
// a non-nil http.Handler (it serves the embedded static files).
func TestHandler_ReturnsNonNil(t *testing.T) {
h := admin.Handler()
if h == nil {
t.Fatal("Handler() returned nil")
}
}
// TestHandler_ServesEmbeddedFiles verifies that the deprecated Handler() serves
// a response (the embedded static FS) without panicking.
func TestHandler_ServesEmbeddedFiles(t *testing.T) {
h := admin.Handler()
req := httptest.NewRequest(http.MethodGet, "/index.html", nil)
w := httptest.NewRecorder()
h.ServeHTTP(w, req)
// http.FileServer returns 200 for a found file or 301/404 for others;
// the important thing is it doesn't panic and returns a valid HTTP status.
if w.Code == 0 {
t.Error("Handler() response has zero status code")
}
}
// ─── ownerOnlyMiddleware (tested via API endpoints that use it) ───────────────
// TestOwnerOnlyMiddleware_OwnerAllowed verifies that a user with Owner role
@@ -186,9 +160,9 @@ func TestOwnerOnlyMiddleware_AdminDenied(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
// Create admin user (role_id=2, position=80)
adminUID, _ := database.CreateUser("middlewareadmin", "hash", 2)
adminUID, _ := database.CreateUser(context.Background(), "middlewareadmin", "hash", 2)
token := "mw-admin-token"
_, _ = database.CreateSession(adminUID, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), adminUID, auth.HashToken(token), "test", "127.0.0.1")
w := doRequest(t, handler, http.MethodPost, "/backup", token, nil)
+18 -18
View File
@@ -41,8 +41,8 @@ func TestAdminAPI_PatchUser_UnbanUser(t *testing.T) {
token := createAdminUser(t, database)
// Create and ban a target user first.
targetUID, _ := database.CreateUser("unbanme", "hash", 3)
_ = database.BanUser(targetUID, "test ban", nil)
targetUID, _ := database.CreateUser(context.Background(), "unbanme", "hash", 3)
_ = database.BanUser(context.Background(), targetUID, "test ban", nil)
body := map[string]any{"banned": false}
w := doRequest(t, handler, http.MethodPatch, "/users/"+itoa(targetUID), token, body)
@@ -52,7 +52,7 @@ func TestAdminAPI_PatchUser_UnbanUser(t *testing.T) {
}
// Verify the user is now unbanned.
user, _ := database.GetUserByID(targetUID)
user, _ := database.GetUserByID(context.Background(), targetUID)
if user.Banned {
t.Error("user is still banned after unban request")
}
@@ -64,7 +64,7 @@ func TestAdminAPI_PatchUser_InvalidBody(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", nil, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
targetUID, _ := database.CreateUser("invalidbody", "hash", 3)
targetUID, _ := database.CreateUser(context.Background(), "invalidbody", "hash", 3)
req := httptest.NewRequest(http.MethodPatch, "/users/"+itoa(targetUID), bytes.NewReader([]byte("not-json")))
req.Header.Set("Authorization", "Bearer "+token)
@@ -147,7 +147,7 @@ func TestAdminAPI_PatchChannel_InvalidBody(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", nil, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
chID, _ := database.AdminCreateChannel("malformed", "text", "", "", 0)
chID, _ := database.AdminCreateChannel(context.Background(), "malformed", "text", "", "", 0)
req := httptest.NewRequest(http.MethodPatch, "/channels/"+itoa(chID), bytes.NewReader([]byte("not-json")))
req.Header.Set("Authorization", "Bearer "+token)
@@ -235,9 +235,9 @@ func TestAdminAPI_AuditLog_Pagination(t *testing.T) {
token := createAdminUser(t, database)
// Create several audit entries.
uid, _ := database.CreateUser("auditpager", "hash", 1)
uid, _ := database.CreateUser(context.Background(), "auditpager", "hash", 1)
for i := 0; i < 5; i++ {
_ = database.LogAudit(uid, "TEST", "test", int64(i), "")
_ = database.LogAudit(context.Background(), uid, "TEST", "test", int64(i), "")
}
// Fetch page 2 with limit=2, offset=2 — should return 2 entries.
@@ -324,7 +324,7 @@ func TestAdminAPI_PatchUser_BanNilHubDoesNotPanic(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", nil, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
targetUID, _ := database.CreateUser("ban-nohub", "hash", 3)
targetUID, _ := database.CreateUser(context.Background(), "ban-nohub", "hash", 3)
body := map[string]any{"banned": true, "ban_reason": "nil hub test"}
w := doRequest(t, handler, http.MethodPatch, "/users/"+itoa(targetUID), token, body)
@@ -334,7 +334,7 @@ func TestAdminAPI_PatchUser_BanNilHubDoesNotPanic(t *testing.T) {
}
// Verify ban was still applied despite nil hub.
user, _ := database.GetUserByID(targetUID)
user, _ := database.GetUserByID(context.Background(), targetUID)
if !user.Banned {
t.Error("user should be banned even with nil hub")
}
@@ -363,7 +363,7 @@ func TestAdminAPI_LogStreamTicketFlow(t *testing.T) {
if payload.Ticket == "" {
t.Fatal("expected non-empty log stream ticket")
}
if err := database.DeleteSession(auth.HashToken(token)); err != nil {
if err := database.DeleteSession(context.Background(), auth.HashToken(token)); err != nil {
t.Fatalf("DeleteSession: %v", err)
}
@@ -412,7 +412,7 @@ func TestAdminAPI_LogStreamTicketFlow(t *testing.T) {
t.Fatalf("legacy token stream status = %d, want 401; body: %s", legacyResp.StatusCode, string(body))
}
if _, err := database.CreateSession(1, auth.HashToken(token), "test", "127.0.0.1"); err != nil {
if _, err := database.CreateSession(context.Background(), 1, auth.HashToken(token), "test", "127.0.0.1"); err != nil {
t.Fatalf("CreateSession: %v", err)
}
ticketResp = doRequest(t, handler, http.MethodPost, "/logs/ticket", token, nil)
@@ -422,7 +422,7 @@ func TestAdminAPI_LogStreamTicketFlow(t *testing.T) {
if err := json.Unmarshal(ticketResp.Body.Bytes(), &payload); err != nil {
t.Fatalf("unmarshal restored ticket response: %v", err)
}
if err := database.UpdateUserRole(1, 3); err != nil {
if err := database.UpdateUserRole(context.Background(), 1, 3); err != nil {
t.Fatalf("UpdateUserRole: %v", err)
}
demotedResp, err := http.Get(srv.URL + "/logs/stream?ticket=" + payload.Ticket)
@@ -444,7 +444,7 @@ func TestAdminAPI_PatchUser_RoleChangeNilHubDoesNotPanic(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", nil, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
targetUID, _ := database.CreateUser("role-nohub", "hash", 3)
targetUID, _ := database.CreateUser(context.Background(), "role-nohub", "hash", 3)
body := map[string]any{"role_id": float64(2)}
w := doRequest(t, handler, http.MethodPatch, "/users/"+itoa(targetUID), token, body)
@@ -454,7 +454,7 @@ func TestAdminAPI_PatchUser_RoleChangeNilHubDoesNotPanic(t *testing.T) {
}
// Verify role was still changed despite nil hub.
user, _ := database.GetUserByID(targetUID)
user, _ := database.GetUserByID(context.Background(), targetUID)
if user.RoleID != 2 {
t.Errorf("RoleID = %d, want 2", user.RoleID)
}
@@ -469,7 +469,7 @@ func TestAdminAPI_PatchUser_BanWithoutReason(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
targetUID, _ := database.CreateUser("banwithout", "hash", 3)
targetUID, _ := database.CreateUser(context.Background(), "banwithout", "hash", 3)
// No ban_reason in body — the nil check in handlePatchUser uses empty string.
body := map[string]any{"banned": true}
@@ -490,7 +490,7 @@ func TestAdminAPI_PatchUser_RoleChangeBroadcast(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", hub, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
targetUID, _ := database.CreateUser("rolebroadcast", "hash", 3)
targetUID, _ := database.CreateUser(context.Background(), "rolebroadcast", "hash", 3)
body := map[string]any{"role_id": float64(2)}
w := doRequest(t, handler, http.MethodPatch, "/users/"+itoa(targetUID), token, body)
@@ -531,7 +531,7 @@ func TestAdminAPI_SetupStatus_AlreadySetup(t *testing.T) {
database := openAdminTestDB(t)
handler := admin.NewAdminAPI(database, "1.0.0", nil, nil, nil, nil, nil, newTestModService(database))
_, _ = database.CreateUser("existing", "hash", 1)
_, _ = database.CreateUser(context.Background(), "existing", "hash", 1)
w := doRequest(t, handler, http.MethodGet, "/setup/status", "", nil)
@@ -583,7 +583,7 @@ func TestAdminAPI_Setup_AlreadyCompleted(t *testing.T) {
database := openAdminTestDB(t)
handler := admin.NewAdminAPI(database, "1.0.0", nil, nil, nil, nil, nil, newTestModService(database))
_, _ = database.CreateUser("existing", "hash", 1)
_, _ = database.CreateUser(context.Background(), "existing", "hash", 1)
body := map[string]string{
"username": "hacker",
+43 -42
View File
@@ -2,6 +2,7 @@ package admin_test
import (
"bytes"
"context"
"encoding/json"
"fmt"
"net/http"
@@ -160,14 +161,14 @@ func openAdminTestDB(t *testing.T) *db.DB {
func createAdminUser(t *testing.T, database *db.DB) string {
t.Helper()
// Owner role has permissions = 2147483647 (includes ADMINISTRATOR bit 0x40000000)
uid, err := database.CreateUser("adminuser", "$2a$12$placeholder", 1)
uid, err := database.CreateUser(context.Background(), "adminuser", "$2a$12$placeholder", 1)
if err != nil {
t.Fatalf("CreateUser admin: %v", err)
}
token := "test-admin-token-" + t.Name()
tokenHash := auth.HashToken(token)
if _, err := database.CreateSession(uid, tokenHash, "test", "127.0.0.1"); err != nil {
if _, err := database.CreateSession(context.Background(), uid, tokenHash, "test", "127.0.0.1"); err != nil {
t.Fatalf("CreateSession: %v", err)
}
return token
@@ -177,14 +178,14 @@ func createAdminUser(t *testing.T, database *db.DB) string {
func createMemberUser(t *testing.T, database *db.DB) string {
t.Helper()
// Member role (id=3) has limited permissions, not ADMINISTRATOR
uid, err := database.CreateUser("memberuser", "$2a$12$placeholder", 3)
uid, err := database.CreateUser(context.Background(), "memberuser", "$2a$12$placeholder", 3)
if err != nil {
t.Fatalf("CreateUser member: %v", err)
}
token := "test-member-token-" + t.Name()
tokenHash := auth.HashToken(token)
if _, err := database.CreateSession(uid, tokenHash, "test", "127.0.0.1"); err != nil {
if _, err := database.CreateSession(context.Background(), uid, tokenHash, "test", "127.0.0.1"); err != nil {
t.Fatalf("CreateSession: %v", err)
}
return token
@@ -322,7 +323,7 @@ func TestAdminAPI_PatchUser_BanHierarchy(t *testing.T) {
ownerToken := createAdminUser(t, database) // Owner role (pos 100)
// A second owner-rank user: equal position, cannot be banned.
peerUID, err := database.CreateUser("peerowner", "$2a$12$placeholder", 1)
peerUID, err := database.CreateUser(context.Background(), "peerowner", "$2a$12$placeholder", 1)
if err != nil {
t.Fatalf("CreateUser peerowner: %v", err)
}
@@ -331,27 +332,27 @@ func TestAdminAPI_PatchUser_BanHierarchy(t *testing.T) {
if w.Code != http.StatusForbidden {
t.Fatalf("equal-rank ban: status = %d, want 403; body: %s", w.Code, w.Body.String())
}
if u, _ := database.GetUserByID(peerUID); u.Banned {
if u, _ := database.GetUserByID(context.Background(), peerUID); u.Banned {
t.Fatal("equal-rank target must not be banned")
}
// A lower-positioned role that still holds ADMINISTRATOR (panel access):
// its holder must not be able to ban the higher-ranked owner.
if _, err := database.Exec(
if _, err := database.ExecContext(context.Background(),
`INSERT INTO roles (id, name, permissions, position, is_default) VALUES (9, 'JuniorAdmin', ?, 50, 0)`,
permissions.Administrator,
); err != nil {
t.Fatalf("inserting junior admin role: %v", err)
}
juniorUID, err := database.CreateUser("junioradmin", "$2a$12$placeholder", 9)
juniorUID, err := database.CreateUser(context.Background(), "junioradmin", "$2a$12$placeholder", 9)
if err != nil {
t.Fatalf("CreateUser junioradmin: %v", err)
}
juniorToken := "junior-token-" + t.Name()
if _, err := database.CreateSession(juniorUID, auth.HashToken(juniorToken), "test", "127.0.0.1"); err != nil {
if _, err := database.CreateSession(context.Background(), juniorUID, auth.HashToken(juniorToken), "test", "127.0.0.1"); err != nil {
t.Fatalf("CreateSession junior: %v", err)
}
ownerUser, err := database.GetUserByUsername("adminuser")
ownerUser, err := database.GetUserByUsername(context.Background(), "adminuser")
if err != nil || ownerUser == nil {
t.Fatalf("GetUserByUsername adminuser: %v", err)
}
@@ -360,12 +361,12 @@ func TestAdminAPI_PatchUser_BanHierarchy(t *testing.T) {
if w.Code != http.StatusForbidden {
t.Fatalf("junior bans owner: status = %d, want 403; body: %s", w.Code, w.Body.String())
}
if u, _ := database.GetUserByID(ownerUser.ID); u.Banned {
if u, _ := database.GetUserByID(context.Background(), ownerUser.ID); u.Banned {
t.Fatal("owner must not be banned by a lower rank")
}
// Downward ban still works: junior admin (pos 50) bans a member (pos 40).
memberUID, err := database.CreateUser("banme", "$2a$12$placeholder", 3)
memberUID, err := database.CreateUser(context.Background(), "banme", "$2a$12$placeholder", 3)
if err != nil {
t.Fatalf("CreateUser banme: %v", err)
}
@@ -374,7 +375,7 @@ func TestAdminAPI_PatchUser_BanHierarchy(t *testing.T) {
if w.Code != http.StatusOK {
t.Fatalf("junior bans member: status = %d, want 200; body: %s", w.Code, w.Body.String())
}
if u, _ := database.GetUserByID(memberUID); !u.Banned {
if u, _ := database.GetUserByID(context.Background(), memberUID); !u.Banned {
t.Fatal("member should be banned by higher-ranked actor")
}
}
@@ -385,7 +386,7 @@ func TestAdminAPI_PatchUser_BanUser(t *testing.T) {
token := createAdminUser(t, database)
// Create a target user
targetUID, _ := database.CreateUser("target", "hash", 3)
targetUID, _ := database.CreateUser(context.Background(), "target", "hash", 3)
body := map[string]any{
"banned": true,
@@ -398,7 +399,7 @@ func TestAdminAPI_PatchUser_BanUser(t *testing.T) {
}
// Verify user is banned in DB
user, err := database.GetUserByID(targetUID)
user, err := database.GetUserByID(context.Background(), targetUID)
if err != nil {
t.Fatalf("GetUserByID: %v", err)
}
@@ -412,7 +413,7 @@ func TestAdminAPI_PatchUser_ChangeRole(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
targetUID, _ := database.CreateUser("rolechange", "hash", 3)
targetUID, _ := database.CreateUser(context.Background(), "rolechange", "hash", 3)
body := map[string]any{
"role_id": float64(2),
@@ -423,7 +424,7 @@ func TestAdminAPI_PatchUser_ChangeRole(t *testing.T) {
t.Errorf("status = %d, want 200; body: %s", w.Code, w.Body.String())
}
user, _ := database.GetUserByID(targetUID)
user, _ := database.GetUserByID(context.Background(), targetUID)
if user.RoleID != 2 {
t.Errorf("RoleID = %d, want 2", user.RoleID)
}
@@ -461,8 +462,8 @@ func TestAdminAPI_ForceLogout_OK(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
targetUID, _ := database.CreateUser("logoutme", "hash", 3)
_, _ = database.CreateSession(targetUID, "victim-token-hash", "web", "1.2.3.4")
targetUID, _ := database.CreateUser(context.Background(), "logoutme", "hash", 3)
_, _ = database.CreateSession(context.Background(), targetUID, "victim-token-hash", "web", "1.2.3.4")
w := doRequest(t, handler, http.MethodDelete, "/users/"+itoa(targetUID)+"/sessions", token, nil)
@@ -470,7 +471,7 @@ func TestAdminAPI_ForceLogout_OK(t *testing.T) {
t.Errorf("status = %d, want 204", w.Code)
}
sessions, _ := database.GetUserSessions(targetUID)
sessions, _ := database.GetUserSessions(context.Background(), targetUID)
if len(sessions) != 0 {
t.Errorf("expected 0 sessions after force logout, got %d", len(sessions))
}
@@ -494,7 +495,7 @@ func TestAdminAPI_ListChannels_OK(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
_, _ = database.AdminCreateChannel("general", "text", "", "", 0)
_, _ = database.AdminCreateChannel(context.Background(), "general", "text", "", "", 0)
w := doRequest(t, handler, http.MethodGet, "/channels", token, nil)
@@ -562,7 +563,7 @@ func TestAdminAPI_UpdateChannel_OK(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
chID, _ := database.AdminCreateChannel("old", "text", "", "", 0)
chID, _ := database.AdminCreateChannel(context.Background(), "old", "text", "", "", 0)
body := map[string]any{
"name": "updated",
@@ -598,7 +599,7 @@ func TestAdminAPI_DeleteChannel_OK(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
chID, _ := database.AdminCreateChannel("del-me", "text", "", "", 0)
chID, _ := database.AdminCreateChannel(context.Background(), "del-me", "text", "", "", 0)
w := doRequest(t, handler, http.MethodDelete, "/channels/"+itoa(chID), token, nil)
@@ -626,8 +627,8 @@ func TestAdminAPI_AuditLog_OK(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
uid, _ := database.CreateUser("actor", "hash", 1)
_ = database.LogAudit(uid, "TEST_ACTION", "user", uid, "detail")
uid, _ := database.CreateUser(context.Background(), "actor", "hash", 1)
_ = database.LogAudit(context.Background(), uid, "TEST_ACTION", "user", uid, "detail")
w := doRequest(t, handler, http.MethodGet, "/audit-log?limit=10&offset=0", token, nil)
@@ -702,7 +703,7 @@ func TestAdminAPI_PatchSettings_OK(t *testing.T) {
}
// Verify the change was persisted
val, err := database.GetSetting("server_name")
val, err := database.GetSetting(context.Background(), "server_name")
if err != nil {
t.Fatalf("GetSetting: %v", err)
}
@@ -733,9 +734,9 @@ func TestAdminAPI_Backup_RequiresOwner(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
// Admin (role 2) can authenticate but is not Owner (role 1, position 100)
adminUID, _ := database.CreateUser("adminonly", "hash", 2)
adminUID, _ := database.CreateUser(context.Background(), "adminonly", "hash", 2)
token := "admin-only-token"
_, _ = database.CreateSession(adminUID, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), adminUID, auth.HashToken(token), "test", "127.0.0.1")
w := doRequest(t, handler, http.MethodPost, "/backup", token, nil)
@@ -768,7 +769,7 @@ func TestAdminAPI_ActorFromContext_AuditEntry(t *testing.T) {
token := createAdminUser(t, database)
// Create a target user to act on.
targetUID, _ := database.CreateUser("ctxtarget", "hash", 3)
targetUID, _ := database.CreateUser(context.Background(), "ctxtarget", "hash", 3)
body := map[string]any{"banned": true, "ban_reason": "context test"}
w := doRequest(t, handler, http.MethodPatch, "/users/"+itoa(targetUID), token, body)
@@ -779,7 +780,7 @@ func TestAdminAPI_ActorFromContext_AuditEntry(t *testing.T) {
// The audit log should have a non-zero actor_id showing the actor was
// resolved (not 0, which would indicate a failed context lookup).
entries, err := database.GetAuditLog(10, 0)
entries, err := database.GetAuditLog(context.Background(), 10, 0)
if err != nil {
t.Fatalf("GetAuditLog: %v", err)
}
@@ -801,8 +802,8 @@ func TestAdminAPI_ActorFromContext_ForceLogout(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
targetUID, _ := database.CreateUser("logoutctx", "hash", 3)
_, _ = database.CreateSession(targetUID, "victim-hash-ctx", "web", "1.2.3.4")
targetUID, _ := database.CreateUser(context.Background(), "logoutctx", "hash", 3)
_, _ = database.CreateSession(context.Background(), targetUID, "victim-hash-ctx", "web", "1.2.3.4")
w := doRequest(t, handler, http.MethodDelete, "/users/"+itoa(targetUID)+"/sessions", token, nil)
@@ -810,7 +811,7 @@ func TestAdminAPI_ActorFromContext_ForceLogout(t *testing.T) {
t.Fatalf("status = %d, want 204; body: %s", w.Code, w.Body.String())
}
entries, err := database.GetAuditLog(10, 0)
entries, err := database.GetAuditLog(context.Background(), 10, 0)
if err != nil {
t.Fatalf("GetAuditLog: %v", err)
}
@@ -870,7 +871,7 @@ func TestAdminAPI_PatchSettings_RejectsMixedKeys(t *testing.T) {
}
// The valid key must NOT have been written because the request was rejected.
val, err := database.GetSetting("server_name")
val, err := database.GetSetting(context.Background(), "server_name")
if err != nil {
t.Fatalf("GetSetting: %v", err)
}
@@ -953,7 +954,7 @@ func TestAdminAPI_PatchSettings_AllowsRequire2FAWhenAllUsersEnrolledAndRegistrat
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
if _, err := database.Exec(`UPDATE users SET totp_secret = ? WHERE id = 1`, "JBSWY3DPEHPK3PXP"); err != nil {
if _, err := database.ExecContext(context.Background(), `UPDATE users SET totp_secret = ? WHERE id = 1`, "JBSWY3DPEHPK3PXP"); err != nil {
t.Fatalf("enroll admin user: %v", err)
}
@@ -993,7 +994,7 @@ func TestAdminAPI_ListUsers_NoPasswordHash(t *testing.T) {
token := createAdminUser(t, database)
// Create a second user so the list is non-trivial.
_, _ = database.CreateUser("plainuser", "supersecretbcrypthash", 3)
_, _ = database.CreateUser(context.Background(), "plainuser", "supersecretbcrypthash", 3)
w := doRequest(t, handler, http.MethodGet, "/users", token, nil)
@@ -1067,7 +1068,7 @@ func TestAdminAPI_PatchUser_NoPasswordHash(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
targetUID, _ := database.CreateUser("patchvictim", "topsecretbcrypt", 3)
targetUID, _ := database.CreateUser(context.Background(), "patchvictim", "topsecretbcrypt", 3)
body := map[string]any{
"banned": true,
@@ -1095,7 +1096,7 @@ func TestAdminAPI_PatchUser_NoTOTPSecret(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
targetUID, _ := database.CreateUser("patchtotp", "hash", 3)
targetUID, _ := database.CreateUser(context.Background(), "patchtotp", "hash", 3)
w := doRequest(t, handler, http.MethodPatch, "/users/"+itoa(targetUID), token, map[string]any{
"banned": false,
@@ -1210,7 +1211,7 @@ func TestAdminAPI_UpdateChannel_BroadcastsChannelUpdate(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", hub, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
chID, _ := database.AdminCreateChannel("before", "text", "", "", 0)
chID, _ := database.AdminCreateChannel(context.Background(), "before", "text", "", "", 0)
body := map[string]any{"name": "after"}
w := doRequest(t, handler, http.MethodPatch, "/channels/"+itoa(chID), token, body)
@@ -1231,7 +1232,7 @@ func TestAdminAPI_UpdateChannel_NilHubDoesNotPanic(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", nil, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
chID, _ := database.AdminCreateChannel("patchme", "text", "", "", 0)
chID, _ := database.AdminCreateChannel(context.Background(), "patchme", "text", "", "", 0)
body := map[string]any{"name": "patched"}
w := doRequest(t, handler, http.MethodPatch, "/channels/"+itoa(chID), token, body)
@@ -1246,7 +1247,7 @@ func TestAdminAPI_DeleteChannel_BroadcastsChannelDelete(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", hub, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
chID, _ := database.AdminCreateChannel("delete-me", "text", "", "", 0)
chID, _ := database.AdminCreateChannel(context.Background(), "delete-me", "text", "", "", 0)
w := doRequest(t, handler, http.MethodDelete, "/channels/"+itoa(chID), token, nil)
@@ -1266,7 +1267,7 @@ func TestAdminAPI_DeleteChannel_NilHubDoesNotPanic(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", nil, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
chID, _ := database.AdminCreateChannel("del-no-hub", "text", "", "", 0)
chID, _ := database.AdminCreateChannel(context.Background(), "del-no-hub", "text", "", "", 0)
w := doRequest(t, handler, http.MethodDelete, "/channels/"+itoa(chID), token, nil)
if w.Code != http.StatusNoContent {
+5
View File
@@ -0,0 +1,5 @@
package admin
// SetBackupBaseDir overrides backupBaseDir so tests can point backup handlers
// at a temp dir. Lives here so it stays out of the production binary.
func SetBackupBaseDir(dir string) { backupBaseDir = dir }
+13 -9
View File
@@ -1,6 +1,7 @@
package admin
import (
"context"
"fmt"
"io"
"log/slog"
@@ -28,9 +29,6 @@ func init() {
}
}
// SetBackupBaseDir overrides backupBaseDir. Intended for tests only.
func SetBackupBaseDir(dir string) { backupBaseDir = dir }
// ─── Backup Handlers ─────────────────────────────────────────────────────────
func handleBackup(database *db.DB) http.Handler {
@@ -44,7 +42,10 @@ func handleBackup(database *db.DB) http.Handler {
timestamp := time.Now().UTC().Format("20060102_150405")
backupPath := filepath.Join(backupDir, "chatserver_"+timestamp+".db")
if err := database.BackupTo(backupPath); err != nil {
// Detached like the restore path's safety backup: an interrupted
// VACUUM INTO leaves a truncated .db that handleListBackups would
// present as restorable.
if err := database.BackupTo(context.WithoutCancel(r.Context()), backupPath); err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "backup failed")
return
}
@@ -52,7 +53,7 @@ func handleBackup(database *db.DB) http.Handler {
actor := actorFromContext(r)
backupName := filepath.Base(backupPath)
slog.Info("database backup created", "actor_id", actor, "name", backupName)
db.WriteAudit(database, actor, "backup_create", "server", 0,
db.WriteAudit(context.WithoutCancel(r.Context()), database, actor, "backup_create", "server", 0,
fmt.Sprintf("backup saved: %s", backupName))
writeJSON(w, http.StatusOK, map[string]string{
@@ -136,7 +137,7 @@ func handleDeleteBackup(database *db.DB) http.Handler {
actor := actorFromContext(r)
slog.Info("backup deleted", "actor_id", actor, "name", name)
db.WriteAudit(database, actor, "backup_delete", "server", 0, "deleted backup "+name)
db.WriteAudit(context.WithoutCancel(r.Context()), database, actor, "backup_delete", "server", 0, "deleted backup "+name)
w.WriteHeader(http.StatusNoContent)
})
@@ -163,9 +164,12 @@ func handleRestoreBackup(database *db.DB, hub HubBroadcaster) http.Handler {
dbPath := filepath.Join("data", "chatserver.db")
// Safety: create a pre-restore backup before overwriting.
// Safety: create a pre-restore backup before overwriting. WithoutCancel:
// the restore proceeds regardless of client disconnect (Close/copyFile
// below are not ctx-aware), so the safety backup must not be skippable
// by a canceled request ctx.
preRestore := filepath.Join("data", "backups", "pre_restore_"+time.Now().UTC().Format("20060102_150405")+".db")
if err := database.BackupTo(preRestore); err != nil {
if err := database.BackupTo(context.WithoutCancel(r.Context()), preRestore); err != nil {
slog.Warn("pre-restore backup failed", "err", err)
}
@@ -174,7 +178,7 @@ func handleRestoreBackup(database *db.DB, hub HubBroadcaster) http.Handler {
// Checkpoint the WAL and close the database connection before overwriting
// to prevent corruption from concurrent writes (BUG-096).
if _, checkpointErr := database.SQLDb().Exec("PRAGMA wal_checkpoint(TRUNCATE)"); checkpointErr != nil {
if _, checkpointErr := database.SQLDb().ExecContext(context.WithoutCancel(r.Context()), "PRAGMA wal_checkpoint(TRUNCATE)"); checkpointErr != nil {
slog.Warn("pre-restore WAL checkpoint failed", "err", checkpointErr)
}
+7 -6
View File
@@ -1,6 +1,7 @@
package admin_test
import (
"context"
"encoding/json"
"net/http"
"os"
@@ -77,9 +78,9 @@ func TestHandleBackup_RequiresOwner(t *testing.T) {
database := openAdminTestDB(t)
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
adminUID, _ := database.CreateUser("backupadmin", "hash", 2)
adminUID, _ := database.CreateUser(context.Background(), "backupadmin", "hash", 2)
token := "backup-admin-token"
_, _ = database.CreateSession(adminUID, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), adminUID, auth.HashToken(token), "test", "127.0.0.1")
w := doRequest(t, handler, http.MethodPost, "/backup", token, nil)
@@ -226,9 +227,9 @@ func TestHandleDeleteBackup_RequiresOwner(t *testing.T) {
database := openAdminTestDB(t)
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
adminUID, _ := database.CreateUser("deladmin", "hash", 2)
adminUID, _ := database.CreateUser(context.Background(), "deladmin", "hash", 2)
token := "del-admin-token"
_, _ = database.CreateSession(adminUID, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), adminUID, auth.HashToken(token), "test", "127.0.0.1")
// Create the file so path validation doesn't return 404 before the 403.
backupDir := filepath.Join(tmpDir, "data", "backups")
@@ -354,9 +355,9 @@ func TestHandleRestoreBackup_RequiresOwner(t *testing.T) {
database := openAdminTestDB(t)
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
adminUID, _ := database.CreateUser("restoreadmin", "hash", 2)
adminUID, _ := database.CreateUser(context.Background(), "restoreadmin", "hash", 2)
token := "restore-admin-token"
_, _ = database.CreateSession(adminUID, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), adminUID, auth.HashToken(token), "test", "127.0.0.1")
// Create files so path checks pass before auth check.
backupDir := filepath.Join(tmpDir, "data", "backups")
+8 -7
View File
@@ -1,6 +1,7 @@
package admin
import (
"context"
"encoding/json"
"fmt"
"log/slog"
@@ -27,7 +28,7 @@ func getPermChannel(database *db.DB, w http.ResponseWriter, r *http.Request) *db
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid channel id")
return nil
}
ch, err := database.GetChannel(id)
ch, err := database.GetChannel(r.Context(), id)
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch channel")
return nil
@@ -55,7 +56,7 @@ func handleGetChannelPermissions(database *db.DB) http.HandlerFunc {
if ch == nil {
return
}
overrides, err := database.ListChannelRoleOverrides(ch.ID)
overrides, err := database.ListChannelRoleOverrides(r.Context(), ch.ID)
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to list channel permissions")
return
@@ -81,7 +82,7 @@ func handlePutChannelPermission(database *db.DB, hub HubBroadcaster, permInvalid
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid role id")
return
}
role, err := database.GetRoleByID(roleID)
role, err := database.GetRoleByID(r.Context(), roleID)
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch role")
return
@@ -100,7 +101,7 @@ func handlePutChannelPermission(database *db.DB, hub HubBroadcaster, permInvalid
allow := req.Allow & permissions.AllPerms
deny := req.Deny & permissions.AllPerms
if err := database.UpsertChannelOverride(ch.ID, roleID, allow, deny); err != nil {
if err := database.UpsertChannelOverride(r.Context(), ch.ID, roleID, allow, deny); err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to save channel permission")
return
}
@@ -108,7 +109,7 @@ func handlePutChannelPermission(database *db.DB, hub HubBroadcaster, permInvalid
actor := actorFromContext(r)
slog.Info("channel permissions updated", "actor_id", actor, "channel_id", ch.ID,
"role_id", roleID, "allow", allow, "deny", deny)
db.WriteAudit(database, actor, "channel_perms_update", "channel", ch.ID,
db.WriteAudit(context.WithoutCancel(r.Context()), database, actor, "channel_perms_update", "channel", ch.ID,
fmt.Sprintf("set overrides for role %s on #%s (allow=%#x deny=%#x)", role.Name, ch.Name, allow, deny))
if permInvalidator != nil {
@@ -140,14 +141,14 @@ func handleDeleteChannelPermission(database *db.DB, hub HubBroadcaster, permInva
return
}
if err := database.DeleteChannelOverride(ch.ID, roleID); err != nil {
if err := database.DeleteChannelOverride(r.Context(), ch.ID, roleID); err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to delete channel permission")
return
}
actor := actorFromContext(r)
slog.Info("channel permissions cleared", "actor_id", actor, "channel_id", ch.ID, "role_id", roleID)
db.WriteAudit(database, actor, "channel_perms_clear", "channel", ch.ID,
db.WriteAudit(context.WithoutCancel(r.Context()), database, actor, "channel_perms_clear", "channel", ch.ID,
fmt.Sprintf("cleared overrides for role %d on #%s", roleID, ch.Name))
if permInvalidator != nil {
+13 -12
View File
@@ -1,6 +1,7 @@
package admin_test
import (
"context"
"encoding/json"
"net/http"
"testing"
@@ -31,7 +32,7 @@ func TestGetChannelPermissions_ReturnsAllRoles(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
chID, err := database.CreateChannel("secret", "text", "", "", 0)
chID, err := database.CreateChannel(context.Background(), "secret", "text", "", "", 0)
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
@@ -80,7 +81,7 @@ func TestGetChannelPermissions_DMRejected(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
chID, err := database.CreateChannel("dm-chan", "dm", "", "", 0)
chID, err := database.CreateChannel(context.Background(), "dm-chan", "dm", "", "", 0)
if err != nil {
t.Fatalf("CreateChannel dm: %v", err)
}
@@ -101,7 +102,7 @@ func TestPutChannelPermission_PersistsAndPropagates(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", hub, nil, nil, nil, inv, newTestModService(database))
token := createAdminUser(t, database)
chID, err := database.CreateChannel("secret", "text", "", "", 0)
chID, err := database.CreateChannel(context.Background(), "secret", "text", "", "", 0)
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
@@ -114,7 +115,7 @@ func TestPutChannelPermission_PersistsAndPropagates(t *testing.T) {
t.Fatalf("status = %d, want 200; body: %s", w.Code, w.Body.String())
}
allow, deny, err := database.GetChannelPermissions(chID, 3)
allow, deny, err := database.GetChannelPermissions(context.Background(), chID, 3)
if err != nil {
t.Fatalf("GetChannelPermissions: %v", err)
}
@@ -129,7 +130,7 @@ func TestPutChannelPermission_PersistsAndPropagates(t *testing.T) {
t.Errorf("RefreshChannelVisibility not called for channel %d", chID)
}
entries, err := database.GetAuditLog(10, 0)
entries, err := database.GetAuditLog(context.Background(), 10, 0)
if err != nil {
t.Fatalf("GetAuditLog: %v", err)
}
@@ -149,7 +150,7 @@ func TestPutChannelPermission_MasksUnknownBits(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
chID, err := database.CreateChannel("secret2", "text", "", "", 0)
chID, err := database.CreateChannel(context.Background(), "secret2", "text", "", "", 0)
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
@@ -162,7 +163,7 @@ func TestPutChannelPermission_MasksUnknownBits(t *testing.T) {
t.Fatalf("status = %d, want 200; body: %s", w.Code, w.Body.String())
}
allow, deny, err := database.GetChannelPermissions(chID, 3)
allow, deny, err := database.GetChannelPermissions(context.Background(), chID, 3)
if err != nil {
t.Fatalf("GetChannelPermissions: %v", err)
}
@@ -179,7 +180,7 @@ func TestPutChannelPermission_UnknownRole(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", &mockHub{}, nil, nil, nil, nil, newTestModService(database))
token := createAdminUser(t, database)
chID, err := database.CreateChannel("secret3", "text", "", "", 0)
chID, err := database.CreateChannel(context.Background(), "secret3", "text", "", "", 0)
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
@@ -197,7 +198,7 @@ func TestPutChannelPermission_NonAdminForbidden(t *testing.T) {
_ = createAdminUser(t, database)
memberToken := createMemberUser(t, database)
chID, err := database.CreateChannel("secret4", "text", "", "", 0)
chID, err := database.CreateChannel(context.Background(), "secret4", "text", "", "", 0)
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
@@ -218,11 +219,11 @@ func TestDeleteChannelPermission_ClearsOverride(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", hub, nil, nil, nil, inv, newTestModService(database))
token := createAdminUser(t, database)
chID, err := database.CreateChannel("secret5", "text", "", "", 0)
chID, err := database.CreateChannel(context.Background(), "secret5", "text", "", "", 0)
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
if err := database.UpsertChannelOverride(chID, 3, 0, permissions.ReadMessages); err != nil {
if err := database.UpsertChannelOverride(context.Background(), chID, 3, 0, permissions.ReadMessages); err != nil {
t.Fatalf("UpsertChannelOverride: %v", err)
}
@@ -232,7 +233,7 @@ func TestDeleteChannelPermission_ClearsOverride(t *testing.T) {
t.Fatalf("status = %d, want 204; body: %s", w.Code, w.Body.String())
}
allow, deny, err := database.GetChannelPermissions(chID, 3)
allow, deny, err := database.GetChannelPermissions(context.Background(), chID, 3)
if err != nil {
t.Fatalf("GetChannelPermissions: %v", err)
}
+13 -12
View File
@@ -1,6 +1,7 @@
package admin
import (
"context"
"encoding/json"
"fmt"
"log/slog"
@@ -63,7 +64,7 @@ func validateCategoryType(channelType, category string) string {
func handleListChannels(database *db.DB) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
channels, err := database.ListChannels()
channels, err := database.ListChannels(r.Context())
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to list channels")
return
@@ -102,20 +103,20 @@ func handleCreateChannel(database *db.DB, hub HubBroadcaster) http.HandlerFunc {
return
}
id, err := database.AdminCreateChannel(req.Name, req.Type, req.Category, req.Topic, req.Position)
id, err := database.AdminCreateChannel(r.Context(), req.Name, req.Type, req.Category, req.Topic, req.Position)
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to create channel")
return
}
ch, err := database.GetChannel(id)
ch, err := database.GetChannel(r.Context(), id)
if err != nil || ch == nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch created channel")
return
}
actor := actorFromContext(r)
slog.Info("channel created", "actor_id", actor, "channel", req.Name, "type", req.Type)
db.WriteAudit(database, actor, "channel_create", "channel", id,
db.WriteAudit(context.WithoutCancel(r.Context()), database, actor, "channel_create", "channel", id,
fmt.Sprintf("created #%s (%s)", req.Name, req.Type))
if hub != nil {
hub.BroadcastChannelCreate(ch)
@@ -141,7 +142,7 @@ func handlePatchChannel(database *db.DB, hub HubBroadcaster) http.HandlerFunc {
return
}
existing, err := database.GetChannel(id)
existing, err := database.GetChannel(r.Context(), id)
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch channel")
return
@@ -164,17 +165,17 @@ func handlePatchChannel(database *db.DB, hub HubBroadcaster) http.HandlerFunc {
return
}
if err := database.AdminUpdateChannel(id, req.Name, req.Topic, req.SlowMode, req.Position, req.Archived); err != nil {
if err := database.AdminUpdateChannel(r.Context(), id, req.Name, req.Topic, req.SlowMode, req.Position, req.Archived); err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to update channel")
return
}
actor := actorFromContext(r)
slog.Info("channel updated", "actor_id", actor, "channel_id", id, "name", req.Name)
db.WriteAudit(database, actor, "channel_update", "channel", id,
db.WriteAudit(context.WithoutCancel(r.Context()), database, actor, "channel_update", "channel", id,
fmt.Sprintf("updated #%s", req.Name))
updated, err := database.GetChannel(id)
updated, err := database.GetChannel(r.Context(), id)
if err != nil || updated == nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch updated channel")
return
@@ -194,7 +195,7 @@ func handleDeleteChannel(database *db.DB, hub HubBroadcaster) http.HandlerFunc {
return
}
existing, err := database.GetChannel(id)
existing, err := database.GetChannel(r.Context(), id)
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch channel")
return
@@ -204,13 +205,13 @@ func handleDeleteChannel(database *db.DB, hub HubBroadcaster) http.HandlerFunc {
return
}
if err := database.AdminDeleteChannel(id); err != nil {
if err := database.AdminDeleteChannel(r.Context(), id); err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to delete channel")
return
}
actor := actorFromContext(r)
slog.Warn("channel deleted", "actor_id", actor, "channel_id", id, "name", existing.Name)
db.WriteAudit(database, actor, "channel_delete", "channel", id,
db.WriteAudit(context.WithoutCancel(r.Context()), database, actor, "channel_delete", "channel", id,
fmt.Sprintf("deleted #%s", existing.Name))
if hub != nil {
hub.BroadcastChannelDelete(id)
@@ -224,7 +225,7 @@ func handleGetAuditLog(database *db.DB) http.HandlerFunc {
limit := queryInt(r, "limit", 50, 1)
offset := queryInt(r, "offset", 0, 0)
entries, err := database.GetAuditLog(limit, offset)
entries, err := database.GetAuditLog(r.Context(), limit, offset)
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to get audit log")
return
+13 -12
View File
@@ -1,6 +1,7 @@
package admin
import (
"context"
"encoding/json"
"errors"
"fmt"
@@ -15,7 +16,7 @@ import (
func handleGetSettings(database *db.DB) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
settings, err := database.GetAllSettings()
settings, err := database.GetAllSettings(r.Context())
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to get settings")
return
@@ -48,7 +49,7 @@ func handlePatchSettings(database *db.DB) http.HandlerFunc {
return
}
if err := validateRequire2FAUpdate(database, normalizedUpdates); err != nil {
if err := validateRequire2FAUpdate(r.Context(), database, normalizedUpdates); err != nil {
writeErr(w, http.StatusBadRequest, "BAD_REQUEST", err.Error())
return
}
@@ -57,13 +58,13 @@ func handlePatchSettings(database *db.DB) http.HandlerFunc {
// Apply all settings atomically so a mid-loop failure doesn't leave
// partial updates.
tx, err := database.Begin()
tx, err := database.BeginTx(r.Context(), nil)
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to start transaction")
return
}
for key, value := range normalizedUpdates {
if _, txErr := tx.Exec(
if _, txErr := tx.ExecContext(r.Context(),
`INSERT INTO settings (key, value) VALUES (?, ?)
ON CONFLICT(key) DO UPDATE SET value = excluded.value`,
key, value,
@@ -79,11 +80,11 @@ func handlePatchSettings(database *db.DB) http.HandlerFunc {
}
for key := range normalizedUpdates {
slog.Info("setting changed", "actor_id", actor, "key", key)
db.WriteAudit(database, actor, "setting_change", "setting", 0,
db.WriteAudit(context.WithoutCancel(r.Context()), database, actor, "setting_change", "setting", 0,
fmt.Sprintf("%s updated", key))
}
settings, err := database.GetAllSettings()
settings, err := database.GetAllSettings(r.Context())
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch settings")
return
@@ -112,8 +113,8 @@ func normalizeSettingUpdates(updates map[string]string) (map[string]string, erro
return normalized, nil
}
func validateRequire2FAUpdate(database *db.DB, updates map[string]string) error {
targetRequire2FA, err := targetBoolSetting(database, updates, "require_2fa")
func validateRequire2FAUpdate(ctx context.Context, database *db.DB, updates map[string]string) error {
targetRequire2FA, err := targetBoolSetting(ctx, database, updates, "require_2fa")
if err != nil {
return err
}
@@ -121,7 +122,7 @@ func validateRequire2FAUpdate(database *db.DB, updates map[string]string) error
return nil
}
registrationOpen, err := targetBoolSetting(database, updates, "registration_open")
registrationOpen, err := targetBoolSetting(ctx, database, updates, "registration_open")
if err != nil {
return err
}
@@ -129,7 +130,7 @@ func validateRequire2FAUpdate(database *db.DB, updates map[string]string) error
return fmt.Errorf("require_2fa cannot be enabled while registration is open")
}
count, err := database.CountUsersWithoutTOTP()
count, err := database.CountUsersWithoutTOTP(ctx)
if err != nil {
return fmt.Errorf("failed to validate 2FA enrollment")
}
@@ -139,11 +140,11 @@ func validateRequire2FAUpdate(database *db.DB, updates map[string]string) error
return nil
}
func targetBoolSetting(database *db.DB, updates map[string]string, key string) (bool, error) {
func targetBoolSetting(ctx context.Context, database *db.DB, updates map[string]string, key string) (bool, error) {
if value, ok := updates[key]; ok {
return parseBooleanSettingValue(value)
}
value, err := database.GetSetting(key)
value, err := database.GetSetting(ctx, key)
if errors.Is(err, db.ErrNotFound) {
return false, nil
}
+11 -10
View File
@@ -1,6 +1,7 @@
package admin
import (
"context"
"encoding/json"
"errors"
"fmt"
@@ -15,7 +16,7 @@ import (
func handleGetStats(database *db.DB, hub HubBroadcaster) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
stats, err := database.GetServerStats()
stats, err := database.GetServerStats(r.Context())
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to get stats")
return
@@ -32,7 +33,7 @@ func handleListUsers(database *db.DB) http.HandlerFunc {
limit := queryInt(r, "limit", 50, 1)
offset := queryInt(r, "offset", 0, 0)
users, err := database.ListAllUsers(limit, offset)
users, err := database.ListAllUsers(r.Context(), limit, offset)
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to list users")
return
@@ -81,7 +82,7 @@ func handlePatchUser(database *db.DB, hub HubBroadcaster, permInvalidator Permis
return
}
user, err := database.GetUserByID(id)
user, err := database.GetUserByID(r.Context(), id)
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch user")
return
@@ -131,7 +132,7 @@ func handlePatchUser(database *db.DB, hub HubBroadcaster, permInvalidator Permis
}
if req.RoleID != nil {
if _, err := database.Exec(`UPDATE users SET role_id = ? WHERE id = ?`, *req.RoleID, id); err != nil {
if _, err := database.ExecContext(r.Context(), `UPDATE users SET role_id = ? WHERE id = ?`, *req.RoleID, id); err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to update role")
return
}
@@ -139,21 +140,21 @@ func handlePatchUser(database *db.DB, hub HubBroadcaster, permInvalidator Permis
if permInvalidator != nil {
permInvalidator.InvalidateUser(id)
}
db.WriteAudit(database, actor, "role_change", "user", id,
db.WriteAudit(context.WithoutCancel(r.Context()), database, actor, "role_change", "user", id,
fmt.Sprintf("changed %s role to %d", user.Username, *req.RoleID))
if role, err := database.GetRoleByID(*req.RoleID); err == nil && role != nil {
if role, err := database.GetRoleByID(r.Context(), *req.RoleID); err == nil && role != nil {
if hub != nil {
hub.BroadcastMemberUpdate(id, role.Name)
}
}
}
updated, err := database.GetUserByID(id)
updated, err := database.GetUserByID(r.Context(), id)
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch updated user")
return
}
writeJSON(w, http.StatusOK, toAdminUserResponseFromUser(database, updated))
writeJSON(w, http.StatusOK, toAdminUserResponseFromUser(r.Context(), database, updated))
}
}
@@ -165,13 +166,13 @@ func handleForceLogout(database *db.DB) http.HandlerFunc {
return
}
if err := database.ForceLogoutUser(id); err != nil {
if err := database.ForceLogoutUser(r.Context(), id); err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to logout user")
return
}
actor := actorFromContext(r)
slog.Info("force logout", "actor_id", actor, "target_user_id", id)
db.WriteAudit(database, actor, "force_logout", "user", id, "all sessions terminated")
db.WriteAudit(context.WithoutCancel(r.Context()), database, actor, "force_logout", "user", id, "all sessions terminated")
w.WriteHeader(http.StatusNoContent)
}
}
+7 -5
View File
@@ -341,7 +341,10 @@ func handleLogStream(database *db.DB, ringBuf *RingBuffer) http.HandlerFunc {
http.Error(w, string(errResp), http.StatusUnauthorized)
return
}
sess, err := database.GetSessionByTokenHash(entry.tokenHash)
// Stream lifetime == request lifetime, so all session re-checks below
// use the stream request's context.
ctx := r.Context()
sess, err := database.GetSessionByTokenHash(ctx, entry.tokenHash)
if err != nil || sess == nil || auth.IsSessionExpired(sess.ExpiresAt) {
errResp, _ := json.Marshal(map[string]string{
"error": "UNAUTHORIZED",
@@ -351,15 +354,15 @@ func handleLogStream(database *db.DB, ringBuf *RingBuffer) http.HandlerFunc {
return
}
sessionStillAuthorized := func() bool {
current, currentErr := database.GetSessionByTokenHash(entry.tokenHash)
current, currentErr := database.GetSessionByTokenHash(ctx, entry.tokenHash)
if currentErr != nil || current == nil || auth.IsSessionExpired(current.ExpiresAt) {
return false
}
user, userErr := database.GetUserByID(current.UserID)
user, userErr := database.GetUserByID(ctx, current.UserID)
if userErr != nil || user == nil {
return false
}
role, roleErr := database.GetRoleByID(user.RoleID)
role, roleErr := database.GetRoleByID(ctx, user.RoleID)
if roleErr != nil || role == nil {
return false
}
@@ -408,7 +411,6 @@ func handleLogStream(database *db.DB, ringBuf *RingBuffer) http.HandlerFunc {
keepalive := time.NewTicker(15 * time.Second)
defer keepalive.Stop()
ctx := r.Context()
for {
select {
case entry := <-ch:
+3 -3
View File
@@ -73,7 +73,7 @@ func TestHandleLogStream_BackfillStopsAfterSessionRevocation(t *testing.T) {
logBuf.Write(LogEntry{Timestamp: "2026-03-29T10:00:00Z", Level: "info", Message: "first", Source: "test"})
logBuf.Write(LogEntry{Timestamp: "2026-03-29T10:00:01Z", Level: "info", Message: "second", Source: "test"})
userID, err := database.CreateUser("owner", "hash", 1)
userID, err := database.CreateUser(context.Background(), "owner", "hash", 1)
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
@@ -83,7 +83,7 @@ func TestHandleLogStream_BackfillStopsAfterSessionRevocation(t *testing.T) {
t.Fatalf("GenerateToken: %v", err)
}
tokenHash := auth.HashToken(token)
if _, err := database.CreateSession(userID, tokenHash, "test", "127.0.0.1"); err != nil {
if _, err := database.CreateSession(context.Background(), userID, tokenHash, "test", "127.0.0.1"); err != nil {
t.Fatalf("CreateSession: %v", err)
}
@@ -99,7 +99,7 @@ func TestHandleLogStream_BackfillStopsAfterSessionRevocation(t *testing.T) {
writer := &revokingSSEWriter{
header: make(http.Header),
revoke: func() {
_ = database.DeleteSession(tokenHash)
_ = database.DeleteSession(context.Background(), tokenHash)
},
cancel: cancel,
}
+4 -4
View File
@@ -31,7 +31,7 @@ func adminAuthMiddleware(database *db.DB) func(http.Handler) http.Handler {
}
hash := auth.HashToken(token)
sess, err := database.GetSessionByTokenHash(hash)
sess, err := database.GetSessionByTokenHash(r.Context(), hash)
if err != nil || sess == nil {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "invalid or expired session")
return
@@ -42,13 +42,13 @@ func adminAuthMiddleware(database *db.DB) func(http.Handler) http.Handler {
return
}
user, err := database.GetUserByID(sess.UserID)
user, err := database.GetUserByID(r.Context(), sess.UserID)
if err != nil || user == nil {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "user not found")
return
}
role, err := database.GetRoleByID(user.RoleID)
role, err := database.GetRoleByID(r.Context(), user.RoleID)
if err != nil || role == nil {
writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "role not found")
return
@@ -77,7 +77,7 @@ func ownerOnlyMiddleware(database *db.DB, next http.Handler) http.Handler {
return
}
role, err := database.GetRoleByID(user.RoleID)
role, err := database.GetRoleByID(r.Context(), user.RoleID)
if err != nil || role == nil {
writeErr(w, http.StatusForbidden, "FORBIDDEN", "role not found")
return
+12 -12
View File
@@ -157,23 +157,23 @@ func TestOwnerOnlyMiddleware_RoleNotFound(t *testing.T) {
// Create a user initially with a valid role, then mutate role_id to a
// nonexistent value (disabling FK checks temporarily so SQLite allows it).
uid, err := database.CreateUser("orphanuser", "$2a$12$x", 1)
uid, err := database.CreateUser(context.Background(), "orphanuser", "$2a$12$x", 1)
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
user, err := database.GetUserByID(uid)
user, err := database.GetUserByID(context.Background(), uid)
if err != nil || user == nil {
t.Fatalf("GetUserByID: %v", err)
}
// Disable FK enforcement, update role_id, re-enable.
if _, err := database.Exec(`PRAGMA foreign_keys=OFF`); err != nil {
if _, err := database.ExecContext(context.Background(), `PRAGMA foreign_keys=OFF`); err != nil {
t.Fatalf("disable FK: %v", err)
}
if _, err := database.Exec(`UPDATE users SET role_id = 9999 WHERE id = ?`, uid); err != nil {
if _, err := database.ExecContext(context.Background(), `UPDATE users SET role_id = 9999 WHERE id = ?`, uid); err != nil {
t.Fatalf("UPDATE role_id: %v", err)
}
if _, err := database.Exec(`PRAGMA foreign_keys=ON`); err != nil {
if _, err := database.ExecContext(context.Background(), `PRAGMA foreign_keys=ON`); err != nil {
t.Fatalf("re-enable FK: %v", err)
}
user.RoleID = 9999 // mirror the DB value in our in-memory struct
@@ -213,11 +213,11 @@ func TestOwnerOnlyMiddleware_RoleNotFound(t *testing.T) {
func TestOwnerOnlyMiddleware_OwnerPassesThrough(t *testing.T) {
database := openWhiteboxTestDB(t)
uid, err := database.CreateUser("ownerpass", "$2a$12$x", 1)
uid, err := database.CreateUser(context.Background(), "ownerpass", "$2a$12$x", 1)
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
user, err := database.GetUserByID(uid)
user, err := database.GetUserByID(context.Background(), uid)
if err != nil || user == nil {
t.Fatalf("GetUserByID: %v", err)
}
@@ -251,23 +251,23 @@ func TestAdminAuthMiddleware_RoleNotFound(t *testing.T) {
database := openWhiteboxTestDB(t)
handler := NewAdminAPI(database, "1.0.0", nil, nil, nil, nil, nil, nil)
uid, err := database.CreateUser("noroleuser", "$2a$12$x", 1)
uid, err := database.CreateUser(context.Background(), "noroleuser", "$2a$12$x", 1)
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
token := "norole-token"
if _, err := database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1"); err != nil {
if _, err := database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1"); err != nil {
t.Fatalf("CreateSession: %v", err)
}
// Disable FK enforcement, assign a non-existent role_id, re-enable.
if _, err := database.Exec(`PRAGMA foreign_keys=OFF`); err != nil {
if _, err := database.ExecContext(context.Background(), `PRAGMA foreign_keys=OFF`); err != nil {
t.Fatalf("disable FK: %v", err)
}
if _, err := database.Exec(`UPDATE users SET role_id = 9999 WHERE id = ?`, uid); err != nil {
if _, err := database.ExecContext(context.Background(), `UPDATE users SET role_id = 9999 WHERE id = ?`, uid); err != nil {
t.Fatalf("UPDATE role_id: %v", err)
}
if _, err := database.Exec(`PRAGMA foreign_keys=ON`); err != nil {
if _, err := database.ExecContext(context.Background(), `PRAGMA foreign_keys=ON`); err != nil {
t.Fatalf("re-enable FK: %v", err)
}
+4 -3
View File
@@ -4,6 +4,7 @@ package admin_test
// ownerOnlyMiddleware, and related helpers.
import (
"context"
"net/http"
"testing"
"time"
@@ -22,19 +23,19 @@ func TestAdminAuthMiddleware_ExpiredSession(t *testing.T) {
// Create a user and session, then manually expire the session by setting
// expires_at to a past timestamp via the exported Exec helper.
uid, err := database.CreateUser("expireduser", "$2a$12$x", 1)
uid, err := database.CreateUser(context.Background(), "expireduser", "$2a$12$x", 1)
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
token := "expired-session-token"
tokenHash := auth.HashToken(token)
if _, err := database.CreateSession(uid, tokenHash, "test", "127.0.0.1"); err != nil {
if _, err := database.CreateSession(context.Background(), uid, tokenHash, "test", "127.0.0.1"); err != nil {
t.Fatalf("CreateSession: %v", err)
}
// Set expires_at to yesterday so the session is treated as expired.
pastTime := time.Now().Add(-24 * time.Hour).UTC().Format("2006-01-02T15:04:05Z")
if _, err := database.Exec(
if _, err := database.ExecContext(context.Background(),
`UPDATE sessions SET expires_at = ? WHERE token = ?`,
pastTime, tokenHash,
); err != nil {
+8 -7
View File
@@ -1,6 +1,7 @@
package admin
import (
"context"
"encoding/json"
"errors"
"log/slog"
@@ -42,7 +43,7 @@ type setupResponse struct {
// handleSetupStatus returns whether initial setup is needed (no users exist).
func handleSetupStatus(database *db.DB) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
count, err := database.UserCount()
count, err := database.UserCount(r.Context())
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to check user count")
return
@@ -110,7 +111,7 @@ func handleSetup(database *db.DB, limiter *auth.RateLimiter, allowedOrigins []st
// Atomically check no users exist and create the owner (BUG-119).
// This closes the TOCTOU race between UserCount() and CreateUser().
uid, err := database.CreateOwnerIfEmpty(req.Username, hash, ownerRoleID)
uid, err := database.CreateOwnerIfEmpty(r.Context(), req.Username, hash, ownerRoleID)
if errors.Is(err, db.ErrConflict) {
writeErr(w, http.StatusForbidden, "FORBIDDEN", "setup has already been completed")
return
@@ -132,27 +133,27 @@ func handleSetup(database *db.DB, limiter *auth.RateLimiter, allowedOrigins []st
if len(device) > maxDeviceLen {
device = device[:maxDeviceLen]
}
if _, err := database.CreateSession(uid, auth.HashToken(token), device, host); err != nil {
if _, err := database.CreateSession(r.Context(), uid, auth.HashToken(token), device, host); err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to create session")
return
}
// Create default channels under canonical categories.
_, _ = database.CreateChannel("general", "text", "Text Channels", "Welcome to the server!", 0)
_, _ = database.CreateChannel("General", "voice", "Voice Channels", "", 0)
_, _ = database.CreateChannel(r.Context(), "general", "text", "Text Channels", "Welcome to the server!", 0)
_, _ = database.CreateChannel(r.Context(), "General", "voice", "Voice Channels", "", 0)
// Generate a bootstrap invite code so the owner can invite others.
// Bound it (5 uses / 24h) rather than minting an unlimited, non-expiring
// invite — the owner can create fresh invites once logged in.
bootstrapInviteExpiry := time.Now().Add(24 * time.Hour)
inviteCode, err := database.CreateInvite(uid, 5, &bootstrapInviteExpiry)
inviteCode, err := database.CreateInvite(r.Context(), uid, 5, &bootstrapInviteExpiry)
if err != nil {
writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to generate invite code")
return
}
slog.Info("server setup completed", "owner", req.Username, "user_id", uid)
db.WriteAudit(database, uid, "server_setup", "server", 0,
db.WriteAudit(context.WithoutCancel(r.Context()), database, uid, "server_setup", "server", 0,
"initial setup: owner account created, default channel and invite generated")
writeJSON(w, http.StatusCreated, setupResponse{
+3 -2
View File
@@ -1,6 +1,7 @@
package admin_test
import (
"context"
"encoding/json"
"fmt"
"net/http"
@@ -86,7 +87,7 @@ func TestSetup_CreatesOwner(t *testing.T) {
}
// Verify user was created with Owner role.
user, err := database.GetUserByUsername("myadmin")
user, err := database.GetUserByUsername(context.Background(), "myadmin")
if err != nil || user == nil {
t.Fatal("user not found in database after setup")
}
@@ -185,7 +186,7 @@ func TestSetup_ConcurrentRace(t *testing.T) {
}
// Verify only one user exists in the database.
count, err := database.UserCount()
count, err := database.UserCount(context.Background())
if err != nil {
t.Fatalf("UserCount: %v", err)
}
File diff suppressed because it is too large Load Diff
+7 -3
View File
@@ -1,6 +1,10 @@
package admin
import "github.com/owncord/server/db"
import (
"context"
"github.com/owncord/server/db"
)
// ─── Context keys ─────────────────────────────────────────────────────────────
@@ -93,9 +97,9 @@ func toAdminUserResponse(u db.UserWithRole) adminUserResponse {
// toAdminUserResponseFromUser converts a plain db.User to the safe response
// shape, resolving the role name via the database.
func toAdminUserResponseFromUser(database *db.DB, u *db.User) adminUserResponse {
func toAdminUserResponseFromUser(ctx context.Context, database *db.DB, u *db.User) adminUserResponse {
roleName := ""
if role, err := database.GetRoleByID(u.RoleID); err == nil && role != nil {
if role, err := database.GetRoleByID(ctx, u.RoleID); err == nil && role != nil {
roleName = role.Name
}
return adminUserResponse{
+3 -2
View File
@@ -1,6 +1,7 @@
package admin_test
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
@@ -136,9 +137,9 @@ func TestAdminAPI_ApplyUpdate_RequiresOwner(t *testing.T) {
handler := admin.NewAdminAPI(database, "1.0.0", nil, nil, nil, nil, nil, newTestModService(database))
// Create admin user (not owner - role 2)
adminUID, _ := database.CreateUser("adminonly2", "hash", 2)
adminUID, _ := database.CreateUser(context.Background(), "adminonly2", "hash", 2)
token := "admin-role-token"
_, _ = database.CreateSession(adminUID, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), adminUID, auth.HashToken(token), "test", "127.0.0.1")
w := doRequest(t, handler, http.MethodPost, "/updates/apply", token, nil)
if w.Code != http.StatusForbidden {
+37 -30
View File
@@ -1,6 +1,7 @@
package api
import (
"context"
"encoding/json"
"errors"
"fmt"
@@ -107,7 +108,7 @@ func MountAuthRoutes(r chi.Router, database *db.DB, limiter *auth.RateLimiter, t
// handleRegister processes POST /api/v1/auth/register.
func handleRegister(database *db.DB) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
registrationOpen, err := isRegistrationOpen(database)
registrationOpen, err := isRegistrationOpen(r.Context(), database)
if err != nil {
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL_ERROR",
@@ -123,7 +124,7 @@ func handleRegister(database *db.DB) http.HandlerFunc {
return
}
require2FA, err := isRequire2FAEnabled(database)
require2FA, err := isRequire2FAEnabled(r.Context(), database)
if err != nil {
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL_ERROR",
@@ -190,7 +191,7 @@ func handleRegister(database *db.DB) http.HandlerFunc {
// Atomically consume the invite and create the user so failed
// registrations do not burn a valid invite code.
uid, err := database.CreateUserWithInvite(req.Username, hash, int(permissions.MemberRoleID), req.InviteCode)
uid, err := database.CreateUserWithInvite(r.Context(), req.Username, hash, int(permissions.MemberRoleID), req.InviteCode)
if err != nil {
// UNIQUE constraint violation → duplicate username → 400.
// Any other DB error → 500.
@@ -211,7 +212,7 @@ func handleRegister(database *db.DB) http.HandlerFunc {
ip := clientIP(r)
slog.Info("user registered", "username", req.Username, "user_id", uid, "ip", ip)
db.WriteAudit(database, uid, "user_register", "user", uid,
db.WriteAudit(context.WithoutCancel(r.Context()), database, uid, "user_register", "user", uid,
"new account created via invite")
// Issue session.
@@ -225,7 +226,7 @@ func handleRegister(database *db.DB) http.HandlerFunc {
}
device := truncateDevice(r.Header.Get("User-Agent"))
if _, err := database.CreateSession(uid, auth.HashToken(token), device, ip); err != nil {
if _, err := database.CreateSession(r.Context(), uid, auth.HashToken(token), device, ip); err != nil {
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL_ERROR",
Message: "failed to create session",
@@ -233,7 +234,7 @@ func handleRegister(database *db.DB) http.HandlerFunc {
return
}
user, err := database.GetUserByID(uid)
user, err := database.GetUserByID(r.Context(), uid)
if err != nil || user == nil {
slog.Error("failed to fetch user after registration", "user_id", uid, "error", err)
writeJSON(w, http.StatusInternalServerError, errorResponse{
@@ -287,7 +288,11 @@ func handleLogin(database *db.DB, limiter *auth.RateLimiter, partialStore *auth.
}
// BUG-110: Also check per-username lockout to prevent distributed brute force.
userLockKey := "login_user_lock:" + req.Username
// F1: canonicalize the username the same way GetUserByUsername does (COLLATE
// NOCASE) before keying the lockout, so case variants of one account
// (admin/Admin/ADMIN) share a single bucket instead of each getting its own.
unameKey := strings.ToLower(req.Username)
userLockKey := "login_user_lock:" + unameKey
if limiter.IsLockedOut(userLockKey) {
writeJSON(w, http.StatusTooManyRequests, errorResponse{
Error: "RATE_LIMITED",
@@ -298,7 +303,7 @@ func handleLogin(database *db.DB, limiter *auth.RateLimiter, partialStore *auth.
// Constant-time lookup: always attempt bcrypt compare even when user
// does not exist to prevent timing-based username enumeration.
user, err := database.GetUserByUsername(req.Username)
user, err := database.GetUserByUsername(r.Context(), req.Username)
// Distinguish DB errors from authentication failures. DB errors
// should NOT increment the rate limiter — otherwise a transient
@@ -316,7 +321,7 @@ func handleLogin(database *db.DB, limiter *auth.RateLimiter, partialStore *auth.
}
failKey := "login_fail:" + ip
userFailKey := "login_user_fail:" + req.Username
userFailKey := "login_user_fail:" + unameKey
// Always run the password check — with an empty hash when the user does
// not exist. auth.CheckPassword performs a dummy bcrypt comparison for an
// empty hash, so bcrypt executes on every path and response time stays
@@ -330,11 +335,11 @@ func handleLogin(database *db.DB, limiter *auth.RateLimiter, partialStore *auth.
if !auth.CheckPassword(storedHash, req.Password) {
// Track failures per-IP; lockout on threshold.
if !limiter.Allow(failKey, loginFailureThreshold, loginFailureWindow) {
limiter.Lockout(lockKey, loginLockoutDuration)
limiter.Lockout(r.Context(), lockKey, loginLockoutDuration)
}
// BUG-110: Track failures per-username; lockout on threshold.
if !limiter.Allow(userFailKey, loginUserFailureThreshold, loginUserFailureWindow) {
limiter.Lockout(userLockKey, loginUserLockoutDuration)
limiter.Lockout(r.Context(), userLockKey, loginUserLockoutDuration)
}
slog.Info("login failed", "ip", ip, "username_len", len(req.Username))
writeJSON(w, http.StatusUnauthorized, errorResponse{
@@ -345,12 +350,12 @@ func handleLogin(database *db.DB, limiter *auth.RateLimiter, partialStore *auth.
}
// Reset failure counters on success.
limiter.Reset(failKey)
limiter.Reset(userFailKey)
limiter.Reset(r.Context(), failKey)
limiter.Reset(r.Context(), userFailKey)
if auth.IsEffectivelyBanned(user) {
slog.Warn("banned user login attempt", "username", user.Username, "user_id", user.ID, "ip", ip)
db.WriteAudit(database, user.ID, "login_blocked_banned", "user", user.ID,
db.WriteAudit(context.WithoutCancel(r.Context()), database, user.ID, "login_blocked_banned", "user", user.ID,
"banned user attempted login from "+ip)
writeJSON(w, http.StatusForbidden, errorResponse{
Error: "FORBIDDEN",
@@ -359,7 +364,7 @@ func handleLogin(database *db.DB, limiter *auth.RateLimiter, partialStore *auth.
return
}
require2FA, err := isRequire2FAEnabled(database)
require2FA, err := isRequire2FAEnabled(r.Context(), database)
if err != nil {
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL_ERROR",
@@ -391,7 +396,7 @@ func handleLogin(database *db.DB, limiter *auth.RateLimiter, partialStore *auth.
}
// Issue session.
token, err := issueSession(database, user.ID, truncateDevice(r.Header.Get("User-Agent")), ip)
token, err := issueSession(r.Context(), database, user.ID, truncateDevice(r.Header.Get("User-Agent")), ip)
if err != nil {
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL_ERROR",
@@ -405,7 +410,7 @@ func handleLogin(database *db.DB, limiter *auth.RateLimiter, partialStore *auth.
// would leave the user permanently "online" if they never open a WS
// connection or if the client crashes before connecting.
slog.Info("user logged in", "username", user.Username, "user_id", user.ID, "ip", ip)
db.WriteAudit(database, user.ID, "user_login", "user", user.ID,
db.WriteAudit(context.WithoutCancel(r.Context()), database, user.ID, "user_login", "user", user.ID,
"logged in from "+ip)
writeJSON(w, http.StatusOK, authSuccessResponse{
Token: token,
@@ -427,7 +432,9 @@ func handleLogout(database *db.DB) http.HandlerFunc {
return
}
if err := database.DeleteSession(sess.TokenHash); err != nil {
// The client clears its token optimistically — once logout reaches the
// server, the revocation must not die with a dropped connection.
if err := database.DeleteSession(context.WithoutCancel(r.Context()), sess.TokenHash); err != nil {
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL_ERROR",
Message: "failed to logout",
@@ -436,7 +443,7 @@ func handleLogout(database *db.DB) http.HandlerFunc {
}
slog.Info("user logged out", "user_id", sess.UserID)
db.WriteAudit(database, sess.UserID, "user_logout", "user", sess.UserID, "")
db.WriteAudit(context.WithoutCancel(r.Context()), database, sess.UserID, "user_logout", "user", sess.UserID, "")
w.WriteHeader(http.StatusNoContent)
}
@@ -507,7 +514,7 @@ func handleDeleteAccount(database *db.DB, limiter *auth.RateLimiter) http.Handle
failKey := fmt.Sprintf("delete_fail:%d", user.ID)
if !auth.CheckPassword(user.PasswordHash, req.Password) {
if !limiter.Allow(failKey, deleteAccountFailureThreshold, deleteAccountFailureWindow) {
limiter.Lockout(lockKey, deleteAccountLockoutDuration)
limiter.Lockout(r.Context(), lockKey, deleteAccountLockoutDuration)
}
writeJSON(w, http.StatusBadRequest, errorResponse{
Error: "INVALID_INPUT",
@@ -515,7 +522,7 @@ func handleDeleteAccount(database *db.DB, limiter *auth.RateLimiter) http.Handle
})
return
}
limiter.Reset(failKey)
limiter.Reset(r.Context(), failKey)
if err := database.DeleteAccount(r.Context(), user.ID); err != nil {
if errors.Is(err, db.ErrLastAdmin) {
@@ -535,7 +542,7 @@ func handleDeleteAccount(database *db.DB, limiter *auth.RateLimiter) http.Handle
ip := clientIP(r)
slog.Info("account deleted", "username", user.Username, "user_id", user.ID, "ip", ip)
db.WriteAudit(database, user.ID, "account_deleted", "user", user.ID,
db.WriteAudit(context.WithoutCancel(r.Context()), database, user.ID, "account_deleted", "user", user.ID,
"account self-deleted from "+ip)
w.WriteHeader(http.StatusNoContent)
@@ -570,27 +577,27 @@ func truncateDevice(ua string) string {
return ua
}
func issueSession(database *db.DB, userID int64, device, ip string) (string, error) {
func issueSession(ctx context.Context, database *db.DB, userID int64, device, ip string) (string, error) {
token, err := auth.GenerateToken()
if err != nil {
return "", err
}
if _, err := database.CreateSession(userID, auth.HashToken(token), device, ip); err != nil {
if _, err := database.CreateSession(ctx, userID, auth.HashToken(token), device, ip); err != nil {
return "", err
}
return token, nil
}
func isRequire2FAEnabled(database *db.DB) (bool, error) {
return getBooleanSetting(database, "require_2fa", false)
func isRequire2FAEnabled(ctx context.Context, database *db.DB) (bool, error) {
return getBooleanSetting(ctx, database, "require_2fa", false)
}
func isRegistrationOpen(database *db.DB) (bool, error) {
return getBooleanSetting(database, "registration_open", true)
func isRegistrationOpen(ctx context.Context, database *db.DB) (bool, error) {
return getBooleanSetting(ctx, database, "registration_open", true)
}
func getBooleanSetting(database *db.DB, key string, defaultValue bool) (bool, error) {
value, err := database.GetSetting(key)
func getBooleanSetting(ctx context.Context, database *db.DB, key string, defaultValue bool) (bool, error) {
value, err := database.GetSetting(ctx, key)
if err != nil {
if errors.Is(err, db.ErrNotFound) {
return defaultValue, nil
+122 -87
View File
@@ -2,6 +2,7 @@ package api_test
import (
"bytes"
"context"
"encoding/json"
"fmt"
"net/http"
@@ -104,8 +105,8 @@ func TestRegister_Success(t *testing.T) {
router := buildAuthRouter(database, limiter)
// Create an invite first.
ownerID, _ := database.CreateUser("owner", "hash", 1)
code, _ := database.CreateInvite(ownerID, 1, nil)
ownerID, _ := database.CreateUser(context.Background(), "owner", "hash", 1)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil)
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "newuser",
@@ -132,12 +133,12 @@ func TestRegister_RegistrationClosed(t *testing.T) {
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
if _, err := database.Exec(`UPDATE settings SET value = '0' WHERE key = 'registration_open'`); err != nil {
if _, err := database.ExecContext(context.Background(), `UPDATE settings SET value = '0' WHERE key = 'registration_open'`); err != nil {
t.Fatalf("close registration: %v", err)
}
ownerID, _ := database.CreateUser("owner", "hash", 1)
code, _ := database.CreateInvite(ownerID, 1, nil)
ownerID, _ := database.CreateUser(context.Background(), "owner", "hash", 1)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil)
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "closeduser",
@@ -171,8 +172,8 @@ func TestRegister_WeakPassword(t *testing.T) {
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
ownerID, _ := database.CreateUser("owner2", "hash", 1)
code, _ := database.CreateInvite(ownerID, 1, nil)
ownerID, _ := database.CreateUser(context.Background(), "owner2", "hash", 1)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil)
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "newuser",
@@ -190,8 +191,8 @@ func TestRegister_InviteUsedUp(t *testing.T) {
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
ownerID, _ := database.CreateUser("owner3", "hash", 1)
code, _ := database.CreateInvite(ownerID, 1, nil) // max 1 use
ownerID, _ := database.CreateUser(context.Background(), "owner3", "hash", 1)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil) // max 1 use
// First registration should succeed.
postJSON(t, router, "/api/v1/auth/register", map[string]string{
@@ -217,9 +218,9 @@ func TestRegister_DuplicateUsername_DoesNotConsumeInvite(t *testing.T) {
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
ownerID, _ := database.CreateUser("owner4", "hash", 1)
_, _ = database.CreateUser("takenuser", "hash", 4)
code, _ := database.CreateInvite(ownerID, 1, nil)
ownerID, _ := database.CreateUser(context.Background(), "owner4", "hash", 1)
_, _ = database.CreateUser(context.Background(), "takenuser", "hash", 4)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil)
duplicate := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "takenuser",
@@ -277,7 +278,7 @@ func TestLogin_Success(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser("loginuser", hash, 4)
_, _ = database.CreateUser(context.Background(), "loginuser", hash, 4)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "loginuser",
@@ -301,7 +302,7 @@ func TestLogin_WrongPassword(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser("loginuser2", hash, 4)
_, _ = database.CreateUser(context.Background(), "loginuser2", hash, 4)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "loginuser2",
@@ -361,7 +362,7 @@ func TestLogin_UsernameLockoutAcrossDifferentIPs(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser("lockoutuser", hash, 4)
_, _ = database.CreateUser(context.Background(), "lockoutuser", hash, 4)
for i := 0; i < 10; i++ {
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
@@ -388,7 +389,7 @@ func TestLogin_UsernameLockoutBlocksCorrectPasswordFromFreshIP(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser("lockoutcorrect", hash, 4)
_, _ = database.CreateUser(context.Background(), "lockoutcorrect", hash, 4)
for i := 0; i < 10; i++ {
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
@@ -409,13 +410,47 @@ func TestLogin_UsernameLockoutBlocksCorrectPasswordFromFreshIP(t *testing.T) {
}
}
// TestLogin_UsernameLockoutIgnoresCasing locks F1: the per-username lockout key
// must be case-folded so it matches the DB's COLLATE NOCASE username lookup.
// Otherwise an attacker splits the 9-attempt lockout budget across case variants
// of one account (admin, Admin, ADMIN, …), all of which authenticate the same row.
func TestLogin_UsernameLockoutIgnoresCasing(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser(context.Background(), "casehunt", hash, 4)
// Trip the per-username lockout using the lowercase spelling, from many IPs
// so the per-IP limiter is never the binding cap.
for i := 0; i < 10; i++ {
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "casehunt",
"password": "wrongpassword",
}, fmt.Sprintf("198.51.100.%d", i+1))
if rr.Code != http.StatusUnauthorized {
t.Fatalf("setup attempt %d status = %d, want 401; body = %s", i+1, rr.Code, rr.Body.String())
}
}
// A different casing of the SAME account must land in the same lockout bucket.
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
"username": "CASEHUNT",
"password": "wrongpassword",
}, "198.51.100.250")
if rr.Code != http.StatusTooManyRequests {
t.Fatalf("case-variant username bypassed the per-username lockout: status = %d, want 429; body = %s", rr.Code, rr.Body.String())
}
}
func TestLogin_SuccessResetsUsernameFailureCounter(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser("resetuser", hash, 4)
_, _ = database.CreateUser(context.Background(), "resetuser", hash, 4)
for i := 0; i < 8; i++ {
rr := postJSONFromIP(t, router, "/api/v1/auth/login", map[string]string{
@@ -469,8 +504,8 @@ func TestLogin_RequiresTOTPChallenge(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
userID, _ := database.CreateUser("totpuser", hash, 4)
if _, err := database.Exec(`UPDATE users SET totp_secret = ? WHERE id = ?`, "JBSWY3DPEHPK3PXP", userID); err != nil {
userID, _ := database.CreateUser(context.Background(), "totpuser", hash, 4)
if _, err := database.ExecContext(context.Background(), `UPDATE users SET totp_secret = ? WHERE id = ?`, "JBSWY3DPEHPK3PXP", userID); err != nil {
t.Fatalf("set totp secret: %v", err)
}
@@ -504,8 +539,8 @@ func TestLogin_UsernameLockoutBlocksTOTPChallenge(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
userID, _ := database.CreateUser("totplocked", hash, 4)
if _, err := database.Exec(`UPDATE users SET totp_secret = ? WHERE id = ?`, "JBSWY3DPEHPK3PXP", userID); err != nil {
userID, _ := database.CreateUser(context.Background(), "totplocked", hash, 4)
if _, err := database.ExecContext(context.Background(), `UPDATE users SET totp_secret = ? WHERE id = ?`, "JBSWY3DPEHPK3PXP", userID); err != nil {
t.Fatalf("set totp secret: %v", err)
}
@@ -542,9 +577,9 @@ func TestVerifyTotp_Success(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
userID, _ := database.CreateUser("totpverify", hash, 4)
userID, _ := database.CreateUser(context.Background(), "totpverify", hash, 4)
secret := "JBSWY3DPEHPK3PXP"
if _, err := database.Exec(`UPDATE users SET totp_secret = ? WHERE id = ?`, secret, userID); err != nil {
if _, err := database.ExecContext(context.Background(), `UPDATE users SET totp_secret = ? WHERE id = ?`, secret, userID); err != nil {
t.Fatalf("set totp secret: %v", err)
}
@@ -592,9 +627,9 @@ func TestEnableConfirmDisableTotp(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
userID, _ := database.CreateUser("enrolltotp", hash, 4)
userID, _ := database.CreateUser(context.Background(), "enrolltotp", hash, 4)
token, _ := auth.GenerateToken()
if _, err := database.CreateSession(userID, auth.HashToken(token), "test", "127.0.0.1"); err != nil {
if _, err := database.CreateSession(context.Background(), userID, auth.HashToken(token), "test", "127.0.0.1"); err != nil {
t.Fatalf("CreateSession: %v", err)
}
@@ -612,7 +647,7 @@ func TestEnableConfirmDisableTotp(t *testing.T) {
t.Fatal("expected qr_uri from enable response")
}
userBeforeConfirm, err := database.GetUserByID(userID)
userBeforeConfirm, err := database.GetUserByID(context.Background(), userID)
if err != nil {
t.Fatalf("GetUserByID before confirm: %v", err)
}
@@ -638,7 +673,7 @@ func TestEnableConfirmDisableTotp(t *testing.T) {
t.Fatalf("confirm status = %d, want 204; body = %s", confirm.Code, confirm.Body.String())
}
userAfterConfirm, err := database.GetUserByID(userID)
userAfterConfirm, err := database.GetUserByID(context.Background(), userID)
if err != nil {
t.Fatalf("GetUserByID after confirm: %v", err)
}
@@ -660,7 +695,7 @@ func TestEnableConfirmDisableTotp(t *testing.T) {
t.Fatalf("disable status = %d, want 204; body = %s", deleteRec.Code, deleteRec.Body.String())
}
userAfterDelete, err := database.GetUserByID(userID)
userAfterDelete, err := database.GetUserByID(context.Background(), userID)
if err != nil {
t.Fatalf("GetUserByID after delete: %v", err)
}
@@ -675,9 +710,9 @@ func TestTOTPManagement_RequiresPasswordConfirmation(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
userID, _ := database.CreateUser("totppassword", hash, 4)
userID, _ := database.CreateUser(context.Background(), "totppassword", hash, 4)
token, _ := auth.GenerateToken()
if _, err := database.CreateSession(userID, auth.HashToken(token), "test", "127.0.0.1"); err != nil {
if _, err := database.CreateSession(context.Background(), userID, auth.HashToken(token), "test", "127.0.0.1"); err != nil {
t.Fatalf("CreateSession: %v", err)
}
@@ -686,7 +721,7 @@ func TestTOTPManagement_RequiresPasswordConfirmation(t *testing.T) {
t.Fatalf("enable status = %d, want 400; body = %s", enable.Code, enable.Body.String())
}
userAfterEnable, err := database.GetUserByID(userID)
userAfterEnable, err := database.GetUserByID(context.Background(), userID)
if err != nil {
t.Fatalf("GetUserByID after failed enable: %v", err)
}
@@ -715,9 +750,9 @@ func TestVerifyTotp_ConsumesChallengeAfterRepeatedFailures(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
userID, _ := database.CreateUser("totplockout", hash, 4)
userID, _ := database.CreateUser(context.Background(), "totplockout", hash, 4)
secret := "JBSWY3DPEHPK3PXP"
if _, err := database.Exec(`UPDATE users SET totp_secret = ? WHERE id = ?`, secret, userID); err != nil {
if _, err := database.ExecContext(context.Background(), `UPDATE users SET totp_secret = ? WHERE id = ?`, secret, userID); err != nil {
t.Fatalf("set totp secret: %v", err)
}
@@ -760,15 +795,15 @@ func TestLogin_Require2FASettingRejectsUsersWithoutEnrollment(t *testing.T) {
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
if _, err := database.Exec(`UPDATE settings SET value = 'true' WHERE key = 'require_2fa'`); err != nil {
if _, err := database.ExecContext(context.Background(), `UPDATE settings SET value = 'true' WHERE key = 'require_2fa'`); err != nil {
t.Fatalf("enable require_2fa: %v", err)
}
if _, err := database.Exec(`UPDATE settings SET value = 'false' WHERE key = 'registration_open'`); err != nil {
if _, err := database.ExecContext(context.Background(), `UPDATE settings SET value = 'false' WHERE key = 'registration_open'`); err != nil {
t.Fatalf("disable registration_open: %v", err)
}
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser("needsenrollment", hash, 4)
_, _ = database.CreateUser(context.Background(), "needsenrollment", hash, 4)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "needsenrollment",
@@ -786,8 +821,8 @@ func TestLogin_BannedUser(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
id, _ := database.CreateUser("banned", hash, 4)
_ = database.BanUser(id, "violated rules", nil)
id, _ := database.CreateUser(context.Background(), "banned", hash, 4)
_ = database.BanUser(context.Background(), id, "violated rules", nil)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "banned",
@@ -818,10 +853,10 @@ func TestLogout_Success(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser("logoutuser", hash, 4)
uid, _ := database.CreateUser(context.Background(), "logoutuser", hash, 4)
token, _ := auth.GenerateToken()
tokenHash := auth.HashToken(token)
_, _ = database.CreateSession(uid, tokenHash, "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, tokenHash, "test", "127.0.0.1")
rr := postJSONWithToken(t, router, "/api/v1/auth/logout", token, nil)
@@ -830,7 +865,7 @@ func TestLogout_Success(t *testing.T) {
}
// Session should be gone.
sess, _ := database.GetSessionByTokenHash(tokenHash)
sess, _ := database.GetSessionByTokenHash(context.Background(), tokenHash)
if sess != nil {
t.Error("Session still exists after logout")
}
@@ -859,9 +894,9 @@ func TestMe_Success(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser("meuser", hash, 4)
uid, _ := database.CreateUser(context.Background(), "meuser", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
rr := getWithToken(t, router, "/api/v1/auth/me", token)
@@ -906,7 +941,7 @@ func TestLogin_PasswordWithLeadingSpaceIsPreserved(t *testing.T) {
// Hash the password WITH the leading space — this is what was registered.
hash, _ := auth.HashPassword(" securePass1")
_, _ = database.CreateUser("spacepassuser", hash, 4)
_, _ = database.CreateUser(context.Background(), "spacepassuser", hash, 4)
// Login with the exact same password (including space) must succeed.
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
@@ -928,7 +963,7 @@ func TestLogin_PasswordWithLeadingSpaceTrimmedFails(t *testing.T) {
// Register with password that has a leading space.
hash, _ := auth.HashPassword(" securePass1")
_, _ = database.CreateUser("spacepassuser2", hash, 4)
_, _ = database.CreateUser(context.Background(), "spacepassuser2", hash, 4)
// Login without the leading space must fail.
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
@@ -949,7 +984,7 @@ func TestLogin_PasswordWithTrailingSpaceIsPreserved(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("securePass1 ")
_, _ = database.CreateUser("trailingspaceuser", hash, 4)
_, _ = database.CreateUser(context.Background(), "trailingspaceuser", hash, 4)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "trailingspaceuser",
@@ -969,7 +1004,7 @@ func TestLogin_UsernameIsStillTrimmed(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
_, _ = database.CreateUser("trimuser", hash, 4)
_, _ = database.CreateUser(context.Background(), "trimuser", hash, 4)
// Username with surrounding spaces should resolve to "trimuser".
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
@@ -989,12 +1024,12 @@ func TestRegister_RateLimit(t *testing.T) {
limiter := auth.NewRateLimiter()
router := buildAuthRouter(database, limiter)
ownerID, _ := database.CreateUser("rl_owner", "hash", 1)
ownerID, _ := database.CreateUser(context.Background(), "rl_owner", "hash", 1)
// Attempt register 4 times (limit=3) — 4th should be rate-limited.
var lastCode int
for i := range 4 {
code, _ := database.CreateInvite(ownerID, 1, nil)
code, _ := database.CreateInvite(context.Background(), ownerID, 1, nil)
rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{
"username": "rl_user" + string(rune('0'+i)),
"password": "securePass1",
@@ -1030,10 +1065,10 @@ func TestDeleteAccount_Success(t *testing.T) {
hash, _ := auth.HashPassword("correctPass1")
// Create as Member (role_id=4) so the last-admin check does not block deletion.
uid, _ := database.CreateUser("deleteuser", hash, 4)
uid, _ := database.CreateUser(context.Background(), "deleteuser", hash, 4)
token, _ := auth.GenerateToken()
tokenHash := auth.HashToken(token)
_, _ = database.CreateSession(uid, tokenHash, "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, tokenHash, "test", "127.0.0.1")
rr := deleteJSONWithToken(t, router, "/api/v1/auth/account", token, map[string]string{
"password": "correctPass1",
@@ -1044,7 +1079,7 @@ func TestDeleteAccount_Success(t *testing.T) {
}
// User should be anonymised (banned, username changed).
user, err := database.GetUserByID(uid)
user, err := database.GetUserByID(context.Background(), uid)
if err != nil {
t.Fatalf("GetUserByID after delete: %v", err)
}
@@ -1059,7 +1094,7 @@ func TestDeleteAccount_Success(t *testing.T) {
}
// Session should be gone.
sess, _ := database.GetSessionByTokenHash(tokenHash)
sess, _ := database.GetSessionByTokenHash(context.Background(), tokenHash)
if sess != nil {
t.Error("session should be deleted after account deletion")
}
@@ -1071,9 +1106,9 @@ func TestDeleteAccount_MissingPassword(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser("delnopass", hash, 4)
uid, _ := database.CreateUser(context.Background(), "delnopass", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
rr := deleteJSONWithToken(t, router, "/api/v1/auth/account", token, map[string]string{})
@@ -1088,9 +1123,9 @@ func TestDeleteAccount_WrongPassword(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser("delwrong", hash, 4)
uid, _ := database.CreateUser(context.Background(), "delwrong", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
rr := deleteJSONWithToken(t, router, "/api/v1/auth/account", token, map[string]string{
"password": "wrongPassword1",
@@ -1101,7 +1136,7 @@ func TestDeleteAccount_WrongPassword(t *testing.T) {
}
// Verify user is NOT deleted.
user, _ := database.GetUserByID(uid)
user, _ := database.GetUserByID(context.Background(), uid)
if user == nil || user.Banned {
t.Error("user should not be deleted after wrong password")
}
@@ -1114,9 +1149,9 @@ func TestDeleteAccount_LastAdmin(t *testing.T) {
hash, _ := auth.HashPassword("correctPass1")
// Create as Owner (role_id=1) — the only admin-class user.
uid, _ := database.CreateUser("lastadmin", hash, 1)
uid, _ := database.CreateUser(context.Background(), "lastadmin", hash, 1)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
rr := deleteJSONWithToken(t, router, "/api/v1/auth/account", token, map[string]string{
"password": "correctPass1",
@@ -1127,7 +1162,7 @@ func TestDeleteAccount_LastAdmin(t *testing.T) {
}
// User should still be intact.
user, _ := database.GetUserByID(uid)
user, _ := database.GetUserByID(context.Background(), uid)
if user == nil || user.Banned {
t.Error("last admin should not be deleted")
}
@@ -1155,9 +1190,9 @@ func TestDeleteAccount_LockoutAfterRepeatedFailures(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser("dellockout", hash, 4)
uid, _ := database.CreateUser(context.Background(), "dellockout", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// 3 failures should trigger lockout on the 4th attempt.
for i := 0; i < 4; i++ {
@@ -1184,9 +1219,9 @@ func TestConfirmTOTP_InvalidCode(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser("totpbadcode", hash, 4)
uid, _ := database.CreateUser(context.Background(), "totpbadcode", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// Enable TOTP first to get a pending secret.
enable := postJSONWithToken(t, router, "/api/v1/users/me/totp/enable", token, map[string]string{"password": "correctPass1"})
@@ -1205,7 +1240,7 @@ func TestConfirmTOTP_InvalidCode(t *testing.T) {
}
// Secret should NOT be persisted.
user, _ := database.GetUserByID(uid)
user, _ := database.GetUserByID(context.Background(), uid)
if user.TOTPSecret != nil {
t.Error("TOTP secret should not be persisted after invalid code")
}
@@ -1217,9 +1252,9 @@ func TestConfirmTOTP_NoPendingSecret(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser("totpnopending", hash, 4)
uid, _ := database.CreateUser(context.Background(), "totpnopending", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// Confirm without enabling first — no pending secret.
confirm := postJSONWithToken(t, router, "/api/v1/users/me/totp/confirm", token, map[string]string{
@@ -1238,9 +1273,9 @@ func TestConfirmTOTP_MissingPassword(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser("totpnoconfirmpass", hash, 4)
uid, _ := database.CreateUser(context.Background(), "totpnoconfirmpass", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// Enable TOTP first.
postJSONWithToken(t, router, "/api/v1/users/me/totp/enable", token, map[string]string{"password": "correctPass1"})
@@ -1261,9 +1296,9 @@ func TestConfirmTOTP_WrongPassword(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser("totpwrongconfirm", hash, 4)
uid, _ := database.CreateUser(context.Background(), "totpwrongconfirm", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// Enable TOTP first.
postJSONWithToken(t, router, "/api/v1/users/me/totp/enable", token, map[string]string{"password": "correctPass1"})
@@ -1303,12 +1338,12 @@ func TestDisableTOTP_WrongPassword(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser("disabletotpwrong", hash, 4)
uid, _ := database.CreateUser(context.Background(), "disabletotpwrong", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// Set TOTP secret directly.
if _, err := database.Exec(`UPDATE users SET totp_secret = 'JBSWY3DPEHPK3PXP' WHERE id = ?`, uid); err != nil {
if _, err := database.ExecContext(context.Background(), `UPDATE users SET totp_secret = 'JBSWY3DPEHPK3PXP' WHERE id = ?`, uid); err != nil {
t.Fatalf("set totp secret: %v", err)
}
@@ -1325,7 +1360,7 @@ func TestDisableTOTP_WrongPassword(t *testing.T) {
}
// TOTP should still be enabled.
user, _ := database.GetUserByID(uid)
user, _ := database.GetUserByID(context.Background(), uid)
if user.TOTPSecret == nil {
t.Error("TOTP secret should still be set after wrong password")
}
@@ -1337,17 +1372,17 @@ func TestDisableTOTP_Require2FABlocksDisable(t *testing.T) {
router := buildAuthRouter(database, limiter)
// Enable require_2fa setting.
if _, err := database.Exec(`UPDATE settings SET value = 'true' WHERE key = 'require_2fa'`); err != nil {
if _, err := database.ExecContext(context.Background(), `UPDATE settings SET value = 'true' WHERE key = 'require_2fa'`); err != nil {
t.Fatalf("enable require_2fa: %v", err)
}
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser("disabletotpreq", hash, 4)
uid, _ := database.CreateUser(context.Background(), "disabletotpreq", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
// Set TOTP secret directly.
if _, err := database.Exec(`UPDATE users SET totp_secret = 'JBSWY3DPEHPK3PXP' WHERE id = ?`, uid); err != nil {
if _, err := database.ExecContext(context.Background(), `UPDATE users SET totp_secret = 'JBSWY3DPEHPK3PXP' WHERE id = ?`, uid); err != nil {
t.Fatalf("set totp secret: %v", err)
}
@@ -1364,7 +1399,7 @@ func TestDisableTOTP_Require2FABlocksDisable(t *testing.T) {
}
// TOTP should still be enabled.
user, _ := database.GetUserByID(uid)
user, _ := database.GetUserByID(context.Background(), uid)
if user.TOTPSecret == nil {
t.Error("TOTP secret should still be set when require_2fa is enabled")
}
@@ -1406,10 +1441,10 @@ func TestLogout_SessionGoneAfterLogout(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser("logoutsess", hash, 4)
uid, _ := database.CreateUser(context.Background(), "logoutsess", hash, 4)
token, _ := auth.GenerateToken()
tokenHash := auth.HashToken(token)
_, _ = database.CreateSession(uid, tokenHash, "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, tokenHash, "test", "127.0.0.1")
// First logout should succeed.
rr := postJSONWithToken(t, router, "/api/v1/auth/logout", token, nil)
@@ -1432,9 +1467,9 @@ func TestMe_ReturnsCorrectUserFields(t *testing.T) {
router := buildAuthRouter(database, limiter)
hash, _ := auth.HashPassword("correctPass1")
uid, _ := database.CreateUser("medetailed", hash, 4)
uid, _ := database.CreateUser(context.Background(), "medetailed", hash, 4)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
rr := getWithToken(t, router, "/api/v1/auth/me", token)
@@ -1492,9 +1527,9 @@ func containsStr(s, sub string) bool {
func expiredInviteDB(t *testing.T) (*db.DB, string) {
t.Helper()
database := newAuthTestDB(t)
ownerID, _ := database.CreateUser("expowner", "hash", 1)
ownerID, _ := database.CreateUser(context.Background(), "expowner", "hash", 1)
past := time.Now().Add(-time.Hour)
code, _ := database.CreateInvite(ownerID, 0, &past)
code, _ := database.CreateInvite(context.Background(), ownerID, 0, &past)
return database, code
}
+23 -22
View File
@@ -1,6 +1,7 @@
package api_test
import (
"context"
"encoding/json"
"fmt"
"net/http"
@@ -18,7 +19,7 @@ import (
// given role on the given channel.
func denyReadMessages(t *testing.T, database *db.DB, channelID, roleID int64) {
t.Helper()
_, err := database.Exec(
_, err := database.ExecContext(context.Background(),
`INSERT INTO channel_overrides (channel_id, role_id, allow, deny) VALUES (?, ?, 0, ?)`,
channelID, roleID, permissions.ReadMessages,
)
@@ -36,8 +37,8 @@ func TestChannelList_FiltersOutDeniedChannels(t *testing.T) {
// Create member user (roleID=4, has READ_MESSAGES by default).
token := chTestCreateToken(t, database, "authz-member1", 4)
chVisible, _ := database.CreateChannel("visible", "text", "", "", 0)
chHidden, _ := database.CreateChannel("hidden", "text", "", "", 1)
chVisible, _ := database.CreateChannel(context.Background(), "visible", "text", "", "", 0)
chHidden, _ := database.CreateChannel(context.Background(), "hidden", "text", "", "", 1)
_ = chVisible // used implicitly in response
// Deny READ_MESSAGES on the hidden channel for the Member role.
@@ -70,8 +71,8 @@ func TestChannelList_AdminSeesAllChannels(t *testing.T) {
// Owner (roleID=1) has Administrator bit — bypasses all checks.
token := chTestCreateToken(t, database, "authz-owner1", 1)
chA, _ := database.CreateChannel("a", "text", "", "", 0)
chB, _ := database.CreateChannel("b", "text", "", "", 1)
chA, _ := database.CreateChannel(context.Background(), "a", "text", "", "", 0)
chB, _ := database.CreateChannel(context.Background(), "b", "text", "", "", 1)
// Deny READ_MESSAGES on both channels for all roles.
denyReadMessages(t, database, chA, permissions.MemberRoleID)
@@ -96,7 +97,7 @@ func TestChannelMessages_DeniedByPermission(t *testing.T) {
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "authz-member2", 4)
chID, _ := database.CreateChannel("restricted", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "restricted", "text", "", "", 0)
// Deny READ_MESSAGES for Member role on this channel.
denyReadMessages(t, database, chID, permissions.MemberRoleID)
@@ -112,7 +113,7 @@ func TestChannelMessages_AdminBypassesDeny(t *testing.T) {
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "authz-owner2", 1)
chID, _ := database.CreateChannel("restricted", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "restricted", "text", "", "", 0)
// Deny READ_MESSAGES for Member role — should not affect Owner.
denyReadMessages(t, database, chID, permissions.MemberRoleID)
@@ -131,17 +132,17 @@ func TestSearch_FiltersResultsByPermission(t *testing.T) {
// Create an owner to insert messages (owner can write anywhere).
_ = chTestCreateToken(t, database, "authz-owner3", 1)
owner, _ := database.GetUserByUsername("authz-owner3")
owner, _ := database.GetUserByUsername(context.Background(), "authz-owner3")
// Member user for search.
memberToken := chTestCreateToken(t, database, "authz-member3", 4)
chVisible, _ := database.CreateChannel("pub", "text", "", "", 0)
chHidden, _ := database.CreateChannel("priv", "text", "", "", 1)
chVisible, _ := database.CreateChannel(context.Background(), "pub", "text", "", "", 0)
chHidden, _ := database.CreateChannel(context.Background(), "priv", "text", "", "", 1)
// Insert messages in both channels with a common keyword.
_, _ = database.CreateMessage(chVisible, owner.ID, "searchable keyword public", nil)
_, _ = database.CreateMessage(chHidden, owner.ID, "searchable keyword private", nil)
_, _ = database.CreateMessage(context.Background(), chVisible, owner.ID, "searchable keyword public", nil)
_, _ = database.CreateMessage(context.Background(), chHidden, owner.ID, "searchable keyword private", nil)
// Deny READ_MESSAGES on the hidden channel for members.
denyReadMessages(t, database, chHidden, permissions.MemberRoleID)
@@ -167,13 +168,13 @@ func TestSearch_AdminSeesAllResults(t *testing.T) {
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "authz-owner4", 1)
owner, _ := database.GetUserByUsername("authz-owner4")
owner, _ := database.GetUserByUsername(context.Background(), "authz-owner4")
chA, _ := database.CreateChannel("a", "text", "", "", 0)
chB, _ := database.CreateChannel("b", "text", "", "", 1)
chA, _ := database.CreateChannel(context.Background(), "a", "text", "", "", 0)
chB, _ := database.CreateChannel(context.Background(), "b", "text", "", "", 1)
_, _ = database.CreateMessage(chA, owner.ID, "findme alpha", nil)
_, _ = database.CreateMessage(chB, owner.ID, "findme beta", nil)
_, _ = database.CreateMessage(context.Background(), chA, owner.ID, "findme alpha", nil)
_, _ = database.CreateMessage(context.Background(), chB, owner.ID, "findme beta", nil)
// Deny READ_MESSAGES on both for member role — admin bypasses.
denyReadMessages(t, database, chA, permissions.MemberRoleID)
@@ -200,8 +201,8 @@ func TestChannelList_ExcludesDMChannels_Member(t *testing.T) {
token := chTestCreateToken(t, database, "dm-excl-member", 4)
// Create a normal text channel and a DM channel.
database.CreateChannel("general", "text", "", "", 0)
database.Exec(`INSERT INTO channels (name, type, position) VALUES ('dm-1', 'dm', 0)`)
database.CreateChannel(context.Background(), "general", "text", "", "", 0)
database.ExecContext(context.Background(), `INSERT INTO channels (name, type, position) VALUES ('dm-1', 'dm', 0)`)
rr := chGet(t, router, "/api/v1/channels", token)
if rr.Code != http.StatusOK {
@@ -227,9 +228,9 @@ func TestChannelList_ExcludesDMChannels_Admin(t *testing.T) {
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "dm-excl-admin", 1) // Owner
database.CreateChannel("general", "text", "", "", 0)
database.CreateChannel("voice", "voice", "", "", 1)
database.Exec(`INSERT INTO channels (name, type, position) VALUES ('dm-1', 'dm', 0)`)
database.CreateChannel(context.Background(), "general", "text", "", "", 0)
database.CreateChannel(context.Background(), "voice", "voice", "", "", 1)
database.ExecContext(context.Background(), `INSERT INTO channels (name, type, position) VALUES ('dm-1', 'dm', 0)`)
rr := chGet(t, router, "/api/v1/channels", token)
if rr.Code != http.StatusOK {
+4 -4
View File
@@ -131,7 +131,7 @@ func handleGetMessages(svc *service.Services) http.HandlerFunc {
limit = v
}
msgs, hasMore, err := svc.Messages.GetMessages(user.ID, channelID, before, limit)
msgs, hasMore, err := svc.Messages.GetMessages(r.Context(), user.ID, channelID, before, limit)
if err != nil {
writeServiceError(w, err)
return
@@ -191,7 +191,7 @@ func handleSearch(svc *service.Services) http.HandlerFunc {
limit = v
}
results, err := svc.Messages.SearchMessages(user.ID, q, channelID, limit)
results, err := svc.Messages.SearchMessages(r.Context(), user.ID, q, channelID, limit)
if err != nil {
if isInvalidSearchQueryError(err) {
writeJSON(w, http.StatusBadRequest, errorResponse{
@@ -229,7 +229,7 @@ func handleGetPins(svc *service.Services) http.HandlerFunc {
return
}
msgs, err := svc.Messages.GetPinnedMessages(user.ID, channelID)
msgs, err := svc.Messages.GetPinnedMessages(r.Context(), user.ID, channelID)
if err != nil {
writeServiceError(w, err)
return
@@ -263,7 +263,7 @@ func handleSetPinned(svc *service.Services, pinned bool) http.HandlerFunc {
return
}
if err := svc.Messages.SetMessagePinned(user.ID, channelID, messageID, pinned); err != nil {
if err := svc.Messages.SetMessagePinned(r.Context(), user.ID, channelID, messageID, pinned); err != nil {
writeServiceError(w, err)
return
}
+68 -67
View File
@@ -1,6 +1,7 @@
package api_test
import (
"context"
"encoding/json"
"fmt"
"net/http"
@@ -197,13 +198,13 @@ func buildChannelRouter(database *db.DB) http.Handler {
// chTestCreateToken creates a user+session and returns the plaintext token.
func chTestCreateToken(t *testing.T, database *db.DB, username string, roleID int) string {
t.Helper()
_, err := database.CreateUser(username, "$2a$12$fake", roleID)
_, err := database.CreateUser(context.Background(), username, "$2a$12$fake", roleID)
if err != nil {
t.Fatalf("CreateUser %q: %v", username, err)
}
token := "chtest-token-" + username
hash := auth.HashToken(token)
_, err = database.Exec(
_, err = database.ExecContext(context.Background(),
`INSERT INTO sessions (user_id, token, device, ip_address, expires_at)
SELECT id, ?, 'test', '127.0.0.1', '2099-01-01T00:00:00Z' FROM users WHERE username = ?`,
hash, username,
@@ -257,8 +258,8 @@ func TestChannelList_WithChannels(t *testing.T) {
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "bob", 1)
_, _ = database.CreateChannel("general", "text", "", "", 0)
_, _ = database.CreateChannel("random", "text", "", "", 1)
_, _ = database.CreateChannel(context.Background(), "general", "text", "", "", 0)
_, _ = database.CreateChannel(context.Background(), "random", "text", "", "", 1)
rr := chGet(t, router, "/api/v1/channels", token)
if rr.Code != http.StatusOK {
@@ -307,7 +308,7 @@ func TestChannelMessages_EmptyChannel(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "eve", 1)
chID, _ := database.CreateChannel("general", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "general", "text", "", "", 0)
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages", chID), token)
if rr.Code != http.StatusOK {
@@ -325,11 +326,11 @@ func TestChannelMessages_ReturnsMessages(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "frank", 1)
user, _ := database.GetUserByUsername("frank")
chID, _ := database.CreateChannel("ch", "text", "", "", 0)
user, _ := database.GetUserByUsername(context.Background(), "frank")
chID, _ := database.CreateChannel(context.Background(), "ch", "text", "", "", 0)
for i := range 3 {
_, _ = database.CreateMessage(chID, user.ID, fmt.Sprintf("msg%d", i), nil)
_, _ = database.CreateMessage(context.Background(), chID, user.ID, fmt.Sprintf("msg%d", i), nil)
}
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages", chID), token)
@@ -348,7 +349,7 @@ func TestChannelMessages_LimitCappedAt100(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "grace", 1)
chID, _ := database.CreateChannel("ch", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "ch", "text", "", "", 0)
// limit=200 should succeed (capped internally).
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages?limit=200", chID), token)
@@ -361,11 +362,11 @@ func TestChannelMessages_HasMore(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "henry", 1)
user, _ := database.GetUserByUsername("henry")
chID, _ := database.CreateChannel("ch", "text", "", "", 0)
user, _ := database.GetUserByUsername(context.Background(), "henry")
chID, _ := database.CreateChannel(context.Background(), "ch", "text", "", "", 0)
for i := range 60 {
_, _ = database.CreateMessage(chID, user.ID, fmt.Sprintf("m%d", i), nil)
_, _ = database.CreateMessage(context.Background(), chID, user.ID, fmt.Sprintf("m%d", i), nil)
}
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages?limit=50", chID), token)
@@ -383,11 +384,11 @@ func TestChannelMessages_HasMoreFalse(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "ivan", 1)
user, _ := database.GetUserByUsername("ivan")
chID, _ := database.CreateChannel("ch", "text", "", "", 0)
user, _ := database.GetUserByUsername(context.Background(), "ivan")
chID, _ := database.CreateChannel(context.Background(), "ch", "text", "", "", 0)
for i := range 5 {
_, _ = database.CreateMessage(chID, user.ID, fmt.Sprintf("m%d", i), nil)
_, _ = database.CreateMessage(context.Background(), chID, user.ID, fmt.Sprintf("m%d", i), nil)
}
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages?limit=50", chID), token)
@@ -426,9 +427,9 @@ func TestSearch_ReturnsResults(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "kim", 1)
user, _ := database.GetUserByUsername("kim")
chID, _ := database.CreateChannel("searchable", "text", "", "", 0)
_, _ = database.CreateMessage(chID, user.ID, "uniqueterm in message", nil)
user, _ := database.GetUserByUsername(context.Background(), "kim")
chID, _ := database.CreateChannel(context.Background(), "searchable", "text", "", "", 0)
_, _ = database.CreateMessage(context.Background(), chID, user.ID, "uniqueterm in message", nil)
rr := chGet(t, router, "/api/v1/search?q=uniqueterm", token)
if rr.Code != http.StatusOK {
@@ -463,9 +464,9 @@ func TestSearch_WithChannelID(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "searchch", 1)
user, _ := database.GetUserByUsername("searchch")
chID, _ := database.CreateChannel("filtered", "text", "", "", 0)
_, _ = database.CreateMessage(chID, user.ID, "filtered message here", nil)
user, _ := database.GetUserByUsername(context.Background(), "searchch")
chID, _ := database.CreateChannel(context.Background(), "filtered", "text", "", "", 0)
_, _ = database.CreateMessage(context.Background(), chID, user.ID, "filtered message here", nil)
rr := chGet(t, router, fmt.Sprintf("/api/v1/search?q=filtered&channel_id=%d", chID), token)
if rr.Code != http.StatusOK {
@@ -521,9 +522,9 @@ func TestSearch_InvalidFTSQuery(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "badfts", 1)
user, _ := database.GetUserByUsername("badfts")
chID, _ := database.CreateChannel("fts", "text", "", "", 0)
_, _ = database.CreateMessage(chID, user.ID, "search seed", nil)
user, _ := database.GetUserByUsername(context.Background(), "badfts")
chID, _ := database.CreateChannel(context.Background(), "fts", "text", "", "", 0)
_, _ = database.CreateMessage(context.Background(), chID, user.ID, "search seed", nil)
// FTS5 operator characters are now stripped by sanitizeFTSQuery, so a
// bare quote becomes an empty query which returns 200 with no results.
@@ -560,22 +561,22 @@ func TestSearch_ChannelTypeLookupFailure_FailsClosed(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "searchfailclosed", 1)
user, _ := database.GetUserByUsername("searchfailclosed")
chID, _ := database.CreateChannel("searchable", "text", "", "", 0)
_, _ = database.CreateMessage(chID, user.ID, "closedlookupterm", nil)
user, _ := database.GetUserByUsername(context.Background(), "searchfailclosed")
chID, _ := database.CreateChannel(context.Background(), "searchable", "text", "", "", 0)
_, _ = database.CreateMessage(context.Background(), chID, user.ID, "closedlookupterm", nil)
_, err := database.Exec(`ALTER TABLE channels RENAME TO channels_with_type`)
_, err := database.ExecContext(context.Background(), `ALTER TABLE channels RENAME TO channels_with_type`)
if err != nil {
t.Fatalf("rename channels: %v", err)
}
_, err = database.Exec(`CREATE TABLE channels (
_, err = database.ExecContext(context.Background(), `CREATE TABLE channels (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL
)`)
if err != nil {
t.Fatalf("recreate channels without type: %v", err)
}
_, err = database.Exec(`INSERT INTO channels (id, name) SELECT id, name FROM channels_with_type`)
_, err = database.ExecContext(context.Background(), `INSERT INTO channels (id, name) SELECT id, name FROM channels_with_type`)
if err != nil {
t.Fatalf("copy channels: %v", err)
}
@@ -590,11 +591,11 @@ func TestSearch_ChannelOverrideLookupFailure_ReturnsError(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "searchoverridefail", 4)
user, _ := database.GetUserByUsername("searchoverridefail")
chID, _ := database.CreateChannel("searchable", "text", "", "", 0)
_, _ = database.CreateMessage(chID, user.ID, "overridefailterm", nil)
user, _ := database.GetUserByUsername(context.Background(), "searchoverridefail")
chID, _ := database.CreateChannel(context.Background(), "searchable", "text", "", "", 0)
_, _ = database.CreateMessage(context.Background(), chID, user.ID, "overridefailterm", nil)
_, err := database.Exec(`DROP TABLE channel_overrides`)
_, err := database.ExecContext(context.Background(), `DROP TABLE channel_overrides`)
if err != nil {
t.Fatalf("drop channel_overrides: %v", err)
}
@@ -642,12 +643,12 @@ func TestChannelMessages_BeforeCursor(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "cursoruser", 1)
user, _ := database.GetUserByUsername("cursoruser")
chID, _ := database.CreateChannel("cursor", "text", "", "", 0)
user, _ := database.GetUserByUsername(context.Background(), "cursoruser")
chID, _ := database.CreateChannel(context.Background(), "cursor", "text", "", "", 0)
var lastID int64
for i := range 5 {
lastID, _ = database.CreateMessage(chID, user.ID, fmt.Sprintf("msg%d", i), nil)
lastID, _ = database.CreateMessage(context.Background(), chID, user.ID, fmt.Sprintf("msg%d", i), nil)
}
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages?before=%d", chID, lastID), token)
@@ -660,7 +661,7 @@ func TestChannelMessages_InvalidLimit(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "badlimituser", 1)
chID, _ := database.CreateChannel("lim", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "lim", "text", "", "", 0)
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages?limit=abc", chID), token)
if rr.Code != http.StatusBadRequest {
@@ -675,7 +676,7 @@ func newPinTestDB(t *testing.T) *db.DB {
t.Helper()
database := newChannelTestDB(t)
// Add DM tables required by pin handlers for DM authorization.
_, err := database.Exec(`
_, err := database.ExecContext(context.Background(), `
CREATE TABLE IF NOT EXISTS dm_participants (
channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE,
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
@@ -718,7 +719,7 @@ func TestGetPins_EmptyPins(t *testing.T) {
database := newPinTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "pinuser2", 1)
chID, _ := database.CreateChannel("general", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "general", "text", "", "", 0)
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/pins", chID), token)
if rr.Code != http.StatusOK {
@@ -742,13 +743,13 @@ func TestGetPins_ReturnsPinnedMessages(t *testing.T) {
database := newPinTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "pinuser3", 1)
user, _ := database.GetUserByUsername("pinuser3")
chID, _ := database.CreateChannel("general", "text", "", "", 0)
user, _ := database.GetUserByUsername(context.Background(), "pinuser3")
chID, _ := database.CreateChannel(context.Background(), "general", "text", "", "", 0)
msgID, _ := database.CreateMessage(chID, user.ID, "pinned message", nil)
_ = database.SetMessagePinned(msgID, true)
msgID, _ := database.CreateMessage(context.Background(), chID, user.ID, "pinned message", nil)
_ = database.SetMessagePinned(context.Background(), msgID, true)
// Also create an unpinned message — should not appear.
_, _ = database.CreateMessage(chID, user.ID, "not pinned", nil)
_, _ = database.CreateMessage(context.Background(), chID, user.ID, "not pinned", nil)
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/pins", chID), token)
if rr.Code != http.StatusOK {
@@ -773,11 +774,11 @@ func TestGetPins_DMChannel_NonParticipantForbidden(t *testing.T) {
chTestCreateToken(t, database, "dmuser2", 4)
outsiderToken := chTestCreateToken(t, database, "outsider", 4)
user1, _ := database.GetUserByUsername("dmuser1")
user2, _ := database.GetUserByUsername("dmuser2")
user1, _ := database.GetUserByUsername(context.Background(), "dmuser1")
user2, _ := database.GetUserByUsername(context.Background(), "dmuser2")
// Create a DM channel manually.
dmCh, _, _ := database.GetOrCreateDMChannel(user1.ID, user2.ID)
dmCh, _, _ := database.GetOrCreateDMChannel(context.Background(), user1.ID, user2.ID)
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/pins", dmCh.ID), outsiderToken)
if rr.Code != http.StatusNotFound {
@@ -791,10 +792,10 @@ func TestGetPins_MemberNoReadPermission(t *testing.T) {
// Role 4 = Member with permissions 1635 (0x663).
// Deny READ_MESSAGES on a specific channel via override.
token := chTestCreateToken(t, database, "nopermuser", 4)
chID, _ := database.CreateChannel("restricted", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "restricted", "text", "", "", 0)
// Deny all permissions for role 4 on this channel.
_, _ = database.Exec(
_, _ = database.ExecContext(context.Background(),
`INSERT INTO channel_overrides (channel_id, role_id, allow, deny) VALUES (?, 4, 0, 2147483647)`,
chID,
)
@@ -835,9 +836,9 @@ func TestSetPinned_PinSuccessfully(t *testing.T) {
database := newPinTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "pinner1", 1)
user, _ := database.GetUserByUsername("pinner1")
chID, _ := database.CreateChannel("general", "text", "", "", 0)
msgID, _ := database.CreateMessage(chID, user.ID, "pin me", nil)
user, _ := database.GetUserByUsername(context.Background(), "pinner1")
chID, _ := database.CreateChannel(context.Background(), "general", "text", "", "", 0)
msgID, _ := database.CreateMessage(context.Background(), chID, user.ID, "pin me", nil)
rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", chID, msgID), token)
if rr.Code != http.StatusNoContent {
@@ -845,7 +846,7 @@ func TestSetPinned_PinSuccessfully(t *testing.T) {
}
// Verify the message is actually pinned.
msg, _ := database.GetMessage(msgID)
msg, _ := database.GetMessage(context.Background(), msgID)
if !msg.Pinned {
t.Error("message should be pinned after POST")
}
@@ -855,17 +856,17 @@ func TestSetPinned_UnpinSuccessfully(t *testing.T) {
database := newPinTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "unpinner1", 1)
user, _ := database.GetUserByUsername("unpinner1")
chID, _ := database.CreateChannel("general", "text", "", "", 0)
msgID, _ := database.CreateMessage(chID, user.ID, "unpin me", nil)
_ = database.SetMessagePinned(msgID, true)
user, _ := database.GetUserByUsername(context.Background(), "unpinner1")
chID, _ := database.CreateChannel(context.Background(), "general", "text", "", "", 0)
msgID, _ := database.CreateMessage(context.Background(), chID, user.ID, "unpin me", nil)
_ = database.SetMessagePinned(context.Background(), msgID, true)
rr := chDelete(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", chID, msgID), token)
if rr.Code != http.StatusNoContent {
t.Errorf("unpin status = %d, want 204; body: %s", rr.Code, rr.Body.String())
}
msg, _ := database.GetMessage(msgID)
msg, _ := database.GetMessage(context.Background(), msgID)
if msg.Pinned {
t.Error("message should not be pinned after DELETE")
}
@@ -875,7 +876,7 @@ func TestSetPinned_MessageNotFound(t *testing.T) {
database := newPinTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "pinner2", 1)
chID, _ := database.CreateChannel("general", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "general", "text", "", "", 0)
rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/9999", chID), token)
if rr.Code != http.StatusNotFound {
@@ -899,9 +900,9 @@ func TestSetPinned_NoPermission(t *testing.T) {
router := buildChannelRouter(database)
// Member role (4) has permissions 1635 — does not include MANAGE_MESSAGES (0x2000).
token := chTestCreateToken(t, database, "noperm", 4)
user, _ := database.GetUserByUsername("noperm")
chID, _ := database.CreateChannel("general", "text", "", "", 0)
msgID, _ := database.CreateMessage(chID, user.ID, "try to pin", nil)
user, _ := database.GetUserByUsername(context.Background(), "noperm")
chID, _ := database.CreateChannel(context.Background(), "general", "text", "", "", 0)
msgID, _ := database.CreateMessage(context.Background(), chID, user.ID, "try to pin", nil)
rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", chID, msgID), token)
if rr.Code != http.StatusForbidden {
@@ -913,10 +914,10 @@ func TestSetPinned_Idempotent(t *testing.T) {
database := newPinTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "pinner4", 1)
user, _ := database.GetUserByUsername("pinner4")
chID, _ := database.CreateChannel("general", "text", "", "", 0)
msgID, _ := database.CreateMessage(chID, user.ID, "already pinned", nil)
_ = database.SetMessagePinned(msgID, true)
user, _ := database.GetUserByUsername(context.Background(), "pinner4")
chID, _ := database.CreateChannel(context.Background(), "general", "text", "", "", 0)
msgID, _ := database.CreateMessage(context.Background(), chID, user.ID, "already pinned", nil)
_ = database.SetMessagePinned(context.Background(), msgID, true)
// Pinning again should still succeed (idempotent).
rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", chID, msgID), token)
+11 -10
View File
@@ -1,6 +1,7 @@
package api_test
import (
"context"
"encoding/json"
"fmt"
"net/http"
@@ -17,9 +18,9 @@ func TestContract_Messages_HasRequiredFields(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "contract-msg1", 1)
user, _ := database.GetUserByUsername("contract-msg1")
chID, _ := database.CreateChannel("contract-ch", "text", "", "", 0)
_, _ = database.CreateMessage(chID, user.ID, "contract test message", nil)
user, _ := database.GetUserByUsername(context.Background(), "contract-msg1")
chID, _ := database.CreateChannel(context.Background(), "contract-ch", "text", "", "", 0)
_, _ = database.CreateMessage(context.Background(), chID, user.ID, "contract test message", nil)
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages", chID), token)
if rr.Code != http.StatusOK {
@@ -85,10 +86,10 @@ func TestContract_Messages_ReactionsHaveMeFlag(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "contract-react1", 1)
user, _ := database.GetUserByUsername("contract-react1")
chID, _ := database.CreateChannel("react-ch", "text", "", "", 0)
msgID, _ := database.CreateMessage(chID, user.ID, "reaction target", nil)
_ = database.AddReaction(msgID, user.ID, "👍")
user, _ := database.GetUserByUsername(context.Background(), "contract-react1")
chID, _ := database.CreateChannel(context.Background(), "react-ch", "text", "", "", 0)
msgID, _ := database.CreateMessage(context.Background(), chID, user.ID, "reaction target", nil)
_ = database.AddReaction(context.Background(), msgID, user.ID, "👍")
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages", chID), token)
if rr.Code != http.StatusOK {
@@ -133,9 +134,9 @@ func TestContract_Search_HasRequiredFields(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "contract-search1", 1)
user, _ := database.GetUserByUsername("contract-search1")
chID, _ := database.CreateChannel("search-ch", "text", "", "", 0)
_, _ = database.CreateMessage(chID, user.ID, "contractsearchterm in body", nil)
user, _ := database.GetUserByUsername(context.Background(), "contract-search1")
chID, _ := database.CreateChannel(context.Background(), "search-ch", "text", "", "", 0)
_, _ = database.CreateMessage(context.Background(), chID, user.ID, "contractsearchterm in body", nil)
rr := chGet(t, router, "/api/v1/search?q=contractsearchterm", token)
if rr.Code != http.StatusOK {
+47 -47
View File
@@ -485,7 +485,7 @@ func TestCloseDM_BroadcasterUserOffline(t *testing.T) {
tokenAlice := dmCreateToken(t, database, "offline_alice", 4)
_ = dmCreateToken(t, database, "offline_bob", 4)
bob, _ := database.GetUserByUsername("offline_bob")
bob, _ := database.GetUserByUsername(context.Background(), "offline_bob")
// Create a DM.
rr := dmPost(t, router, "/api/v1/dms", tokenAlice, map[string]any{
@@ -565,7 +565,7 @@ func TestGetMessages_InvalidLimit(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "msglimit", 1)
chID, _ := database.CreateChannel("limit-ch", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "limit-ch", "text", "", "", 0)
// Negative limit.
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages?limit=-1", chID), token)
@@ -592,7 +592,7 @@ func TestGetPins_Success(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "pinuser", 1)
chID, _ := database.CreateChannel("pin-ch", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "pin-ch", "text", "", "", 0)
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/pins", chID), token)
if rr.Code != http.StatusOK {
@@ -603,7 +603,7 @@ func TestGetPins_Success(t *testing.T) {
func TestGetPins_Unauthorized(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
chID, _ := database.CreateChannel("pin-unauth-ch", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "pin-unauth-ch", "text", "", "", 0)
req := httptest.NewRequest(http.MethodGet, fmt.Sprintf("/api/v1/channels/%d/pins", chID), nil)
req.RemoteAddr = "127.0.0.1:9999"
@@ -620,7 +620,7 @@ func TestGetPins_Unauthorized(t *testing.T) {
func TestSetPinned_Unauthorized(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
chID, _ := database.CreateChannel("setpin-ch", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "setpin-ch", "text", "", "", 0)
req := httptest.NewRequest(http.MethodPut,
fmt.Sprintf("/api/v1/channels/%d/messages/1/pin", chID),
@@ -712,8 +712,8 @@ func TestListSessions_MultipleSessions(t *testing.T) {
token := profileCreateToken(t, database, "multisess", 4)
// Create additional session.
user, _ := database.GetUserByUsername("multisess")
_, _ = database.CreateSession(user.ID, auth.HashToken("extra-token"), "Chrome", "1.2.3.4")
user, _ := database.GetUserByUsername(context.Background(), "multisess")
_, _ = database.CreateSession(context.Background(), user.ID, auth.HashToken("extra-token"), "Chrome", "1.2.3.4")
rr := getWithToken(t, router, "/api/v1/users/me/sessions", token)
if rr.Code != http.StatusOK {
@@ -784,7 +784,7 @@ func TestSetPinned_MessageNotFound_Push(t *testing.T) {
database := newPinTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "pinmissmsg", 1)
chID, _ := database.CreateChannel("pinmiss-ch", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "pinmiss-ch", "text", "", "", 0)
rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", chID, 99999), token)
if rr.Code != http.StatusNotFound {
@@ -818,7 +818,7 @@ func TestSetPinned_InvalidMessageID(t *testing.T) {
database := newPinTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "pinbadmsg", 1)
chID, _ := database.CreateChannel("badmsgid-ch", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "badmsgid-ch", "text", "", "", 0)
rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/abc", chID), token)
if rr.Code != http.StatusBadRequest {
@@ -830,10 +830,10 @@ func TestUnpin_Success(t *testing.T) {
database := newPinTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "unpinner", 1)
user, _ := database.GetUserByUsername("unpinner")
chID, _ := database.CreateChannel("unpin-ch", "text", "", "", 0)
msgID, _ := database.CreateMessage(chID, user.ID, "to unpin", nil)
_ = database.SetMessagePinned(msgID, true)
user, _ := database.GetUserByUsername(context.Background(), "unpinner")
chID, _ := database.CreateChannel(context.Background(), "unpin-ch", "text", "", "", 0)
msgID, _ := database.CreateMessage(context.Background(), chID, user.ID, "to unpin", nil)
_ = database.SetMessagePinned(context.Background(), msgID, true)
rr := chDelete(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", chID, msgID), token)
if rr.Code != http.StatusNoContent {
@@ -845,9 +845,9 @@ func TestSetPinned_MemberForbidden(t *testing.T) {
database := newPinTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "pinmember", 4)
user, _ := database.GetUserByUsername("pinmember")
chID, _ := database.CreateChannel("pinforbid-ch", "text", "", "", 0)
msgID, _ := database.CreateMessage(chID, user.ID, "cant pin", nil)
user, _ := database.GetUserByUsername(context.Background(), "pinmember")
chID, _ := database.CreateChannel(context.Background(), "pinforbid-ch", "text", "", "", 0)
msgID, _ := database.CreateMessage(context.Background(), chID, user.ID, "cant pin", nil)
rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", chID, msgID), token)
if rr.Code != http.StatusForbidden {
@@ -859,10 +859,10 @@ func TestSetPinned_WrongChannel(t *testing.T) {
database := newPinTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "pinwrongch", 1)
user, _ := database.GetUserByUsername("pinwrongch")
chID1, _ := database.CreateChannel("pin-ch1", "text", "", "", 0)
chID2, _ := database.CreateChannel("pin-ch2", "text", "", "", 0)
msgID, _ := database.CreateMessage(chID1, user.ID, "wrong channel", nil)
user, _ := database.GetUserByUsername(context.Background(), "pinwrongch")
chID1, _ := database.CreateChannel(context.Background(), "pin-ch1", "text", "", "", 0)
chID2, _ := database.CreateChannel(context.Background(), "pin-ch2", "text", "", "", 0)
msgID, _ := database.CreateMessage(context.Background(), chID1, user.ID, "wrong channel", nil)
// Try to pin a message from chID1 using chID2.
rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", chID2, msgID), token)
@@ -878,11 +878,11 @@ func TestSetPinned_DMChannel_ParticipantSuccess(t *testing.T) {
router := buildChannelRouter(database)
tokenAlice := chTestCreateToken(t, database, "dmpin_alice", 4)
_ = chTestCreateToken(t, database, "dmpin_bob", 4)
alice, _ := database.GetUserByUsername("dmpin_alice")
bob, _ := database.GetUserByUsername("dmpin_bob")
alice, _ := database.GetUserByUsername(context.Background(), "dmpin_alice")
bob, _ := database.GetUserByUsername(context.Background(), "dmpin_bob")
dmCh, _, _ := database.GetOrCreateDMChannel(alice.ID, bob.ID)
msgID, _ := database.CreateMessage(dmCh.ID, alice.ID, "pin this dm msg", nil)
dmCh, _, _ := database.GetOrCreateDMChannel(context.Background(), alice.ID, bob.ID)
msgID, _ := database.CreateMessage(context.Background(), dmCh.ID, alice.ID, "pin this dm msg", nil)
rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", dmCh.ID, msgID), tokenAlice)
if rr.Code != http.StatusNoContent {
@@ -896,11 +896,11 @@ func TestSetPinned_DMChannel_NonParticipantForbidden(t *testing.T) {
_ = chTestCreateToken(t, database, "dmpinforbid_alice", 4)
_ = chTestCreateToken(t, database, "dmpinforbid_bob", 4)
tokenCharlie := chTestCreateToken(t, database, "dmpinforbid_charlie", 4)
alice, _ := database.GetUserByUsername("dmpinforbid_alice")
bob, _ := database.GetUserByUsername("dmpinforbid_bob")
alice, _ := database.GetUserByUsername(context.Background(), "dmpinforbid_alice")
bob, _ := database.GetUserByUsername(context.Background(), "dmpinforbid_bob")
dmCh, _, _ := database.GetOrCreateDMChannel(alice.ID, bob.ID)
msgID, _ := database.CreateMessage(dmCh.ID, alice.ID, "secret msg", nil)
dmCh, _, _ := database.GetOrCreateDMChannel(context.Background(), alice.ID, bob.ID)
msgID, _ := database.CreateMessage(context.Background(), dmCh.ID, alice.ID, "secret msg", nil)
rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", dmCh.ID, msgID), tokenCharlie)
if rr.Code != http.StatusNotFound {
@@ -914,9 +914,9 @@ func TestSearch_WithChannelID_Push(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "searchch", 1)
user, _ := database.GetUserByUsername("searchch")
chID, _ := database.CreateChannel("search-ch1", "text", "", "", 0)
_, _ = database.CreateMessage(chID, user.ID, "findable in channel", nil)
user, _ := database.GetUserByUsername(context.Background(), "searchch")
chID, _ := database.CreateChannel(context.Background(), "search-ch1", "text", "", "", 0)
_, _ = database.CreateMessage(context.Background(), chID, user.ID, "findable in channel", nil)
rr := chGet(t, router, fmt.Sprintf("/api/v1/search?q=findable&channel_id=%d", chID), token)
if rr.Code != http.StatusOK {
@@ -1028,10 +1028,10 @@ func TestGetMessages_WithBeforeParam(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "msgbefore", 1)
user, _ := database.GetUserByUsername("msgbefore")
chID, _ := database.CreateChannel("before-ch", "text", "", "", 0)
_, _ = database.CreateMessage(chID, user.ID, "msg one", nil)
msgID2, _ := database.CreateMessage(chID, user.ID, "msg two", nil)
user, _ := database.GetUserByUsername(context.Background(), "msgbefore")
chID, _ := database.CreateChannel(context.Background(), "before-ch", "text", "", "", 0)
_, _ = database.CreateMessage(context.Background(), chID, user.ID, "msg one", nil)
msgID2, _ := database.CreateMessage(context.Background(), chID, user.ID, "msg two", nil)
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages?before=%d", chID, msgID2), token)
if rr.Code != http.StatusOK {
@@ -1043,7 +1043,7 @@ func TestGetMessages_WithCustomLimit(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "msglimitcust", 1)
chID, _ := database.CreateChannel("limitch", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "limitch", "text", "", "", 0)
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages?limit=5", chID), token)
if rr.Code != http.StatusOK {
@@ -1057,7 +1057,7 @@ func TestListChannels_MemberRole(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "memberchanlist", 4)
_, _ = database.CreateChannel("visible-ch", "text", "", "", 0)
_, _ = database.CreateChannel(context.Background(), "visible-ch", "text", "", "", 0)
rr := chGet(t, router, "/api/v1/channels", token)
if rr.Code != http.StatusOK {
@@ -1069,7 +1069,7 @@ func TestListChannels_AdminSeesAll(t *testing.T) {
database := newChannelTestDB(t)
router := buildChannelRouter(database)
token := chTestCreateToken(t, database, "adminchanlist", 2)
_, _ = database.CreateChannel("admin-visible-ch", "text", "", "", 0)
_, _ = database.CreateChannel(context.Background(), "admin-visible-ch", "text", "", "", 0)
rr := chGet(t, router, "/api/v1/channels", token)
if rr.Code != http.StatusOK {
@@ -1091,10 +1091,10 @@ func TestGetMessages_DMChannel_NonParticipant(t *testing.T) {
_ = chTestCreateToken(t, database, "dmmsg_alice", 4)
_ = chTestCreateToken(t, database, "dmmsg_bob", 4)
tokenCharlie := chTestCreateToken(t, database, "dmmsg_charlie", 4)
alice, _ := database.GetUserByUsername("dmmsg_alice")
bob, _ := database.GetUserByUsername("dmmsg_bob")
alice, _ := database.GetUserByUsername(context.Background(), "dmmsg_alice")
bob, _ := database.GetUserByUsername(context.Background(), "dmmsg_bob")
dmCh, _, _ := database.GetOrCreateDMChannel(alice.ID, bob.ID)
dmCh, _, _ := database.GetOrCreateDMChannel(context.Background(), alice.ID, bob.ID)
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages", dmCh.ID), tokenCharlie)
if rr.Code != http.StatusNotFound {
@@ -1107,10 +1107,10 @@ func TestGetMessages_DMChannel_ParticipantSuccess(t *testing.T) {
router := buildChannelRouter(database)
tokenAlice := chTestCreateToken(t, database, "dmmsgok_alice", 4)
_ = chTestCreateToken(t, database, "dmmsgok_bob", 4)
alice, _ := database.GetUserByUsername("dmmsgok_alice")
bob, _ := database.GetUserByUsername("dmmsgok_bob")
alice, _ := database.GetUserByUsername(context.Background(), "dmmsgok_alice")
bob, _ := database.GetUserByUsername(context.Background(), "dmmsgok_bob")
dmCh, _, _ := database.GetOrCreateDMChannel(alice.ID, bob.ID)
dmCh, _, _ := database.GetOrCreateDMChannel(context.Background(), alice.ID, bob.ID)
rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/messages", dmCh.ID), tokenAlice)
if rr.Code != http.StatusOK {
@@ -1126,10 +1126,10 @@ func TestSearch_DMChannelFilter_NonParticipant(t *testing.T) {
_ = chTestCreateToken(t, database, "dmsearch_alice", 4)
_ = chTestCreateToken(t, database, "dmsearch_bob", 4)
tokenCharlie := chTestCreateToken(t, database, "dmsearch_charlie", 4)
alice, _ := database.GetUserByUsername("dmsearch_alice")
bob, _ := database.GetUserByUsername("dmsearch_bob")
alice, _ := database.GetUserByUsername(context.Background(), "dmsearch_alice")
bob, _ := database.GetUserByUsername(context.Background(), "dmsearch_bob")
dmCh, _, _ := database.GetOrCreateDMChannel(alice.ID, bob.ID)
dmCh, _, _ := database.GetOrCreateDMChannel(context.Background(), alice.ID, bob.ID)
rr := chGet(t, router, fmt.Sprintf("/api/v1/search?q=test&channel_id=%d", dmCh.ID), tokenCharlie)
if rr.Code != http.StatusForbidden {
+39 -6
View File
@@ -1,6 +1,7 @@
package api_test
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
@@ -10,11 +11,12 @@ import (
"github.com/owncord/server/auth"
"github.com/owncord/server/config"
"github.com/owncord/server/db"
"github.com/owncord/server/permissions"
)
// setupDiagnosticsRouter creates a full router with an authenticated user for
// diagnostics testing.
func setupDiagnosticsRouter(t *testing.T) (http.Handler, string) {
func setupDiagnosticsRouter(t *testing.T) (http.Handler, string, *db.DB) {
t.Helper()
database, err := db.Open(":memory:")
@@ -37,20 +39,20 @@ func setupDiagnosticsRouter(t *testing.T) (http.Handler, string) {
t.Cleanup(cleanup)
// Create a user and session for authenticated requests.
uid, _ := database.CreateUser("diaguser", "$2a$12$fake", 1)
uid, _ := database.CreateUser(context.Background(), "diaguser", "$2a$12$fake", 1)
token := "diagtest-token-123"
hash := auth.HashToken(token)
_, _ = database.Exec(
_, _ = database.ExecContext(context.Background(),
`INSERT INTO sessions (user_id, token, device, ip_address, expires_at)
VALUES (?, ?, 'test', '127.0.0.1', '2099-01-01T00:00:00Z')`,
uid, hash,
)
return handler, token
return handler, token, database
}
func TestDiagnosticsConnectivity_ReturnsData(t *testing.T) {
router, token := setupDiagnosticsRouter(t)
router, token, _ := setupDiagnosticsRouter(t)
req := httptest.NewRequest(http.MethodGet, "/api/v1/diagnostics/connectivity", nil)
req.Header.Set("Authorization", "Bearer "+token)
@@ -82,7 +84,7 @@ func TestDiagnosticsConnectivity_ReturnsData(t *testing.T) {
}
func TestDiagnosticsConnectivity_Unauthenticated(t *testing.T) {
router, _ := setupDiagnosticsRouter(t)
router, _, _ := setupDiagnosticsRouter(t)
req := httptest.NewRequest(http.MethodGet, "/api/v1/diagnostics/connectivity", nil)
req.RemoteAddr = "127.0.0.1:9999"
@@ -94,6 +96,37 @@ func TestDiagnosticsConnectivity_Unauthenticated(t *testing.T) {
}
}
// TestDiagnosticsConnectivity_MemberForbidden locks the RequirePermission gate
// on the route. Without it, only 200-for-owner and 401-unauthenticated were
// covered, so deleting the ADMINISTRATOR gate broke no test while exposing the
// server's network topology to every member.
func TestDiagnosticsConnectivity_MemberForbidden(t *testing.T) {
router, _, database := setupDiagnosticsRouter(t)
uid, err := database.CreateUser(context.Background(), "diagmember", "$2a$12$fake", int(permissions.MemberRoleID))
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
token := "diagtest-member-token"
if _, err := database.ExecContext(context.Background(),
`INSERT INTO sessions (user_id, token, device, ip_address, expires_at)
VALUES (?, ?, 'test', '127.0.0.1', '2099-01-01T00:00:00Z')`,
uid, auth.HashToken(token),
); err != nil {
t.Fatalf("insert session: %v", err)
}
req := httptest.NewRequest(http.MethodGet, "/api/v1/diagnostics/connectivity", nil)
req.Header.Set("Authorization", "Bearer "+token)
req.RemoteAddr = "127.0.0.1:9999"
rr := httptest.NewRecorder()
router.ServeHTTP(rr, req)
if rr.Code != http.StatusForbidden {
t.Errorf("status = %d, want 403; body: %s", rr.Code, rr.Body.String())
}
}
// ─── isPrivateIP tests ──────────────────────────────────────────────────────
func TestIsPrivateIP(t *testing.T) {
+4 -4
View File
@@ -113,7 +113,7 @@ func handleListDMs(svc *service.Services) http.HandlerFunc {
return
}
channels, err := svc.DMs.ListDMs(user.ID)
channels, err := svc.DMs.ListDMs(r.Context(), user.ID)
if err != nil {
writeServiceError(w, err)
return
@@ -138,7 +138,7 @@ func handleCloseDM(svc *service.Services, broadcaster DMBroadcaster) http.Handle
return
}
if err := svc.DMs.CloseDM(user.ID, channelID); err != nil {
if err := svc.DMs.CloseDM(r.Context(), user.ID, channelID); err != nil {
writeServiceError(w, err)
return
}
@@ -191,7 +191,7 @@ func handleUnblockUser(svc *service.Services) http.HandlerFunc {
return
}
if err := svc.Blocks.UnblockUser(user.ID, targetID); err != nil {
if err := svc.Blocks.UnblockUser(r.Context(), user.ID, targetID); err != nil {
writeServiceError(w, err)
return
}
@@ -208,7 +208,7 @@ func handleListBlocks(svc *service.Services) http.HandlerFunc {
return
}
ids, err := svc.Blocks.ListBlocked(user.ID)
ids, err := svc.Blocks.ListBlocked(r.Context(), user.ID)
if err != nil {
writeServiceError(w, err)
return
+11 -10
View File
@@ -2,6 +2,7 @@ package api_test
import (
"bytes"
"context"
"encoding/json"
"fmt"
"net/http"
@@ -154,13 +155,13 @@ func (m *mockBroadcaster) SendToUser(userID int64, msg []byte) bool {
// dmCreateToken creates a user+session and returns the plaintext token.
func dmCreateToken(t *testing.T, database *db.DB, username string, roleID int) string {
t.Helper()
_, err := database.CreateUser(username, "$2a$12$fake", roleID)
_, err := database.CreateUser(context.Background(), username, "$2a$12$fake", roleID)
if err != nil {
t.Fatalf("CreateUser %q: %v", username, err)
}
token := "dmtest-token-" + username
hash := auth.HashToken(token)
_, err = database.Exec(
_, err = database.ExecContext(context.Background(),
`INSERT INTO sessions (user_id, token, device, ip_address, expires_at)
SELECT id, ?, 'test', '127.0.0.1', '2099-01-01T00:00:00Z' FROM users WHERE username = ?`,
hash, username,
@@ -225,7 +226,7 @@ func TestCreateDM_Success_NewDM(t *testing.T) {
tokenAlice := dmCreateToken(t, database, "alice", 4)
_ = dmCreateToken(t, database, "bob", 4)
bob, _ := database.GetUserByUsername("bob")
bob, _ := database.GetUserByUsername(context.Background(), "bob")
rr := dmPost(t, router, "/api/v1/dms", tokenAlice, map[string]any{
"recipient_id": bob.ID,
@@ -256,7 +257,7 @@ func TestCreateDM_Success_ExistingDM(t *testing.T) {
tokenAlice := dmCreateToken(t, database, "alice2", 4)
_ = dmCreateToken(t, database, "bob2", 4)
bob, _ := database.GetUserByUsername("bob2")
bob, _ := database.GetUserByUsername(context.Background(), "bob2")
// First call creates the DM.
rr1 := dmPost(t, router, "/api/v1/dms", tokenAlice, map[string]any{
@@ -328,7 +329,7 @@ func TestCreateDM_BadRequest_SelfDM(t *testing.T) {
database := newDMTestDB(t)
router := buildDMRouter(database, nil)
token := dmCreateToken(t, database, "selfuser", 4)
self, _ := database.GetUserByUsername("selfuser")
self, _ := database.GetUserByUsername(context.Background(), "selfuser")
rr := dmPost(t, router, "/api/v1/dms", token, map[string]any{
"recipient_id": self.ID,
@@ -372,7 +373,7 @@ func TestListDMs_ReturnsOpenDMs(t *testing.T) {
tokenAlice := dmCreateToken(t, database, "list_alice", 4)
_ = dmCreateToken(t, database, "list_bob", 4)
bob, _ := database.GetUserByUsername("list_bob")
bob, _ := database.GetUserByUsername(context.Background(), "list_bob")
// Create a DM.
rr1 := dmPost(t, router, "/api/v1/dms", tokenAlice, map[string]any{
@@ -436,7 +437,7 @@ func TestCloseDM_Success(t *testing.T) {
tokenAlice := dmCreateToken(t, database, "close_alice", 4)
_ = dmCreateToken(t, database, "close_bob", 4)
bob, _ := database.GetUserByUsername("close_bob")
bob, _ := database.GetUserByUsername(context.Background(), "close_bob")
// Create a DM.
rr1 := dmPost(t, router, "/api/v1/dms", tokenAlice, map[string]any{
@@ -468,7 +469,7 @@ func TestCloseDM_Success_VerifyRemovedFromList(t *testing.T) {
tokenAlice := dmCreateToken(t, database, "closelist_alice", 4)
_ = dmCreateToken(t, database, "closelist_bob", 4)
bob, _ := database.GetUserByUsername("closelist_bob")
bob, _ := database.GetUserByUsername(context.Background(), "closelist_bob")
// Create a DM.
rr1 := dmPost(t, router, "/api/v1/dms", tokenAlice, map[string]any{
@@ -502,7 +503,7 @@ func TestCloseDM_Forbidden_NotParticipant(t *testing.T) {
tokenAlice := dmCreateToken(t, database, "forbid_alice", 4)
_ = dmCreateToken(t, database, "forbid_bob", 4)
tokenCharlie := dmCreateToken(t, database, "forbid_charlie", 4)
bob, _ := database.GetUserByUsername("forbid_bob")
bob, _ := database.GetUserByUsername(context.Background(), "forbid_bob")
// Alice creates DM with Bob.
rr1 := dmPost(t, router, "/api/v1/dms", tokenAlice, map[string]any{
@@ -549,7 +550,7 @@ func TestCloseDM_NilBroadcaster(t *testing.T) {
token := dmCreateToken(t, database, "nilbc_alice", 4)
_ = dmCreateToken(t, database, "nilbc_bob", 4)
bob, _ := database.GetUserByUsername("nilbc_bob")
bob, _ := database.GetUserByUsername(context.Background(), "nilbc_bob")
// Create a DM.
rr1 := dmPost(t, router, "/api/v1/dms", token, map[string]any{
+6
View File
@@ -44,3 +44,9 @@ func SetGIFUpstreamForTest(baseURL string, client *http.Client) func() {
gifAPIBase, gifClient = baseURL, client
return func() { gifAPIBase, gifClient = prevBase, prevClient }
}
// SecurityHeaders is SecurityHeadersWithTLS with TLS disabled (no HSTS).
// Test-only convenience — production always goes through SecurityHeadersWithTLS.
func SecurityHeaders(next http.Handler) http.Handler {
return SecurityHeadersWithTLS("")(next)
}
+2 -2
View File
@@ -85,7 +85,7 @@ func handleCreateInvite(svc *service.Services) http.HandlerFunc {
// handleListInvites processes GET /api/v1/invites.
func handleListInvites(svc *service.Services) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
invites, err := svc.Invites.ListInvites()
invites, err := svc.Invites.ListInvites(r.Context())
if err != nil {
writeServiceError(w, err)
return
@@ -103,7 +103,7 @@ func handleListInvites(svc *service.Services) http.HandlerFunc {
func handleRevokeInvite(svc *service.Services) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
code := chi.URLParam(r, "code")
if err := svc.Invites.RevokeInvite(code); err != nil {
if err := svc.Invites.RevokeInvite(r.Context(), code); err != nil {
writeServiceError(w, err)
return
}
+43 -7
View File
@@ -1,6 +1,7 @@
package api_test
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
@@ -10,6 +11,7 @@ import (
"github.com/owncord/server/api"
"github.com/owncord/server/auth"
"github.com/owncord/server/db"
"github.com/owncord/server/permissions"
"github.com/owncord/server/service"
)
@@ -26,9 +28,9 @@ func buildInviteRouter(database *db.DB, limiter *auth.RateLimiter) http.Handler
func loginAndGetToken(t *testing.T, _ http.Handler, database *db.DB, username string, roleID int) string {
t.Helper()
hash, _ := auth.HashPassword("Password1!")
uid, _ := database.CreateUser(username, hash, roleID)
uid, _ := database.CreateUser(context.Background(), username, hash, roleID)
token, _ := auth.GenerateToken()
_, _ = database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, auth.HashToken(token), "test", "127.0.0.1")
return token
}
@@ -89,6 +91,40 @@ func TestCreateInvite_MemberForbidden(t *testing.T) {
}
}
// TestCreateInvite_ChannelAllowOverrideDoesNotGrant pins the scope boundary of
// RequirePermission: it gates on SERVER-WIDE bits, so a per-channel allow must
// never open it. The state is reachable — the admin channel-permission handler
// masks override input with permissions.AllPerms, which includes ManageInvites.
// This kills the plausible-looking "just route RequirePermission through
// Checker.HasChannelPerm" refactor, which would pass naive review because
// GetChannelPermissions returns (0, 0, nil) when no override row exists.
func TestCreateInvite_ChannelAllowOverrideDoesNotGrant(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
router := buildInviteRouter(database, limiter)
token := loginAndGetToken(t, router, database, "overrideuser", 4)
if _, err := database.ExecContext(context.Background(),
`INSERT INTO channels (id, name, type) VALUES (1, 'general', 'text')`); err != nil {
t.Fatalf("insert channel: %v", err)
}
if _, err := database.ExecContext(context.Background(),
`INSERT INTO channel_overrides (channel_id, role_id, allow, deny) VALUES (1, 4, ?, 0)`,
permissions.ManageInvites,
); err != nil {
t.Fatalf("insert channel override: %v", err)
}
rr := postJSONWithToken(t, router, "/api/v1/invites", token, map[string]any{
"max_uses": 1,
})
if rr.Code != http.StatusForbidden {
t.Errorf("CreateInvite with channel allow override status = %d, want 403", rr.Code)
}
}
func TestCreateInvite_Unlimited(t *testing.T) {
database := newAuthTestDB(t)
limiter := auth.NewRateLimiter()
@@ -141,7 +177,7 @@ func TestCreateInvite_CreateInviteFailure(t *testing.T) {
router := buildInviteRouter(database, limiter)
token := loginAndGetToken(t, router, database, "invitecreatefail", 2)
if _, err := database.Exec(`DROP TABLE invites`); err != nil {
if _, err := database.ExecContext(context.Background(), `DROP TABLE invites`); err != nil {
t.Fatalf("drop invites table: %v", err)
}
@@ -165,7 +201,7 @@ func TestCreateInvite_GetInviteFailure(t *testing.T) {
router := buildInviteRouter(database, limiter)
token := loginAndGetToken(t, router, database, "invitegetfail", 2)
if _, err := database.Exec(`
if _, err := database.ExecContext(context.Background(), `
CREATE TRIGGER delete_invite_after_insert
AFTER INSERT ON invites
BEGIN
@@ -187,7 +223,7 @@ func TestCreateInvite_GetInviteFailure(t *testing.T) {
if resp["message"] != "an internal error occurred" {
t.Errorf("message = %v, want an internal error occurred", resp["message"])
}
if _, err := database.Exec(`DROP TRIGGER delete_invite_after_insert`); err != nil {
if _, err := database.ExecContext(context.Background(), `DROP TRIGGER delete_invite_after_insert`); err != nil {
t.Fatalf("drop trigger: %v", err)
}
}
@@ -296,7 +332,7 @@ func TestRevokeInvite_Success(t *testing.T) {
}
// Verify invite is revoked.
inv, _ := database.GetInvite(code)
inv, _ := database.GetInvite(context.Background(), code)
if inv == nil || !inv.Revoked {
t.Error("Invite not revoked in database after DELETE")
}
@@ -362,7 +398,7 @@ func TestRevokeInvite_RevokeFailure(t *testing.T) {
}
code := created["code"].(string)
if _, err := database.Exec(`
if _, err := database.ExecContext(context.Background(), `
CREATE TRIGGER block_revoke_invite
BEFORE UPDATE OF revoked ON invites
BEGIN
+27 -22
View File
@@ -42,7 +42,7 @@ func AuthMiddleware(database *db.DB) func(http.Handler) http.Handler {
}
hash := auth.HashToken(token)
sess, err := database.GetSessionByTokenHash(hash)
sess, err := database.GetSessionByTokenHash(r.Context(), hash)
if err != nil || sess == nil {
writeJSON(w, http.StatusUnauthorized, errorResponse{
Error: "UNAUTHORIZED",
@@ -54,8 +54,11 @@ func AuthMiddleware(database *db.DB) func(http.Handler) http.Handler {
// Check expiry.
if auth.IsSessionExpired(sess.ExpiresAt) {
// Clean up expired session in background to prevent accumulation.
// The request ctx is cancelled as soon as the 401 below is
// written, so detach cancellation: the deletion must complete.
cleanupCtx := context.WithoutCancel(r.Context())
go func(h string) {
_ = database.DeleteSession(h)
_ = database.DeleteSession(cleanupCtx, h)
}(hash)
writeJSON(w, http.StatusUnauthorized, errorResponse{
Error: "UNAUTHORIZED",
@@ -65,7 +68,7 @@ func AuthMiddleware(database *db.DB) func(http.Handler) http.Handler {
}
// Load user.
user, err := database.GetUserByID(sess.UserID)
user, err := database.GetUserByID(r.Context(), sess.UserID)
if err != nil || user == nil {
writeJSON(w, http.StatusUnauthorized, errorResponse{
Error: "UNAUTHORIZED",
@@ -84,8 +87,11 @@ func AuthMiddleware(database *db.DB) func(http.Handler) http.Handler {
}
// Load role for permission checks.
role, err := database.GetRoleByID(user.RoleID)
if err != nil {
// A dangling role_id returns (nil, nil) from GetRoleByID, so the nil
// check is load-bearing: without it a nil role reaches the context
// and every downstream permission check has to re-guard it.
role, err := database.GetRoleByID(r.Context(), user.RoleID)
if err != nil || role == nil {
writeJSON(w, http.StatusUnauthorized, errorResponse{
Error: "UNAUTHORIZED",
Message: "role not found",
@@ -94,7 +100,7 @@ func AuthMiddleware(database *db.DB) func(http.Handler) http.Handler {
}
// Touch session in background — non-fatal if it fails.
if err := database.TouchSession(hash); err != nil {
if err := database.TouchSession(r.Context(), hash); err != nil {
slog.Warn("failed to touch session", "error", err, "user_id", user.ID)
}
@@ -106,9 +112,20 @@ func AuthMiddleware(database *db.DB) func(http.Handler) http.Handler {
}
}
// RequirePermission returns middleware that checks the authenticated user's
// role permissions. Returns 403 if the user lacks the required permission.
// The ADMINISTRATOR bit (0x40000000) bypasses all checks.
// RequirePermission returns middleware gating a route on SERVER-WIDE role
// permissions. Returns 403 if the user lacks them.
//
// Scope contract — this is the whole reason the middleware and the service
// layer look like two permission systems:
// - It consults the role bitfield only. Channel overrides are NOT applied,
// because a route reaching this middleware has no channel id to resolve
// them against, and a per-channel allow must never open a server-wide gate.
// - Anything channel-scoped belongs in the service layer behind
// permissions.Checker (via svc.Permissions), which resolves overrides.
// - ADMINISTRATOR bypasses; multi-bit masks require ALL bits.
//
// The rule itself lives in permissions.HasServerPerm so no call site can
// re-derive it.
func RequirePermission(perm int64) func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
@@ -121,13 +138,7 @@ func RequirePermission(perm int64) func(http.Handler) http.Handler {
return
}
// ADMINISTRATOR bypasses all permission checks.
if permissions.HasAdmin(role.Permissions) {
next.ServeHTTP(w, r)
return
}
if role.Permissions&perm == 0 {
if !permissions.HasServerPerm(role.Permissions, perm) {
writeJSON(w, http.StatusForbidden, errorResponse{
Error: "FORBIDDEN",
Message: "insufficient permissions",
@@ -367,12 +378,6 @@ func SecurityHeadersWithTLS(tlsMode string) func(http.Handler) http.Handler {
}
}
// SecurityHeaders is a convenience wrapper for SecurityHeadersWithTLS with TLS
// disabled (no HSTS header). Kept for backwards compatibility with tests.
func SecurityHeaders(next http.Handler) http.Handler {
return SecurityHeadersWithTLS("")(next)
}
// MaxBodySize wraps r.Body with http.MaxBytesReader so that reads beyond
// maxBytes return an error. This prevents clients from exhausting server memory
// by sending arbitrarily large request bodies.
+102 -24
View File
@@ -14,6 +14,7 @@ import (
"github.com/owncord/server/api"
"github.com/owncord/server/auth"
"github.com/owncord/server/db"
"github.com/owncord/server/permissions"
)
// ─── Helpers ─────────────────────────────────────────────────────────────────
@@ -51,10 +52,10 @@ func withBearer(req *http.Request, token string) *http.Request {
func TestAuthMiddleware_ValidToken(t *testing.T) {
database := newAPITestDB(t)
uid, _ := database.CreateUser("alice", "hash", 4)
uid, _ := database.CreateUser(context.Background(), "alice", "hash", 4)
token, _ := auth.GenerateToken()
hash := auth.HashToken(token)
_, _ = database.CreateSession(uid, hash, "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, hash, "test", "127.0.0.1")
h := api.AuthMiddleware(database)(http.HandlerFunc(ok))
req := httptest.NewRequest(http.MethodGet, "/", nil)
@@ -99,13 +100,13 @@ func TestAuthMiddleware_InvalidToken(t *testing.T) {
func TestAuthMiddleware_ExpiredSession(t *testing.T) {
database := newAPITestDB(t)
uid, _ := database.CreateUser("bob", "hash", 4)
uid, _ := database.CreateUser(context.Background(), "bob", "hash", 4)
token, _ := auth.GenerateToken()
hash := auth.HashToken(token)
// Insert an already-expired session.
pastTime := time.Now().Add(-time.Hour).UTC().Format("2006-01-02 15:04:05")
_, _ = database.Exec(
_, _ = database.ExecContext(context.Background(),
`INSERT INTO sessions (user_id, token, device, ip_address, expires_at) VALUES (?, ?, ?, ?, ?)`,
uid, hash, "test", "127.0.0.1", pastTime,
)
@@ -143,18 +144,62 @@ func TestAuthMiddleware_MalformedAuthHeader(t *testing.T) {
}
}
// TestAuthMiddleware_DanglingRoleUnauthorized pins the `role == nil` guard:
// GetRoleByID returns (nil, nil) for a role_id with no roles row, so without
// the guard a nil role reached the request context and the request only died
// later, at RequirePermission's own nil check (403) — or not at all on routes
// that have no RequirePermission.
func TestAuthMiddleware_DanglingRoleUnauthorized(t *testing.T) {
database := newAPITestDB(t)
// users.role_id has a FK to roles(id), so the dangling row can only be
// created with FK enforcement momentarily off (db.Open pins the pool to a
// single connection, so the pragma applies to the inserts that follow).
if _, err := database.ExecContext(context.Background(), `PRAGMA foreign_keys=OFF`); err != nil {
t.Fatalf("disable foreign keys: %v", err)
}
res, err := database.ExecContext(context.Background(),
`INSERT INTO users (username, password, role_id) VALUES ('dangling', '$2a$12$fake', 999)`)
if err != nil {
t.Fatalf("insert dangling user: %v", err)
}
uid, _ := res.LastInsertId()
if _, err := database.ExecContext(context.Background(), `PRAGMA foreign_keys=ON`); err != nil {
t.Fatalf("re-enable foreign keys: %v", err)
}
token, _ := auth.GenerateToken()
if _, err := database.ExecContext(context.Background(),
`INSERT INTO sessions (user_id, token, device, ip_address, expires_at)
VALUES (?, ?, 'test', '127.0.0.1', '2099-01-01T00:00:00Z')`,
uid, auth.HashToken(token),
); err != nil {
t.Fatalf("insert session: %v", err)
}
h := api.AuthMiddleware(database)(http.HandlerFunc(ok))
req := httptest.NewRequest(http.MethodGet, "/", nil)
withBearer(req, token)
rr := httptest.NewRecorder()
h.ServeHTTP(rr, req)
if rr.Code != http.StatusUnauthorized {
t.Errorf("AuthMiddleware dangling role status = %d, want 401", rr.Code)
}
}
// ─── RequirePermission tests ──────────────────────────────────────────────────
func TestRequirePermission_Allowed(t *testing.T) {
database := newAPITestDB(t)
uid, _ := database.CreateUser("carol", "hash", 4) // Member role = 0x663
uid, _ := database.CreateUser(context.Background(), "carol", "hash", 4) // Member role = 0x663
token, _ := auth.GenerateToken()
hash := auth.HashToken(token)
_, _ = database.CreateSession(uid, hash, "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, hash, "test", "127.0.0.1")
// SEND_MESSAGES = 0x1 — Member role has this bit
h := api.AuthMiddleware(database)(
api.RequirePermission(0x1)(http.HandlerFunc(ok)),
api.RequirePermission(permissions.SendMessages)(http.HandlerFunc(ok)),
)
req := httptest.NewRequest(http.MethodGet, "/", nil)
withBearer(req, token)
@@ -169,14 +214,13 @@ func TestRequirePermission_Allowed(t *testing.T) {
func TestRequirePermission_Forbidden(t *testing.T) {
database := newAPITestDB(t)
uid, _ := database.CreateUser("dave", "hash", 4) // Member role = 0x663
uid, _ := database.CreateUser(context.Background(), "dave", "hash", 4) // Member role = 0x663
token, _ := auth.GenerateToken()
hash := auth.HashToken(token)
_, _ = database.CreateSession(uid, hash, "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, hash, "test", "127.0.0.1")
// MANAGE_ROLES = 0x1000000 — Member does not have this
h := api.AuthMiddleware(database)(
api.RequirePermission(0x1000000)(http.HandlerFunc(ok)),
api.RequirePermission(permissions.ManageRoles)(http.HandlerFunc(ok)),
)
req := httptest.NewRequest(http.MethodGet, "/", nil)
withBearer(req, token)
@@ -192,14 +236,14 @@ func TestRequirePermission_Forbidden(t *testing.T) {
func TestRequirePermission_Administrator_Bypass(t *testing.T) {
database := newAPITestDB(t)
// Owner role (id=1) has permissions 0x7FFFFFFF which includes ADMINISTRATOR (0x40000000)
uid, _ := database.CreateUser("owner", "hash", 1)
uid, _ := database.CreateUser(context.Background(), "owner", "hash", 1)
token, _ := auth.GenerateToken()
hash := auth.HashToken(token)
_, _ = database.CreateSession(uid, hash, "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, hash, "test", "127.0.0.1")
// Any permission should pass for ADMINISTRATOR
h := api.AuthMiddleware(database)(
api.RequirePermission(0x1000000)(http.HandlerFunc(ok)),
api.RequirePermission(permissions.ManageRoles)(http.HandlerFunc(ok)),
)
req := httptest.NewRequest(http.MethodGet, "/", nil)
withBearer(req, token)
@@ -212,6 +256,31 @@ func TestRequirePermission_Administrator_Bypass(t *testing.T) {
}
}
// TestRequirePermission_MultiBitRequiresAllBits pins the one behaviour the
// HasServerPerm consolidation changed: a multi-bit mask is ALL-of, not any-of.
// The previous raw `role.Permissions&perm == 0` test returned 200 here because
// Member holds SendMessages, which was enough to make the mask non-zero.
func TestRequirePermission_MultiBitRequiresAllBits(t *testing.T) {
database := newAPITestDB(t)
uid, _ := database.CreateUser(context.Background(), "multibit", "hash", 4) // Member role = 1635, has SendMessages, not ManageRoles
token, _ := auth.GenerateToken()
hash := auth.HashToken(token)
_, _ = database.CreateSession(context.Background(), uid, hash, "test", "127.0.0.1")
h := api.AuthMiddleware(database)(
api.RequirePermission(permissions.SendMessages | permissions.ManageRoles)(http.HandlerFunc(ok)),
)
req := httptest.NewRequest(http.MethodGet, "/", nil)
withBearer(req, token)
rr := httptest.NewRecorder()
h.ServeHTTP(rr, req)
if rr.Code != http.StatusForbidden {
t.Errorf("RequirePermission partial multi-bit mask status = %d, want 403", rr.Code)
}
}
// ─── RateLimitMiddleware tests ────────────────────────────────────────────────
func TestRateLimitMiddleware_UnderLimit(t *testing.T) {
@@ -342,11 +411,11 @@ func TestRateLimitMiddleware_XRealIPHonouredFromTrustedProxy(t *testing.T) {
// with no expiry cannot pass the auth middleware.
func TestAuthMiddleware_BannedUserBlocked(t *testing.T) {
database := newAPITestDB(t)
uid, _ := database.CreateUser("banneduser", "hash", 4)
_ = database.BanUser(uid, "rule violation", nil) // permanent ban
uid, _ := database.CreateUser(context.Background(), "banneduser", "hash", 4)
_ = database.BanUser(context.Background(), uid, "rule violation", nil) // permanent ban
token, _ := auth.GenerateToken()
hash := auth.HashToken(token)
_, _ = database.CreateSession(uid, hash, "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, hash, "test", "127.0.0.1")
h := api.AuthMiddleware(database)(http.HandlerFunc(ok))
req := httptest.NewRequest(http.MethodGet, "/", nil)
@@ -364,15 +433,15 @@ func TestAuthMiddleware_BannedUserBlocked(t *testing.T) {
// expired in the past can pass the auth middleware.
func TestAuthMiddleware_ExpiredBanAllowed(t *testing.T) {
database := newAPITestDB(t)
uid, _ := database.CreateUser("expbanned", "hash", 4)
uid, _ := database.CreateUser(context.Background(), "expbanned", "hash", 4)
// Set ban with an expiry time in the past.
past := time.Now().UTC().Add(-time.Hour)
_ = database.BanUser(uid, "temp ban", &past)
_ = database.BanUser(context.Background(), uid, "temp ban", &past)
token, _ := auth.GenerateToken()
hash := auth.HashToken(token)
_, _ = database.CreateSession(uid, hash, "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, hash, "test", "127.0.0.1")
h := api.AuthMiddleware(database)(http.HandlerFunc(ok))
req := httptest.NewRequest(http.MethodGet, "/", nil)
@@ -390,15 +459,15 @@ func TestAuthMiddleware_ExpiredBanAllowed(t *testing.T) {
// temporary ban whose expiry is in the future is still blocked.
func TestAuthMiddleware_ActiveTemporaryBanBlocked(t *testing.T) {
database := newAPITestDB(t)
uid, _ := database.CreateUser("tempbanned", "hash", 4)
uid, _ := database.CreateUser(context.Background(), "tempbanned", "hash", 4)
// Set ban with an expiry time in the future.
future := time.Now().UTC().Add(time.Hour)
_ = database.BanUser(uid, "temp ban", &future)
_ = database.BanUser(context.Background(), uid, "temp ban", &future)
token, _ := auth.GenerateToken()
hash := auth.HashToken(token)
_, _ = database.CreateSession(uid, hash, "test", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, hash, "test", "127.0.0.1")
h := api.AuthMiddleware(database)(http.HandlerFunc(ok))
req := httptest.NewRequest(http.MethodGet, "/", nil)
@@ -940,6 +1009,15 @@ CREATE TABLE IF NOT EXISTS channels (
voice_max_video INTEGER NOT NULL DEFAULT 0
);
CREATE TABLE IF NOT EXISTS channel_overrides (
id INTEGER PRIMARY KEY AUTOINCREMENT,
channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE,
role_id INTEGER NOT NULL REFERENCES roles(id) ON DELETE CASCADE,
allow INTEGER NOT NULL DEFAULT 0,
deny INTEGER NOT NULL DEFAULT 0,
UNIQUE(channel_id, role_id)
);
CREATE TABLE IF NOT EXISTS messages (
id INTEGER PRIMARY KEY AUTOINCREMENT,
channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE,
+5 -5
View File
@@ -187,14 +187,14 @@ func handleChangePassword(svc *service.Services, limiter *auth.RateLimiter) http
failKey := fmt.Sprintf("pw_confirm_fail:%d", user.ID)
if !auth.CheckPassword(user.PasswordHash, req.OldPassword) {
if !limiter.Allow(failKey, pwConfirmFailureThreshold, pwConfirmFailureWindow) {
limiter.Lockout(lockKey, pwConfirmLockoutDuration)
limiter.Lockout(r.Context(), lockKey, pwConfirmLockoutDuration)
}
writeJSON(w, http.StatusForbidden, errorResponse{
Error: "FORBIDDEN", Message: "incorrect password",
})
return
}
limiter.Reset(failKey)
limiter.Reset(r.Context(), failKey)
// Reject same old/new password.
if req.OldPassword == req.NewPassword {
@@ -228,7 +228,7 @@ func handleChangePassword(svc *service.Services, limiter *auth.RateLimiter) http
keepSessionID = sess.ID
}
res, err := svc.Users.ChangePassword(user.ID, hash, keepSessionID)
res, err := svc.Users.ChangePassword(r.Context(), user.ID, hash, keepSessionID)
if err != nil {
// Only reachable when the password itself failed to commit.
writeServiceError(w, err)
@@ -268,7 +268,7 @@ func handleListSessions(svc *service.Services) http.HandlerFunc {
return
}
sessions, err := svc.Users.ListSessions(user.ID)
sessions, err := svc.Users.ListSessions(r.Context(), user.ID)
if err != nil {
writeServiceError(w, err)
return
@@ -308,7 +308,7 @@ func handleRevokeSession(svc *service.Services) http.HandlerFunc {
return
}
if err := svc.Users.RevokeSession(user.ID, sessionID); err != nil {
if err := svc.Users.RevokeSession(r.Context(), user.ID, sessionID); err != nil {
writeServiceError(w, err)
return
}
+13 -12
View File
@@ -2,6 +2,7 @@ package api_test
import (
"bytes"
"context"
"encoding/json"
"fmt"
"net/http"
@@ -28,7 +29,7 @@ func buildProfileRouter(database *db.DB) http.Handler {
// profileCreateToken creates a user and session, returning the raw token.
func profileCreateToken(t *testing.T, database *db.DB, username string, roleID int) string {
t.Helper()
uid, err := database.CreateUser(username, mustHash(t), roleID)
uid, err := database.CreateUser(context.Background(), username, mustHash(t), roleID)
if err != nil {
t.Fatalf("CreateUser(%s): %v", username, err)
}
@@ -37,7 +38,7 @@ func profileCreateToken(t *testing.T, database *db.DB, username string, roleID i
t.Fatalf("GenerateToken: %v", err)
}
expiresAt := time.Now().Add(24 * time.Hour).UTC().Format("2006-01-02T15:04:05Z")
_, err = database.Exec(
_, err = database.ExecContext(context.Background(),
"INSERT INTO sessions (user_id, token, device, ip_address, expires_at) VALUES (?, ?, ?, ?, ?)",
uid, auth.HashToken(token), "TestAgent", "127.0.0.1", expiresAt,
)
@@ -183,12 +184,12 @@ func TestChangePassword_RevokesOtherSessions(t *testing.T) {
// Create user with two sessions.
token1 := profileCreateToken(t, database, "pw-revoke", 4)
user, _ := database.GetUserByUsername("pw-revoke")
user, _ := database.GetUserByUsername(context.Background(), "pw-revoke")
// Create a second session for the same user.
token2, _ := auth.GenerateToken()
expiresAt := time.Now().Add(24 * time.Hour).UTC().Format("2006-01-02T15:04:05Z")
_, _ = database.Exec(
_, _ = database.ExecContext(context.Background(),
"INSERT INTO sessions (user_id, token, device, ip_address, expires_at) VALUES (?, ?, ?, ?, ?)",
user.ID, auth.HashToken(token2), "OtherDevice", "10.0.0.1", expiresAt,
)
@@ -203,13 +204,13 @@ func TestChangePassword_RevokesOtherSessions(t *testing.T) {
}
// token1 (current session) should still work.
sess1, _ := database.GetSessionByTokenHash(auth.HashToken(token1))
sess1, _ := database.GetSessionByTokenHash(context.Background(), auth.HashToken(token1))
if sess1 == nil {
t.Error("current session should survive password change")
}
// token2 (other session) should be revoked.
sess2, _ := database.GetSessionByTokenHash(auth.HashToken(token2))
sess2, _ := database.GetSessionByTokenHash(context.Background(), auth.HashToken(token2))
if sess2 != nil {
t.Error("other session should be revoked after password change")
}
@@ -314,8 +315,8 @@ func TestRevokeSession_Success(t *testing.T) {
token := profileCreateToken(t, database, "revoke", 4)
// Create a second session to revoke.
user, _ := database.GetUserByUsername("revoke")
secondSessID, _ := database.CreateSession(user.ID, auth.HashToken("second-tok"), "Firefox", "1.2.3.4")
user, _ := database.GetUserByUsername(context.Background(), "revoke")
secondSessID, _ := database.CreateSession(context.Background(), user.ID, auth.HashToken("second-tok"), "Firefox", "1.2.3.4")
rr := profileDelete(t, router, fmt.Sprintf("/api/v1/users/me/sessions/%d", secondSessID), token)
@@ -342,8 +343,8 @@ func TestRevokeSession_OtherUsersSession(t *testing.T) {
token := profileCreateToken(t, database, "revokeother", 4)
// Create another user with a session.
otherUID, _ := database.CreateUser("victim", mustHash(t), 4)
otherSessID, _ := database.CreateSession(otherUID, auth.HashToken("victim-tok"), "Safari", "9.8.7.6")
otherUID, _ := database.CreateUser(context.Background(), "victim", mustHash(t), 4)
otherSessID, _ := database.CreateSession(context.Background(), otherUID, auth.HashToken("victim-tok"), "Safari", "9.8.7.6")
rr := profileDelete(t, router, fmt.Sprintf("/api/v1/users/me/sessions/%d", otherSessID), token)
@@ -358,8 +359,8 @@ func TestRevokeSession_CurrentSession(t *testing.T) {
token := profileCreateToken(t, database, "revokeself", 4)
// Find the current session ID.
user, _ := database.GetUserByUsername("revokeself")
sessions, _ := database.ListUserSessions(user.ID)
user, _ := database.GetUserByUsername(context.Background(), "revokeself")
sessions, _ := database.ListUserSessions(context.Background(), user.ID)
if len(sessions) == 0 {
t.Fatal("expected at least 1 session")
}
+22 -17
View File
@@ -1,6 +1,7 @@
package api
import (
"context"
"encoding/json"
"errors"
"fmt"
@@ -81,7 +82,7 @@ func handleVerifyTOTP(database *db.DB, partialStore *auth.PartialAuthStore, limi
return
}
user, err := database.GetUserByID(challenge.UserID)
user, err := database.GetUserByID(r.Context(), challenge.UserID)
if err != nil || user == nil || user.TOTPSecret == nil {
writeJSON(w, http.StatusUnauthorized, errorResponse{
Error: "UNAUTHORIZED",
@@ -111,7 +112,7 @@ func handleVerifyTOTP(database *db.DB, partialStore *auth.PartialAuthStore, limi
return
}
limiter.Reset(totpRateLimitKey)
limiter.Reset(r.Context(), totpRateLimitKey)
if _, ok := partialStore.Consume(partialToken); !ok {
writeJSON(w, http.StatusUnauthorized, errorResponse{
@@ -121,7 +122,7 @@ func handleVerifyTOTP(database *db.DB, partialStore *auth.PartialAuthStore, limi
return
}
token, err := issueSession(database, user.ID, challenge.Device, challenge.IP)
token, err := issueSession(r.Context(), database, user.ID, challenge.Device, challenge.IP)
if err != nil {
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL_ERROR",
@@ -131,7 +132,7 @@ func handleVerifyTOTP(database *db.DB, partialStore *auth.PartialAuthStore, limi
}
slog.Info("totp verified", "user_id", user.ID, "ip", challenge.IP)
db.WriteAudit(database, user.ID, "totp_verified", "user", user.ID,
db.WriteAudit(context.WithoutCancel(r.Context()), database, user.ID, "totp_verified", "user", user.ID,
"two-factor verification completed from "+challenge.IP)
writeJSON(w, http.StatusOK, authSuccessResponse{
@@ -182,7 +183,7 @@ func handleEnableTOTP(pendingStore *auth.PendingTOTPStore, limiter *auth.RateLim
failKey := fmt.Sprintf("pw_confirm_fail:%d", user.ID)
if err := requirePasswordConfirmation(user, req.Password); err != nil {
if !limiter.Allow(failKey, pwConfirmFailureThreshold, pwConfirmFailureWindow) {
limiter.Lockout(lockKey, pwConfirmLockoutDuration)
limiter.Lockout(r.Context(), lockKey, pwConfirmLockoutDuration)
}
writeJSON(w, http.StatusBadRequest, errorResponse{
Error: "INVALID_INPUT",
@@ -190,7 +191,7 @@ func handleEnableTOTP(pendingStore *auth.PendingTOTPStore, limiter *auth.RateLim
})
return
}
limiter.Reset(failKey)
limiter.Reset(r.Context(), failKey)
secret, err := auth.GenerateTOTPSecret()
if err != nil {
@@ -241,7 +242,7 @@ func handleConfirmTOTP(database *db.DB, pendingStore *auth.PendingTOTPStore, use
failKey := fmt.Sprintf("pw_confirm_fail:%d", user.ID)
if err := requirePasswordConfirmation(user, req.Password); err != nil {
if !limiter.Allow(failKey, pwConfirmFailureThreshold, pwConfirmFailureWindow) {
limiter.Lockout(lockKey, pwConfirmLockoutDuration)
limiter.Lockout(r.Context(), lockKey, pwConfirmLockoutDuration)
}
writeJSON(w, http.StatusBadRequest, errorResponse{
Error: "INVALID_INPUT",
@@ -249,7 +250,7 @@ func handleConfirmTOTP(database *db.DB, pendingStore *auth.PendingTOTPStore, use
})
return
}
limiter.Reset(failKey)
limiter.Reset(r.Context(), failKey)
secret, ok := pendingStore.Lookup(user.ID)
if !ok {
@@ -278,7 +279,7 @@ func handleConfirmTOTP(database *db.DB, pendingStore *auth.PendingTOTPStore, use
return
}
if err := database.UpdateUserTOTPSecret(user.ID, &encryptedSecret); err != nil {
if err := database.UpdateUserTOTPSecret(r.Context(), user.ID, &encryptedSecret); err != nil {
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL_ERROR",
Message: "failed to enable two-factor authentication",
@@ -289,14 +290,16 @@ func handleConfirmTOTP(database *db.DB, pendingStore *auth.PendingTOTPStore, use
// BUG-108: Revoke all other sessions after 2FA state change.
if sess, ok := r.Context().Value(SessionKey).(*db.Session); ok && sess != nil {
n, _ := database.DeleteOtherSessions(user.ID, sess.ID)
// Security tail of the 2FA change: once the secret update committed,
// revoking the other sessions must not be aborted by a dead request.
n, _ := database.DeleteOtherSessions(context.WithoutCancel(r.Context()), user.ID, sess.ID)
if n > 0 {
slog.Info("revoked other sessions after totp enable", "user_id", user.ID, "revoked", n)
}
}
slog.Info("totp enabled", "user_id", user.ID)
db.WriteAudit(database, user.ID, "totp_enabled", "user", user.ID,
db.WriteAudit(context.WithoutCancel(r.Context()), database, user.ID, "totp_enabled", "user", user.ID,
"two-factor authentication enrolled")
w.WriteHeader(http.StatusNoContent)
@@ -335,7 +338,7 @@ func handleDisableTOTP(database *db.DB, pendingStore *auth.PendingTOTPStore, lim
failKey := fmt.Sprintf("pw_confirm_fail:%d", user.ID)
if err := requirePasswordConfirmation(user, req.Password); err != nil {
if !limiter.Allow(failKey, pwConfirmFailureThreshold, pwConfirmFailureWindow) {
limiter.Lockout(lockKey, pwConfirmLockoutDuration)
limiter.Lockout(r.Context(), lockKey, pwConfirmLockoutDuration)
}
writeJSON(w, http.StatusBadRequest, errorResponse{
Error: "INVALID_INPUT",
@@ -343,9 +346,9 @@ func handleDisableTOTP(database *db.DB, pendingStore *auth.PendingTOTPStore, lim
})
return
}
limiter.Reset(failKey)
limiter.Reset(r.Context(), failKey)
require2FA, err := isRequire2FAEnabled(database)
require2FA, err := isRequire2FAEnabled(r.Context(), database)
if err != nil {
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL_ERROR",
@@ -362,7 +365,7 @@ func handleDisableTOTP(database *db.DB, pendingStore *auth.PendingTOTPStore, lim
}
pendingStore.Delete(user.ID)
if err := database.UpdateUserTOTPSecret(user.ID, nil); err != nil {
if err := database.UpdateUserTOTPSecret(r.Context(), user.ID, nil); err != nil {
writeJSON(w, http.StatusInternalServerError, errorResponse{
Error: "INTERNAL_ERROR",
Message: "failed to disable two-factor authentication",
@@ -372,14 +375,16 @@ func handleDisableTOTP(database *db.DB, pendingStore *auth.PendingTOTPStore, lim
// BUG-108: Revoke all other sessions after 2FA state change.
if sess, ok := r.Context().Value(SessionKey).(*db.Session); ok && sess != nil {
n, _ := database.DeleteOtherSessions(user.ID, sess.ID)
// Security tail of the 2FA change: once the secret update committed,
// revoking the other sessions must not be aborted by a dead request.
n, _ := database.DeleteOtherSessions(context.WithoutCancel(r.Context()), user.ID, sess.ID)
if n > 0 {
slog.Info("revoked other sessions after totp disable", "user_id", user.ID, "revoked", n)
}
}
slog.Info("totp disabled", "user_id", user.ID)
db.WriteAudit(database, user.ID, "totp_disabled", "user", user.ID,
db.WriteAudit(context.WithoutCancel(r.Context()), database, user.ID, "totp_disabled", "user", user.ID,
"two-factor authentication disabled")
w.WriteHeader(http.StatusNoContent)
+11 -10
View File
@@ -2,6 +2,7 @@ package api_test
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
@@ -22,8 +23,8 @@ func TestVerifyTOTP_Success(t *testing.T) {
// Create user with TOTP enabled.
secret, _ := auth.GenerateTOTPSecret()
hash, _ := auth.HashPassword("Password1!")
uid, _ := database.CreateUser("totpuser", hash, 4)
_ = database.UpdateUserTOTPSecret(uid, &secret)
uid, _ := database.CreateUser(context.Background(), "totpuser", hash, 4)
_ = database.UpdateUserTOTPSecret(context.Background(), uid, &secret)
// Login should return requires_2fa + partial_token.
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
@@ -70,8 +71,8 @@ func TestVerifyTOTP_InvalidCode(t *testing.T) {
secret, _ := auth.GenerateTOTPSecret()
hash, _ := auth.HashPassword("Password1!")
uid, _ := database.CreateUser("totpuser2", hash, 4)
_ = database.UpdateUserTOTPSecret(uid, &secret)
uid, _ := database.CreateUser(context.Background(), "totpuser2", hash, 4)
_ = database.UpdateUserTOTPSecret(context.Background(), uid, &secret)
// Login to get partial token.
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
@@ -122,8 +123,8 @@ func TestVerifyTOTP_MalformedBody(t *testing.T) {
// Need a valid partial token to get past the token check.
secret, _ := auth.GenerateTOTPSecret()
hash, _ := auth.HashPassword("Password1!")
uid, _ := database.CreateUser("totpuser3", hash, 4)
_ = database.UpdateUserTOTPSecret(uid, &secret)
uid, _ := database.CreateUser(context.Background(), "totpuser3", hash, 4)
_ = database.UpdateUserTOTPSecret(context.Background(), uid, &secret)
rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{
"username": "totpuser3",
@@ -154,8 +155,8 @@ func TestVerifyTOTP_ReplayProtection(t *testing.T) {
secret, _ := auth.GenerateTOTPSecret()
hash, _ := auth.HashPassword("Password1!")
uid, _ := database.CreateUser("totpuser4", hash, 4)
_ = database.UpdateUserTOTPSecret(uid, &secret)
uid, _ := database.CreateUser(context.Background(), "totpuser4", hash, 4)
_ = database.UpdateUserTOTPSecret(context.Background(), uid, &secret)
code, _ := auth.GenerateTOTPCode(secret, time.Now().UTC())
@@ -271,7 +272,7 @@ func TestConfirmTOTP_Success(t *testing.T) {
}
// Verify TOTP is now stored on user.
user, _ := database.GetUserByUsername("confirmuser")
user, _ := database.GetUserByUsername(context.Background(), "confirmuser")
if user == nil {
t.Fatal("user not found after confirm")
}
@@ -392,7 +393,7 @@ func TestDisableTOTP_BlockedByServerPolicy(t *testing.T) {
token := loginAndGetToken(t, router, database, "disableuser3", 4)
// Enable require_2fa server policy.
_, _ = database.Exec(`INSERT OR REPLACE INTO settings (key, value) VALUES ('require_2fa', '1')`)
_, _ = database.ExecContext(context.Background(), `INSERT OR REPLACE INTO settings (key, value) VALUES ('require_2fa', '1')`)
rr := deleteWithToken(t, router, "/api/v1/users/me/totp", token,
map[string]string{"password": "Password1!"})
+4 -4
View File
@@ -186,7 +186,7 @@ func handleUpload(database *db.DB, store *storage.Storage, limiter *auth.RateLim
// Insert attachment record in DB (unlinked — message_id is NULL).
user, _ = r.Context().Value(UserKey).(*db.User)
safeFilename := sanitizeUploadFilename(header.Filename)
if err := database.CreateAttachment(fileID, user.ID, safeFilename, fileID, mime, writtenBytes, width, height); err != nil {
if err := database.CreateAttachment(r.Context(), fileID, user.ID, safeFilename, fileID, mime, writtenBytes, width, height); err != nil {
// Clean up stored file on DB failure.
_ = store.Delete(fileID)
slog.Error("failed to create attachment record", "error", err)
@@ -223,7 +223,7 @@ func handleServeFile(database *db.DB, store *storage.Storage, allowedOrigins []s
role, _ := r.Context().Value(RoleKey).(*db.Role)
// Look up attachment metadata with channel context.
aa, err := database.GetAttachmentWithChannel(fileID)
aa, err := database.GetAttachmentWithChannel(r.Context(), fileID)
if err != nil {
slog.Error("failed to look up attachment", "id", fileID, "error", err)
writeJSON(w, http.StatusInternalServerError, errorResponse{
@@ -269,7 +269,7 @@ func handleServeFile(database *db.DB, store *storage.Storage, allowedOrigins []s
})
return
}
ok, dmErr := database.IsDMParticipant(user.ID, *aa.ChannelID)
ok, dmErr := database.IsDMParticipant(r.Context(), user.ID, *aa.ChannelID)
if dmErr != nil || !ok {
writeJSON(w, http.StatusForbidden, errorResponse{
Error: "FORBIDDEN",
@@ -277,7 +277,7 @@ func handleServeFile(database *db.DB, store *storage.Storage, allowedOrigins []s
})
return
}
} else if user == nil || !permSvc.HasChannelPerm(user.ID, *aa.ChannelID, permissions.ReadMessages) {
} else if user == nil || !permSvc.HasChannelPerm(r.Context(), user.ID, *aa.ChannelID, permissions.ReadMessages) {
writeJSON(w, http.StatusForbidden, errorResponse{
Error: "FORBIDDEN",
Message: "you do not have access to this file",
+30 -29
View File
@@ -2,6 +2,7 @@ package api_test
import (
"bytes"
"context"
"encoding/json"
"fmt"
"image"
@@ -177,13 +178,13 @@ func buildUploadRouterWithLimiter(database *db.DB, store *storage.Storage, limit
// uploadCreateToken creates a user+session and returns the plaintext token.
func uploadCreateToken(t *testing.T, database *db.DB, username string, roleID int) string {
t.Helper()
_, err := database.CreateUser(username, "$2a$12$fake", roleID)
_, err := database.CreateUser(context.Background(), username, "$2a$12$fake", roleID)
if err != nil {
t.Fatalf("CreateUser %q: %v", username, err)
}
token := "upload-test-token-" + username
hash := auth.HashToken(token)
_, err = database.Exec(
_, err = database.ExecContext(context.Background(),
`INSERT INTO sessions (user_id, token, device, ip_address, expires_at)
SELECT id, ?, 'test', '127.0.0.1', '2099-01-01T00:00:00Z' FROM users WHERE username = ?`,
hash, username,
@@ -320,7 +321,7 @@ func TestUpload_Success_TextFile(t *testing.T) {
}
// Verify attachment record was created in DB.
att, err := database.GetAttachmentByID(resp["id"].(string))
att, err := database.GetAttachmentByID(context.Background(), resp["id"].(string))
if err != nil {
t.Fatalf("GetAttachmentByID: %v", err)
}
@@ -557,7 +558,7 @@ func TestUpload_DBCreateAttachmentFailureDeletesStoredFile(t *testing.T) {
router := buildUploadRouter(database, store, nil)
token := uploadCreateToken(t, database, "dbfailupload", 1)
if _, err := database.Exec(`DROP TABLE attachments`); err != nil {
if _, err := database.ExecContext(context.Background(), `DROP TABLE attachments`); err != nil {
t.Fatalf("drop attachments table: %v", err)
}
@@ -604,7 +605,7 @@ func TestUpload_SanitizesReservedFilenameToUnnamed(t *testing.T) {
t.Fatalf("filename = %v, want unnamed", resp["filename"])
}
att, err := database.GetAttachmentByID(resp["id"].(string))
att, err := database.GetAttachmentByID(context.Background(), resp["id"].(string))
if err != nil {
t.Fatalf("GetAttachmentByID: %v", err)
}
@@ -633,7 +634,7 @@ func TestUpload_SuccessfulUploadCreatesDBRecord(t *testing.T) {
_ = json.NewDecoder(rr.Body).Decode(&resp)
fileID := resp["id"].(string)
att, err := database.GetAttachmentByID(fileID)
att, err := database.GetAttachmentByID(context.Background(), fileID)
if err != nil {
t.Fatalf("GetAttachmentByID: %v", err)
}
@@ -1159,20 +1160,20 @@ func TestServeFile_LinkedToGuildChannel_MemberWithPerm(t *testing.T) {
fileID := resp["id"].(string)
// Create a guild channel and link the attachment via a message.
_, err := database.Exec(`INSERT INTO channels (id, name, type) VALUES (1, 'general', 'text')`)
_, err := database.ExecContext(context.Background(), `INSERT INTO channels (id, name, type) VALUES (1, 'general', 'text')`)
if err != nil {
t.Fatalf("insert channel: %v", err)
}
// Get the uploader's user ID.
var userID int64
if err := database.QueryRow(`SELECT id FROM users WHERE username = 'guildmember'`).Scan(&userID); err != nil {
if err := database.QueryRowContext(context.Background(), `SELECT id FROM users WHERE username = 'guildmember'`).Scan(&userID); err != nil {
t.Fatalf("get user id: %v", err)
}
_, err = database.Exec(`INSERT INTO messages (id, channel_id, user_id, content) VALUES (1, 1, ?, 'test')`, userID)
_, err = database.ExecContext(context.Background(), `INSERT INTO messages (id, channel_id, user_id, content) VALUES (1, 1, ?, 'test')`, userID)
if err != nil {
t.Fatalf("insert message: %v", err)
}
_, err = database.Exec(`UPDATE attachments SET message_id = 1 WHERE id = ?`, fileID)
_, err = database.ExecContext(context.Background(), `UPDATE attachments SET message_id = 1 WHERE id = ?`, fileID)
if err != nil {
t.Fatalf("link attachment: %v", err)
}
@@ -1202,24 +1203,24 @@ func TestServeFile_LinkedToGuildChannel_MemberWithoutPerm(t *testing.T) {
fileID := resp["id"].(string)
// Create channel and link.
_, err := database.Exec(`INSERT INTO channels (id, name, type) VALUES (1, 'secret', 'text')`)
_, err := database.ExecContext(context.Background(), `INSERT INTO channels (id, name, type) VALUES (1, 'secret', 'text')`)
if err != nil {
t.Fatalf("insert channel: %v", err)
}
var uploaderID int64
if err := database.QueryRow(`SELECT id FROM users WHERE username = 'guilduploader2'`).Scan(&uploaderID); err != nil {
if err := database.QueryRowContext(context.Background(), `SELECT id FROM users WHERE username = 'guilduploader2'`).Scan(&uploaderID); err != nil {
t.Fatalf("get user id: %v", err)
}
_, err = database.Exec(`INSERT INTO messages (id, channel_id, user_id, content) VALUES (1, 1, ?, 'test')`, uploaderID)
_, err = database.ExecContext(context.Background(), `INSERT INTO messages (id, channel_id, user_id, content) VALUES (1, 1, ?, 'test')`, uploaderID)
if err != nil {
t.Fatalf("insert message: %v", err)
}
_, err = database.Exec(`UPDATE attachments SET message_id = 1 WHERE id = ?`, fileID)
_, err = database.ExecContext(context.Background(), `UPDATE attachments SET message_id = 1 WHERE id = ?`, fileID)
if err != nil {
t.Fatalf("link attachment: %v", err)
}
// Deny ReadMessages (0x0002) for role 4 (Member) on channel 1.
_, err = database.Exec(`INSERT INTO channel_overrides (channel_id, role_id, allow, deny) VALUES (1, 4, 0, 2)`)
_, err = database.ExecContext(context.Background(), `INSERT INTO channel_overrides (channel_id, role_id, allow, deny) VALUES (1, 4, 0, 2)`)
if err != nil {
t.Fatalf("insert channel_override: %v", err)
}
@@ -1249,17 +1250,17 @@ func TestServeFile_LinkedToDM_ParticipantAllowed(t *testing.T) {
fileID := resp["id"].(string)
// Create DM channel, add participants, link attachment.
_, err := database.Exec(`INSERT INTO channels (id, name, type) VALUES (1, 'dm-1', 'dm')`)
_, err := database.ExecContext(context.Background(), `INSERT INTO channels (id, name, type) VALUES (1, 'dm-1', 'dm')`)
if err != nil {
t.Fatalf("insert channel: %v", err)
}
var aliceID, bobID int64
_ = database.QueryRow(`SELECT id FROM users WHERE username = 'dmalice'`).Scan(&aliceID)
_ = database.QueryRow(`SELECT id FROM users WHERE username = 'dmbob'`).Scan(&bobID)
_, _ = database.Exec(`INSERT INTO dm_participants (user_id, channel_id) VALUES (?, 1)`, aliceID)
_, _ = database.Exec(`INSERT INTO dm_participants (user_id, channel_id) VALUES (?, 1)`, bobID)
_, _ = database.Exec(`INSERT INTO messages (id, channel_id, user_id, content) VALUES (1, 1, ?, 'hi')`, aliceID)
_, _ = database.Exec(`UPDATE attachments SET message_id = 1 WHERE id = ?`, fileID)
_ = database.QueryRowContext(context.Background(), `SELECT id FROM users WHERE username = 'dmalice'`).Scan(&aliceID)
_ = database.QueryRowContext(context.Background(), `SELECT id FROM users WHERE username = 'dmbob'`).Scan(&bobID)
_, _ = database.ExecContext(context.Background(), `INSERT INTO dm_participants (user_id, channel_id) VALUES (?, 1)`, aliceID)
_, _ = database.ExecContext(context.Background(), `INSERT INTO dm_participants (user_id, channel_id) VALUES (?, 1)`, bobID)
_, _ = database.ExecContext(context.Background(), `INSERT INTO messages (id, channel_id, user_id, content) VALUES (1, 1, ?, 'hi')`, aliceID)
_, _ = database.ExecContext(context.Background(), `UPDATE attachments SET message_id = 1 WHERE id = ?`, fileID)
// DM participant can access.
rr2 := doServeFile(t, router, fileID, token1, nil)
@@ -1287,17 +1288,17 @@ func TestServeFile_LinkedToDM_NonParticipantForbidden(t *testing.T) {
fileID := resp["id"].(string)
// Create DM channel with two participants (not the outsider).
_, err := database.Exec(`INSERT INTO channels (id, name, type) VALUES (1, 'dm-1', 'dm')`)
_, err := database.ExecContext(context.Background(), `INSERT INTO channels (id, name, type) VALUES (1, 'dm-1', 'dm')`)
if err != nil {
t.Fatalf("insert channel: %v", err)
}
var ownerID, partnerID int64
_ = database.QueryRow(`SELECT id FROM users WHERE username = 'dmowner'`).Scan(&ownerID)
_ = database.QueryRow(`SELECT id FROM users WHERE username = 'dmpartner'`).Scan(&partnerID)
_, _ = database.Exec(`INSERT INTO dm_participants (user_id, channel_id) VALUES (?, 1)`, ownerID)
_, _ = database.Exec(`INSERT INTO dm_participants (user_id, channel_id) VALUES (?, 1)`, partnerID)
_, _ = database.Exec(`INSERT INTO messages (id, channel_id, user_id, content) VALUES (1, 1, ?, 'hi')`, ownerID)
_, _ = database.Exec(`UPDATE attachments SET message_id = 1 WHERE id = ?`, fileID)
_ = database.QueryRowContext(context.Background(), `SELECT id FROM users WHERE username = 'dmowner'`).Scan(&ownerID)
_ = database.QueryRowContext(context.Background(), `SELECT id FROM users WHERE username = 'dmpartner'`).Scan(&partnerID)
_, _ = database.ExecContext(context.Background(), `INSERT INTO dm_participants (user_id, channel_id) VALUES (?, 1)`, ownerID)
_, _ = database.ExecContext(context.Background(), `INSERT INTO dm_participants (user_id, channel_id) VALUES (?, 1)`, partnerID)
_, _ = database.ExecContext(context.Background(), `INSERT INTO messages (id, channel_id, user_id, content) VALUES (1, 1, ?, 'hi')`, ownerID)
_, _ = database.ExecContext(context.Background(), `UPDATE attachments SET message_id = 1 WHERE id = ?`, fileID)
// Non-participant gets 403.
rr2 := doServeFile(t, router, fileID, outsiderToken, nil)
+18 -12
View File
@@ -1,6 +1,7 @@
package auth
import (
"context"
"time"
"github.com/owncord/server/syncutil"
@@ -20,11 +21,11 @@ type lockoutEntry struct {
// When provided, lockouts survive server restarts. The interface uses only
// stdlib types to avoid circular dependencies between packages.
type LockoutPersister interface {
UpsertLockout(key string, expiresAt time.Time) error
DeleteLockout(key string) error
CleanupExpiredLockouts() error
UpsertLockout(ctx context.Context, key string, expiresAt time.Time) error
DeleteLockout(ctx context.Context, key string) error
CleanupExpiredLockouts(ctx context.Context) error
// LoadActiveLockouts returns (keys, expiresAt) slices of equal length.
LoadActiveLockouts() (keys []string, expiresAt []time.Time, err error)
LoadActiveLockouts(ctx context.Context) (keys []string, expiresAt []time.Time, err error)
}
// RateLimiter is an in-memory, thread-safe sliding-window rate limiter with
@@ -58,8 +59,9 @@ func NewPersistentRateLimiter(store LockoutPersister) *RateLimiter {
lockouts: make(map[string]*lockoutEntry),
store: store,
}
// Load surviving lockouts from the store.
if keys, expiresAt, err := store.LoadActiveLockouts(); err == nil {
// Load surviving lockouts from the store. Constructor runs at startup
// with no request in flight, so background context.
if keys, expiresAt, err := store.LoadActiveLockouts(context.Background()); err == nil {
for i, key := range keys {
rl.lockouts[key] = &lockoutEntry{expiresAt: expiresAt[i]}
}
@@ -111,14 +113,16 @@ func (r *RateLimiter) Allow(key string, limit int, window time.Duration) bool {
// Lockout prevents any requests from key for duration regardless of the
// sliding-window counter. When a LockoutStore is configured, the lockout
// is persisted so it survives server restarts.
func (r *RateLimiter) Lockout(key string, duration time.Duration) {
// is persisted so it survives server restarts. The persist write must land
// once the lockout is decided, so the caller's cancellation is detached
// (WithoutCancel) rather than aborting the write mid-request.
func (r *RateLimiter) Lockout(ctx context.Context, key string, duration time.Duration) {
r.mu.Lock()
defer r.mu.Unlock()
expiresAt := time.Now().Add(duration)
r.lockouts[key] = &lockoutEntry{expiresAt: expiresAt}
if r.store != nil {
_ = r.store.UpsertLockout(key, expiresAt)
_ = r.store.UpsertLockout(context.WithoutCancel(ctx), key, expiresAt)
}
}
@@ -170,13 +174,14 @@ func (r *RateLimiter) Check(key string, limit int, window time.Duration) bool {
}
// Reset clears all rate-limit state (timestamps and lockout) for key.
func (r *RateLimiter) Reset(key string) {
// Like Lockout, the store delete must complete once decided (WithoutCancel).
func (r *RateLimiter) Reset(ctx context.Context, key string) {
r.mu.Lock()
defer r.mu.Unlock()
delete(r.windows, key)
delete(r.lockouts, key)
if r.store != nil {
_ = r.store.DeleteLockout(key)
_ = r.store.DeleteLockout(context.WithoutCancel(ctx), key)
}
}
@@ -217,7 +222,8 @@ func (r *RateLimiter) Cleanup(maxWindow time.Duration) {
}
if r.store != nil {
_ = r.store.CleanupExpiredLockouts()
// Runs from the StartCleanup background goroutine — no request ctx.
_ = r.store.CleanupExpiredLockouts(context.Background())
}
}
+7 -6
View File
@@ -1,6 +1,7 @@
package auth_test
import (
"context"
"testing"
"time"
@@ -36,7 +37,7 @@ func TestCleanup_RemovesExpiredWindows(t *testing.T) {
func TestCleanup_RemovesExpiredLockouts(t *testing.T) {
rl := auth.NewRateLimiter()
rl.Lockout("stale-lockout", 20*time.Millisecond)
rl.Lockout(context.Background(), "stale-lockout", 20*time.Millisecond)
time.Sleep(40 * time.Millisecond)
rl.Cleanup(15 * time.Minute)
@@ -69,7 +70,7 @@ func TestCleanup_PreservesActiveWindows(t *testing.T) {
func TestCleanup_PreservesActiveLockouts(t *testing.T) {
rl := auth.NewRateLimiter()
rl.Lockout("live-lockout", time.Hour)
rl.Lockout(context.Background(), "live-lockout", time.Hour)
rl.Cleanup(15 * time.Minute)
@@ -92,14 +93,14 @@ func TestCleanup_MixedEntries(t *testing.T) {
// Stale window entry — its timestamp will be older than shortWindow.
rl.Allow("stale", 10, shortWindow)
// Stale lockout — expires in shortWindow.
rl.Lockout("stale-lock", shortWindow)
rl.Lockout(context.Background(), "stale-lock", shortWindow)
// Wait until the stale timestamps fall outside shortWindow.
time.Sleep(shortWindow + 10*time.Millisecond)
// Active entries added AFTER the sleep — their timestamps are fresh.
rl.Allow("active", 10, time.Hour)
rl.Lockout("live-lock", time.Hour)
rl.Lockout(context.Background(), "live-lock", time.Hour)
// Cleanup with shortWindow: "stale" was recorded before the cutoff, so it
// is evicted. "active" was just recorded, so it is kept.
@@ -143,8 +144,8 @@ func TestLen_AfterAllows(t *testing.T) {
// active lockout entries.
func TestLen_AfterLockouts(t *testing.T) {
rl := auth.NewRateLimiter()
rl.Lockout("x", time.Hour)
rl.Lockout("y", time.Hour)
rl.Lockout(context.Background(), "x", time.Hour)
rl.Lockout(context.Background(), "y", time.Hour)
_, locks := rl.Len()
if locks != 2 {
+9 -8
View File
@@ -1,6 +1,7 @@
package auth_test
import (
"context"
"testing"
"time"
@@ -69,7 +70,7 @@ func TestRateLimiter_DifferentKeysIndependent(t *testing.T) {
func TestRateLimiter_LockoutEnforced(t *testing.T) {
rl := auth.NewRateLimiter()
rl.Lockout("keyLock", time.Hour)
rl.Lockout(context.Background(), "keyLock", time.Hour)
if !rl.IsLockedOut("keyLock") {
t.Error("IsLockedOut() = false after Lockout(), want true")
}
@@ -77,7 +78,7 @@ func TestRateLimiter_LockoutEnforced(t *testing.T) {
func TestRateLimiter_LockoutExpires(t *testing.T) {
rl := auth.NewRateLimiter()
rl.Lockout("keyExp", 30*time.Millisecond)
rl.Lockout(context.Background(), "keyExp", 30*time.Millisecond)
time.Sleep(50 * time.Millisecond)
if rl.IsLockedOut("keyExp") {
t.Error("IsLockedOut() = true after lockout expired, want false")
@@ -95,7 +96,7 @@ func TestRateLimiter_Reset(t *testing.T) {
rl := auth.NewRateLimiter()
rl.Allow("keyR", 1, time.Second)
rl.Allow("keyR", 1, time.Second) // now blocked
rl.Reset("keyR")
rl.Reset(context.Background(), "keyR")
if !rl.Allow("keyR", 1, time.Second) {
t.Error("Allow() = false after Reset(), want true")
}
@@ -103,7 +104,7 @@ func TestRateLimiter_Reset(t *testing.T) {
func TestRateLimiter_LockoutBlocksAllow(t *testing.T) {
rl := auth.NewRateLimiter()
rl.Lockout("keyLB", time.Hour)
rl.Lockout(context.Background(), "keyLB", time.Hour)
// Even under normal limit, lockout should block
if rl.Allow("keyLB", 100, time.Second) {
t.Error("Allow() = true for locked-out key, want false")
@@ -161,7 +162,7 @@ func TestRateLimiter_Check_AtLimit(t *testing.T) {
func TestRateLimiter_Check_RespectsLockout(t *testing.T) {
rl := auth.NewRateLimiter()
rl.Lockout("checkLocked", time.Hour)
rl.Lockout(context.Background(), "checkLocked", time.Hour)
if rl.Check("checkLocked", 100, time.Second) {
t.Error("Check() = true for locked-out key, want false")
}
@@ -169,7 +170,7 @@ func TestRateLimiter_Check_RespectsLockout(t *testing.T) {
func TestRateLimiter_Check_LockoutExpired(t *testing.T) {
rl := auth.NewRateLimiter()
rl.Lockout("checkExpLock", 10*time.Millisecond)
rl.Lockout(context.Background(), "checkExpLock", 10*time.Millisecond)
time.Sleep(30 * time.Millisecond)
if !rl.Check("checkExpLock", 5, time.Second) {
t.Error("Check() = false after lockout expired, want true")
@@ -221,11 +222,11 @@ func TestRateLimiter_ConcurrentHammering(t *testing.T) {
func TestRateLimiter_ResetClearsLockout(t *testing.T) {
rl := auth.NewRateLimiter()
rl.Lockout("resetLock", time.Hour)
rl.Lockout(context.Background(), "resetLock", time.Hour)
if !rl.IsLockedOut("resetLock") {
t.Fatal("precondition: key should be locked out")
}
rl.Reset("resetLock")
rl.Reset(context.Background(), "resetLock")
if rl.IsLockedOut("resetLock") {
t.Error("Reset() should clear lockout, but key is still locked out")
}
+10 -10
View File
@@ -83,7 +83,7 @@ func TestDeleteAccount_AnonymisesUsername(t *testing.T) {
t.Fatalf("DeleteAccount: %v", err)
}
user, err := database.GetUserByID(userID)
user, err := database.GetUserByID(context.Background(), userID)
if err != nil {
t.Fatalf("GetUserByID after delete: %v", err)
}
@@ -100,7 +100,7 @@ func TestDeleteAccount_ClearsPassword(t *testing.T) {
database.DeleteAccount(context.Background(), userID) //nolint:errcheck
user, _ := database.GetUserByID(userID)
user, _ := database.GetUserByID(context.Background(), userID)
if user.PasswordHash != "" {
t.Errorf("PasswordHash = %q, want empty", user.PasswordHash)
}
@@ -111,11 +111,11 @@ func TestDeleteAccount_ClearsAvatarAndTOTP(t *testing.T) {
userID := seedUser(t, database, "charlie")
// Set avatar and TOTP before deletion.
database.Exec("UPDATE users SET avatar = 'pic.png', totp_secret = 'SECRET' WHERE id = ?", userID) //nolint:errcheck
database.ExecContext(context.Background(), "UPDATE users SET avatar = 'pic.png', totp_secret = 'SECRET' WHERE id = ?", userID) //nolint:errcheck
database.DeleteAccount(context.Background(), userID) //nolint:errcheck
user, _ := database.GetUserByID(userID)
user, _ := database.GetUserByID(context.Background(), userID)
if user.Avatar != nil {
t.Errorf("Avatar = %v, want nil", user.Avatar)
}
@@ -130,7 +130,7 @@ func TestDeleteAccount_SetsBannedAndOffline(t *testing.T) {
database.DeleteAccount(context.Background(), userID) //nolint:errcheck
user, _ := database.GetUserByID(userID)
user, _ := database.GetUserByID(context.Background(), userID)
if !user.Banned {
t.Error("Banned should be true after deletion")
}
@@ -146,7 +146,7 @@ func TestDeleteAccount_DeletesSessions(t *testing.T) {
userID := seedUser(t, database, "eve")
// Insert a session directly.
database.Exec(
database.ExecContext(context.Background(),
"INSERT INTO sessions (user_id, token, expires_at) VALUES (?, 'tok123', datetime('now', '+1 day'))",
userID,
) //nolint:errcheck
@@ -154,7 +154,7 @@ func TestDeleteAccount_DeletesSessions(t *testing.T) {
database.DeleteAccount(context.Background(), userID) //nolint:errcheck
var count int
database.QueryRow("SELECT COUNT(*) FROM sessions WHERE user_id = ?", userID).Scan(&count) //nolint:errcheck
database.QueryRowContext(context.Background(), "SELECT COUNT(*) FROM sessions WHERE user_id = ?", userID).Scan(&count) //nolint:errcheck
if count != 0 {
t.Errorf("sessions count = %d, want 0", count)
}
@@ -165,11 +165,11 @@ func TestDeleteAccount_SoftDeletesMessages(t *testing.T) {
userID := seedUser(t, database, "frank")
chID := seedChannel(t, database, "general")
msgID, _ := database.CreateMessage(chID, userID, "hello world", nil)
msgID, _ := database.CreateMessage(context.Background(), chID, userID, "hello world", nil)
database.DeleteAccount(context.Background(), userID) //nolint:errcheck
msg, err := database.GetMessage(msgID)
msg, err := database.GetMessage(context.Background(), msgID)
if err != nil {
t.Fatalf("GetMessage after delete: %v", err)
}
@@ -194,7 +194,7 @@ func TestDeleteAccount_NonexistentUser(t *testing.T) {
func setRole(t *testing.T, database *db.DB, userID, roleID int64) {
t.Helper()
if _, err := database.Exec("UPDATE users SET role_id = ? WHERE id = ?", roleID, userID); err != nil {
if _, err := database.ExecContext(context.Background(), "UPDATE users SET role_id = ? WHERE id = ?", roleID, userID); err != nil {
t.Fatalf("setRole(%d, %d): %v", userID, roleID, err)
}
}
+40 -39
View File
@@ -1,6 +1,7 @@
package db
import (
"context"
"database/sql"
"errors"
"fmt"
@@ -13,8 +14,8 @@ import (
// ─── Setup ───────────────────────────────────────────────────────────────────
// UserCount returns the total number of registered users.
func (d *DB) UserCount() (int64, error) {
count, err := d.q.UserCount(dbCtx())
func (d *DB) UserCount(ctx context.Context) (int64, error) {
count, err := d.q.UserCount(ctx)
if err != nil {
return 0, fmt.Errorf("UserCount: %w", err)
}
@@ -26,20 +27,20 @@ func (d *DB) UserCount() (int64, error) {
// GetServerStats returns aggregate counts for the admin dashboard.
// DBSizeBytes is 0 for in-memory databases (page_count * page_size returns
// a meaningful value only for file-backed databases).
func (d *DB) GetServerStats() (*ServerStats, error) {
func (d *DB) GetServerStats(ctx context.Context) (*ServerStats, error) {
stats := &ServerStats{}
var err error
if stats.UserCount, err = d.q.CountUsers(dbCtx()); err != nil {
if stats.UserCount, err = d.q.CountUsers(ctx); err != nil {
return nil, fmt.Errorf("GetServerStats users: %w", err)
}
if stats.MessageCount, err = d.q.CountActiveMessages(dbCtx()); err != nil {
if stats.MessageCount, err = d.q.CountActiveMessages(ctx); err != nil {
return nil, fmt.Errorf("GetServerStats messages: %w", err)
}
if stats.ChannelCount, err = d.q.CountChannels(dbCtx()); err != nil {
if stats.ChannelCount, err = d.q.CountChannels(ctx); err != nil {
return nil, fmt.Errorf("GetServerStats channels: %w", err)
}
if stats.InviteCount, err = d.q.CountActiveInvites(dbCtx()); err != nil {
if stats.InviteCount, err = d.q.CountActiveInvites(ctx); err != nil {
return nil, fmt.Errorf("GetServerStats invites: %w", err)
}
@@ -47,10 +48,10 @@ func (d *DB) GetServerStats() (*ServerStats, error) {
// expressible as sqlc queries, so they stay on the raw connection.
// For :memory: databases this still works (returns the in-memory size).
var pageCount, pageSize int64
if err := d.sqlDB.QueryRow(`PRAGMA page_count`).Scan(&pageCount); err != nil {
if err := d.sqlDB.QueryRowContext(ctx, `PRAGMA page_count`).Scan(&pageCount); err != nil {
return nil, fmt.Errorf("GetServerStats page_count: %w", err)
}
if err := d.sqlDB.QueryRow(`PRAGMA page_size`).Scan(&pageSize); err != nil {
if err := d.sqlDB.QueryRowContext(ctx, `PRAGMA page_size`).Scan(&pageSize); err != nil {
return nil, fmt.Errorf("GetServerStats page_size: %w", err)
}
stats.DBSizeBytes = pageCount * pageSize
@@ -62,8 +63,8 @@ func (d *DB) GetServerStats() (*ServerStats, error) {
// ListAllUsers returns users joined with their role name, ordered by ID.
// limit=0 returns no rows.
func (d *DB) ListAllUsers(limit, offset int) ([]UserWithRole, error) {
rows, err := d.q.ListAllUsers(dbCtx(), dbgen.ListAllUsersParams{
func (d *DB) ListAllUsers(ctx context.Context, limit, offset int) ([]UserWithRole, error) {
rows, err := d.q.ListAllUsers(ctx, dbgen.ListAllUsersParams{
Limit: int64(limit),
Offset: int64(offset),
})
@@ -92,8 +93,8 @@ func (d *DB) ListAllUsers(limit, offset int) ([]UserWithRole, error) {
}
// UpdateUserRole changes the role_id of a user.
func (d *DB) UpdateUserRole(userID, roleID int64) error {
if err := d.q.UpdateUserRole(dbCtx(), dbgen.UpdateUserRoleParams{
func (d *DB) UpdateUserRole(ctx context.Context, userID, roleID int64) error {
if err := d.q.UpdateUserRole(ctx, dbgen.UpdateUserRoleParams{
RoleID: roleID,
ID: userID,
}); err != nil {
@@ -103,16 +104,16 @@ func (d *DB) UpdateUserRole(userID, roleID int64) error {
}
// ForceLogoutUser deletes all sessions for the given user ID.
func (d *DB) ForceLogoutUser(userID int64) error {
if err := d.q.ForceLogoutUser(dbCtx(), userID); err != nil {
func (d *DB) ForceLogoutUser(ctx context.Context, userID int64) error {
if err := d.q.ForceLogoutUser(ctx, userID); err != nil {
return fmt.Errorf("ForceLogoutUser: %w", err)
}
return nil
}
// GetUserSessions returns all active sessions for the given user ID.
func (d *DB) GetUserSessions(userID int64) ([]Session, error) {
rows, err := d.q.GetUserSessions(dbCtx(), userID)
func (d *DB) GetUserSessions(ctx context.Context, userID int64) ([]Session, error) {
rows, err := d.q.GetUserSessions(ctx, userID)
if err != nil {
return nil, fmt.Errorf("GetUserSessions: %w", err)
}
@@ -127,8 +128,8 @@ func (d *DB) GetUserSessions(userID int64) ([]Session, error) {
// AdminCreateChannel creates a channel with full field control including position.
// No sqlc query covers this exact INSERT shape, so it stays on raw SQL.
func (d *DB) AdminCreateChannel(name, chanType, category, topic string, position int) (int64, error) {
res, err := d.sqlDB.Exec(
func (d *DB) AdminCreateChannel(ctx context.Context, name, chanType, category, topic string, position int) (int64, error) {
res, err := d.sqlDB.ExecContext(ctx,
`INSERT INTO channels (name, type, category, topic, position)
VALUES (?, ?, ?, ?, ?)`,
name, chanType, strToNullPtr(category), strToNullPtr(topic), position,
@@ -140,8 +141,8 @@ func (d *DB) AdminCreateChannel(name, chanType, category, topic string, position
}
// AdminUpdateChannel updates all mutable channel fields.
func (d *DB) AdminUpdateChannel(id int64, name, topic string, slowMode, position int, archived bool) error {
if err := d.q.AdminUpdateChannel(dbCtx(), dbgen.AdminUpdateChannelParams{
func (d *DB) AdminUpdateChannel(ctx context.Context, id int64, name, topic string, slowMode, position int, archived bool) error {
if err := d.q.AdminUpdateChannel(ctx, dbgen.AdminUpdateChannelParams{
Name: name,
Topic: strToNullPtr(topic),
SlowMode: int64(slowMode),
@@ -155,8 +156,8 @@ func (d *DB) AdminUpdateChannel(id int64, name, topic string, slowMode, position
}
// AdminDeleteChannel removes a channel by ID (cascades to messages, etc.).
func (d *DB) AdminDeleteChannel(id int64) error {
if err := d.q.DeleteChannel(dbCtx(), id); err != nil {
func (d *DB) AdminDeleteChannel(ctx context.Context, id int64) error {
if err := d.q.DeleteChannel(ctx, id); err != nil {
return fmt.Errorf("AdminDeleteChannel: %w", err)
}
return nil
@@ -165,8 +166,8 @@ func (d *DB) AdminDeleteChannel(id int64) error {
// ─── Audit Log ────────────────────────────────────────────────────────────────
// LogAudit inserts an audit log entry.
func (d *DB) LogAudit(actorID int64, action, targetType string, targetID int64, detail string) error {
if err := d.q.LogAudit(dbCtx(), dbgen.LogAuditParams{
func (d *DB) LogAudit(ctx context.Context, actorID int64, action, targetType string, targetID int64, detail string) error {
if err := d.q.LogAudit(ctx, dbgen.LogAuditParams{
ActorID: actorID,
Action: action,
TargetType: targetType,
@@ -179,8 +180,8 @@ func (d *DB) LogAudit(actorID int64, action, targetType string, targetID int64,
}
// GetAuditLog returns audit log entries ordered newest-first with pagination.
func (d *DB) GetAuditLog(limit, offset int) ([]AuditEntry, error) {
rows, err := d.q.GetAuditLog(dbCtx(), dbgen.GetAuditLogParams{
func (d *DB) GetAuditLog(ctx context.Context, limit, offset int) ([]AuditEntry, error) {
rows, err := d.q.GetAuditLog(ctx, dbgen.GetAuditLogParams{
Limit: int64(limit),
Offset: int64(offset),
})
@@ -207,8 +208,8 @@ func (d *DB) GetAuditLog(limit, offset int) ([]AuditEntry, error) {
// GetSetting returns the value for the given settings key.
// Returns an error (wrapping sql.ErrNoRows) when the key does not exist.
func (d *DB) GetSetting(key string) (string, error) {
value, err := d.q.GetSetting(dbCtx(), key)
func (d *DB) GetSetting(ctx context.Context, key string) (string, error) {
value, err := d.q.GetSetting(ctx, key)
if errors.Is(err, sql.ErrNoRows) {
return "", fmt.Errorf("GetSetting: key %q: %w", key, ErrNotFound)
}
@@ -219,8 +220,8 @@ func (d *DB) GetSetting(key string) (string, error) {
}
// SetSetting upserts a setting value for the given key.
func (d *DB) SetSetting(key, value string) error {
if err := d.q.SetSetting(dbCtx(), dbgen.SetSettingParams{
func (d *DB) SetSetting(ctx context.Context, key, value string) error {
if err := d.q.SetSetting(ctx, dbgen.SetSettingParams{
Key: key,
Value: value,
}); err != nil {
@@ -230,8 +231,8 @@ func (d *DB) SetSetting(key, value string) error {
}
// GetAllSettings returns all settings as a key→value map.
func (d *DB) GetAllSettings() (map[string]string, error) {
rows, err := d.q.GetAllSettings(dbCtx())
func (d *DB) GetAllSettings(ctx context.Context) (map[string]string, error) {
rows, err := d.q.GetAllSettings(ctx)
if err != nil {
return nil, fmt.Errorf("GetAllSettings: %w", err)
}
@@ -244,8 +245,8 @@ func (d *DB) GetAllSettings() (map[string]string, error) {
// CountUsersWithoutTOTP returns the number of non-banned users that do not
// currently have a confirmed TOTP secret.
func (d *DB) CountUsersWithoutTOTP() (int, error) {
count, err := d.q.CountUsersWithoutTOTP(dbCtx())
func (d *DB) CountUsersWithoutTOTP(ctx context.Context) (int, error) {
count, err := d.q.CountUsersWithoutTOTP(ctx)
if err != nil {
return 0, fmt.Errorf("CountUsersWithoutTOTP: %w", err)
}
@@ -266,13 +267,13 @@ func (d *DB) CountUsersWithoutTOTP() (int, error) {
//
// The caller in handleBackup constructs the path from a hardcoded directory
// and a timestamp — no user input reaches this function.
func (d *DB) BackupTo(path string) error {
return d.BackupToSafe(path, filepath.Join("data", "backups"))
func (d *DB) BackupTo(ctx context.Context, path string) error {
return d.BackupToSafe(ctx, path, filepath.Join("data", "backups"))
}
// BackupToSafe is the internal implementation that accepts an explicit safe
// root directory. Exported for testing with isolated directories.
func (d *DB) BackupToSafe(path, safeRoot string) error {
func (d *DB) BackupToSafe(ctx context.Context, path, safeRoot string) error {
clean := filepath.Clean(path)
absRoot, err := filepath.Abs(safeRoot)
@@ -310,7 +311,7 @@ func (d *DB) BackupToSafe(path, safeRoot string) error {
return fmt.Errorf("BackupToSafe: path contains forbidden sequence %q", "--")
}
_, err = d.sqlDB.Exec(fmt.Sprintf("VACUUM INTO '%s'", absClean))
_, err = d.sqlDB.ExecContext(ctx, fmt.Sprintf("VACUUM INTO '%s'", absClean))
if err != nil {
return fmt.Errorf("BackupToSafe: %w", err)
}
+77 -76
View File
@@ -1,6 +1,7 @@
package db_test
import (
"context"
"fmt"
"os"
"path/filepath"
@@ -87,7 +88,7 @@ func newAdminTestDB(t *testing.T) *db.DB {
func TestGetServerStats_EmptyDB(t *testing.T) {
database := newAdminTestDB(t)
stats, err := database.GetServerStats()
stats, err := database.GetServerStats(context.Background())
if err != nil {
t.Fatalf("GetServerStats() error: %v", err)
}
@@ -114,17 +115,17 @@ func TestGetServerStats_EmptyDB(t *testing.T) {
func TestGetServerStats_WithData(t *testing.T) {
database := newAdminTestDB(t)
_, err := database.CreateUser("statuser", "hash", 4)
_, err := database.CreateUser(context.Background(), "statuser", "hash", 4)
if err != nil {
t.Fatalf("CreateUser error: %v", err)
}
_, err = database.CreateChannel("general", "text", "", "", 0)
_, err = database.CreateChannel(context.Background(), "general", "text", "", "", 0)
if err != nil {
t.Fatalf("CreateChannel error: %v", err)
}
stats, err := database.GetServerStats()
stats, err := database.GetServerStats(context.Background())
if err != nil {
t.Fatalf("GetServerStats() error: %v", err)
}
@@ -141,7 +142,7 @@ func TestGetServerStats_WithData(t *testing.T) {
func TestListAllUsers_Empty(t *testing.T) {
database := newAdminTestDB(t)
users, err := database.ListAllUsers(50, 0)
users, err := database.ListAllUsers(context.Background(), 50, 0)
if err != nil {
t.Fatalf("ListAllUsers() error: %v", err)
}
@@ -153,12 +154,12 @@ func TestListAllUsers_Empty(t *testing.T) {
func TestListAllUsers_WithRoleName(t *testing.T) {
database := newAdminTestDB(t)
_, err := database.CreateUser("alice", "hash", 4)
_, err := database.CreateUser(context.Background(), "alice", "hash", 4)
if err != nil {
t.Fatalf("CreateUser error: %v", err)
}
users, err := database.ListAllUsers(50, 0)
users, err := database.ListAllUsers(context.Background(), 50, 0)
if err != nil {
t.Fatalf("ListAllUsers() error: %v", err)
}
@@ -178,7 +179,7 @@ func TestListAllUsers_Pagination(t *testing.T) {
database := newAdminTestDB(t)
for i := range 5 {
_, err := database.CreateUser(
_, err := database.CreateUser(context.Background(),
strings.Repeat("u", i+1),
"hash",
4,
@@ -188,7 +189,7 @@ func TestListAllUsers_Pagination(t *testing.T) {
}
}
page1, err := database.ListAllUsers(3, 0)
page1, err := database.ListAllUsers(context.Background(), 3, 0)
if err != nil {
t.Fatalf("ListAllUsers page1 error: %v", err)
}
@@ -196,7 +197,7 @@ func TestListAllUsers_Pagination(t *testing.T) {
t.Errorf("page1 len = %d, want 3", len(page1))
}
page2, err := database.ListAllUsers(3, 3)
page2, err := database.ListAllUsers(context.Background(), 3, 3)
if err != nil {
t.Fatalf("ListAllUsers page2 error: %v", err)
}
@@ -207,9 +208,9 @@ func TestListAllUsers_Pagination(t *testing.T) {
func TestListAllUsers_ZeroLimit(t *testing.T) {
database := newAdminTestDB(t)
_, _ = database.CreateUser("zerotest", "hash", 4)
_, _ = database.CreateUser(context.Background(), "zerotest", "hash", 4)
users, err := database.ListAllUsers(0, 0)
users, err := database.ListAllUsers(context.Background(), 0, 0)
if err != nil {
t.Fatalf("ListAllUsers(0, 0) error: %v", err)
}
@@ -224,16 +225,16 @@ func TestListAllUsers_ZeroLimit(t *testing.T) {
func TestUpdateUserRole(t *testing.T) {
database := newAdminTestDB(t)
uid, err := database.CreateUser("roleuser", "hash", 4)
uid, err := database.CreateUser(context.Background(), "roleuser", "hash", 4)
if err != nil {
t.Fatalf("CreateUser error: %v", err)
}
if err := database.UpdateUserRole(uid, 2); err != nil {
if err := database.UpdateUserRole(context.Background(), uid, 2); err != nil {
t.Fatalf("UpdateUserRole() error: %v", err)
}
user, err := database.GetUserByID(uid)
user, err := database.GetUserByID(context.Background(), uid)
if err != nil {
t.Fatalf("GetUserByID error: %v", err)
}
@@ -246,7 +247,7 @@ func TestUpdateUserRole_NonexistentUser(t *testing.T) {
database := newAdminTestDB(t)
// UPDATE with no matching rows is not an error
err := database.UpdateUserRole(99999, 2)
err := database.UpdateUserRole(context.Background(), 99999, 2)
if err != nil {
t.Errorf("UpdateUserRole() for nonexistent user returned unexpected error: %v", err)
}
@@ -257,15 +258,15 @@ func TestUpdateUserRole_NonexistentUser(t *testing.T) {
func TestForceLogoutUser_DeletesSessions(t *testing.T) {
database := newAdminTestDB(t)
uid, err := database.CreateUser("logoutuser", "hash", 4)
uid, err := database.CreateUser(context.Background(), "logoutuser", "hash", 4)
if err != nil {
t.Fatalf("CreateUser error: %v", err)
}
_, _ = database.CreateSession(uid, "token1hash", "device1", "127.0.0.1")
_, _ = database.CreateSession(uid, "token2hash", "device2", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, "token1hash", "device1", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, "token2hash", "device2", "127.0.0.1")
sessions, err := database.GetUserSessions(uid)
sessions, err := database.GetUserSessions(context.Background(), uid)
if err != nil {
t.Fatalf("GetUserSessions error: %v", err)
}
@@ -273,11 +274,11 @@ func TestForceLogoutUser_DeletesSessions(t *testing.T) {
t.Fatalf("expected 2 sessions before logout, got %d", len(sessions))
}
if err := database.ForceLogoutUser(uid); err != nil {
if err := database.ForceLogoutUser(context.Background(), uid); err != nil {
t.Fatalf("ForceLogoutUser() error: %v", err)
}
sessions, err = database.GetUserSessions(uid)
sessions, err = database.GetUserSessions(context.Background(), uid)
if err != nil {
t.Fatalf("GetUserSessions after logout error: %v", err)
}
@@ -289,12 +290,12 @@ func TestForceLogoutUser_DeletesSessions(t *testing.T) {
func TestForceLogoutUser_NoSessions(t *testing.T) {
database := newAdminTestDB(t)
uid, err := database.CreateUser("nosessions", "hash", 4)
uid, err := database.CreateUser(context.Background(), "nosessions", "hash", 4)
if err != nil {
t.Fatalf("CreateUser error: %v", err)
}
if err := database.ForceLogoutUser(uid); err != nil {
if err := database.ForceLogoutUser(context.Background(), uid); err != nil {
t.Errorf("ForceLogoutUser() on user with no sessions returned error: %v", err)
}
}
@@ -304,12 +305,12 @@ func TestForceLogoutUser_NoSessions(t *testing.T) {
func TestGetUserSessions_Empty(t *testing.T) {
database := newAdminTestDB(t)
uid, err := database.CreateUser("sessionuser", "hash", 4)
uid, err := database.CreateUser(context.Background(), "sessionuser", "hash", 4)
if err != nil {
t.Fatalf("CreateUser error: %v", err)
}
sessions, err := database.GetUserSessions(uid)
sessions, err := database.GetUserSessions(context.Background(), uid)
if err != nil {
t.Fatalf("GetUserSessions() error: %v", err)
}
@@ -321,14 +322,14 @@ func TestGetUserSessions_Empty(t *testing.T) {
func TestGetUserSessions_IsolatedByUser(t *testing.T) {
database := newAdminTestDB(t)
uid1, _ := database.CreateUser("user1sess", "hash", 4)
uid2, _ := database.CreateUser("user2sess", "hash", 4)
uid1, _ := database.CreateUser(context.Background(), "user1sess", "hash", 4)
uid2, _ := database.CreateUser(context.Background(), "user2sess", "hash", 4)
_, _ = database.CreateSession(uid1, "u1t1", "web", "1.2.3.4")
_, _ = database.CreateSession(uid1, "u1t2", "mobile", "1.2.3.5")
_, _ = database.CreateSession(uid2, "u2t1", "web", "1.2.3.6")
_, _ = database.CreateSession(context.Background(), uid1, "u1t1", "web", "1.2.3.4")
_, _ = database.CreateSession(context.Background(), uid1, "u1t2", "mobile", "1.2.3.5")
_, _ = database.CreateSession(context.Background(), uid2, "u2t1", "web", "1.2.3.6")
sessions, err := database.GetUserSessions(uid1)
sessions, err := database.GetUserSessions(context.Background(), uid1)
if err != nil {
t.Fatalf("GetUserSessions() error: %v", err)
}
@@ -347,7 +348,7 @@ func TestGetUserSessions_IsolatedByUser(t *testing.T) {
func TestAdminCreateChannel(t *testing.T) {
database := newAdminTestDB(t)
id, err := database.AdminCreateChannel("announce", "text", "General", "Announcements", 1)
id, err := database.AdminCreateChannel(context.Background(), "announce", "text", "General", "Announcements", 1)
if err != nil {
t.Fatalf("AdminCreateChannel() error: %v", err)
}
@@ -355,7 +356,7 @@ func TestAdminCreateChannel(t *testing.T) {
t.Errorf("AdminCreateChannel() id = %d, want > 0", id)
}
ch, err := database.GetChannel(id)
ch, err := database.GetChannel(context.Background(), id)
if err != nil {
t.Fatalf("GetChannel() error: %v", err)
}
@@ -382,12 +383,12 @@ func TestAdminCreateChannel(t *testing.T) {
func TestAdminCreateChannel_EmptyOptionals(t *testing.T) {
database := newAdminTestDB(t)
id, err := database.AdminCreateChannel("simple", "voice", "", "", 0)
id, err := database.AdminCreateChannel(context.Background(), "simple", "voice", "", "", 0)
if err != nil {
t.Fatalf("AdminCreateChannel() error: %v", err)
}
ch, err := database.GetChannel(id)
ch, err := database.GetChannel(context.Background(), id)
if err != nil {
t.Fatalf("GetChannel() error: %v", err)
}
@@ -404,16 +405,16 @@ func TestAdminCreateChannel_EmptyOptionals(t *testing.T) {
func TestAdminUpdateChannel(t *testing.T) {
database := newAdminTestDB(t)
id, err := database.AdminCreateChannel("old-name", "text", "", "", 0)
id, err := database.AdminCreateChannel(context.Background(), "old-name", "text", "", "", 0)
if err != nil {
t.Fatalf("AdminCreateChannel() error: %v", err)
}
if err := database.AdminUpdateChannel(id, "new-name", "new topic", 5, 2, true); err != nil {
if err := database.AdminUpdateChannel(context.Background(), id, "new-name", "new topic", 5, 2, true); err != nil {
t.Fatalf("AdminUpdateChannel() error: %v", err)
}
ch, err := database.GetChannel(id)
ch, err := database.GetChannel(context.Background(), id)
if err != nil {
t.Fatalf("GetChannel() error: %v", err)
}
@@ -437,17 +438,17 @@ func TestAdminUpdateChannel(t *testing.T) {
func TestAdminUpdateChannel_Unarchive(t *testing.T) {
database := newAdminTestDB(t)
id, _ := database.AdminCreateChannel("arch-ch", "text", "", "", 0)
_ = database.AdminUpdateChannel(id, "arch-ch", "", 0, 0, true)
id, _ := database.AdminCreateChannel(context.Background(), "arch-ch", "text", "", "", 0)
_ = database.AdminUpdateChannel(context.Background(), id, "arch-ch", "", 0, 0, true)
ch, _ := database.GetChannel(id)
ch, _ := database.GetChannel(context.Background(), id)
if !ch.Archived {
t.Fatal("channel should be archived")
}
// Unarchive
_ = database.AdminUpdateChannel(id, "arch-ch", "", 0, 0, false)
ch, _ = database.GetChannel(id)
_ = database.AdminUpdateChannel(context.Background(), id, "arch-ch", "", 0, 0, false)
ch, _ = database.GetChannel(context.Background(), id)
if ch.Archived {
t.Error("Archived = true after unarchiving, want false")
}
@@ -458,16 +459,16 @@ func TestAdminUpdateChannel_Unarchive(t *testing.T) {
func TestAdminDeleteChannel(t *testing.T) {
database := newAdminTestDB(t)
id, err := database.AdminCreateChannel("to-delete", "text", "", "", 0)
id, err := database.AdminCreateChannel(context.Background(), "to-delete", "text", "", "", 0)
if err != nil {
t.Fatalf("AdminCreateChannel() error: %v", err)
}
if err := database.AdminDeleteChannel(id); err != nil {
if err := database.AdminDeleteChannel(context.Background(), id); err != nil {
t.Fatalf("AdminDeleteChannel() error: %v", err)
}
ch, err := database.GetChannel(id)
ch, err := database.GetChannel(context.Background(), id)
if err != nil {
t.Fatalf("GetChannel() after delete error: %v", err)
}
@@ -480,7 +481,7 @@ func TestAdminDeleteChannel_NonExistent(t *testing.T) {
database := newAdminTestDB(t)
// Deleting nonexistent channel should not error
if err := database.AdminDeleteChannel(99999); err != nil {
if err := database.AdminDeleteChannel(context.Background(), 99999); err != nil {
t.Errorf("AdminDeleteChannel(nonexistent) error: %v", err)
}
}
@@ -490,16 +491,16 @@ func TestAdminDeleteChannel_NonExistent(t *testing.T) {
func TestLogAudit_AndRetrieve(t *testing.T) {
database := newAdminTestDB(t)
uid, err := database.CreateUser("auditor", "hash", 1)
uid, err := database.CreateUser(context.Background(), "auditor", "hash", 1)
if err != nil {
t.Fatalf("CreateUser error: %v", err)
}
if err := database.LogAudit(uid, "USER_BANNED", "user", 42, "banned for spam"); err != nil {
if err := database.LogAudit(context.Background(), uid, "USER_BANNED", "user", 42, "banned for spam"); err != nil {
t.Fatalf("LogAudit() error: %v", err)
}
entries, err := database.GetAuditLog(10, 0)
entries, err := database.GetAuditLog(context.Background(), 10, 0)
if err != nil {
t.Fatalf("GetAuditLog() error: %v", err)
}
@@ -534,7 +535,7 @@ func TestLogAudit_AndRetrieve(t *testing.T) {
func TestGetAuditLog_Empty(t *testing.T) {
database := newAdminTestDB(t)
entries, err := database.GetAuditLog(10, 0)
entries, err := database.GetAuditLog(context.Background(), 10, 0)
if err != nil {
t.Fatalf("GetAuditLog() error: %v", err)
}
@@ -546,12 +547,12 @@ func TestGetAuditLog_Empty(t *testing.T) {
func TestGetAuditLog_Pagination(t *testing.T) {
database := newAdminTestDB(t)
uid, _ := database.CreateUser("auditpager", "hash", 1)
uid, _ := database.CreateUser(context.Background(), "auditpager", "hash", 1)
for i := range 5 {
_ = database.LogAudit(uid, "ACTION", "target", int64(i), "detail")
_ = database.LogAudit(context.Background(), uid, "ACTION", "target", int64(i), "detail")
}
page1, err := database.GetAuditLog(3, 0)
page1, err := database.GetAuditLog(context.Background(), 3, 0)
if err != nil {
t.Fatalf("GetAuditLog page1 error: %v", err)
}
@@ -559,7 +560,7 @@ func TestGetAuditLog_Pagination(t *testing.T) {
t.Errorf("page1 len = %d, want 3", len(page1))
}
page2, err := database.GetAuditLog(3, 3)
page2, err := database.GetAuditLog(context.Background(), 3, 3)
if err != nil {
t.Fatalf("GetAuditLog page2 error: %v", err)
}
@@ -571,11 +572,11 @@ func TestGetAuditLog_Pagination(t *testing.T) {
func TestGetAuditLog_NewestFirst(t *testing.T) {
database := newAdminTestDB(t)
uid, _ := database.CreateUser("auditorder", "hash", 1)
_ = database.LogAudit(uid, "FIRST", "", 0, "")
_ = database.LogAudit(uid, "SECOND", "", 0, "")
uid, _ := database.CreateUser(context.Background(), "auditorder", "hash", 1)
_ = database.LogAudit(context.Background(), uid, "FIRST", "", 0, "")
_ = database.LogAudit(context.Background(), uid, "SECOND", "", 0, "")
entries, err := database.GetAuditLog(10, 0)
entries, err := database.GetAuditLog(context.Background(), 10, 0)
if err != nil {
t.Fatalf("GetAuditLog() error: %v", err)
}
@@ -592,7 +593,7 @@ func TestGetAuditLog_NewestFirst(t *testing.T) {
func TestGetSetting_Exists(t *testing.T) {
database := newAdminTestDB(t)
val, err := database.GetSetting("server_name")
val, err := database.GetSetting(context.Background(), "server_name")
if err != nil {
t.Fatalf("GetSetting() error: %v", err)
}
@@ -604,7 +605,7 @@ func TestGetSetting_Exists(t *testing.T) {
func TestGetSetting_NotFound(t *testing.T) {
database := newAdminTestDB(t)
_, err := database.GetSetting("nonexistent_key_xyz")
_, err := database.GetSetting(context.Background(), "nonexistent_key_xyz")
if err == nil {
t.Error("GetSetting() for nonexistent key should return error")
}
@@ -613,11 +614,11 @@ func TestGetSetting_NotFound(t *testing.T) {
func TestSetSetting_NewKey(t *testing.T) {
database := newAdminTestDB(t)
if err := database.SetSetting("custom_key", "custom_val"); err != nil {
if err := database.SetSetting(context.Background(), "custom_key", "custom_val"); err != nil {
t.Fatalf("SetSetting() error: %v", err)
}
val, err := database.GetSetting("custom_key")
val, err := database.GetSetting(context.Background(), "custom_key")
if err != nil {
t.Fatalf("GetSetting() after SetSetting error: %v", err)
}
@@ -629,11 +630,11 @@ func TestSetSetting_NewKey(t *testing.T) {
func TestSetSetting_UpdateExisting(t *testing.T) {
database := newAdminTestDB(t)
if err := database.SetSetting("server_name", "My Custom Server"); err != nil {
if err := database.SetSetting(context.Background(), "server_name", "My Custom Server"); err != nil {
t.Fatalf("SetSetting() update error: %v", err)
}
val, err := database.GetSetting("server_name")
val, err := database.GetSetting(context.Background(), "server_name")
if err != nil {
t.Fatalf("GetSetting() error: %v", err)
}
@@ -645,7 +646,7 @@ func TestSetSetting_UpdateExisting(t *testing.T) {
func TestGetAllSettings_ReturnsMap(t *testing.T) {
database := newAdminTestDB(t)
settings, err := database.GetAllSettings()
settings, err := database.GetAllSettings(context.Background())
if err != nil {
t.Fatalf("GetAllSettings() error: %v", err)
}
@@ -660,9 +661,9 @@ func TestGetAllSettings_ReturnsMap(t *testing.T) {
func TestGetAllSettings_AfterClearing(t *testing.T) {
database := newAdminTestDB(t)
_, _ = database.Exec("DELETE FROM settings")
_, _ = database.ExecContext(context.Background(), "DELETE FROM settings")
settings, err := database.GetAllSettings()
settings, err := database.GetAllSettings(context.Background())
if err != nil {
t.Fatalf("GetAllSettings() after clearing error: %v", err)
}
@@ -693,7 +694,7 @@ func TestBackupToSafe_AdminQueries(t *testing.T) {
backupDir := filepath.Join(tmpDir, "backups")
_ = os.MkdirAll(backupDir, 0o755)
backupPath := filepath.Join(backupDir, "backup.db")
if err := database.BackupToSafe(backupPath, backupDir); err != nil {
if err := database.BackupToSafe(context.Background(), backupPath, backupDir); err != nil {
t.Fatalf("BackupToSafe() error: %v", err)
}
@@ -725,7 +726,7 @@ func TestBackupToSafe_CreatesDirectoryFile(t *testing.T) {
_ = os.MkdirAll(backupDir, 0o755)
backupPath := filepath.Join(backupDir, "chatserver_20260314_120000.db")
if err := database.BackupToSafe(backupPath, backupDir); err != nil {
if err := database.BackupToSafe(context.Background(), backupPath, backupDir); err != nil {
t.Fatalf("BackupToSafe() error: %v", err)
}
@@ -739,7 +740,7 @@ func TestBackupToSafe_CreatesDirectoryFile(t *testing.T) {
func TestUserCount_Empty(t *testing.T) {
database := newAdminTestDB(t)
count, err := database.UserCount()
count, err := database.UserCount(context.Background())
if err != nil {
t.Fatalf("UserCount() error: %v", err)
}
@@ -752,7 +753,7 @@ func TestUserCount_WithUsers(t *testing.T) {
database := newAdminTestDB(t)
for i := range 3 {
_, err := database.CreateUser(
_, err := database.CreateUser(context.Background(),
fmt.Sprintf("countuser%d", i),
"hash",
4,
@@ -762,7 +763,7 @@ func TestUserCount_WithUsers(t *testing.T) {
}
}
count, err := database.UserCount()
count, err := database.UserCount(context.Background())
if err != nil {
t.Fatalf("UserCount() error: %v", err)
}
@@ -793,7 +794,7 @@ func TestBackupToSafe_DirectCall(t *testing.T) {
backupDir := filepath.Join(tmpDir, "backups")
_ = os.MkdirAll(backupDir, 0o755)
backupPath := filepath.Join(backupDir, "backup_direct.db")
if err := database.BackupToSafe(backupPath, backupDir); err != nil {
if err := database.BackupToSafe(context.Background(), backupPath, backupDir); err != nil {
t.Fatalf("BackupToSafe() error: %v", err)
}
@@ -825,7 +826,7 @@ func TestBackupToSafe_RejectsTraversal(t *testing.T) {
_ = os.MkdirAll(safeRoot, 0o755)
unsafePath := filepath.Join(tmpDir, "outside", "evil.db")
err = database.BackupToSafe(unsafePath, safeRoot)
err = database.BackupToSafe(context.Background(), unsafePath, safeRoot)
if err == nil {
t.Error("BackupToSafe should reject path outside safe root")
}
+13 -12
View File
@@ -1,6 +1,7 @@
package db
import (
"context"
"database/sql"
"errors"
"fmt"
@@ -33,8 +34,8 @@ type AttachmentAccess struct {
// CreateAttachment inserts a new attachment record (initially unlinked to any message).
// uploaderID records who uploaded the file for ownership checks on unlinked files.
// width and height are optional image dimensions (pass nil for non-image files).
func (d *DB) CreateAttachment(id string, uploaderID int64, filename, storedAs, mimeType string, size int64, width, height *int) error {
if err := d.q.CreateAttachment(dbCtx(), dbgen.CreateAttachmentParams{
func (d *DB) CreateAttachment(ctx context.Context, id string, uploaderID int64, filename, storedAs, mimeType string, size int64, width, height *int) error {
if err := d.q.CreateAttachment(ctx, dbgen.CreateAttachmentParams{
ID: id,
UploaderID: &uploaderID,
Filename: filename,
@@ -50,8 +51,8 @@ func (d *DB) CreateAttachment(id string, uploaderID int64, filename, storedAs, m
}
// GetAttachmentByID returns the attachment with the given ID, or nil if not found.
func (d *DB) GetAttachmentByID(id string) (*Attachment, error) {
r, err := d.q.GetAttachmentByID(dbCtx(), id)
func (d *DB) GetAttachmentByID(ctx context.Context, id string) (*Attachment, error) {
r, err := d.q.GetAttachmentByID(ctx, id)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
@@ -74,8 +75,8 @@ func (d *DB) GetAttachmentByID(id string) (*Attachment, error) {
// (channel ID and type) for access-control checks. Returns nil if the
// attachment does not exist. ChannelID/ChannelType are nil/empty when the
// attachment is unlinked or its message/channel was deleted.
func (d *DB) GetAttachmentWithChannel(id string) (*AttachmentAccess, error) {
r, err := d.q.GetAttachmentWithChannel(dbCtx(), id)
func (d *DB) GetAttachmentWithChannel(ctx context.Context, id string) (*AttachmentAccess, error) {
r, err := d.q.GetAttachmentWithChannel(ctx, id)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
@@ -107,7 +108,7 @@ func (d *DB) GetAttachmentWithChannel(id string) (*AttachmentAccess, error) {
// is the atomic attachment-IDOR guard for message sends: ownership is
// enforced in the same statement that links, so there is no check-then-link
// race. Returns the number of rows updated.
func (d *DB) LinkAttachmentsToMessage(messageID, uploaderID int64, attachmentIDs []string) (int64, error) {
func (d *DB) LinkAttachmentsToMessage(ctx context.Context, messageID, uploaderID int64, attachmentIDs []string) (int64, error) {
if len(attachmentIDs) == 0 {
return 0, nil
}
@@ -127,7 +128,7 @@ func (d *DB) LinkAttachmentsToMessage(messageID, uploaderID int64, attachmentIDs
AND (uploader_id = ? OR uploader_id IS NULL)`,
strings.Join(placeholders, ","),
)
res, err := d.sqlDB.Exec(query, args...)
res, err := d.sqlDB.ExecContext(ctx, query, args...)
if err != nil {
return 0, fmt.Errorf("LinkAttachmentsToMessage: %w", err)
}
@@ -135,7 +136,7 @@ func (d *DB) LinkAttachmentsToMessage(messageID, uploaderID int64, attachmentIDs
}
// GetAttachmentsByMessageIDs returns attachments grouped by message ID.
func (d *DB) GetAttachmentsByMessageIDs(msgIDs []int64) (map[int64][]AttachmentInfo, error) {
func (d *DB) GetAttachmentsByMessageIDs(ctx context.Context, msgIDs []int64) (map[int64][]AttachmentInfo, error) {
if len(msgIDs) == 0 {
return map[int64][]AttachmentInfo{}, nil
}
@@ -152,7 +153,7 @@ func (d *DB) GetAttachmentsByMessageIDs(msgIDs []int64) (map[int64][]AttachmentI
FROM attachments WHERE message_id IN (%s)`,
strings.Join(placeholders, ","),
)
rows, err := d.sqlDB.Query(query, args...)
rows, err := d.sqlDB.QueryContext(ctx, query, args...)
if err != nil {
return nil, fmt.Errorf("GetAttachmentsByMessageIDs: %w", err)
}
@@ -184,8 +185,8 @@ func (d *DB) GetAttachmentsByMessageIDs(msgIDs []int64) (map[int64][]AttachmentI
// BUG-132: Uses DELETE ... RETURNING to make select+delete atomic,
// preventing a race where an attachment linked between SELECT and DELETE
// would have its file deleted while the DB row survives.
func (d *DB) DeleteOrphanedAttachments(cutoff string) ([]string, error) {
files, err := d.q.DeleteOrphanedAttachments(dbCtx(), cutoff)
func (d *DB) DeleteOrphanedAttachments(ctx context.Context, cutoff string) ([]string, error) {
files, err := d.q.DeleteOrphanedAttachments(ctx, cutoff)
if err != nil {
return nil, fmt.Errorf("DeleteOrphanedAttachments: %w", err)
}
+26 -25
View File
@@ -1,6 +1,7 @@
package db_test
import (
"context"
"testing"
)
@@ -9,7 +10,7 @@ import (
func TestGetAttachmentByID_NotFound(t *testing.T) {
database := openMigratedMemory(t)
att, err := database.GetAttachmentByID("nonexistent-id")
att, err := database.GetAttachmentByID(context.Background(), "nonexistent-id")
if err != nil {
t.Errorf("GetAttachmentByID for nonexistent ID should return nil error, got %v", err)
}
@@ -22,7 +23,7 @@ func TestGetAttachmentByID_Found(t *testing.T) {
database := openMigratedMemory(t)
// Insert an attachment directly.
_, err := database.Exec(
_, err := database.ExecContext(context.Background(),
`INSERT INTO attachments (id, filename, stored_as, mime_type, size)
VALUES (?, ?, ?, ?, ?)`,
"att-001", "photo.png", "stored-photo.png", "image/png", 12345,
@@ -31,7 +32,7 @@ func TestGetAttachmentByID_Found(t *testing.T) {
t.Fatalf("inserting attachment: %v", err)
}
att, err := database.GetAttachmentByID("att-001")
att, err := database.GetAttachmentByID(context.Background(), "att-001")
if err != nil {
t.Fatalf("GetAttachmentByID: %v", err)
}
@@ -57,7 +58,7 @@ func TestGetAttachmentByID_Found(t *testing.T) {
func TestLinkAttachmentsToMessage_Empty(t *testing.T) {
database := openMigratedMemory(t)
n, err := database.LinkAttachmentsToMessage(1, 1, nil)
n, err := database.LinkAttachmentsToMessage(context.Background(), 1, 1, nil)
if err != nil {
t.Fatalf("LinkAttachmentsToMessage(nil): %v", err)
}
@@ -70,11 +71,11 @@ func TestLinkAttachmentsToMessage_LinksUnlinked(t *testing.T) {
database := openMigratedMemory(t)
userID := seedUser(t, database, "linkuser")
chID := seedChannel(t, database, "linkchan")
msgID, _ := database.CreateMessage(chID, userID, "with attachment", nil)
msgID, _ := database.CreateMessage(context.Background(), chID, userID, "with attachment", nil)
// Insert two unlinked attachments.
for _, id := range []string{"att-a", "att-b"} {
_, err := database.Exec(
_, err := database.ExecContext(context.Background(),
`INSERT INTO attachments (id, filename, stored_as, mime_type, size)
VALUES (?, ?, ?, ?, ?)`,
id, "file.txt", "stored.txt", "text/plain", 100,
@@ -84,7 +85,7 @@ func TestLinkAttachmentsToMessage_LinksUnlinked(t *testing.T) {
}
}
n, err := database.LinkAttachmentsToMessage(msgID, userID, []string{"att-a", "att-b"})
n, err := database.LinkAttachmentsToMessage(context.Background(), msgID, userID, []string{"att-a", "att-b"})
if err != nil {
t.Fatalf("LinkAttachmentsToMessage: %v", err)
}
@@ -93,7 +94,7 @@ func TestLinkAttachmentsToMessage_LinksUnlinked(t *testing.T) {
}
// Verify linkage.
att, _ := database.GetAttachmentByID("att-a")
att, _ := database.GetAttachmentByID(context.Background(), "att-a")
if att.MessageID == nil || *att.MessageID != msgID {
t.Errorf("att-a MessageID = %v, want %d", att.MessageID, msgID)
}
@@ -103,17 +104,17 @@ func TestLinkAttachmentsToMessage_SkipsAlreadyLinked(t *testing.T) {
database := openMigratedMemory(t)
userID := seedUser(t, database, "linkuser2")
chID := seedChannel(t, database, "linkchan2")
msg1, _ := database.CreateMessage(chID, userID, "msg1", nil)
msg2, _ := database.CreateMessage(chID, userID, "msg2", nil)
msg1, _ := database.CreateMessage(context.Background(), chID, userID, "msg1", nil)
msg2, _ := database.CreateMessage(context.Background(), chID, userID, "msg2", nil)
_, _ = database.Exec(
_, _ = database.ExecContext(context.Background(),
`INSERT INTO attachments (id, filename, stored_as, mime_type, size, message_id)
VALUES (?, ?, ?, ?, ?, ?)`,
"att-linked", "file.txt", "stored.txt", "text/plain", 100, msg1,
)
// Try to re-link to a different message — should skip (WHERE message_id IS NULL).
n, err := database.LinkAttachmentsToMessage(msg2, userID, []string{"att-linked"})
n, err := database.LinkAttachmentsToMessage(context.Background(), msg2, userID, []string{"att-linked"})
if err != nil {
t.Fatalf("LinkAttachmentsToMessage: %v", err)
}
@@ -131,23 +132,23 @@ func TestLinkAttachmentsToMessage_OwnershipGuard(t *testing.T) {
owner := seedUser(t, database, "att-owner")
other := seedUser(t, database, "att-other")
chID := seedChannel(t, database, "att-owner-ch")
msgID, _ := database.CreateMessage(chID, owner, "attachment carrier", nil)
msgID, _ := database.CreateMessage(context.Background(), chID, owner, "attachment carrier", nil)
if err := database.CreateAttachment("att-owned", owner, "o.txt", "s-o.txt", "text/plain", 1, nil, nil); err != nil {
if err := database.CreateAttachment(context.Background(), "att-owned", owner, "o.txt", "s-o.txt", "text/plain", 1, nil, nil); err != nil {
t.Fatalf("CreateAttachment att-owned: %v", err)
}
if err := database.CreateAttachment("att-foreign", other, "f.txt", "s-f.txt", "text/plain", 1, nil, nil); err != nil {
if err := database.CreateAttachment(context.Background(), "att-foreign", other, "f.txt", "s-f.txt", "text/plain", 1, nil, nil); err != nil {
t.Fatalf("CreateAttachment att-foreign: %v", err)
}
// Legacy row from before uploader tracking: uploader_id IS NULL.
if _, err := database.Exec(
if _, err := database.ExecContext(context.Background(),
`INSERT INTO attachments (id, filename, stored_as, mime_type, size)
VALUES ('att-legacy', 'l.txt', 's-l.txt', 'text/plain', 1)`,
); err != nil {
t.Fatalf("inserting legacy attachment: %v", err)
}
n, err := database.LinkAttachmentsToMessage(msgID, owner,
n, err := database.LinkAttachmentsToMessage(context.Background(), msgID, owner,
[]string{"att-owned", "att-foreign", "att-legacy", "att-missing"})
if err != nil {
t.Fatalf("LinkAttachmentsToMessage: %v", err)
@@ -155,13 +156,13 @@ func TestLinkAttachmentsToMessage_OwnershipGuard(t *testing.T) {
if n != 2 {
t.Errorf("expected 2 linked (owned + legacy), got %d", n)
}
if att, _ := database.GetAttachmentByID("att-owned"); att.MessageID == nil || *att.MessageID != msgID {
if att, _ := database.GetAttachmentByID(context.Background(), "att-owned"); att.MessageID == nil || *att.MessageID != msgID {
t.Error("owner's unlinked attachment should link")
}
if att, _ := database.GetAttachmentByID("att-foreign"); att.MessageID != nil {
if att, _ := database.GetAttachmentByID(context.Background(), "att-foreign"); att.MessageID != nil {
t.Error("another user's attachment must never link (IDOR guard)")
}
if att, _ := database.GetAttachmentByID("att-legacy"); att.MessageID == nil {
if att, _ := database.GetAttachmentByID(context.Background(), "att-legacy"); att.MessageID == nil {
t.Error("legacy NULL-uploader attachment should be claimable")
}
}
@@ -171,7 +172,7 @@ func TestLinkAttachmentsToMessage_OwnershipGuard(t *testing.T) {
func TestGetAttachmentsByMessageIDs_Empty(t *testing.T) {
database := openMigratedMemory(t)
result, err := database.GetAttachmentsByMessageIDs(nil)
result, err := database.GetAttachmentsByMessageIDs(context.Background(), nil)
if err != nil {
t.Fatalf("GetAttachmentsByMessageIDs(nil): %v", err)
}
@@ -184,8 +185,8 @@ func TestGetAttachmentsByMessageIDs_GroupsByMessage(t *testing.T) {
database := openMigratedMemory(t)
userID := seedUser(t, database, "attuser")
chID := seedChannel(t, database, "attchan")
msg1, _ := database.CreateMessage(chID, userID, "msg1", nil)
msg2, _ := database.CreateMessage(chID, userID, "msg2", nil)
msg1, _ := database.CreateMessage(context.Background(), chID, userID, "msg1", nil)
msg2, _ := database.CreateMessage(context.Background(), chID, userID, "msg2", nil)
// Two attachments on msg1, one on msg2.
for _, row := range []struct {
@@ -196,7 +197,7 @@ func TestGetAttachmentsByMessageIDs_GroupsByMessage(t *testing.T) {
{"att-1b", msg1},
{"att-2a", msg2},
} {
_, err := database.Exec(
_, err := database.ExecContext(context.Background(),
`INSERT INTO attachments (id, filename, stored_as, mime_type, size, message_id)
VALUES (?, ?, ?, ?, ?, ?)`,
row.id, "f.txt", "s.txt", "text/plain", 50, row.msgID,
@@ -206,7 +207,7 @@ func TestGetAttachmentsByMessageIDs_GroupsByMessage(t *testing.T) {
}
}
result, err := database.GetAttachmentsByMessageIDs([]int64{msg1, msg2})
result, err := database.GetAttachmentsByMessageIDs(context.Background(), []int64{msg1, msg2})
if err != nil {
t.Fatalf("GetAttachmentsByMessageIDs: %v", err)
}
+7 -4
View File
@@ -1,13 +1,16 @@
package db
import "log/slog"
import (
"context"
"log/slog"
)
// Auditor is the minimal audit-write surface WriteAudit needs. *DB satisfies
// it directly, and the service layer's Store interface does too, so every
// caller — api, admin, ws, service — can route its audit writes through this
// one helper regardless of whether it holds a *DB or a narrower interface.
type Auditor interface {
LogAudit(actorID int64, action, targetType string, targetID int64, detail string) error
LogAudit(ctx context.Context, actorID int64, action, targetType string, targetID int64, detail string) error
}
// WriteAudit records an audit entry best-effort.
@@ -19,8 +22,8 @@ type Auditor interface {
// gap is visible in the logs. The detail string is intentionally not logged;
// it can carry request-specific or sensitive text and the structured fields
// already identify what was attempted.
func WriteAudit(a Auditor, actorID int64, action, targetType string, targetID int64, detail string) {
if err := a.LogAudit(actorID, action, targetType, targetID, detail); err != nil {
func WriteAudit(ctx context.Context, a Auditor, actorID int64, action, targetType string, targetID int64, detail string) {
if err := a.LogAudit(ctx, actorID, action, targetType, targetID, detail); err != nil {
slog.Error("audit log write failed",
"action", action,
"actor_id", actorID,
+4 -3
View File
@@ -1,6 +1,7 @@
package db_test
import (
"context"
"errors"
"log/slog"
"strings"
@@ -16,7 +17,7 @@ type fakeAuditor struct {
called bool
}
func (f *fakeAuditor) LogAudit(_ int64, _, _ string, _ int64, _ string) error {
func (f *fakeAuditor) LogAudit(_ context.Context, _ int64, _, _ string, _ int64, _ string) error {
f.called = true
return f.err
}
@@ -39,7 +40,7 @@ func TestWriteAudit_LogsFailureButDoesNotPropagate(t *testing.T) {
// WriteAudit returns nothing, so "never propagated" is structural — the
// call simply must not panic and must record the failure.
out := captureLogs(t, func() {
db.WriteAudit(a, 7, "user_ban", "user", 42, "spam")
db.WriteAudit(context.Background(), a, 7, "user_ban", "user", 42, "spam")
})
if !a.called {
@@ -67,7 +68,7 @@ func TestWriteAudit_SuccessLogsNothing(t *testing.T) {
a := &fakeAuditor{err: nil}
out := captureLogs(t, func() {
db.WriteAudit(a, 1, "user_login", "user", 1, "")
db.WriteAudit(context.Background(), a, 1, "user_login", "user", 1, "")
})
if !a.called {
+46 -45
View File
@@ -1,6 +1,7 @@
package db
import (
"context"
"crypto/rand"
"database/sql"
"encoding/hex"
@@ -14,8 +15,8 @@ import (
// ─── User Operations ──────────────────────────────────────────────────────────
// CreateUser inserts a new user record and returns the assigned ID.
func (d *DB) CreateUser(username, passwordHash string, roleID int) (int64, error) {
res, err := d.sqlDB.Exec(
func (d *DB) CreateUser(ctx context.Context, username, passwordHash string, roleID int) (int64, error) {
res, err := d.sqlDB.ExecContext(ctx,
`INSERT INTO users (username, password, role_id) VALUES (?, ?, ?)`,
username, passwordHash, roleID,
)
@@ -28,8 +29,8 @@ func (d *DB) CreateUser(username, passwordHash string, roleID int) (int64, error
// CreateOwnerIfEmpty atomically checks that no users exist and inserts the
// first owner in a single transaction. Returns ErrConflict if any user already
// exists, closing the TOCTOU race in the setup endpoint (BUG-119).
func (d *DB) CreateOwnerIfEmpty(username, passwordHash string, roleID int) (int64, error) {
tx, err := d.sqlDB.Begin()
func (d *DB) CreateOwnerIfEmpty(ctx context.Context, username, passwordHash string, roleID int) (int64, error) {
tx, err := d.sqlDB.BeginTx(ctx, nil)
if err != nil {
return 0, fmt.Errorf("CreateOwnerIfEmpty begin: %w", err)
}
@@ -70,8 +71,8 @@ func (d *DB) CreateOwnerIfEmpty(username, passwordHash string, roleID int) (int6
// CreateUserWithInvite atomically consumes an invite and creates the user in
// the same transaction so a failed registration does not burn the invite.
func (d *DB) CreateUserWithInvite(username, passwordHash string, roleID int, inviteCode string) (int64, error) {
tx, err := d.sqlDB.Begin()
func (d *DB) CreateUserWithInvite(ctx context.Context, username, passwordHash string, roleID int, inviteCode string) (int64, error) {
tx, err := d.sqlDB.BeginTx(ctx, nil)
if err != nil {
return 0, fmt.Errorf("CreateUserWithInvite begin: %w", err)
}
@@ -120,8 +121,8 @@ func (d *DB) CreateUserWithInvite(username, passwordHash string, roleID int, inv
// GetUserByUsername returns the user with the given username (case-insensitive),
// or nil if not found.
func (d *DB) GetUserByUsername(username string) (*User, error) {
u, err := d.q.GetUserByUsername(dbCtx(), username)
func (d *DB) GetUserByUsername(ctx context.Context, username string) (*User, error) {
u, err := d.q.GetUserByUsername(ctx, username)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
@@ -132,8 +133,8 @@ func (d *DB) GetUserByUsername(username string) (*User, error) {
}
// GetUserByID returns the user with the given ID, or nil if not found.
func (d *DB) GetUserByID(id int64) (*User, error) {
u, err := d.q.GetUserByID(dbCtx(), id)
func (d *DB) GetUserByID(ctx context.Context, id int64) (*User, error) {
u, err := d.q.GetUserByID(ctx, id)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
@@ -144,8 +145,8 @@ func (d *DB) GetUserByID(id int64) (*User, error) {
}
// UpdateUserStatus sets the status column for the given user ID.
func (d *DB) UpdateUserStatus(id int64, status string) error {
if err := d.q.UpdateUserStatus(dbCtx(), dbgen.UpdateUserStatusParams{
func (d *DB) UpdateUserStatus(ctx context.Context, id int64, status string) error {
if err := d.q.UpdateUserStatus(ctx, dbgen.UpdateUserStatusParams{
Status: status,
ID: id,
}); err != nil {
@@ -155,8 +156,8 @@ func (d *DB) UpdateUserStatus(id int64, status string) error {
}
// UpdateUserTOTPSecret sets or clears the TOTP secret for a user.
func (d *DB) UpdateUserTOTPSecret(id int64, secret *string) error {
if err := d.q.UpdateUserTOTPSecret(dbCtx(), dbgen.UpdateUserTOTPSecretParams{
func (d *DB) UpdateUserTOTPSecret(ctx context.Context, id int64, secret *string) error {
if err := d.q.UpdateUserTOTPSecret(ctx, dbgen.UpdateUserTOTPSecretParams{
TotpSecret: secret,
ID: id,
}); err != nil {
@@ -167,8 +168,8 @@ func (d *DB) UpdateUserTOTPSecret(id int64, secret *string) error {
// ResetAllUserStatuses sets all users to "offline". Called on server startup
// to clear stale statuses from a previous run or crash.
func (d *DB) ResetAllUserStatuses() error {
if err := d.q.ResetAllUserStatuses(dbCtx()); err != nil {
func (d *DB) ResetAllUserStatuses(ctx context.Context) error {
if err := d.q.ResetAllUserStatuses(ctx); err != nil {
return fmt.Errorf("ResetAllUserStatuses: %w", err)
}
return nil
@@ -176,14 +177,14 @@ func (d *DB) ResetAllUserStatuses() error {
// BanUser marks a user as banned with an optional expiry. Pass nil for a
// permanent ban.
func (d *DB) BanUser(id int64, reason string, expires *time.Time) error {
func (d *DB) BanUser(ctx context.Context, id int64, reason string, expires *time.Time) error {
var expiresStr *string
if expires != nil {
s := expires.UTC().Format("2006-01-02T15:04:05Z")
expiresStr = &s
}
reasonCopy := reason
if err := d.q.BanUser(dbCtx(), dbgen.BanUserParams{
if err := d.q.BanUser(ctx, dbgen.BanUserParams{
BanReason: &reasonCopy,
BanExpires: expiresStr,
ID: id,
@@ -194,8 +195,8 @@ func (d *DB) BanUser(id int64, reason string, expires *time.Time) error {
}
// UnbanUser removes the ban from a user.
func (d *DB) UnbanUser(id int64) error {
if err := d.q.UnbanUser(dbCtx(), id); err != nil {
func (d *DB) UnbanUser(ctx context.Context, id int64) error {
if err := d.q.UnbanUser(ctx, id); err != nil {
return fmt.Errorf("UnbanUser: %w", err)
}
return nil
@@ -212,16 +213,16 @@ const maxSessionsPerUser = 25
// tokenHash must already be hashed (never store plaintext tokens).
// H-6: Enforces a per-user session cap by evicting the oldest session when
// the limit is reached.
func (d *DB) CreateSession(userID int64, tokenHash, device, ip string) (int64, error) {
func (d *DB) CreateSession(ctx context.Context, userID int64, tokenHash, device, ip string) (int64, error) {
// Evict oldest sessions if at or above the cap.
_ = d.q.EvictOldestSessions(dbCtx(), dbgen.EvictOldestSessionsParams{
_ = d.q.EvictOldestSessions(ctx, dbgen.EvictOldestSessionsParams{
UserID: userID,
Offset: maxSessionsPerUser - 1,
})
expiresAt := time.Now().Add(sessionTTL).UTC().Format("2006-01-02T15:04:05Z")
deviceCopy, ipCopy := device, ip
res, err := d.q.InsertSession(dbCtx(), dbgen.InsertSessionParams{
res, err := d.q.InsertSession(ctx, dbgen.InsertSessionParams{
UserID: userID,
Token: tokenHash,
Device: &deviceCopy,
@@ -236,8 +237,8 @@ func (d *DB) CreateSession(userID int64, tokenHash, device, ip string) (int64, e
// GetSessionByTokenHash retrieves a session by its hashed token, or nil if
// not found.
func (d *DB) GetSessionByTokenHash(tokenHash string) (*Session, error) {
s, err := d.q.GetSessionByTokenHash(dbCtx(), tokenHash)
func (d *DB) GetSessionByTokenHash(ctx context.Context, tokenHash string) (*Session, error) {
s, err := d.q.GetSessionByTokenHash(ctx, tokenHash)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
@@ -259,8 +260,8 @@ type SessionWithBanStatus struct {
// GetSessionWithBanStatus returns the session joined with the user's ban
// status in a single query. Returns nil, nil when not found.
func (d *DB) GetSessionWithBanStatus(tokenHash string) (*SessionWithBanStatus, error) {
row, err := d.q.GetSessionWithBanStatus(dbCtx(), tokenHash)
func (d *DB) GetSessionWithBanStatus(ctx context.Context, tokenHash string) (*SessionWithBanStatus, error) {
row, err := d.q.GetSessionWithBanStatus(ctx, tokenHash)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
@@ -285,8 +286,8 @@ func (d *DB) GetSessionWithBanStatus(tokenHash string) (*SessionWithBanStatus, e
}
// DeleteSession removes the session with the given token hash.
func (d *DB) DeleteSession(tokenHash string) error {
if err := d.q.DeleteSessionByToken(dbCtx(), tokenHash); err != nil {
func (d *DB) DeleteSession(ctx context.Context, tokenHash string) error {
if err := d.q.DeleteSessionByToken(ctx, tokenHash); err != nil {
return fmt.Errorf("DeleteSession: %w", err)
}
return nil
@@ -295,8 +296,8 @@ func (d *DB) DeleteSession(tokenHash string) error {
// DeleteOtherSessions removes all sessions for the given user except the one
// with keepSessionID. Used after password change or 2FA state change to
// invalidate all other sessions (BUG-108).
func (d *DB) DeleteOtherSessions(userID, keepSessionID int64) (int64, error) {
result, err := d.q.DeleteOtherSessions(dbCtx(), dbgen.DeleteOtherSessionsParams{
func (d *DB) DeleteOtherSessions(ctx context.Context, userID, keepSessionID int64) (int64, error) {
result, err := d.q.DeleteOtherSessions(ctx, dbgen.DeleteOtherSessionsParams{
UserID: userID,
ID: keepSessionID,
})
@@ -309,16 +310,16 @@ func (d *DB) DeleteOtherSessions(userID, keepSessionID int64) (int64, error) {
// DeleteExpiredSessions removes all sessions whose expires_at is in the past.
// Compares using strftime to handle both ISO-8601 and SQLite datetime formats.
func (d *DB) DeleteExpiredSessions() error {
if err := d.q.DeleteExpiredSessions(dbCtx()); err != nil {
func (d *DB) DeleteExpiredSessions(ctx context.Context) error {
if err := d.q.DeleteExpiredSessions(ctx); err != nil {
return fmt.Errorf("DeleteExpiredSessions: %w", err)
}
return nil
}
// TouchSession updates last_used for the session with the given token hash.
func (d *DB) TouchSession(tokenHash string) error {
if err := d.q.TouchSession(dbCtx(), tokenHash); err != nil {
func (d *DB) TouchSession(ctx context.Context, tokenHash string) error {
if err := d.q.TouchSession(ctx, tokenHash); err != nil {
return fmt.Errorf("TouchSession: %w", err)
}
return nil
@@ -328,7 +329,7 @@ func (d *DB) TouchSession(tokenHash string) error {
// CreateInvite generates a random invite code, persists it, and returns the
// code. maxUses=0 means unlimited. expiresAt=nil means never expires.
func (d *DB) CreateInvite(createdBy int64, maxUses int, expiresAt *time.Time) (string, error) {
func (d *DB) CreateInvite(ctx context.Context, createdBy int64, maxUses int, expiresAt *time.Time) (string, error) {
code, err := generateInviteCode()
if err != nil {
return "", fmt.Errorf("CreateInvite generate code: %w", err)
@@ -344,7 +345,7 @@ func (d *DB) CreateInvite(createdBy int64, maxUses int, expiresAt *time.Time) (s
expiresStr = &s
}
if err := d.q.CreateInvite(dbCtx(), dbgen.CreateInviteParams{
if err := d.q.CreateInvite(ctx, dbgen.CreateInviteParams{
Code: code,
CreatedBy: createdBy,
MaxUses: ptrItoI64(maxUsesVal),
@@ -356,8 +357,8 @@ func (d *DB) CreateInvite(createdBy int64, maxUses int, expiresAt *time.Time) (s
}
// GetInvite returns the invite for the given code, or nil if not found.
func (d *DB) GetInvite(code string) (*Invite, error) {
r, err := d.q.GetInvite(dbCtx(), code)
func (d *DB) GetInvite(ctx context.Context, code string) (*Invite, error) {
r, err := d.q.GetInvite(ctx, code)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
@@ -388,8 +389,8 @@ func (d *DB) GetInvite(code string) (*Invite, error) {
//
// If zero rows are affected the invite is missing, revoked, expired, or
// exhausted — an error is returned in all such cases.
func (d *DB) UseInviteAtomic(code string) error {
result, err := d.q.UseInviteAtomic(dbCtx(), code)
func (d *DB) UseInviteAtomic(ctx context.Context, code string) error {
result, err := d.q.UseInviteAtomic(ctx, code)
if err != nil {
return fmt.Errorf("UseInviteAtomic: %w", err)
}
@@ -404,8 +405,8 @@ func (d *DB) UseInviteAtomic(code string) error {
}
// RevokeInvite marks an invite as revoked.
func (d *DB) RevokeInvite(code string) error {
if err := d.q.RevokeInvite(dbCtx(), code); err != nil {
func (d *DB) RevokeInvite(ctx context.Context, code string) error {
if err := d.q.RevokeInvite(ctx, code); err != nil {
return fmt.Errorf("RevokeInvite: %w", err)
}
return nil
@@ -424,8 +425,8 @@ type MemberSummary struct {
// ListMembers returns non-banned users as lightweight summaries.
// M-12: Limited to 1000 rows to prevent unbounded result sets on large servers.
func (d *DB) ListMembers() ([]MemberSummary, error) {
rows, err := d.q.ListMembers(dbCtx())
func (d *DB) ListMembers(ctx context.Context) ([]MemberSummary, error) {
rows, err := d.q.ListMembers(ctx)
if err != nil {
return nil, fmt.Errorf("ListMembers: %w", err)
}
+114 -113
View File
@@ -1,6 +1,7 @@
package db_test
import (
"context"
"testing"
"testing/fstest"
"time"
@@ -93,7 +94,7 @@ CREATE INDEX IF NOT EXISTS idx_invites_code ON invites(code);
func TestCreateUser_Success(t *testing.T) {
database := newTestDB(t)
id, err := database.CreateUser("alice", "hash123", 4)
id, err := database.CreateUser(context.Background(), "alice", "hash123", 4)
if err != nil {
t.Fatalf("CreateUser: %v", err)
}
@@ -104,10 +105,10 @@ func TestCreateUser_Success(t *testing.T) {
func TestCreateUser_DuplicateUsername(t *testing.T) {
database := newTestDB(t)
if _, err := database.CreateUser("bob", "hash1", 4); err != nil {
if _, err := database.CreateUser(context.Background(), "bob", "hash1", 4); err != nil {
t.Fatalf("first CreateUser: %v", err)
}
_, err := database.CreateUser("bob", "hash2", 4)
_, err := database.CreateUser(context.Background(), "bob", "hash2", 4)
if err == nil {
t.Error("CreateUser() with duplicate username returned nil error, want error")
}
@@ -115,10 +116,10 @@ func TestCreateUser_DuplicateUsername(t *testing.T) {
func TestCreateUser_CaseInsensitiveDuplicate(t *testing.T) {
database := newTestDB(t)
if _, err := database.CreateUser("Charlie", "hash1", 4); err != nil {
if _, err := database.CreateUser(context.Background(), "Charlie", "hash1", 4); err != nil {
t.Fatalf("first CreateUser: %v", err)
}
_, err := database.CreateUser("charlie", "hash2", 4)
_, err := database.CreateUser(context.Background(), "charlie", "hash2", 4)
if err == nil {
t.Error("CreateUser() with case-insensitive duplicate returned nil error, want error")
}
@@ -126,9 +127,9 @@ func TestCreateUser_CaseInsensitiveDuplicate(t *testing.T) {
func TestGetUserByUsername_Found(t *testing.T) {
database := newTestDB(t)
_, _ = database.CreateUser("dave", "hashDave", 4)
_, _ = database.CreateUser(context.Background(), "dave", "hashDave", 4)
user, err := database.GetUserByUsername("dave")
user, err := database.GetUserByUsername(context.Background(), "dave")
if err != nil {
t.Fatalf("GetUserByUsername: %v", err)
}
@@ -142,9 +143,9 @@ func TestGetUserByUsername_Found(t *testing.T) {
func TestGetUserByUsername_CaseInsensitive(t *testing.T) {
database := newTestDB(t)
_, _ = database.CreateUser("Eve", "hashEve", 4)
_, _ = database.CreateUser(context.Background(), "Eve", "hashEve", 4)
user, err := database.GetUserByUsername("EVE")
user, err := database.GetUserByUsername(context.Background(), "EVE")
if err != nil {
t.Fatalf("GetUserByUsername case-insensitive: %v", err)
}
@@ -155,7 +156,7 @@ func TestGetUserByUsername_CaseInsensitive(t *testing.T) {
func TestGetUserByUsername_NotFound(t *testing.T) {
database := newTestDB(t)
user, err := database.GetUserByUsername("nobody")
user, err := database.GetUserByUsername(context.Background(), "nobody")
if err != nil {
t.Fatalf("GetUserByUsername(not found): %v", err)
}
@@ -166,9 +167,9 @@ func TestGetUserByUsername_NotFound(t *testing.T) {
func TestGetUserByID_Found(t *testing.T) {
database := newTestDB(t)
id, _ := database.CreateUser("frank", "hashFrank", 4)
id, _ := database.CreateUser(context.Background(), "frank", "hashFrank", 4)
user, err := database.GetUserByID(id)
user, err := database.GetUserByID(context.Background(), id)
if err != nil {
t.Fatalf("GetUserByID: %v", err)
}
@@ -179,7 +180,7 @@ func TestGetUserByID_Found(t *testing.T) {
func TestGetUserByID_NotFound(t *testing.T) {
database := newTestDB(t)
user, err := database.GetUserByID(999)
user, err := database.GetUserByID(context.Background(), 999)
if err != nil {
t.Fatalf("GetUserByID(not found): %v", err)
}
@@ -190,12 +191,12 @@ func TestGetUserByID_NotFound(t *testing.T) {
func TestUpdateUserStatus(t *testing.T) {
database := newTestDB(t)
id, _ := database.CreateUser("grace", "hash", 4)
id, _ := database.CreateUser(context.Background(), "grace", "hash", 4)
if err := database.UpdateUserStatus(id, "online"); err != nil {
if err := database.UpdateUserStatus(context.Background(), id, "online"); err != nil {
t.Fatalf("UpdateUserStatus: %v", err)
}
user, _ := database.GetUserByID(id)
user, _ := database.GetUserByID(context.Background(), id)
if user.Status != "online" {
t.Errorf("Status = %q, want %q", user.Status, "online")
}
@@ -203,12 +204,12 @@ func TestUpdateUserStatus(t *testing.T) {
func TestBanUser_Permanent(t *testing.T) {
database := newTestDB(t)
id, _ := database.CreateUser("hank", "hash", 4)
id, _ := database.CreateUser(context.Background(), "hank", "hash", 4)
if err := database.BanUser(id, "spam", nil); err != nil {
if err := database.BanUser(context.Background(), id, "spam", nil); err != nil {
t.Fatalf("BanUser: %v", err)
}
user, _ := database.GetUserByID(id)
user, _ := database.GetUserByID(context.Background(), id)
if !user.Banned {
t.Error("Banned = false after BanUser, want true")
}
@@ -219,13 +220,13 @@ func TestBanUser_Permanent(t *testing.T) {
func TestBanUser_Temporary(t *testing.T) {
database := newTestDB(t)
id, _ := database.CreateUser("ivan", "hash", 4)
id, _ := database.CreateUser(context.Background(), "ivan", "hash", 4)
expires := time.Now().Add(24 * time.Hour)
if err := database.BanUser(id, "temp ban", &expires); err != nil {
if err := database.BanUser(context.Background(), id, "temp ban", &expires); err != nil {
t.Fatalf("BanUser (temp): %v", err)
}
user, _ := database.GetUserByID(id)
user, _ := database.GetUserByID(context.Background(), id)
if !user.Banned {
t.Error("Banned = false after temp ban")
}
@@ -238,9 +239,9 @@ func TestBanUser_Temporary(t *testing.T) {
func TestCreateSession_Success(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("jack", "hash", 4)
uid, _ := database.CreateUser(context.Background(), "jack", "hash", 4)
id, err := database.CreateSession(uid, "tokenHash1", "GoTest/1.0", "127.0.0.1")
id, err := database.CreateSession(context.Background(), uid, "tokenHash1", "GoTest/1.0", "127.0.0.1")
if err != nil {
t.Fatalf("CreateSession: %v", err)
}
@@ -251,10 +252,10 @@ func TestCreateSession_Success(t *testing.T) {
func TestGetSessionByTokenHash_Found(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("kate", "hash", 4)
_, _ = database.CreateSession(uid, "myTokenHash", "GoTest/1.0", "127.0.0.1")
uid, _ := database.CreateUser(context.Background(), "kate", "hash", 4)
_, _ = database.CreateSession(context.Background(), uid, "myTokenHash", "GoTest/1.0", "127.0.0.1")
sess, err := database.GetSessionByTokenHash("myTokenHash")
sess, err := database.GetSessionByTokenHash(context.Background(), "myTokenHash")
if err != nil {
t.Fatalf("GetSessionByTokenHash: %v", err)
}
@@ -268,7 +269,7 @@ func TestGetSessionByTokenHash_Found(t *testing.T) {
func TestGetSessionByTokenHash_NotFound(t *testing.T) {
database := newTestDB(t)
sess, err := database.GetSessionByTokenHash("nonexistent")
sess, err := database.GetSessionByTokenHash(context.Background(), "nonexistent")
if err != nil {
t.Fatalf("GetSessionByTokenHash(not found): %v", err)
}
@@ -279,10 +280,10 @@ func TestGetSessionByTokenHash_NotFound(t *testing.T) {
func TestGetSessionWithBanStatus_Found(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("zara", "hash", 4)
_, _ = database.CreateSession(uid, "banCheckToken", "GoTest/1.0", "127.0.0.1")
uid, _ := database.CreateUser(context.Background(), "zara", "hash", 4)
_, _ = database.CreateSession(context.Background(), uid, "banCheckToken", "GoTest/1.0", "127.0.0.1")
result, err := database.GetSessionWithBanStatus("banCheckToken")
result, err := database.GetSessionWithBanStatus(context.Background(), "banCheckToken")
if err != nil {
t.Fatalf("GetSessionWithBanStatus: %v", err)
}
@@ -299,13 +300,13 @@ func TestGetSessionWithBanStatus_Found(t *testing.T) {
func TestGetSessionWithBanStatus_BannedUser(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("banned-zara", "hash", 4)
_, _ = database.CreateSession(uid, "bannedToken", "GoTest/1.0", "127.0.0.1")
if err := database.BanUser(uid, "rule violation", nil); err != nil {
uid, _ := database.CreateUser(context.Background(), "banned-zara", "hash", 4)
_, _ = database.CreateSession(context.Background(), uid, "bannedToken", "GoTest/1.0", "127.0.0.1")
if err := database.BanUser(context.Background(), uid, "rule violation", nil); err != nil {
t.Fatalf("BanUser: %v", err)
}
result, err := database.GetSessionWithBanStatus("bannedToken")
result, err := database.GetSessionWithBanStatus(context.Background(), "bannedToken")
if err != nil {
t.Fatalf("GetSessionWithBanStatus: %v", err)
}
@@ -322,7 +323,7 @@ func TestGetSessionWithBanStatus_BannedUser(t *testing.T) {
func TestGetSessionWithBanStatus_NotFound(t *testing.T) {
database := newTestDB(t)
result, err := database.GetSessionWithBanStatus("nonexistent")
result, err := database.GetSessionWithBanStatus(context.Background(), "nonexistent")
if err != nil {
t.Fatalf("GetSessionWithBanStatus(not found): %v", err)
}
@@ -333,13 +334,13 @@ func TestGetSessionWithBanStatus_NotFound(t *testing.T) {
func TestDeleteSession(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("leo", "hash", 4)
_, _ = database.CreateSession(uid, "delToken", "GoTest/1.0", "127.0.0.1")
uid, _ := database.CreateUser(context.Background(), "leo", "hash", 4)
_, _ = database.CreateSession(context.Background(), uid, "delToken", "GoTest/1.0", "127.0.0.1")
if err := database.DeleteSession("delToken"); err != nil {
if err := database.DeleteSession(context.Background(), "delToken"); err != nil {
t.Fatalf("DeleteSession: %v", err)
}
sess, _ := database.GetSessionByTokenHash("delToken")
sess, _ := database.GetSessionByTokenHash(context.Background(), "delToken")
if sess != nil {
t.Error("Session still exists after DeleteSession")
}
@@ -347,12 +348,12 @@ func TestDeleteSession(t *testing.T) {
func TestDeleteExpiredSessions(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("mia", "hash", 4)
uid, _ := database.CreateUser(context.Background(), "mia", "hash", 4)
// Insert an already-expired session directly via Exec.
// Use SQLite datetime format (space separator) to match what datetime('now') produces.
pastTime := time.Now().Add(-time.Hour).UTC().Format("2006-01-02 15:04:05")
_, err := database.Exec(
_, err := database.ExecContext(context.Background(),
`INSERT INTO sessions (user_id, token, device, ip_address, expires_at) VALUES (?, ?, ?, ?, ?)`,
uid, "expiredToken", "test", "127.0.0.1", pastTime,
)
@@ -361,17 +362,17 @@ func TestDeleteExpiredSessions(t *testing.T) {
}
// Insert a valid session through the normal path.
_, _ = database.CreateSession(uid, "validToken", "GoTest/1.0", "127.0.0.1")
_, _ = database.CreateSession(context.Background(), uid, "validToken", "GoTest/1.0", "127.0.0.1")
if err := database.DeleteExpiredSessions(); err != nil {
if err := database.DeleteExpiredSessions(context.Background()); err != nil {
t.Fatalf("DeleteExpiredSessions: %v", err)
}
expired, _ := database.GetSessionByTokenHash("expiredToken")
expired, _ := database.GetSessionByTokenHash(context.Background(), "expiredToken")
if expired != nil {
t.Error("Expired session still exists after DeleteExpiredSessions")
}
valid, _ := database.GetSessionByTokenHash("validToken")
valid, _ := database.GetSessionByTokenHash(context.Background(), "validToken")
if valid == nil {
t.Error("Valid session was deleted by DeleteExpiredSessions")
}
@@ -379,17 +380,17 @@ func TestDeleteExpiredSessions(t *testing.T) {
func TestTouchSession(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("noah", "hash", 4)
_, _ = database.CreateSession(uid, "touchToken", "GoTest/1.0", "127.0.0.1")
uid, _ := database.CreateUser(context.Background(), "noah", "hash", 4)
_, _ = database.CreateSession(context.Background(), uid, "touchToken", "GoTest/1.0", "127.0.0.1")
sess1, _ := database.GetSessionByTokenHash("touchToken")
sess1, _ := database.GetSessionByTokenHash(context.Background(), "touchToken")
time.Sleep(2 * time.Millisecond)
if err := database.TouchSession("touchToken"); err != nil {
if err := database.TouchSession(context.Background(), "touchToken"); err != nil {
t.Fatalf("TouchSession: %v", err)
}
sess2, _ := database.GetSessionByTokenHash("touchToken")
sess2, _ := database.GetSessionByTokenHash(context.Background(), "touchToken")
if sess1.LastUsed == sess2.LastUsed {
// last_used should have advanced; if they're equal the touch had no effect
// (This can be flaky at millisecond resolution, but is a reasonable sanity check.)
@@ -401,9 +402,9 @@ func TestTouchSession(t *testing.T) {
func TestCreateInvite_Success(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("olivia", "hash", 4)
uid, _ := database.CreateUser(context.Background(), "olivia", "hash", 4)
code, err := database.CreateInvite(uid, 0, nil)
code, err := database.CreateInvite(context.Background(), uid, 0, nil)
if err != nil {
t.Fatalf("CreateInvite: %v", err)
}
@@ -414,10 +415,10 @@ func TestCreateInvite_Success(t *testing.T) {
func TestGetInvite_Found(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("pedro", "hash", 4)
code, _ := database.CreateInvite(uid, 5, nil)
uid, _ := database.CreateUser(context.Background(), "pedro", "hash", 4)
code, _ := database.CreateInvite(context.Background(), uid, 5, nil)
inv, err := database.GetInvite(code)
inv, err := database.GetInvite(context.Background(), code)
if err != nil {
t.Fatalf("GetInvite: %v", err)
}
@@ -434,7 +435,7 @@ func TestGetInvite_Found(t *testing.T) {
func TestGetInvite_NotFound(t *testing.T) {
database := newTestDB(t)
inv, err := database.GetInvite("bogus")
inv, err := database.GetInvite(context.Background(), "bogus")
if err != nil {
t.Fatalf("GetInvite(not found): %v", err)
}
@@ -445,14 +446,14 @@ func TestGetInvite_NotFound(t *testing.T) {
func TestRevokeInvite(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("uma", "hash", 4)
code, _ := database.CreateInvite(uid, 0, nil)
uid, _ := database.CreateUser(context.Background(), "uma", "hash", 4)
code, _ := database.CreateInvite(context.Background(), uid, 0, nil)
if err := database.RevokeInvite(code); err != nil {
if err := database.RevokeInvite(context.Background(), code); err != nil {
t.Fatalf("RevokeInvite: %v", err)
}
inv, _ := database.GetInvite(code)
inv, _ := database.GetInvite(context.Background(), code)
if !inv.Revoked {
t.Error("Revoked = false after RevokeInvite, want true")
}
@@ -460,10 +461,10 @@ func TestRevokeInvite(t *testing.T) {
func TestCreateInvite_UnlimitedUses(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("vera", "hash", 4)
code, _ := database.CreateInvite(uid, 0, nil) // 0 = unlimited
uid, _ := database.CreateUser(context.Background(), "vera", "hash", 4)
code, _ := database.CreateInvite(context.Background(), uid, 0, nil) // 0 = unlimited
inv, _ := database.GetInvite(code)
inv, _ := database.GetInvite(context.Background(), code)
if inv.MaxUses != nil {
t.Errorf("MaxUses = %v, want nil for unlimited", inv.MaxUses)
}
@@ -475,14 +476,14 @@ func TestCreateInvite_UnlimitedUses(t *testing.T) {
// its use_count incremented in one operation.
func TestUseInviteAtomic_Success(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("atomic_user1", "hash", 4)
code, _ := database.CreateInvite(uid, 0, nil)
uid, _ := database.CreateUser(context.Background(), "atomic_user1", "hash", 4)
code, _ := database.CreateInvite(context.Background(), uid, 0, nil)
if err := database.UseInviteAtomic(code); err != nil {
if err := database.UseInviteAtomic(context.Background(), code); err != nil {
t.Fatalf("UseInviteAtomic: %v", err)
}
inv, _ := database.GetInvite(code)
inv, _ := database.GetInvite(context.Background(), code)
if inv.Uses != 1 {
t.Errorf("Uses = %d, want 1", inv.Uses)
}
@@ -492,16 +493,16 @@ func TestUseInviteAtomic_Success(t *testing.T) {
// multiple sequential calls.
func TestUseInviteAtomic_IncrementsUses(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("atomic_user2", "hash", 4)
code, _ := database.CreateInvite(uid, 5, nil)
uid, _ := database.CreateUser(context.Background(), "atomic_user2", "hash", 4)
code, _ := database.CreateInvite(context.Background(), uid, 5, nil)
for i := range 3 {
if err := database.UseInviteAtomic(code); err != nil {
if err := database.UseInviteAtomic(context.Background(), code); err != nil {
t.Fatalf("UseInviteAtomic iteration %d: %v", i, err)
}
}
inv, _ := database.GetInvite(code)
inv, _ := database.GetInvite(context.Background(), code)
if inv.Uses != 3 {
t.Errorf("Uses = %d, want 3", inv.Uses)
}
@@ -511,16 +512,16 @@ func TestUseInviteAtomic_IncrementsUses(t *testing.T) {
// modifying the database.
func TestUseInviteAtomic_Revoked(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("atomic_user3", "hash", 4)
code, _ := database.CreateInvite(uid, 0, nil)
_ = database.RevokeInvite(code)
uid, _ := database.CreateUser(context.Background(), "atomic_user3", "hash", 4)
code, _ := database.CreateInvite(context.Background(), uid, 0, nil)
_ = database.RevokeInvite(context.Background(), code)
if err := database.UseInviteAtomic(code); err == nil {
if err := database.UseInviteAtomic(context.Background(), code); err == nil {
t.Error("UseInviteAtomic returned nil error for revoked invite, want error")
}
// use_count must not have changed.
inv, _ := database.GetInvite(code)
inv, _ := database.GetInvite(context.Background(), code)
if inv.Uses != 0 {
t.Errorf("Uses = %d after revoked attempt, want 0", inv.Uses)
}
@@ -529,12 +530,12 @@ func TestUseInviteAtomic_Revoked(t *testing.T) {
// TestUseInviteAtomic_Expired returns an error for an expired invite.
func TestUseInviteAtomic_Expired(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("atomic_user4", "hash", 4)
uid, _ := database.CreateUser(context.Background(), "atomic_user4", "hash", 4)
past := time.Now().Add(-time.Hour)
code, _ := database.CreateInvite(uid, 0, &past)
code, _ := database.CreateInvite(context.Background(), uid, 0, &past)
if err := database.UseInviteAtomic(code); err == nil {
if err := database.UseInviteAtomic(context.Background(), code); err == nil {
t.Error("UseInviteAtomic returned nil error for expired invite, want error")
}
}
@@ -543,13 +544,13 @@ func TestUseInviteAtomic_Expired(t *testing.T) {
// reached its maximum use count.
func TestUseInviteAtomic_ExceedsMaxUses(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("atomic_user5", "hash", 4)
code, _ := database.CreateInvite(uid, 1, nil)
uid, _ := database.CreateUser(context.Background(), "atomic_user5", "hash", 4)
code, _ := database.CreateInvite(context.Background(), uid, 1, nil)
if err := database.UseInviteAtomic(code); err != nil {
if err := database.UseInviteAtomic(context.Background(), code); err != nil {
t.Fatalf("UseInviteAtomic first use: %v", err)
}
if err := database.UseInviteAtomic(code); err == nil {
if err := database.UseInviteAtomic(context.Background(), code); err == nil {
t.Error("UseInviteAtomic returned nil error after exceeding max_uses, want error")
}
}
@@ -558,7 +559,7 @@ func TestUseInviteAtomic_ExceedsMaxUses(t *testing.T) {
func TestUseInviteAtomic_NotFound(t *testing.T) {
database := newTestDB(t)
if err := database.UseInviteAtomic("doesnotexist"); err == nil {
if err := database.UseInviteAtomic(context.Background(), "doesnotexist"); err == nil {
t.Error("UseInviteAtomic returned nil error for unknown code, want error")
}
}
@@ -568,15 +569,15 @@ func TestUseInviteAtomic_NotFound(t *testing.T) {
// fail; the use_count must end up at 1.
func TestUseInviteAtomic_ConcurrentSameCode(t *testing.T) {
database := newTestDB(t)
uid, _ := database.CreateUser("atomic_user6", "hash", 4)
code, _ := database.CreateInvite(uid, 1, nil)
uid, _ := database.CreateUser(context.Background(), "atomic_user6", "hash", 4)
code, _ := database.CreateInvite(context.Background(), uid, 1, nil)
type result struct{ err error }
results := make(chan result, 2)
for range 2 {
go func() {
results <- result{err: database.UseInviteAtomic(code)}
results <- result{err: database.UseInviteAtomic(context.Background(), code)}
}()
}
@@ -592,7 +593,7 @@ func TestUseInviteAtomic_ConcurrentSameCode(t *testing.T) {
t.Errorf("concurrent redemptions: %d succeeded, want exactly 1", successes)
}
inv, _ := database.GetInvite(code)
inv, _ := database.GetInvite(context.Background(), code)
if inv.Uses != 1 {
t.Errorf("use_count = %d after concurrent race, want 1", inv.Uses)
}
@@ -602,22 +603,22 @@ func TestUseInviteAtomic_ConcurrentSameCode(t *testing.T) {
func TestUnbanUser_ClearsBan(t *testing.T) {
database := newTestDB(t)
id, _ := database.CreateUser("unban_target", "hash", 4)
id, _ := database.CreateUser(context.Background(), "unban_target", "hash", 4)
if err := database.BanUser(id, "spam", nil); err != nil {
if err := database.BanUser(context.Background(), id, "spam", nil); err != nil {
t.Fatalf("BanUser: %v", err)
}
user, _ := database.GetUserByID(id)
user, _ := database.GetUserByID(context.Background(), id)
if !user.Banned {
t.Fatal("user should be banned before unban")
}
if err := database.UnbanUser(id); err != nil {
if err := database.UnbanUser(context.Background(), id); err != nil {
t.Fatalf("UnbanUser: %v", err)
}
user, _ = database.GetUserByID(id)
user, _ = database.GetUserByID(context.Background(), id)
if user.Banned {
t.Error("Banned = true after UnbanUser, want false")
}
@@ -633,7 +634,7 @@ func TestUnbanUser_NonexistentUser(t *testing.T) {
database := newTestDB(t)
// Unbanning nonexistent user should not error.
if err := database.UnbanUser(99999); err != nil {
if err := database.UnbanUser(context.Background(), 99999); err != nil {
t.Errorf("UnbanUser(nonexistent) error: %v", err)
}
}
@@ -642,18 +643,18 @@ func TestUnbanUser_NonexistentUser(t *testing.T) {
func TestResetAllUserStatuses(t *testing.T) {
database := newTestDB(t)
id1, _ := database.CreateUser("status_u1", "hash", 4)
id2, _ := database.CreateUser("status_u2", "hash", 4)
id1, _ := database.CreateUser(context.Background(), "status_u1", "hash", 4)
id2, _ := database.CreateUser(context.Background(), "status_u2", "hash", 4)
_ = database.UpdateUserStatus(id1, "online")
_ = database.UpdateUserStatus(id2, "dnd")
_ = database.UpdateUserStatus(context.Background(), id1, "online")
_ = database.UpdateUserStatus(context.Background(), id2, "dnd")
if err := database.ResetAllUserStatuses(); err != nil {
if err := database.ResetAllUserStatuses(context.Background()); err != nil {
t.Fatalf("ResetAllUserStatuses: %v", err)
}
u1, _ := database.GetUserByID(id1)
u2, _ := database.GetUserByID(id2)
u1, _ := database.GetUserByID(context.Background(), id1)
u2, _ := database.GetUserByID(context.Background(), id2)
if u1.Status != "offline" {
t.Errorf("user1 status = %q, want 'offline'", u1.Status)
}
@@ -664,10 +665,10 @@ func TestResetAllUserStatuses(t *testing.T) {
func TestResetAllUserStatuses_AlreadyOffline(t *testing.T) {
database := newTestDB(t)
_, _ = database.CreateUser("offline_user", "hash", 4)
_, _ = database.CreateUser(context.Background(), "offline_user", "hash", 4)
// Should not error when all users are already offline.
if err := database.ResetAllUserStatuses(); err != nil {
if err := database.ResetAllUserStatuses(context.Background()); err != nil {
t.Errorf("ResetAllUserStatuses: %v", err)
}
}
@@ -677,7 +678,7 @@ func TestResetAllUserStatuses_AlreadyOffline(t *testing.T) {
func TestListMembers_Empty(t *testing.T) {
database := newTestDB(t)
members, err := database.ListMembers()
members, err := database.ListMembers(context.Background())
if err != nil {
t.Fatalf("ListMembers: %v", err)
}
@@ -688,12 +689,12 @@ func TestListMembers_Empty(t *testing.T) {
func TestListMembers_ExcludesBanned(t *testing.T) {
database := newTestDB(t)
id1, _ := database.CreateUser("member_visible", "hash", 4)
id2, _ := database.CreateUser("member_banned", "hash", 4)
_ = database.BanUser(id2, "test ban", nil)
id1, _ := database.CreateUser(context.Background(), "member_visible", "hash", 4)
id2, _ := database.CreateUser(context.Background(), "member_banned", "hash", 4)
_ = database.BanUser(context.Background(), id2, "test ban", nil)
_ = id1 // suppress unused
members, err := database.ListMembers()
members, err := database.ListMembers(context.Background())
if err != nil {
t.Fatalf("ListMembers: %v", err)
}
@@ -710,11 +711,11 @@ func TestListMembers_ExcludesBanned(t *testing.T) {
func TestListMembers_SortedByUsername(t *testing.T) {
database := newTestDB(t)
_, _ = database.CreateUser("zeta_user", "hash", 4)
_, _ = database.CreateUser("alpha_user", "hash", 4)
_, _ = database.CreateUser("mid_user", "hash", 4)
_, _ = database.CreateUser(context.Background(), "zeta_user", "hash", 4)
_, _ = database.CreateUser(context.Background(), "alpha_user", "hash", 4)
_, _ = database.CreateUser(context.Background(), "mid_user", "hash", 4)
members, err := database.ListMembers()
members, err := database.ListMembers(context.Background())
if err != nil {
t.Fatalf("ListMembers: %v", err)
}
+8 -7
View File
@@ -1,6 +1,7 @@
package db_test
import (
"context"
"os"
"path/filepath"
"testing"
@@ -43,7 +44,7 @@ func TestBackupToSafe_ValidPath(t *testing.T) {
t.Fatalf("MkdirAll: %v", err)
}
backupPath := filepath.Join(backupDir, "chatserver_20260315_120000.db")
if err := database.BackupToSafe(backupPath, backupDir); err != nil {
if err := database.BackupToSafe(context.Background(), backupPath, backupDir); err != nil {
t.Fatalf("BackupToSafe() with valid path returned error: %v", err)
}
@@ -67,7 +68,7 @@ func TestBackupToSafe_RejectsPathOutsideRoot(t *testing.T) {
}
// Try to write outside backupDir
escapePath := filepath.Join(tmpDir, "escaped.db")
err := database.BackupToSafe(escapePath, backupDir)
err := database.BackupToSafe(context.Background(), escapePath, backupDir)
if err == nil {
t.Error("BackupToSafe() should reject path outside safe root, got nil")
}
@@ -83,7 +84,7 @@ func TestBackupToSafe_RejectsSingleQuote(t *testing.T) {
t.Fatalf("MkdirAll: %v", err)
}
malicious := filepath.Join(backupDir, "evil'.db")
err := database.BackupToSafe(malicious, backupDir)
err := database.BackupToSafe(context.Background(), malicious, backupDir)
if err == nil {
t.Error("BackupToSafe() with single-quote in path should return error, got nil")
}
@@ -98,7 +99,7 @@ func TestBackupToSafe_RejectsSemicolon(t *testing.T) {
t.Fatalf("MkdirAll: %v", err)
}
malicious := filepath.Join(backupDir, "evil;drop.db")
err := database.BackupToSafe(malicious, backupDir)
err := database.BackupToSafe(context.Background(), malicious, backupDir)
if err == nil {
t.Error("BackupToSafe() with semicolon in path should return error, got nil")
}
@@ -113,7 +114,7 @@ func TestBackupToSafe_RejectsSQLComment(t *testing.T) {
t.Fatalf("MkdirAll: %v", err)
}
malicious := filepath.Join(backupDir, "evil--comment.db")
err := database.BackupToSafe(malicious, backupDir)
err := database.BackupToSafe(context.Background(), malicious, backupDir)
if err == nil {
t.Error("BackupToSafe() with '--' in path should return error, got nil")
}
@@ -128,7 +129,7 @@ func TestBackupToSafe_RejectsNullByte(t *testing.T) {
t.Fatalf("MkdirAll: %v", err)
}
malicious := filepath.Join(backupDir, "evil\x00.db") //nolint:gocritic // intentional null byte for security test
err := database.BackupToSafe(malicious, backupDir)
err := database.BackupToSafe(context.Background(), malicious, backupDir)
if err == nil {
t.Error("BackupToSafe() with null byte in path should return error, got nil")
}
@@ -144,7 +145,7 @@ func TestBackupToSafe_RejectsDoubleQuote(t *testing.T) {
t.Fatalf("MkdirAll: %v", err)
}
malicious := filepath.Join(backupDir, `evil".db`)
err := database.BackupToSafe(malicious, backupDir)
err := database.BackupToSafe(context.Background(), malicious, backupDir)
if err == nil {
t.Error("BackupToSafe() with double-quote in path should return error, got nil")
}
+11 -10
View File
@@ -1,6 +1,7 @@
package db
import (
"context"
"errors"
"fmt"
@@ -11,8 +12,8 @@ import (
// BlockUser adds a block from blocker to blocked. Idempotent — re-blocking
// a user that is already blocked is a no-op (INSERT OR IGNORE).
func (d *DB) BlockUser(blockerID, blockedID int64) error {
if err := d.q.BlockUser(dbCtx(), dbgen.BlockUserParams{
func (d *DB) BlockUser(ctx context.Context, blockerID, blockedID int64) error {
if err := d.q.BlockUser(ctx, dbgen.BlockUserParams{
BlockerID: blockerID,
BlockedID: blockedID,
}); err != nil {
@@ -23,8 +24,8 @@ func (d *DB) BlockUser(blockerID, blockedID int64) error {
// UnblockUser removes a block. Idempotent — unblocking a non-blocked user is
// a no-op.
func (d *DB) UnblockUser(blockerID, blockedID int64) error {
if err := d.q.UnblockUser(dbCtx(), dbgen.UnblockUserParams{
func (d *DB) UnblockUser(ctx context.Context, blockerID, blockedID int64) error {
if err := d.q.UnblockUser(ctx, dbgen.UnblockUserParams{
BlockerID: blockerID,
BlockedID: blockedID,
}); err != nil {
@@ -34,8 +35,8 @@ func (d *DB) UnblockUser(blockerID, blockedID int64) error {
}
// IsBlocked returns true if blockerID has blocked blockedID.
func (d *DB) IsBlocked(blockerID, blockedID int64) (bool, error) {
_, err := d.q.IsBlocked(dbCtx(), dbgen.IsBlockedParams{
func (d *DB) IsBlocked(ctx context.Context, blockerID, blockedID int64) (bool, error) {
_, err := d.q.IsBlocked(ctx, dbgen.IsBlockedParams{
BlockerID: blockerID,
BlockedID: blockedID,
})
@@ -51,8 +52,8 @@ func (d *DB) IsBlocked(blockerID, blockedID int64) (bool, error) {
// IsEitherBlocked returns true if either user has blocked the other.
// Used for DM authorization — if either party has blocked the other,
// messaging is denied.
func (d *DB) IsEitherBlocked(userA, userB int64) (bool, error) {
_, err := d.q.IsEitherBlocked(dbCtx(), dbgen.IsEitherBlockedParams{
func (d *DB) IsEitherBlocked(ctx context.Context, userA, userB int64) (bool, error) {
_, err := d.q.IsEitherBlocked(ctx, dbgen.IsEitherBlockedParams{
BlockerID: userA,
BlockedID: userB,
BlockerID_2: userB,
@@ -68,8 +69,8 @@ func (d *DB) IsEitherBlocked(userA, userB int64) (bool, error) {
}
// ListBlockedUsers returns the IDs of all users blocked by the given user.
func (d *DB) ListBlockedUsers(blockerID int64) ([]int64, error) {
ids, err := d.q.ListBlockedUsers(dbCtx(), blockerID)
func (d *DB) ListBlockedUsers(ctx context.Context, blockerID int64) ([]int64, error) {
ids, err := d.q.ListBlockedUsers(ctx, blockerID)
if err != nil {
return nil, fmt.Errorf("ListBlockedUsers: %w", err)
}
+29 -28
View File
@@ -1,6 +1,7 @@
package db
import (
"context"
"database/sql"
"errors"
"fmt"
@@ -47,21 +48,21 @@ func channelFromFields(f channelFields) Channel {
}
// ListChannels returns all channels ordered by position.
func (d *DB) ListChannels() ([]Channel, error) {
rows, err := d.q.ListChannels(dbCtx())
func (d *DB) ListChannels(ctx context.Context) ([]Channel, error) {
rows, err := d.q.ListChannels(ctx)
if err != nil {
return nil, fmt.Errorf("ListChannels: %w", err)
}
channels := make([]Channel, 0, len(rows))
for _, r := range rows {
channels = append(channels, channelFromFields(channelFields(r)))
for i := range rows {
channels = append(channels, channelFromFields(channelFields(rows[i])))
}
return channels, nil
}
// GetChannel returns the channel with the given id, or nil if not found.
func (d *DB) GetChannel(id int64) (*Channel, error) {
r, err := d.q.GetChannel(dbCtx(), id)
func (d *DB) GetChannel(ctx context.Context, id int64) (*Channel, error) {
r, err := d.q.GetChannel(ctx, id)
if errors.Is(err, sql.ErrNoRows) {
return nil, nil
}
@@ -73,8 +74,8 @@ func (d *DB) GetChannel(id int64) (*Channel, error) {
}
// CreateChannel inserts a new channel and returns the assigned ID.
func (d *DB) CreateChannel(name, chanType, category, topic string, position int) (int64, error) {
res, err := d.q.CreateChannel(dbCtx(), dbgen.CreateChannelParams{
func (d *DB) CreateChannel(ctx context.Context, name, chanType, category, topic string, position int) (int64, error) {
res, err := d.q.CreateChannel(ctx, dbgen.CreateChannelParams{
Name: name,
Type: chanType,
Category: strToNullPtr(category),
@@ -88,8 +89,8 @@ func (d *DB) CreateChannel(name, chanType, category, topic string, position int)
}
// UpdateChannel modifies name, topic, and slow_mode for the given channel.
func (d *DB) UpdateChannel(id int64, name, topic string, slowMode int) error {
if err := d.q.UpdateChannel(dbCtx(), dbgen.UpdateChannelParams{
func (d *DB) UpdateChannel(ctx context.Context, id int64, name, topic string, slowMode int) error {
if err := d.q.UpdateChannel(ctx, dbgen.UpdateChannelParams{
Name: name,
Topic: strToNullPtr(topic),
SlowMode: int64(slowMode),
@@ -101,8 +102,8 @@ func (d *DB) UpdateChannel(id int64, name, topic string, slowMode int) error {
}
// SetChannelSlowMode updates only the slow_mode field for the given channel.
func (d *DB) SetChannelSlowMode(id int64, slowMode int) error {
if err := d.q.SetChannelSlowMode(dbCtx(), dbgen.SetChannelSlowModeParams{
func (d *DB) SetChannelSlowMode(ctx context.Context, id int64, slowMode int) error {
if err := d.q.SetChannelSlowMode(ctx, dbgen.SetChannelSlowModeParams{
SlowMode: int64(slowMode),
ID: id,
}); err != nil {
@@ -112,8 +113,8 @@ func (d *DB) SetChannelSlowMode(id int64, slowMode int) error {
}
// SetChannelVoiceMaxUsers updates the voice_max_users field for the given channel.
func (d *DB) SetChannelVoiceMaxUsers(id int64, maxUsers int) error {
if err := d.q.SetChannelVoiceMaxUsers(dbCtx(), dbgen.SetChannelVoiceMaxUsersParams{
func (d *DB) SetChannelVoiceMaxUsers(ctx context.Context, id int64, maxUsers int) error {
if err := d.q.SetChannelVoiceMaxUsers(ctx, dbgen.SetChannelVoiceMaxUsersParams{
VoiceMaxUsers: int64(maxUsers),
ID: id,
}); err != nil {
@@ -123,8 +124,8 @@ func (d *DB) SetChannelVoiceMaxUsers(id int64, maxUsers int) error {
}
// DeleteChannel removes the channel row (cascades to messages, overrides, etc.).
func (d *DB) DeleteChannel(id int64) error {
if err := d.q.DeleteChannel(dbCtx(), id); err != nil {
func (d *DB) DeleteChannel(ctx context.Context, id int64) error {
if err := d.q.DeleteChannel(ctx, id); err != nil {
return fmt.Errorf("DeleteChannel: %w", err)
}
return nil
@@ -132,8 +133,8 @@ func (d *DB) DeleteChannel(id int64) error {
// GetChannelPermissions returns the allow/deny override bits for a role on a
// channel. Returns (0, 0, nil) when no override exists.
func (d *DB) GetChannelPermissions(channelID, roleID int64) (allow, deny int64, err error) {
r, scanErr := d.q.GetChannelPermission(dbCtx(), dbgen.GetChannelPermissionParams{
func (d *DB) GetChannelPermissions(ctx context.Context, channelID, roleID int64) (allow, deny int64, err error) {
r, scanErr := d.q.GetChannelPermission(ctx, dbgen.GetChannelPermissionParams{
ChannelID: channelID,
RoleID: roleID,
})
@@ -155,8 +156,8 @@ type ChannelOverride struct {
// GetAllChannelPermissionsForRole returns all channel permission overrides for
// a role in a single query, keyed by channel ID. Eliminates N+1 queries when
// filtering channels by permission.
func (d *DB) GetAllChannelPermissionsForRole(roleID int64) (map[int64]ChannelOverride, error) {
rows, err := d.q.GetRoleChannelPermissions(dbCtx(), roleID)
func (d *DB) GetAllChannelPermissionsForRole(ctx context.Context, roleID int64) (map[int64]ChannelOverride, error) {
rows, err := d.q.GetRoleChannelPermissions(ctx, roleID)
if err != nil {
return nil, fmt.Errorf("GetAllChannelPermissionsForRole: %w", err)
}
@@ -169,8 +170,8 @@ func (d *DB) GetAllChannelPermissionsForRole(roleID int64) (map[int64]ChannelOve
// UpsertChannelOverride inserts or updates the allow/deny permission override
// for a role on a channel.
func (d *DB) UpsertChannelOverride(channelID, roleID, allow, deny int64) error {
if err := d.q.UpsertChannelPermission(dbCtx(), dbgen.UpsertChannelPermissionParams{
func (d *DB) UpsertChannelOverride(ctx context.Context, channelID, roleID, allow, deny int64) error {
if err := d.q.UpsertChannelPermission(ctx, dbgen.UpsertChannelPermissionParams{
ChannelID: channelID,
RoleID: roleID,
Allow: allow,
@@ -183,8 +184,8 @@ func (d *DB) UpsertChannelOverride(channelID, roleID, allow, deny int64) error {
// DeleteChannelOverride removes the permission override for a role on a
// channel. Deleting a non-existent override is a no-op.
func (d *DB) DeleteChannelOverride(channelID, roleID int64) error {
if err := d.q.DeleteChannelPermission(dbCtx(), dbgen.DeleteChannelPermissionParams{
func (d *DB) DeleteChannelOverride(ctx context.Context, channelID, roleID int64) error {
if err := d.q.DeleteChannelPermission(ctx, dbgen.DeleteChannelPermissionParams{
ChannelID: channelID,
RoleID: roleID,
}); err != nil {
@@ -208,8 +209,8 @@ type ChannelRoleOverride struct {
// ListChannelRoleOverrides returns every role together with its override bits
// on the given channel (zero allow/deny when no override row exists), ordered
// by role position descending.
func (d *DB) ListChannelRoleOverrides(channelID int64) ([]ChannelRoleOverride, error) {
rows, err := d.sqlDB.Query(
func (d *DB) ListChannelRoleOverrides(ctx context.Context, channelID int64) ([]ChannelRoleOverride, error) {
rows, err := d.sqlDB.QueryContext(ctx,
`SELECT r.id, r.name, r.position, r.permissions,
COALESCE(o.allow, 0), COALESCE(o.deny, 0)
FROM roles r
@@ -243,7 +244,7 @@ func (d *DB) ListChannelRoleOverrides(channelID int64) ([]ChannelRoleOverride, e
// GetChannelTypes returns a map of channel ID → type string for the given IDs
// in a single query, avoiding N+1 lookups.
func (d *DB) GetChannelTypes(ids []int64) (map[int64]string, error) {
func (d *DB) GetChannelTypes(ctx context.Context, ids []int64) (map[int64]string, error) {
if len(ids) == 0 {
return map[int64]string{}, nil
}
@@ -261,7 +262,7 @@ func (d *DB) GetChannelTypes(ids []int64) (map[int64]string, error) {
strings.Join(placeholders, ","),
)
rows, err := d.sqlDB.Query(query, args...)
rows, err := d.sqlDB.QueryContext(ctx, query, args...)
if err != nil {
return nil, fmt.Errorf("GetChannelTypes query: %w", err)
}
+53 -52
View File
@@ -1,6 +1,7 @@
package db_test
import (
"context"
"testing"
"github.com/owncord/server/db"
@@ -21,7 +22,7 @@ func openMigratedMemory(t *testing.T) *db.DB {
func TestListChannels_Empty(t *testing.T) {
database := openMigratedMemory(t)
channels, err := database.ListChannels()
channels, err := database.ListChannels(context.Background())
if err != nil {
t.Fatalf("ListChannels() error: %v", err)
}
@@ -33,14 +34,14 @@ func TestListChannels_Empty(t *testing.T) {
func TestListChannels_ReturnsAll(t *testing.T) {
database := openMigratedMemory(t)
if _, err := database.CreateChannel("general", "text", "", "General chat", 0); err != nil {
if _, err := database.CreateChannel(context.Background(), "general", "text", "", "General chat", 0); err != nil {
t.Fatalf("CreateChannel general: %v", err)
}
if _, err := database.CreateChannel("announcements", "text", "", "", 1); err != nil {
if _, err := database.CreateChannel(context.Background(), "announcements", "text", "", "", 1); err != nil {
t.Fatalf("CreateChannel announcements: %v", err)
}
channels, err := database.ListChannels()
channels, err := database.ListChannels(context.Background())
if err != nil {
t.Fatalf("ListChannels() error: %v", err)
}
@@ -54,7 +55,7 @@ func TestListChannels_ReturnsAll(t *testing.T) {
func TestGetChannel_NotFound(t *testing.T) {
database := openMigratedMemory(t)
ch, err := database.GetChannel(9999)
ch, err := database.GetChannel(context.Background(), 9999)
if err != nil {
t.Fatalf("GetChannel() error: %v", err)
}
@@ -66,12 +67,12 @@ func TestGetChannel_NotFound(t *testing.T) {
func TestGetChannel_Found(t *testing.T) {
database := openMigratedMemory(t)
id, err := database.CreateChannel("general", "text", "Public", "hello", 0)
id, err := database.CreateChannel(context.Background(), "general", "text", "Public", "hello", 0)
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
ch, err := database.GetChannel(id)
ch, err := database.GetChannel(context.Background(), id)
if err != nil {
t.Fatalf("GetChannel: %v", err)
}
@@ -100,7 +101,7 @@ func TestGetChannel_Found(t *testing.T) {
func TestCreateChannel_ReturnsID(t *testing.T) {
database := openMigratedMemory(t)
id, err := database.CreateChannel("test", "text", "", "", 0)
id, err := database.CreateChannel(context.Background(), "test", "text", "", "", 0)
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
@@ -112,8 +113,8 @@ func TestCreateChannel_ReturnsID(t *testing.T) {
func TestCreateChannel_UniqueIDs(t *testing.T) {
database := openMigratedMemory(t)
id1, _ := database.CreateChannel("ch1", "text", "", "", 0)
id2, _ := database.CreateChannel("ch2", "text", "", "", 1)
id1, _ := database.CreateChannel(context.Background(), "ch1", "text", "", "", 0)
id2, _ := database.CreateChannel(context.Background(), "ch2", "text", "", "", 1)
if id1 == id2 {
t.Error("expected different IDs for different channels")
}
@@ -122,11 +123,11 @@ func TestCreateChannel_UniqueIDs(t *testing.T) {
func TestCreateChannel_EmptyCategory(t *testing.T) {
database := openMigratedMemory(t)
id, err := database.CreateChannel("nocategory", "text", "", "", 0)
id, err := database.CreateChannel(context.Background(), "nocategory", "text", "", "", 0)
if err != nil {
t.Fatalf("CreateChannel with empty category: %v", err)
}
ch, _ := database.GetChannel(id)
ch, _ := database.GetChannel(context.Background(), id)
if ch.Category != "" {
t.Errorf("Category = %q, want ''", ch.Category)
}
@@ -137,13 +138,13 @@ func TestCreateChannel_EmptyCategory(t *testing.T) {
func TestUpdateChannel_ChangesNameAndTopic(t *testing.T) {
database := openMigratedMemory(t)
id, _ := database.CreateChannel("old", "text", "", "old topic", 0)
id, _ := database.CreateChannel(context.Background(), "old", "text", "", "old topic", 0)
if err := database.UpdateChannel(id, "new", "new topic", 5); err != nil {
if err := database.UpdateChannel(context.Background(), id, "new", "new topic", 5); err != nil {
t.Fatalf("UpdateChannel: %v", err)
}
ch, _ := database.GetChannel(id)
ch, _ := database.GetChannel(context.Background(), id)
if ch.Name != "new" {
t.Errorf("Name = %q, want 'new'", ch.Name)
}
@@ -158,7 +159,7 @@ func TestUpdateChannel_ChangesNameAndTopic(t *testing.T) {
func TestUpdateChannel_NonExistent(t *testing.T) {
database := openMigratedMemory(t)
// Should not error even for non-existent row (0 rows affected is still ok).
err := database.UpdateChannel(9999, "x", "y", 0)
err := database.UpdateChannel(context.Background(), 9999, "x", "y", 0)
if err != nil {
t.Errorf("UpdateChannel non-existent should not error: %v", err)
}
@@ -169,13 +170,13 @@ func TestUpdateChannel_NonExistent(t *testing.T) {
func TestDeleteChannel_RemovesChannel(t *testing.T) {
database := openMigratedMemory(t)
id, _ := database.CreateChannel("todelete", "text", "", "", 0)
id, _ := database.CreateChannel(context.Background(), "todelete", "text", "", "", 0)
if err := database.DeleteChannel(id); err != nil {
if err := database.DeleteChannel(context.Background(), id); err != nil {
t.Fatalf("DeleteChannel: %v", err)
}
ch, err := database.GetChannel(id)
ch, err := database.GetChannel(context.Background(), id)
if err != nil {
t.Fatalf("GetChannel after delete: %v", err)
}
@@ -186,7 +187,7 @@ func TestDeleteChannel_RemovesChannel(t *testing.T) {
func TestDeleteChannel_NonExistent(t *testing.T) {
database := openMigratedMemory(t)
err := database.DeleteChannel(9999)
err := database.DeleteChannel(context.Background(), 9999)
if err != nil {
t.Errorf("DeleteChannel non-existent should not error: %v", err)
}
@@ -197,10 +198,10 @@ func TestDeleteChannel_NonExistent(t *testing.T) {
func TestGetChannelPermissions_Default(t *testing.T) {
database := openMigratedMemory(t)
chID, _ := database.CreateChannel("perms", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "perms", "text", "", "", 0)
// No override set — should return 0, 0.
allow, deny, err := database.GetChannelPermissions(chID, 4)
allow, deny, err := database.GetChannelPermissions(context.Background(), chID, 4)
if err != nil {
t.Fatalf("GetChannelPermissions: %v", err)
}
@@ -212,9 +213,9 @@ func TestGetChannelPermissions_Default(t *testing.T) {
func TestGetChannelPermissions_WithOverride(t *testing.T) {
database := openMigratedMemory(t)
chID, _ := database.CreateChannel("perms2", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "perms2", "text", "", "", 0)
// Insert an override directly.
_, err := database.Exec(
_, err := database.ExecContext(context.Background(),
`INSERT INTO channel_overrides (channel_id, role_id, allow, deny) VALUES (?, ?, ?, ?)`,
chID, 4, int64(0x400), int64(0x200),
)
@@ -222,7 +223,7 @@ func TestGetChannelPermissions_WithOverride(t *testing.T) {
t.Fatalf("insert override: %v", err)
}
allow, deny, err := database.GetChannelPermissions(chID, 4)
allow, deny, err := database.GetChannelPermissions(context.Background(), chID, 4)
if err != nil {
t.Fatalf("GetChannelPermissions: %v", err)
}
@@ -238,12 +239,12 @@ func TestGetChannelPermissions_WithOverride(t *testing.T) {
func TestUpsertChannelOverride_InsertAndUpdate(t *testing.T) {
database := openMigratedMemory(t)
chID, _ := database.CreateChannel("private", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "private", "text", "", "", 0)
if err := database.UpsertChannelOverride(chID, 4, 0, 0x202); err != nil {
if err := database.UpsertChannelOverride(context.Background(), chID, 4, 0, 0x202); err != nil {
t.Fatalf("UpsertChannelOverride insert: %v", err)
}
allow, deny, err := database.GetChannelPermissions(chID, 4)
allow, deny, err := database.GetChannelPermissions(context.Background(), chID, 4)
if err != nil {
t.Fatalf("GetChannelPermissions: %v", err)
}
@@ -252,10 +253,10 @@ func TestUpsertChannelOverride_InsertAndUpdate(t *testing.T) {
}
// Upsert again with different bits — must update, not duplicate.
if err := database.UpsertChannelOverride(chID, 4, 0x2, 0x200); err != nil {
if err := database.UpsertChannelOverride(context.Background(), chID, 4, 0x2, 0x200); err != nil {
t.Fatalf("UpsertChannelOverride update: %v", err)
}
allow, deny, err = database.GetChannelPermissions(chID, 4)
allow, deny, err = database.GetChannelPermissions(context.Background(), chID, 4)
if err != nil {
t.Fatalf("GetChannelPermissions: %v", err)
}
@@ -266,15 +267,15 @@ func TestUpsertChannelOverride_InsertAndUpdate(t *testing.T) {
func TestDeleteChannelOverride(t *testing.T) {
database := openMigratedMemory(t)
chID, _ := database.CreateChannel("private2", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "private2", "text", "", "", 0)
if err := database.UpsertChannelOverride(chID, 4, 0, 0x202); err != nil {
if err := database.UpsertChannelOverride(context.Background(), chID, 4, 0, 0x202); err != nil {
t.Fatalf("UpsertChannelOverride: %v", err)
}
if err := database.DeleteChannelOverride(chID, 4); err != nil {
if err := database.DeleteChannelOverride(context.Background(), chID, 4); err != nil {
t.Fatalf("DeleteChannelOverride: %v", err)
}
allow, deny, err := database.GetChannelPermissions(chID, 4)
allow, deny, err := database.GetChannelPermissions(context.Background(), chID, 4)
if err != nil {
t.Fatalf("GetChannelPermissions: %v", err)
}
@@ -283,20 +284,20 @@ func TestDeleteChannelOverride(t *testing.T) {
}
// Deleting again is a no-op.
if err := database.DeleteChannelOverride(chID, 4); err != nil {
if err := database.DeleteChannelOverride(context.Background(), chID, 4); err != nil {
t.Errorf("DeleteChannelOverride non-existent should not error: %v", err)
}
}
func TestListChannelRoleOverrides(t *testing.T) {
database := openMigratedMemory(t)
chID, _ := database.CreateChannel("private3", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "private3", "text", "", "", 0)
if err := database.UpsertChannelOverride(chID, 4, 0, 0x202); err != nil {
if err := database.UpsertChannelOverride(context.Background(), chID, 4, 0, 0x202); err != nil {
t.Fatalf("UpsertChannelOverride: %v", err)
}
overrides, err := database.ListChannelRoleOverrides(chID)
overrides, err := database.ListChannelRoleOverrides(context.Background(), chID)
if err != nil {
t.Fatalf("ListChannelRoleOverrides: %v", err)
}
@@ -328,13 +329,13 @@ func TestListChannelRoleOverrides(t *testing.T) {
func TestSetChannelSlowMode(t *testing.T) {
database := openMigratedMemory(t)
chID, _ := database.CreateChannel("slowch", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "slowch", "text", "", "", 0)
if err := database.SetChannelSlowMode(chID, 10); err != nil {
if err := database.SetChannelSlowMode(context.Background(), chID, 10); err != nil {
t.Fatalf("SetChannelSlowMode: %v", err)
}
ch, _ := database.GetChannel(chID)
ch, _ := database.GetChannel(context.Background(), chID)
if ch.SlowMode != 10 {
t.Errorf("SlowMode = %d, want 10", ch.SlowMode)
}
@@ -342,12 +343,12 @@ func TestSetChannelSlowMode(t *testing.T) {
func TestSetChannelSlowMode_Zero(t *testing.T) {
database := openMigratedMemory(t)
chID, _ := database.CreateChannel("slowch2", "text", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "slowch2", "text", "", "", 0)
_ = database.SetChannelSlowMode(chID, 30)
_ = database.SetChannelSlowMode(chID, 0)
_ = database.SetChannelSlowMode(context.Background(), chID, 30)
_ = database.SetChannelSlowMode(context.Background(), chID, 0)
ch, _ := database.GetChannel(chID)
ch, _ := database.GetChannel(context.Background(), chID)
if ch.SlowMode != 0 {
t.Errorf("SlowMode = %d, want 0 (disabled)", ch.SlowMode)
}
@@ -357,13 +358,13 @@ func TestSetChannelSlowMode_Zero(t *testing.T) {
func TestSetChannelVoiceMaxUsers(t *testing.T) {
database := openMigratedMemory(t)
chID, _ := database.CreateChannel("voicech", "voice", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "voicech", "voice", "", "", 0)
if err := database.SetChannelVoiceMaxUsers(chID, 25); err != nil {
if err := database.SetChannelVoiceMaxUsers(context.Background(), chID, 25); err != nil {
t.Fatalf("SetChannelVoiceMaxUsers: %v", err)
}
ch, _ := database.GetChannel(chID)
ch, _ := database.GetChannel(context.Background(), chID)
if ch.VoiceMaxUsers != 25 {
t.Errorf("VoiceMaxUsers = %d, want 25", ch.VoiceMaxUsers)
}
@@ -371,12 +372,12 @@ func TestSetChannelVoiceMaxUsers(t *testing.T) {
func TestSetChannelVoiceMaxUsers_Unlimited(t *testing.T) {
database := openMigratedMemory(t)
chID, _ := database.CreateChannel("voicech2", "voice", "", "", 0)
chID, _ := database.CreateChannel(context.Background(), "voicech2", "voice", "", "", 0)
_ = database.SetChannelVoiceMaxUsers(chID, 10)
_ = database.SetChannelVoiceMaxUsers(chID, 0)
_ = database.SetChannelVoiceMaxUsers(context.Background(), chID, 10)
_ = database.SetChannelVoiceMaxUsers(context.Background(), chID, 0)
ch, _ := database.GetChannel(chID)
ch, _ := database.GetChannel(context.Background(), chID)
if ch.VoiceMaxUsers != 0 {
t.Errorf("VoiceMaxUsers = %d, want 0 (unlimited)", ch.VoiceMaxUsers)
}
+115 -114
View File
@@ -1,6 +1,7 @@
package db_test
import (
"context"
"testing"
"time"
@@ -14,12 +15,12 @@ func TestVoice_JoinVoiceChannelIfCapacity_UnderLimit(t *testing.T) {
u1 := seedVoiceUser(t, database, "cap-u1")
chanID := seedVoiceChannel(t, database, "cap-ch")
err := database.JoinVoiceChannelIfCapacity(u1, chanID, 2)
err := database.JoinVoiceChannelIfCapacity(context.Background(), u1, chanID, 2)
if err != nil {
t.Fatalf("JoinVoiceChannelIfCapacity: %v", err)
}
state, err := database.GetVoiceState(u1)
state, err := database.GetVoiceState(context.Background(), u1)
if err != nil {
t.Fatalf("GetVoiceState: %v", err)
}
@@ -39,15 +40,15 @@ func TestVoice_JoinVoiceChannelIfCapacity_AtLimit(t *testing.T) {
chanID := seedVoiceChannel(t, database, "cap-full-ch")
// Fill channel to capacity (max 2).
if err := database.JoinVoiceChannelIfCapacity(u1, chanID, 2); err != nil {
if err := database.JoinVoiceChannelIfCapacity(context.Background(), u1, chanID, 2); err != nil {
t.Fatalf("first join: %v", err)
}
if err := database.JoinVoiceChannelIfCapacity(u2, chanID, 2); err != nil {
if err := database.JoinVoiceChannelIfCapacity(context.Background(), u2, chanID, 2); err != nil {
t.Fatalf("second join: %v", err)
}
// Third join should fail with ErrChannelFull.
err := database.JoinVoiceChannelIfCapacity(u3, chanID, 2)
err := database.JoinVoiceChannelIfCapacity(context.Background(), u3, chanID, 2)
if err == nil {
t.Fatal("expected ErrChannelFull, got nil")
}
@@ -63,14 +64,14 @@ func TestVoice_JoinVoiceChannelIfCapacity_ReplacesOwnState(t *testing.T) {
ch2 := seedVoiceChannel(t, database, "cap-ch2")
// Join ch1, then join ch2 with capacity check — should replace.
if err := database.JoinVoiceChannelIfCapacity(u1, ch1, 5); err != nil {
if err := database.JoinVoiceChannelIfCapacity(context.Background(), u1, ch1, 5); err != nil {
t.Fatalf("join ch1: %v", err)
}
if err := database.JoinVoiceChannelIfCapacity(u1, ch2, 5); err != nil {
if err := database.JoinVoiceChannelIfCapacity(context.Background(), u1, ch2, 5); err != nil {
t.Fatalf("join ch2: %v", err)
}
state, _ := database.GetVoiceState(u1)
state, _ := database.GetVoiceState(context.Background(), u1)
if state == nil || state.ChannelID != ch2 {
t.Errorf("expected channel %d, got %v", ch2, state)
}
@@ -81,7 +82,7 @@ func TestVoice_JoinVoiceChannelIfCapacity_ReplacesOwnState(t *testing.T) {
func TestVoice_GetAllVoiceStates_Empty(t *testing.T) {
database := newVoiceTestDB(t)
states, err := database.GetAllVoiceStates()
states, err := database.GetAllVoiceStates(context.Background())
if err != nil {
t.Fatalf("GetAllVoiceStates: %v", err)
}
@@ -98,11 +99,11 @@ func TestVoice_GetAllVoiceStates_MultipleChannels(t *testing.T) {
ch1 := seedVoiceChannel(t, database, "all-vs-ch1")
ch2 := seedVoiceChannel(t, database, "all-vs-ch2")
_ = database.JoinVoiceChannel(u1, ch1)
_ = database.JoinVoiceChannel(u2, ch1)
_ = database.JoinVoiceChannel(u3, ch2)
_ = database.JoinVoiceChannel(context.Background(), u1, ch1)
_ = database.JoinVoiceChannel(context.Background(), u2, ch1)
_ = database.JoinVoiceChannel(context.Background(), u3, ch2)
states, err := database.GetAllVoiceStates()
states, err := database.GetAllVoiceStates(context.Background())
if err != nil {
t.Fatalf("GetAllVoiceStates: %v", err)
}
@@ -117,7 +118,7 @@ func TestVoice_CountActiveCameras_Zero(t *testing.T) {
database := newVoiceTestDB(t)
chanID := seedVoiceChannel(t, database, "cam-count-empty")
count, err := database.CountActiveCameras(chanID)
count, err := database.CountActiveCameras(context.Background(), chanID)
if err != nil {
t.Fatalf("CountActiveCameras: %v", err)
}
@@ -133,15 +134,15 @@ func TestVoice_CountActiveCameras_SomeCameras(t *testing.T) {
u3 := seedVoiceUser(t, database, "cam-cnt-u3")
chanID := seedVoiceChannel(t, database, "cam-cnt-ch")
_ = database.JoinVoiceChannel(u1, chanID)
_ = database.JoinVoiceChannel(u2, chanID)
_ = database.JoinVoiceChannel(u3, chanID)
_ = database.JoinVoiceChannel(context.Background(), u1, chanID)
_ = database.JoinVoiceChannel(context.Background(), u2, chanID)
_ = database.JoinVoiceChannel(context.Background(), u3, chanID)
_ = database.UpdateVoiceCamera(u1, true)
_ = database.UpdateVoiceCamera(u2, true)
_ = database.UpdateVoiceCamera(context.Background(), u1, true)
_ = database.UpdateVoiceCamera(context.Background(), u2, true)
// u3 camera stays off.
count, err := database.CountActiveCameras(chanID)
count, err := database.CountActiveCameras(context.Background(), chanID)
if err != nil {
t.Fatalf("CountActiveCameras: %v", err)
}
@@ -157,9 +158,9 @@ func TestVoice_EnableCameraIfUnderLimit_Success(t *testing.T) {
u1 := seedVoiceUser(t, database, "cam-limit-ok")
chanID := seedVoiceChannel(t, database, "cam-limit-ch")
_ = database.JoinVoiceChannel(u1, chanID)
_ = database.JoinVoiceChannel(context.Background(), u1, chanID)
ok, err := database.EnableCameraIfUnderLimit(u1, chanID, 2)
ok, err := database.EnableCameraIfUnderLimit(context.Background(), u1, chanID, 2)
if err != nil {
t.Fatalf("EnableCameraIfUnderLimit: %v", err)
}
@@ -167,7 +168,7 @@ func TestVoice_EnableCameraIfUnderLimit_Success(t *testing.T) {
t.Error("expected camera to be enabled")
}
state, _ := database.GetVoiceState(u1)
state, _ := database.GetVoiceState(context.Background(), u1)
if state == nil || !state.Camera {
t.Error("camera should be true after enable")
}
@@ -180,16 +181,16 @@ func TestVoice_EnableCameraIfUnderLimit_AtLimit(t *testing.T) {
u3 := seedVoiceUser(t, database, "cam-lim-u3")
chanID := seedVoiceChannel(t, database, "cam-lim-ch")
_ = database.JoinVoiceChannel(u1, chanID)
_ = database.JoinVoiceChannel(u2, chanID)
_ = database.JoinVoiceChannel(u3, chanID)
_ = database.JoinVoiceChannel(context.Background(), u1, chanID)
_ = database.JoinVoiceChannel(context.Background(), u2, chanID)
_ = database.JoinVoiceChannel(context.Background(), u3, chanID)
// Enable cameras for u1 and u2 (max is 2).
_, _ = database.EnableCameraIfUnderLimit(u1, chanID, 2)
_, _ = database.EnableCameraIfUnderLimit(u2, chanID, 2)
_, _ = database.EnableCameraIfUnderLimit(context.Background(), u1, chanID, 2)
_, _ = database.EnableCameraIfUnderLimit(context.Background(), u2, chanID, 2)
// u3 should be denied.
ok, err := database.EnableCameraIfUnderLimit(u3, chanID, 2)
ok, err := database.EnableCameraIfUnderLimit(context.Background(), u3, chanID, 2)
if err != nil {
t.Fatalf("EnableCameraIfUnderLimit: %v", err)
}
@@ -207,12 +208,12 @@ func TestSearchMessagesInChannels_FindsInAllowedChannels(t *testing.T) {
ch2 := seedChannel(t, database, "srch-ch2")
ch3 := seedChannel(t, database, "srch-ch3")
_, _ = database.CreateMessage(ch1, userID, "alpha keyword here", nil)
_, _ = database.CreateMessage(ch2, userID, "beta keyword here", nil)
_, _ = database.CreateMessage(ch3, userID, "gamma keyword here", nil)
_, _ = database.CreateMessage(context.Background(), ch1, userID, "alpha keyword here", nil)
_, _ = database.CreateMessage(context.Background(), ch2, userID, "beta keyword here", nil)
_, _ = database.CreateMessage(context.Background(), ch3, userID, "gamma keyword here", nil)
// Search only in ch1 and ch2.
results, err := database.SearchMessagesInChannels("keyword", []int64{ch1, ch2}, 10)
results, err := database.SearchMessagesInChannels(context.Background(), "keyword", []int64{ch1, ch2}, 10)
if err != nil {
t.Fatalf("SearchMessagesInChannels: %v", err)
}
@@ -229,7 +230,7 @@ func TestSearchMessagesInChannels_FindsInAllowedChannels(t *testing.T) {
func TestSearchMessagesInChannels_EmptyQuery(t *testing.T) {
database := openMigratedMemory(t)
results, err := database.SearchMessagesInChannels("", []int64{1}, 10)
results, err := database.SearchMessagesInChannels(context.Background(), "", []int64{1}, 10)
if err != nil {
t.Fatalf("SearchMessagesInChannels: %v", err)
}
@@ -241,7 +242,7 @@ func TestSearchMessagesInChannels_EmptyQuery(t *testing.T) {
func TestSearchMessagesInChannels_EmptyChannelIDs(t *testing.T) {
database := openMigratedMemory(t)
results, err := database.SearchMessagesInChannels("test", nil, 10)
results, err := database.SearchMessagesInChannels(context.Background(), "test", nil, 10)
if err != nil {
t.Fatalf("SearchMessagesInChannels: %v", err)
}
@@ -256,10 +257,10 @@ func TestSearchMessagesInChannels_LimitRespected(t *testing.T) {
ch1 := seedChannel(t, database, "srch-lim-ch")
for range 5 {
_, _ = database.CreateMessage(ch1, userID, "findme content here", nil)
_, _ = database.CreateMessage(context.Background(), ch1, userID, "findme content here", nil)
}
results, err := database.SearchMessagesInChannels("findme", []int64{ch1}, 2)
results, err := database.SearchMessagesInChannels(context.Background(), "findme", []int64{ch1}, 2)
if err != nil {
t.Fatalf("SearchMessagesInChannels: %v", err)
}
@@ -271,7 +272,7 @@ func TestSearchMessagesInChannels_LimitRespected(t *testing.T) {
func TestSearchMessagesInChannels_ZeroLimit(t *testing.T) {
database := openMigratedMemory(t)
results, err := database.SearchMessagesInChannels("test", []int64{1}, 0)
results, err := database.SearchMessagesInChannels(context.Background(), "test", []int64{1}, 0)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
@@ -287,7 +288,7 @@ func TestGetPinnedMessages_Empty(t *testing.T) {
userID := seedUser(t, database, "pin-empty-u")
chID := seedChannel(t, database, "pin-empty")
msgs, err := database.GetPinnedMessages(chID, userID)
msgs, err := database.GetPinnedMessages(context.Background(), chID, userID)
if err != nil {
t.Fatalf("GetPinnedMessages: %v", err)
}
@@ -301,11 +302,11 @@ func TestGetPinnedMessages_ReturnsPinnedOnly(t *testing.T) {
userID := seedUser(t, database, "pin-user")
chID := seedChannel(t, database, "pin-ch")
id1, _ := database.CreateMessage(chID, userID, "pinned msg", nil)
_, _ = database.CreateMessage(chID, userID, "not pinned", nil)
_ = database.SetMessagePinned(id1, true)
id1, _ := database.CreateMessage(context.Background(), chID, userID, "pinned msg", nil)
_, _ = database.CreateMessage(context.Background(), chID, userID, "not pinned", nil)
_ = database.SetMessagePinned(context.Background(), id1, true)
msgs, err := database.GetPinnedMessages(chID, userID)
msgs, err := database.GetPinnedMessages(context.Background(), chID, userID)
if err != nil {
t.Fatalf("GetPinnedMessages: %v", err)
}
@@ -327,13 +328,13 @@ func TestSetMessagePinned_Pin(t *testing.T) {
userID := seedUser(t, database, "setpin-u")
chID := seedChannel(t, database, "setpin-ch")
id, _ := database.CreateMessage(chID, userID, "to pin", nil)
id, _ := database.CreateMessage(context.Background(), chID, userID, "to pin", nil)
if err := database.SetMessagePinned(id, true); err != nil {
if err := database.SetMessagePinned(context.Background(), id, true); err != nil {
t.Fatalf("SetMessagePinned(true): %v", err)
}
msg, _ := database.GetMessage(id)
msg, _ := database.GetMessage(context.Background(), id)
if msg == nil || !msg.Pinned {
t.Error("message should be pinned")
}
@@ -344,13 +345,13 @@ func TestSetMessagePinned_Unpin(t *testing.T) {
userID := seedUser(t, database, "unpin-u")
chID := seedChannel(t, database, "unpin-ch")
id, _ := database.CreateMessage(chID, userID, "to unpin", nil)
_ = database.SetMessagePinned(id, true)
if err := database.SetMessagePinned(id, false); err != nil {
id, _ := database.CreateMessage(context.Background(), chID, userID, "to unpin", nil)
_ = database.SetMessagePinned(context.Background(), id, true)
if err := database.SetMessagePinned(context.Background(), id, false); err != nil {
t.Fatalf("SetMessagePinned(false): %v", err)
}
msg, _ := database.GetMessage(id)
msg, _ := database.GetMessage(context.Background(), id)
if msg == nil || msg.Pinned {
t.Error("message should not be pinned")
}
@@ -359,7 +360,7 @@ func TestSetMessagePinned_Unpin(t *testing.T) {
func TestSetMessagePinned_NotFound(t *testing.T) {
database := openMigratedMemory(t)
err := database.SetMessagePinned(99999, true)
err := database.SetMessagePinned(context.Background(), 99999, true)
if err == nil {
t.Error("expected error for non-existent message")
}
@@ -370,10 +371,10 @@ func TestSetMessagePinned_DeletedMessage(t *testing.T) {
userID := seedUser(t, database, "pin-del-u")
chID := seedChannel(t, database, "pin-del-ch")
id, _ := database.CreateMessage(chID, userID, "deleted", nil)
_ = database.DeleteMessage(id, userID, false)
id, _ := database.CreateMessage(context.Background(), chID, userID, "deleted", nil)
_ = database.DeleteMessage(context.Background(), id, userID, false)
err := database.SetMessagePinned(id, true)
err := database.SetMessagePinned(context.Background(), id, true)
if err == nil {
t.Error("expected error when pinning deleted message")
}
@@ -385,12 +386,12 @@ func TestCreateAttachment_Success(t *testing.T) {
database := openMigratedMemory(t)
userID := seedUser(t, database, "att-uploader")
err := database.CreateAttachment("att-001", userID, "photo.png", "stored-001.png", "image/png", 12345, nil, nil)
err := database.CreateAttachment(context.Background(), "att-001", userID, "photo.png", "stored-001.png", "image/png", 12345, nil, nil)
if err != nil {
t.Fatalf("CreateAttachment: %v", err)
}
att, err := database.GetAttachmentByID("att-001")
att, err := database.GetAttachmentByID(context.Background(), "att-001")
if err != nil {
t.Fatalf("GetAttachmentByID: %v", err)
}
@@ -413,12 +414,12 @@ func TestCreateAttachment_WithDimensions(t *testing.T) {
userID := seedUser(t, database, "att-dim-uploader")
w, h := 1920, 1080
err := database.CreateAttachment("att-dim", userID, "photo.jpg", "stored-dim.jpg", "image/jpeg", 54321, &w, &h)
err := database.CreateAttachment(context.Background(), "att-dim", userID, "photo.jpg", "stored-dim.jpg", "image/jpeg", 54321, &w, &h)
if err != nil {
t.Fatalf("CreateAttachment with dims: %v", err)
}
att, _ := database.GetAttachmentByID("att-dim")
att, _ := database.GetAttachmentByID(context.Background(), "att-dim")
if att == nil {
t.Fatal("expected attachment")
}
@@ -431,10 +432,10 @@ func TestDeleteOrphanedAttachments_RemovesOrphans(t *testing.T) {
userID := seedUser(t, database, "orphan-uploader")
// Create an unlinked attachment (message_id IS NULL).
_ = database.CreateAttachment("orphan-1", userID, "file.txt", "stored-orphan.txt", "text/plain", 100, nil, nil)
_ = database.CreateAttachment(context.Background(), "orphan-1", userID, "file.txt", "stored-orphan.txt", "text/plain", 100, nil, nil)
// Use a cutoff far in the future so the attachment is considered old.
files, err := database.DeleteOrphanedAttachments("2099-01-01T00:00:00Z")
files, err := database.DeleteOrphanedAttachments(context.Background(), "2099-01-01T00:00:00Z")
if err != nil {
t.Fatalf("DeleteOrphanedAttachments: %v", err)
}
@@ -446,7 +447,7 @@ func TestDeleteOrphanedAttachments_RemovesOrphans(t *testing.T) {
}
// Should be removed from DB.
att, _ := database.GetAttachmentByID("orphan-1")
att, _ := database.GetAttachmentByID(context.Background(), "orphan-1")
if att != nil {
t.Error("orphaned attachment should be deleted from DB")
}
@@ -458,11 +459,11 @@ func TestDeleteOrphanedAttachments_KeepsLinked(t *testing.T) {
chID := seedChannel(t, database, "orphan-linked-ch")
// Create attachment and link it to a message.
_ = database.CreateAttachment("linked-1", userID, "file.txt", "stored-linked.txt", "text/plain", 100, nil, nil)
msgID, _ := database.CreateMessage(chID, userID, "with attachment", nil)
_, _ = database.LinkAttachmentsToMessage(msgID, userID, []string{"linked-1"})
_ = database.CreateAttachment(context.Background(), "linked-1", userID, "file.txt", "stored-linked.txt", "text/plain", 100, nil, nil)
msgID, _ := database.CreateMessage(context.Background(), chID, userID, "with attachment", nil)
_, _ = database.LinkAttachmentsToMessage(context.Background(), msgID, userID, []string{"linked-1"})
files, err := database.DeleteOrphanedAttachments("2099-01-01T00:00:00Z")
files, err := database.DeleteOrphanedAttachments(context.Background(), "2099-01-01T00:00:00Z")
if err != nil {
t.Fatalf("DeleteOrphanedAttachments: %v", err)
}
@@ -475,10 +476,10 @@ func TestDeleteOrphanedAttachments_CutoffRespected(t *testing.T) {
database := openMigratedMemory(t)
userID := seedUser(t, database, "cutoff-uploader")
_ = database.CreateAttachment("future-1", userID, "file.txt", "stored-future.txt", "text/plain", 100, nil, nil)
_ = database.CreateAttachment(context.Background(), "future-1", userID, "file.txt", "stored-future.txt", "text/plain", 100, nil, nil)
// Cutoff in the past — newly created attachment should NOT be deleted.
files, err := database.DeleteOrphanedAttachments("2000-01-01T00:00:00Z")
files, err := database.DeleteOrphanedAttachments(context.Background(), "2000-01-01T00:00:00Z")
if err != nil {
t.Fatalf("DeleteOrphanedAttachments: %v", err)
}
@@ -492,7 +493,7 @@ func TestDeleteOrphanedAttachments_CutoffRespected(t *testing.T) {
func TestGetAllChannelPermissionsForRole_Empty(t *testing.T) {
database := openMigratedMemory(t)
result, err := database.GetAllChannelPermissionsForRole(4)
result, err := database.GetAllChannelPermissionsForRole(context.Background(), 4)
if err != nil {
t.Fatalf("GetAllChannelPermissionsForRole: %v", err)
}
@@ -504,20 +505,20 @@ func TestGetAllChannelPermissionsForRole_Empty(t *testing.T) {
func TestGetAllChannelPermissionsForRole_WithOverrides(t *testing.T) {
database := openMigratedMemory(t)
ch1, _ := database.CreateChannel("perm-ch1", "text", "", "", 0)
ch2, _ := database.CreateChannel("perm-ch2", "text", "", "", 0)
ch1, _ := database.CreateChannel(context.Background(), "perm-ch1", "text", "", "", 0)
ch2, _ := database.CreateChannel(context.Background(), "perm-ch2", "text", "", "", 0)
// Insert overrides for role 4.
_, _ = database.Exec(
_, _ = database.ExecContext(context.Background(),
`INSERT INTO channel_overrides (channel_id, role_id, allow, deny) VALUES (?, ?, ?, ?)`,
ch1, 4, int64(0x100), int64(0x200),
)
_, _ = database.Exec(
_, _ = database.ExecContext(context.Background(),
`INSERT INTO channel_overrides (channel_id, role_id, allow, deny) VALUES (?, ?, ?, ?)`,
ch2, 4, int64(0x300), int64(0),
)
result, err := database.GetAllChannelPermissionsForRole(4)
result, err := database.GetAllChannelPermissionsForRole(context.Background(), 4)
if err != nil {
t.Fatalf("GetAllChannelPermissionsForRole: %v", err)
}
@@ -534,7 +535,7 @@ func TestGetAllChannelPermissionsForRole_WithOverrides(t *testing.T) {
func TestGetChannelTypes_Empty(t *testing.T) {
database := openMigratedMemory(t)
result, err := database.GetChannelTypes(nil)
result, err := database.GetChannelTypes(context.Background(), nil)
if err != nil {
t.Fatalf("GetChannelTypes: %v", err)
}
@@ -546,10 +547,10 @@ func TestGetChannelTypes_Empty(t *testing.T) {
func TestGetChannelTypes_ReturnsTypes(t *testing.T) {
database := openMigratedMemory(t)
ch1, _ := database.CreateChannel("type-text", "text", "", "", 0)
ch2, _ := database.CreateChannel("type-voice", "voice", "", "", 0)
ch1, _ := database.CreateChannel(context.Background(), "type-text", "text", "", "", 0)
ch2, _ := database.CreateChannel(context.Background(), "type-voice", "voice", "", "", 0)
result, err := database.GetChannelTypes([]int64{ch1, ch2})
result, err := database.GetChannelTypes(context.Background(), []int64{ch1, ch2})
if err != nil {
t.Fatalf("GetChannelTypes: %v", err)
}
@@ -564,7 +565,7 @@ func TestGetChannelTypes_ReturnsTypes(t *testing.T) {
func TestGetChannelTypes_NonExistentIDs(t *testing.T) {
database := openMigratedMemory(t)
result, err := database.GetChannelTypes([]int64{99999})
result, err := database.GetChannelTypes(context.Background(), []int64{99999})
if err != nil {
t.Fatalf("GetChannelTypes: %v", err)
}
@@ -577,10 +578,10 @@ func TestGetChannelTypes_NonExistentIDs(t *testing.T) {
func TestCountUsersWithoutTOTP_AllWithout(t *testing.T) {
database := openMigratedMemory(t)
_, _ = database.CreateUser("totp-u1", "hash", 4)
_, _ = database.CreateUser("totp-u2", "hash", 4)
_, _ = database.CreateUser(context.Background(), "totp-u1", "hash", 4)
_, _ = database.CreateUser(context.Background(), "totp-u2", "hash", 4)
count, err := database.CountUsersWithoutTOTP()
count, err := database.CountUsersWithoutTOTP(context.Background())
if err != nil {
t.Fatalf("CountUsersWithoutTOTP: %v", err)
}
@@ -591,13 +592,13 @@ func TestCountUsersWithoutTOTP_AllWithout(t *testing.T) {
func TestCountUsersWithoutTOTP_WithTOTPSetup(t *testing.T) {
database := openMigratedMemory(t)
uid, _ := database.CreateUser("totp-with", "hash", 4)
_, _ = database.CreateUser("totp-without", "hash", 4)
uid, _ := database.CreateUser(context.Background(), "totp-with", "hash", 4)
_, _ = database.CreateUser(context.Background(), "totp-without", "hash", 4)
secret := "JBSWY3DPEHPK3PXP"
_ = database.UpdateUserTOTPSecret(uid, &secret)
_ = database.UpdateUserTOTPSecret(context.Background(), uid, &secret)
count, err := database.CountUsersWithoutTOTP()
count, err := database.CountUsersWithoutTOTP(context.Background())
if err != nil {
t.Fatalf("CountUsersWithoutTOTP: %v", err)
}
@@ -610,14 +611,14 @@ func TestCountUsersWithoutTOTP_WithTOTPSetup(t *testing.T) {
func TestUpdateUserTOTPSecret_Set(t *testing.T) {
database := openMigratedMemory(t)
uid, _ := database.CreateUser("totp-set", "hash", 4)
uid, _ := database.CreateUser(context.Background(), "totp-set", "hash", 4)
secret := "JBSWY3DPEHPK3PXP"
if err := database.UpdateUserTOTPSecret(uid, &secret); err != nil {
if err := database.UpdateUserTOTPSecret(context.Background(), uid, &secret); err != nil {
t.Fatalf("UpdateUserTOTPSecret(set): %v", err)
}
user, _ := database.GetUserByID(uid)
user, _ := database.GetUserByID(context.Background(), uid)
if user == nil || user.TOTPSecret == nil || *user.TOTPSecret != secret {
t.Error("TOTP secret should be set")
}
@@ -625,15 +626,15 @@ func TestUpdateUserTOTPSecret_Set(t *testing.T) {
func TestUpdateUserTOTPSecret_Clear(t *testing.T) {
database := openMigratedMemory(t)
uid, _ := database.CreateUser("totp-clear", "hash", 4)
uid, _ := database.CreateUser(context.Background(), "totp-clear", "hash", 4)
secret := "JBSWY3DPEHPK3PXP"
_ = database.UpdateUserTOTPSecret(uid, &secret)
if err := database.UpdateUserTOTPSecret(uid, nil); err != nil {
_ = database.UpdateUserTOTPSecret(context.Background(), uid, &secret)
if err := database.UpdateUserTOTPSecret(context.Background(), uid, nil); err != nil {
t.Fatalf("UpdateUserTOTPSecret(clear): %v", err)
}
user, _ := database.GetUserByID(uid)
user, _ := database.GetUserByID(context.Background(), uid)
if user == nil || user.TOTPSecret != nil {
t.Error("TOTP secret should be nil after clear")
}
@@ -644,14 +645,14 @@ func TestUpdateUserTOTPSecret_Clear(t *testing.T) {
func TestCreateUserWithInvite_Success(t *testing.T) {
database := openMigratedMemory(t)
// Create a user who will create the invite.
creatorID, _ := database.CreateUser("invite-creator", "hash", 2)
creatorID, _ := database.CreateUser(context.Background(), "invite-creator", "hash", 2)
code, err := database.CreateInvite(creatorID, 5, nil)
code, err := database.CreateInvite(context.Background(), creatorID, 5, nil)
if err != nil {
t.Fatalf("CreateInvite: %v", err)
}
uid, err := database.CreateUserWithInvite("newuser", "hash", 4, code)
uid, err := database.CreateUserWithInvite(context.Background(), "newuser", "hash", 4, code)
if err != nil {
t.Fatalf("CreateUserWithInvite: %v", err)
}
@@ -660,7 +661,7 @@ func TestCreateUserWithInvite_Success(t *testing.T) {
}
// Verify invite use count incremented.
inv, _ := database.GetInvite(code)
inv, _ := database.GetInvite(context.Background(), code)
if inv == nil || inv.Uses != 1 {
t.Errorf("invite uses = %v, want 1", inv)
}
@@ -669,7 +670,7 @@ func TestCreateUserWithInvite_Success(t *testing.T) {
func TestCreateUserWithInvite_InvalidCode(t *testing.T) {
database := openMigratedMemory(t)
_, err := database.CreateUserWithInvite("baduser", "hash", 4, "nonexistent-code")
_, err := database.CreateUserWithInvite(context.Background(), "baduser", "hash", 4, "nonexistent-code")
if err == nil {
t.Error("expected error for invalid invite code")
}
@@ -677,12 +678,12 @@ func TestCreateUserWithInvite_InvalidCode(t *testing.T) {
func TestCreateUserWithInvite_RevokedInvite(t *testing.T) {
database := openMigratedMemory(t)
creatorID, _ := database.CreateUser("inv-revoke-creator", "hash", 2)
creatorID, _ := database.CreateUser(context.Background(), "inv-revoke-creator", "hash", 2)
code, _ := database.CreateInvite(creatorID, 0, nil)
_ = database.RevokeInvite(code)
code, _ := database.CreateInvite(context.Background(), creatorID, 0, nil)
_ = database.RevokeInvite(context.Background(), code)
_, err := database.CreateUserWithInvite("revokeduser", "hash", 4, code)
_, err := database.CreateUserWithInvite(context.Background(), "revokeduser", "hash", 4, code)
if err == nil {
t.Error("expected error for revoked invite")
}
@@ -690,13 +691,13 @@ func TestCreateUserWithInvite_RevokedInvite(t *testing.T) {
func TestCreateUserWithInvite_ExpiredInvite(t *testing.T) {
database := openMigratedMemory(t)
creatorID, _ := database.CreateUser("inv-expire-creator", "hash", 2)
creatorID, _ := database.CreateUser(context.Background(), "inv-expire-creator", "hash", 2)
// Create an invite that expires in the past.
pastTime := time.Now().Add(-1 * time.Hour)
code, _ := database.CreateInvite(creatorID, 0, &pastTime)
code, _ := database.CreateInvite(context.Background(), creatorID, 0, &pastTime)
_, err := database.CreateUserWithInvite("expireduser", "hash", 4, code)
_, err := database.CreateUserWithInvite(context.Background(), "expireduser", "hash", 4, code)
if err == nil {
t.Error("expected error for expired invite")
}
@@ -707,7 +708,7 @@ func TestCreateUserWithInvite_ExpiredInvite(t *testing.T) {
func TestListInvites_DB_Empty(t *testing.T) {
database := openMigratedMemory(t)
invites, err := database.ListInvites()
invites, err := database.ListInvites(context.Background())
if err != nil {
t.Fatalf("ListInvites: %v", err)
}
@@ -718,12 +719,12 @@ func TestListInvites_DB_Empty(t *testing.T) {
func TestListInvites_DB_ReturnsAll(t *testing.T) {
database := openMigratedMemory(t)
creatorID, _ := database.CreateUser("list-inv-creator", "hash", 2)
creatorID, _ := database.CreateUser(context.Background(), "list-inv-creator", "hash", 2)
_, _ = database.CreateInvite(creatorID, 5, nil)
_, _ = database.CreateInvite(creatorID, 0, nil)
_, _ = database.CreateInvite(context.Background(), creatorID, 5, nil)
_, _ = database.CreateInvite(context.Background(), creatorID, 0, nil)
invites, err := database.ListInvites()
invites, err := database.ListInvites(context.Background())
if err != nil {
t.Fatalf("ListInvites: %v", err)
}
@@ -735,7 +736,7 @@ func TestListInvites_DB_ReturnsAll(t *testing.T) {
func TestUseInviteAtomic_NonExistent(t *testing.T) {
database := openMigratedMemory(t)
err := database.UseInviteAtomic("does-not-exist")
err := database.UseInviteAtomic(context.Background(), "does-not-exist")
if err == nil {
t.Error("expected error for non-existent invite")
}
@@ -746,7 +747,7 @@ func TestUseInviteAtomic_NonExistent(t *testing.T) {
func TestSearchMessages_EmptyQuery(t *testing.T) {
database := openMigratedMemory(t)
results, err := database.SearchMessages("", nil, 10)
results, err := database.SearchMessages(context.Background(), "", nil, 10)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
@@ -758,7 +759,7 @@ func TestSearchMessages_EmptyQuery(t *testing.T) {
func TestSearchMessages_ZeroLimit(t *testing.T) {
database := openMigratedMemory(t)
results, err := database.SearchMessages("test", nil, 0)
results, err := database.SearchMessages(context.Background(), "test", nil, 0)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
@@ -772,10 +773,10 @@ func TestSearchMessages_SpecialCharsStripped(t *testing.T) {
userID := seedUser(t, database, "srch-special")
chID := seedChannel(t, database, "srch-special-ch")
_, _ = database.CreateMessage(chID, userID, "hello world content", nil)
_, _ = database.CreateMessage(context.Background(), chID, userID, "hello world content", nil)
// FTS special chars should be stripped, leaving a valid query.
results, err := database.SearchMessages("hello* \"world\"", nil, 10)
results, err := database.SearchMessages(context.Background(), "hello* \"world\"", nil, 10)
if err != nil {
t.Fatalf("SearchMessages with special chars: %v", err)
}
-25
View File
@@ -25,11 +25,6 @@ type DB struct {
q *dbgen.Queries
}
// dbCtx is the context used for delegated dbgen calls. The public db.DB API is
// context-free today; callers that need cancellation use the *Context helpers
// directly. Using Background here preserves the existing behavior exactly.
func dbCtx() context.Context { return context.Background() }
// Open opens (or creates) a SQLite database at path, enables WAL mode and
// foreign key enforcement, and returns a ready-to-use DB.
func Open(path string) (*DB, error) {
@@ -104,41 +99,21 @@ func (d *DB) Close() error {
return d.sqlDB.Close()
}
// QueryRow executes a query that returns at most one row.
func (d *DB) QueryRow(query string, args ...any) *sql.Row {
return d.sqlDB.QueryRow(query, args...)
}
// QueryRowContext executes a query that returns at most one row, with context.
func (d *DB) QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row {
return d.sqlDB.QueryRowContext(ctx, query, args...)
}
// Exec executes a query that doesn't return rows.
func (d *DB) Exec(query string, args ...any) (sql.Result, error) {
return d.sqlDB.Exec(query, args...)
}
// ExecContext executes a query that doesn't return rows, with context.
func (d *DB) ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) {
return d.sqlDB.ExecContext(ctx, query, args...)
}
// Query executes a query that returns multiple rows.
func (d *DB) Query(query string, args ...any) (*sql.Rows, error) {
return d.sqlDB.Query(query, args...)
}
// QueryContext executes a query that returns multiple rows, with context.
func (d *DB) QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error) {
return d.sqlDB.QueryContext(ctx, query, args...)
}
// Begin starts a database transaction.
func (d *DB) Begin() (*sql.Tx, error) {
return d.sqlDB.Begin()
}
// BeginTx starts a database transaction with context and options.
func (d *DB) BeginTx(ctx context.Context, opts *sql.TxOptions) (*sql.Tx, error) {
return d.sqlDB.BeginTx(ctx, opts)
+17 -16
View File
@@ -1,6 +1,7 @@
package db_test
import (
"context"
"database/sql"
"fmt"
"io"
@@ -59,7 +60,7 @@ func TestWALModeEnabled(t *testing.T) {
database := openMemory(t)
var journalMode string
err := database.QueryRow("PRAGMA journal_mode;").Scan(&journalMode)
err := database.QueryRowContext(context.Background(), "PRAGMA journal_mode;").Scan(&journalMode)
if err != nil {
t.Fatalf("PRAGMA journal_mode query error: %v", err)
}
@@ -82,7 +83,7 @@ func TestWALModeEnabledOnFile(t *testing.T) {
defer database.Close() //nolint:errcheck
var journalMode string
if err := database.QueryRow("PRAGMA journal_mode;").Scan(&journalMode); err != nil {
if err := database.QueryRowContext(context.Background(), "PRAGMA journal_mode;").Scan(&journalMode); err != nil {
t.Fatalf("PRAGMA journal_mode query error: %v", err)
}
if journalMode != "wal" {
@@ -94,7 +95,7 @@ func TestForeignKeysEnabled(t *testing.T) {
database := openMemory(t)
var fkEnabled int
if err := database.QueryRow("PRAGMA foreign_keys;").Scan(&fkEnabled); err != nil {
if err := database.QueryRowContext(context.Background(), "PRAGMA foreign_keys;").Scan(&fkEnabled); err != nil {
t.Fatalf("PRAGMA foreign_keys query error: %v", err)
}
if fkEnabled != 1 {
@@ -118,7 +119,7 @@ func TestMigrateCreatesAllTables(t *testing.T) {
for _, table := range expectedTables {
t.Run(table, func(t *testing.T) {
var name string
err := database.QueryRow(
err := database.QueryRowContext(context.Background(),
"SELECT name FROM sqlite_master WHERE type='table' AND name=?",
table,
).Scan(&name)
@@ -139,7 +140,7 @@ func TestMigrateCreatesFTSTable(t *testing.T) {
}
var name string
err := database.QueryRow(
err := database.QueryRowContext(context.Background(),
"SELECT name FROM sqlite_master WHERE type='table' AND name='messages_fts'",
).Scan(&name)
if err == sql.ErrNoRows {
@@ -169,7 +170,7 @@ func TestMigrateInsertsDefaultRoles(t *testing.T) {
}
var count int
if err := database.QueryRow("SELECT COUNT(*) FROM roles").Scan(&count); err != nil {
if err := database.QueryRowContext(context.Background(), "SELECT COUNT(*) FROM roles").Scan(&count); err != nil {
t.Fatalf("COUNT roles error: %v", err)
}
if count < 4 {
@@ -185,7 +186,7 @@ func TestMigrateInsertsDefaultSettings(t *testing.T) {
}
var value string
err := database.QueryRow("SELECT value FROM settings WHERE key='registration_open'").Scan(&value)
err := database.QueryRowContext(context.Background(), "SELECT value FROM settings WHERE key='registration_open'").Scan(&value)
if err != nil {
t.Fatalf("settings query error: %v", err)
}
@@ -211,7 +212,7 @@ func TestMigrateCreatesIndexes(t *testing.T) {
for _, idx := range expectedIndexes {
t.Run(idx, func(t *testing.T) {
var name string
err := database.QueryRow(
err := database.QueryRowContext(context.Background(),
"SELECT name FROM sqlite_master WHERE type='index' AND name=?",
idx,
).Scan(&name)
@@ -244,7 +245,7 @@ func TestQueryRow(t *testing.T) {
// Verify we can run a simple query via the exposed DB.
var schemaVersion string
err := database.QueryRow("SELECT value FROM settings WHERE key='schema_version'").Scan(&schemaVersion)
err := database.QueryRowContext(context.Background(), "SELECT value FROM settings WHERE key='schema_version'").Scan(&schemaVersion)
if err != nil {
t.Fatalf("QueryRow error: %v", err)
}
@@ -261,13 +262,13 @@ func TestExec(t *testing.T) {
}
// Insert a settings row using Exec.
_, err := database.Exec("INSERT OR REPLACE INTO settings (key, value) VALUES (?, ?)", "test_key", "test_val")
_, err := database.ExecContext(context.Background(), "INSERT OR REPLACE INTO settings (key, value) VALUES (?, ?)", "test_key", "test_val")
if err != nil {
t.Fatalf("Exec() error: %v", err)
}
var val string
if err := database.QueryRow("SELECT value FROM settings WHERE key='test_key'").Scan(&val); err != nil {
if err := database.QueryRowContext(context.Background(), "SELECT value FROM settings WHERE key='test_key'").Scan(&val); err != nil {
t.Fatalf("QueryRow after Exec error: %v", err)
}
if val != "test_val" {
@@ -282,7 +283,7 @@ func TestQuery(t *testing.T) {
t.Fatalf("Migrate() error: %v", err)
}
rows, err := database.Query("SELECT key FROM settings")
rows, err := database.QueryContext(context.Background(), "SELECT key FROM settings")
if err != nil {
t.Fatalf("Query() error: %v", err)
}
@@ -308,7 +309,7 @@ func TestBegin(t *testing.T) {
t.Fatalf("Migrate() error: %v", err)
}
tx, err := database.Begin()
tx, err := database.BeginTx(context.Background(), nil)
if err != nil {
t.Fatalf("Begin() error: %v", err)
}
@@ -325,7 +326,7 @@ func TestBegin(t *testing.T) {
// After rollback, tx_key should not exist.
var val string
err = database.QueryRow("SELECT value FROM settings WHERE key='tx_key'").Scan(&val)
err = database.QueryRowContext(context.Background(), "SELECT value FROM settings WHERE key='tx_key'").Scan(&val)
if err == nil {
t.Error("tx_key should not exist after rollback")
}
@@ -433,7 +434,7 @@ func TestMigrateFSSkipsNonSQL(t *testing.T) {
// The table from the .sql file should exist.
var name string
if err := database.QueryRow(
if err := database.QueryRowContext(context.Background(),
"SELECT name FROM sqlite_master WHERE type='table' AND name='test_skip'",
).Scan(&name); err != nil {
t.Error("table test_skip not found after MigrateFS")
@@ -465,7 +466,7 @@ func TestMigrateWALAndFKOnFile(t *testing.T) {
// Tables should exist.
var name string
if err := database.QueryRow(
if err := database.QueryRowContext(context.Background(),
"SELECT name FROM sqlite_master WHERE type='table' AND name='users'",
).Scan(&name); err != nil {
t.Errorf("users table not found after migration on file db: %v", err)
-23
View File
@@ -7,7 +7,6 @@ package dbgen
import (
"context"
"database/sql"
)
const createAttachment = `-- name: CreateAttachment :exec
@@ -40,15 +39,6 @@ func (q *Queries) CreateAttachment(ctx context.Context, arg CreateAttachmentPara
return err
}
const deleteAttachment = `-- name: DeleteAttachment :exec
DELETE FROM attachments WHERE id = ?
`
func (q *Queries) DeleteAttachment(ctx context.Context, id string) error {
_, err := q.db.ExecContext(ctx, deleteAttachment, id)
return err
}
const deleteOrphanedAttachments = `-- name: DeleteOrphanedAttachments :many
DELETE FROM attachments WHERE message_id IS NULL AND uploaded_at < ? RETURNING stored_as
`
@@ -147,16 +137,3 @@ func (q *Queries) GetAttachmentWithChannel(ctx context.Context, id string) (GetA
)
return i, err
}
const linkAttachmentToMessage = `-- name: LinkAttachmentToMessage :execresult
UPDATE attachments SET message_id = ? WHERE id = ? AND message_id IS NULL
`
type LinkAttachmentToMessageParams struct {
MessageID *int64 `json:"messageId"`
ID string `json:"id"`
}
func (q *Queries) LinkAttachmentToMessage(ctx context.Context, arg LinkAttachmentToMessageParams) (sql.Result, error) {
return q.db.ExecContext(ctx, linkAttachmentToMessage, arg.MessageID, arg.ID)
}
-56
View File
@@ -37,20 +37,6 @@ func (q *Queries) AdminUpdateChannel(ctx context.Context, arg AdminUpdateChannel
return err
}
const archiveChannel = `-- name: ArchiveChannel :exec
UPDATE channels SET archived = ? WHERE id = ?
`
type ArchiveChannelParams struct {
Archived int64 `json:"archived"`
ID int64 `json:"id"`
}
func (q *Queries) ArchiveChannel(ctx context.Context, arg ArchiveChannelParams) error {
_, err := q.db.ExecContext(ctx, archiveChannel, arg.Archived, arg.ID)
return err
}
const createChannel = `-- name: CreateChannel :execresult
INSERT INTO channels (name, type, category, topic, position) VALUES (?, ?, ?, ?, ?)
`
@@ -260,20 +246,6 @@ func (q *Queries) ListChannels(ctx context.Context) ([]ListChannelsRow, error) {
return items, nil
}
const setChannelMixingThreshold = `-- name: SetChannelMixingThreshold :exec
UPDATE channels SET mixing_threshold = ? WHERE id = ?
`
type SetChannelMixingThresholdParams struct {
MixingThreshold *int64 `json:"mixingThreshold"`
ID int64 `json:"id"`
}
func (q *Queries) SetChannelMixingThreshold(ctx context.Context, arg SetChannelMixingThresholdParams) error {
_, err := q.db.ExecContext(ctx, setChannelMixingThreshold, arg.MixingThreshold, arg.ID)
return err
}
const setChannelSlowMode = `-- name: SetChannelSlowMode :exec
UPDATE channels SET slow_mode = ? WHERE id = ?
`
@@ -302,34 +274,6 @@ func (q *Queries) SetChannelVoiceMaxUsers(ctx context.Context, arg SetChannelVoi
return err
}
const setChannelVoiceMaxVideo = `-- name: SetChannelVoiceMaxVideo :exec
UPDATE channels SET voice_max_video = ? WHERE id = ?
`
type SetChannelVoiceMaxVideoParams struct {
VoiceMaxVideo int64 `json:"voiceMaxVideo"`
ID int64 `json:"id"`
}
func (q *Queries) SetChannelVoiceMaxVideo(ctx context.Context, arg SetChannelVoiceMaxVideoParams) error {
_, err := q.db.ExecContext(ctx, setChannelVoiceMaxVideo, arg.VoiceMaxVideo, arg.ID)
return err
}
const setChannelVoiceQuality = `-- name: SetChannelVoiceQuality :exec
UPDATE channels SET voice_quality = ? WHERE id = ?
`
type SetChannelVoiceQualityParams struct {
VoiceQuality *string `json:"voiceQuality"`
ID int64 `json:"id"`
}
func (q *Queries) SetChannelVoiceQuality(ctx context.Context, arg SetChannelVoiceQualityParams) error {
_, err := q.db.ExecContext(ctx, setChannelVoiceQuality, arg.VoiceQuality, arg.ID)
return err
}
const updateChannel = `-- name: UpdateChannel :exec
UPDATE channels SET name = ?, topic = ?, slow_mode = ? WHERE id = ?
`
-72
View File
@@ -7,7 +7,6 @@ package dbgen
import (
"context"
"database/sql"
)
const closeDM = `-- name: CloseDM :exec
@@ -24,27 +23,6 @@ func (q *Queries) CloseDM(ctx context.Context, arg CloseDMParams) error {
return err
}
const findExistingDMChannel = `-- name: FindExistingDMChannel :one
SELECT dp1.channel_id
FROM dm_participants dp1
JOIN dm_participants dp2 ON dp1.channel_id = dp2.channel_id
JOIN channels c ON c.id = dp1.channel_id
WHERE dp1.user_id = ? AND dp2.user_id = ? AND c.type = 'dm'
LIMIT 1
`
type FindExistingDMChannelParams struct {
UserID int64 `json:"userId"`
UserID_2 int64 `json:"userId2"`
}
func (q *Queries) FindExistingDMChannel(ctx context.Context, arg FindExistingDMChannelParams) (int64, error) {
row := q.db.QueryRowContext(ctx, findExistingDMChannel, arg.UserID, arg.UserID_2)
var channel_id int64
err := row.Scan(&channel_id)
return channel_id, err
}
const getDMParticipantIDs = `-- name: GetDMParticipantIDs :many
SELECT user_id FROM dm_participants WHERE channel_id = ?
`
@@ -149,56 +127,6 @@ func (q *Queries) GetUserDMChannels(ctx context.Context, arg GetUserDMChannelsPa
return items, nil
}
const insertDMChannel = `-- name: InsertDMChannel :execresult
INSERT INTO channels (name, type) VALUES ('', 'dm')
`
func (q *Queries) InsertDMChannel(ctx context.Context) (sql.Result, error) {
return q.db.ExecContext(ctx, insertDMChannel)
}
const insertDMOpenState = `-- name: InsertDMOpenState :exec
INSERT OR IGNORE INTO dm_open_state (user_id, channel_id) VALUES (?, ?), (?, ?)
`
type InsertDMOpenStateParams struct {
UserID int64 `json:"userId"`
ChannelID int64 `json:"channelId"`
UserID_2 int64 `json:"userId2"`
ChannelID_2 int64 `json:"channelId2"`
}
func (q *Queries) InsertDMOpenState(ctx context.Context, arg InsertDMOpenStateParams) error {
_, err := q.db.ExecContext(ctx, insertDMOpenState,
arg.UserID,
arg.ChannelID,
arg.UserID_2,
arg.ChannelID_2,
)
return err
}
const insertDMParticipants = `-- name: InsertDMParticipants :exec
INSERT INTO dm_participants (channel_id, user_id) VALUES (?, ?), (?, ?)
`
type InsertDMParticipantsParams struct {
ChannelID int64 `json:"channelId"`
UserID int64 `json:"userId"`
ChannelID_2 int64 `json:"channelId2"`
UserID_2 int64 `json:"userId2"`
}
func (q *Queries) InsertDMParticipants(ctx context.Context, arg InsertDMParticipantsParams) error {
_, err := q.db.ExecContext(ctx, insertDMParticipants,
arg.ChannelID,
arg.UserID,
arg.ChannelID_2,
arg.UserID_2,
)
return err
}
const isDMParticipant = `-- name: IsDMParticipant :one
SELECT user_id FROM dm_participants WHERE user_id = ? AND channel_id = ?
`
-247
View File
@@ -117,133 +117,6 @@ func (q *Queries) GetMessage(ctx context.Context, id int64) (Message, error) {
return i, err
}
const getMessagesByChannel = `-- name: GetMessagesByChannel :many
SELECT m.id, m.channel_id, m.user_id, m.content, m.reply_to,
m.edited_at, m.deleted, m.pinned, m.timestamp,
u.username, u.avatar
FROM messages m JOIN users u ON m.user_id = u.id
WHERE m.channel_id = ? AND m.deleted = 0
ORDER BY m.id DESC LIMIT ?
`
type GetMessagesByChannelParams struct {
ChannelID int64 `json:"channelId"`
Limit int64 `json:"limit"`
}
type GetMessagesByChannelRow struct {
ID int64 `json:"id"`
ChannelID int64 `json:"channelId"`
UserID int64 `json:"userId"`
Content string `json:"content"`
ReplyTo *int64 `json:"replyTo"`
EditedAt *string `json:"editedAt"`
Deleted int64 `json:"deleted"`
Pinned int64 `json:"pinned"`
Timestamp string `json:"timestamp"`
Username string `json:"username"`
Avatar *string `json:"avatar"`
}
func (q *Queries) GetMessagesByChannel(ctx context.Context, arg GetMessagesByChannelParams) ([]GetMessagesByChannelRow, error) {
rows, err := q.db.QueryContext(ctx, getMessagesByChannel, arg.ChannelID, arg.Limit)
if err != nil {
return nil, err
}
defer rows.Close()
items := []GetMessagesByChannelRow{}
for rows.Next() {
var i GetMessagesByChannelRow
if err := rows.Scan(
&i.ID,
&i.ChannelID,
&i.UserID,
&i.Content,
&i.ReplyTo,
&i.EditedAt,
&i.Deleted,
&i.Pinned,
&i.Timestamp,
&i.Username,
&i.Avatar,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const getMessagesByChannelBeforeCursor = `-- name: GetMessagesByChannelBeforeCursor :many
SELECT m.id, m.channel_id, m.user_id, m.content, m.reply_to,
m.edited_at, m.deleted, m.pinned, m.timestamp,
u.username, u.avatar
FROM messages m JOIN users u ON m.user_id = u.id
WHERE m.channel_id = ? AND m.id < ? AND m.deleted = 0
ORDER BY m.id DESC LIMIT ?
`
type GetMessagesByChannelBeforeCursorParams struct {
ChannelID int64 `json:"channelId"`
ID int64 `json:"id"`
Limit int64 `json:"limit"`
}
type GetMessagesByChannelBeforeCursorRow struct {
ID int64 `json:"id"`
ChannelID int64 `json:"channelId"`
UserID int64 `json:"userId"`
Content string `json:"content"`
ReplyTo *int64 `json:"replyTo"`
EditedAt *string `json:"editedAt"`
Deleted int64 `json:"deleted"`
Pinned int64 `json:"pinned"`
Timestamp string `json:"timestamp"`
Username string `json:"username"`
Avatar *string `json:"avatar"`
}
func (q *Queries) GetMessagesByChannelBeforeCursor(ctx context.Context, arg GetMessagesByChannelBeforeCursorParams) ([]GetMessagesByChannelBeforeCursorRow, error) {
rows, err := q.db.QueryContext(ctx, getMessagesByChannelBeforeCursor, arg.ChannelID, arg.ID, arg.Limit)
if err != nil {
return nil, err
}
defer rows.Close()
items := []GetMessagesByChannelBeforeCursorRow{}
for rows.Next() {
var i GetMessagesByChannelBeforeCursorRow
if err := rows.Scan(
&i.ID,
&i.ChannelID,
&i.UserID,
&i.Content,
&i.ReplyTo,
&i.EditedAt,
&i.Deleted,
&i.Pinned,
&i.Timestamp,
&i.Username,
&i.Avatar,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const getMessagesForAPI = `-- name: GetMessagesForAPI :many
SELECT m.id, m.channel_id, m.user_id, u.username, u.avatar,
m.content, m.reply_to, m.edited_at, m.deleted, m.pinned, m.timestamp
@@ -306,126 +179,6 @@ func (q *Queries) GetMessagesForAPI(ctx context.Context, arg GetMessagesForAPIPa
return items, nil
}
const getMessagesForAPIBeforeCursor = `-- name: GetMessagesForAPIBeforeCursor :many
SELECT m.id, m.channel_id, m.user_id, u.username, u.avatar,
m.content, m.reply_to, m.edited_at, m.deleted, m.pinned, m.timestamp
FROM messages m JOIN users u ON m.user_id = u.id
WHERE m.channel_id = ? AND m.id < ? AND m.deleted = 0
ORDER BY m.id DESC LIMIT ?
`
type GetMessagesForAPIBeforeCursorParams struct {
ChannelID int64 `json:"channelId"`
ID int64 `json:"id"`
Limit int64 `json:"limit"`
}
type GetMessagesForAPIBeforeCursorRow struct {
ID int64 `json:"id"`
ChannelID int64 `json:"channelId"`
UserID int64 `json:"userId"`
Username string `json:"username"`
Avatar *string `json:"avatar"`
Content string `json:"content"`
ReplyTo *int64 `json:"replyTo"`
EditedAt *string `json:"editedAt"`
Deleted int64 `json:"deleted"`
Pinned int64 `json:"pinned"`
Timestamp string `json:"timestamp"`
}
func (q *Queries) GetMessagesForAPIBeforeCursor(ctx context.Context, arg GetMessagesForAPIBeforeCursorParams) ([]GetMessagesForAPIBeforeCursorRow, error) {
rows, err := q.db.QueryContext(ctx, getMessagesForAPIBeforeCursor, arg.ChannelID, arg.ID, arg.Limit)
if err != nil {
return nil, err
}
defer rows.Close()
items := []GetMessagesForAPIBeforeCursorRow{}
for rows.Next() {
var i GetMessagesForAPIBeforeCursorRow
if err := rows.Scan(
&i.ID,
&i.ChannelID,
&i.UserID,
&i.Username,
&i.Avatar,
&i.Content,
&i.ReplyTo,
&i.EditedAt,
&i.Deleted,
&i.Pinned,
&i.Timestamp,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const getPinnedMessageRows = `-- name: GetPinnedMessageRows :many
SELECT m.id, m.channel_id, m.user_id, u.username, u.avatar,
m.content, m.reply_to, m.edited_at, m.deleted, m.pinned, m.timestamp
FROM messages m JOIN users u ON m.user_id = u.id
WHERE m.channel_id = ? AND m.pinned = 1 AND m.deleted = 0
ORDER BY m.id DESC
`
type GetPinnedMessageRowsRow struct {
ID int64 `json:"id"`
ChannelID int64 `json:"channelId"`
UserID int64 `json:"userId"`
Username string `json:"username"`
Avatar *string `json:"avatar"`
Content string `json:"content"`
ReplyTo *int64 `json:"replyTo"`
EditedAt *string `json:"editedAt"`
Deleted int64 `json:"deleted"`
Pinned int64 `json:"pinned"`
Timestamp string `json:"timestamp"`
}
func (q *Queries) GetPinnedMessageRows(ctx context.Context, channelID int64) ([]GetPinnedMessageRowsRow, error) {
rows, err := q.db.QueryContext(ctx, getPinnedMessageRows, channelID)
if err != nil {
return nil, err
}
defer rows.Close()
items := []GetPinnedMessageRowsRow{}
for rows.Next() {
var i GetPinnedMessageRowsRow
if err := rows.Scan(
&i.ID,
&i.ChannelID,
&i.UserID,
&i.Username,
&i.Avatar,
&i.Content,
&i.ReplyTo,
&i.EditedAt,
&i.Deleted,
&i.Pinned,
&i.Timestamp,
); err != nil {
return nil, err
}
items = append(items, i)
}
if err := rows.Close(); err != nil {
return nil, err
}
if err := rows.Err(); err != nil {
return nil, err
}
return items, nil
}
const setMessagePinned = `-- name: SetMessagePinned :execresult
UPDATE messages SET pinned = ? WHERE id = ? AND deleted = 0
`

Some files were not shown because too many files have changed in this diff Show More