mirror of
https://github.com/J3vb/OwnCord.git
synced 2026-09-03 03:50:00 +03:00
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:
File diff suppressed because it is too large
Load Diff
@@ -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-058–067) 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"
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -4,7 +4,8 @@
|
||||
"project": ["src/**/*.ts"],
|
||||
"ignore": [
|
||||
"public/**",
|
||||
"src-tauri/**"
|
||||
"src-tauri/**",
|
||||
"src/lib/protocolTypes.ts"
|
||||
],
|
||||
"ignoreDependencies": [
|
||||
"@tauri-apps/cli"
|
||||
|
||||
Generated
+3
-32
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
|
||||
Generated
-67
@@ -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"
|
||||
|
||||
@@ -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",
|
||||
{
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
|
||||
@@ -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 "";
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
@@ -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. */
|
||||
|
||||
@@ -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>;
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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));
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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 }
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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 {
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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
@@ -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{
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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!"})
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
@@ -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())
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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,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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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 = ?
|
||||
`
|
||||
|
||||
@@ -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 = ?
|
||||
`
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user