From ab389764b594007b5d2b28c3f27944b84e2432f5 Mon Sep 17 00:00:00 2001 From: jevb Date: Sat, 14 Mar 2026 21:31:03 +0100 Subject: [PATCH] feat: implement Phase 5 (voice/WebRTC signaling) and Phase 6 (admin panel) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Phase 5 — Voice: - migrations/002_voice_states.sql: voice_states table with FK + index - db/voice_queries: JoinVoiceChannel, LeaveVoiceChannel, GetVoiceState, GetChannelVoiceStates, UpdateVoiceMute, UpdateVoiceDeafen, ClearVoiceState - ws/voice_handlers: handleVoiceJoin (perm check, DB, broadcast existing states), handleVoiceLeave, handleVoiceMute, handleVoiceDeafen, handleVoiceSignal (rate-limited relay, SDP never logged), handleSoundboard (rate-limited, USE_SOUNDBOARD perm check) - ws/handlers: dispatch voice_join/leave/mute/deafen/offer/answer/ice/soundboard - ws/serve: call handleVoiceLeave on disconnect; include voice states in ready payload - ws/messages: buildVoiceState, buildVoiceLeave, buildVoiceSignalRelay - api/voice_handler: GET /api/v1/voice/credentials — HMAC-SHA1 TURN creds - config: VoiceConfig (TURNSecret, STUNPort, TURNPort, TURNEnabled) Phase 6 — Admin Panel: - migrations/003_audit_log.sql: audit_log table with indexes - db/admin_queries: GetServerStats, ListAllUsers, UpdateUserRole, ForceLogoutUser, AdminCreate/Update/DeleteChannel, LogAudit, GetAuditLog, GetSetting, SetSetting, GetAllSettings, BackupTo - admin/api: full REST API — stats, users, channels, audit log, settings, backup; adminAuthMiddleware (ADMINISTRATOR bit), ownerOnlyMiddleware - admin/static/index.html: single-page admin panel (dark theme, vanilla JS, no CDN) — dashboard, users, channels, audit log, settings sections - admin/admin.go: NewHandler wiring go:embed static files + API Fixes: Channel struct json tags (was serializing as "ID" not "id"), duplicate getWithToken helper renamed in voice_handler_test.go Test coverage: admin 59.1%, api 78.2%, auth 90.9%, db 82.0%, ws 37.9% --- Server/admin/admin.go | 37 +- Server/admin/api.go | 496 ++++++++++++++++ Server/admin/api_test.go | 669 +++++++++++++++++++++ Server/admin/static/index.html | 688 +++++++++++++++++++++- Server/api/router.go | 7 + Server/api/voice_handler.go | 121 ++++ Server/api/voice_handler_test.go | 344 +++++++++++ Server/config/config.go | 14 + Server/db/admin_queries.go | 307 ++++++++++ Server/db/admin_queries_test.go | 729 +++++++++++++++++++++++ Server/db/auth_queries.go | 12 + Server/db/models.go | 57 +- Server/db/voice_queries.go | 172 ++++++ Server/db/voice_queries_test.go | 423 ++++++++++++++ Server/migrations/002_voice_states.sql | 12 + Server/migrations/003_audit_log.sql | 45 ++ Server/ws/handlers.go | 12 + Server/ws/messages.go | 49 ++ Server/ws/serve.go | 31 +- Server/ws/voice_handlers.go | 193 +++++++ Server/ws/voice_handlers_test.go | 767 +++++++++++++++++++++++++ 21 files changed, 5161 insertions(+), 24 deletions(-) create mode 100644 Server/admin/api.go create mode 100644 Server/admin/api_test.go create mode 100644 Server/api/voice_handler.go create mode 100644 Server/api/voice_handler_test.go create mode 100644 Server/db/admin_queries.go create mode 100644 Server/db/admin_queries_test.go create mode 100644 Server/db/voice_queries.go create mode 100644 Server/db/voice_queries_test.go create mode 100644 Server/migrations/002_voice_states.sql create mode 100644 Server/migrations/003_audit_log.sql create mode 100644 Server/ws/voice_handlers.go create mode 100644 Server/ws/voice_handlers_test.go diff --git a/Server/admin/admin.go b/Server/admin/admin.go index 550f7d81..bb703272 100644 --- a/Server/admin/admin.go +++ b/Server/admin/admin.go @@ -1,16 +1,47 @@ -// Package admin provides the embedded admin panel static file server. -// Full admin API implementation follows in Phase 6. +// Package admin provides the embedded admin panel static file server and the +// admin REST API for the OwnCord server. package admin import ( "embed" + "io/fs" "net/http" + + "github.com/go-chi/chi/v5" + "github.com/owncord/server/db" ) //go:embed static var staticFiles embed.FS -// Handler returns an http.Handler that serves the embedded admin panel static files. +// NewHandler returns an http.Handler that serves both the admin REST API and +// the embedded admin panel static files. +// +// Routes: +// +// /api/* — admin REST API (all require ADMINISTRATOR permission) +// /* — embedded static files (SPA; index.html for unknown paths) +func NewHandler(database *db.DB) http.Handler { + r := chi.NewRouter() + + // Admin REST API mounted at /api + r.Mount("/api", NewAdminAPI(database)) + + // Static files — serve from the embedded FS sub-tree. + staticFS, err := fs.Sub(staticFiles, "admin/static") + if err != nil { + // This is a programming error (wrong embed path) and should never + // happen in production. Panic so it surfaces immediately in tests. + panic("admin: failed to create static sub-FS: " + err.Error()) + } + r.Handle("/*", http.FileServer(http.FS(staticFS))) + + 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)) } diff --git a/Server/admin/api.go b/Server/admin/api.go new file mode 100644 index 00000000..19764719 --- /dev/null +++ b/Server/admin/api.go @@ -0,0 +1,496 @@ +package admin + +import ( + "encoding/json" + "net/http" + "os" + "path/filepath" + "strconv" + "strings" + "time" + + "github.com/go-chi/chi/v5" + "github.com/owncord/server/auth" + "github.com/owncord/server/db" +) + +// ─── Permission constants ───────────────────────────────────────────────────── + +const ( + permAdministrator = int64(0x40000000) + ownerRolePosition = 100 +) + +// ─── NewAdminAPI ────────────────────────────────────────────────────────────── + +// NewAdminAPI returns a chi router with all /admin/api/* routes. All routes +// are protected by adminAuthMiddleware which requires the ADMINISTRATOR bit. +func NewAdminAPI(database *db.DB) http.Handler { + r := chi.NewRouter() + + // All routes require authentication and ADMINISTRATOR permission. + r.Use(adminAuthMiddleware(database)) + + r.Get("/stats", handleGetStats(database)) + r.Get("/users", handleListUsers(database)) + r.Patch("/users/{id}", handlePatchUser(database)) + r.Delete("/users/{id}/sessions", handleForceLogout(database)) + r.Get("/channels", handleListChannels(database)) + r.Post("/channels", handleCreateChannel(database)) + r.Patch("/channels/{id}", handlePatchChannel(database)) + r.Delete("/channels/{id}", handleDeleteChannel(database)) + r.Get("/audit-log", handleGetAuditLog(database)) + r.Get("/settings", handleGetSettings(database)) + r.Patch("/settings", handlePatchSettings(database)) + r.Post("/backup", http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + ownerOnlyMiddleware(database, handleBackup(database)).ServeHTTP(w, req) + })) + + return r +} + +// ─── Middleware ─────────────────────────────────────────────────────────────── + +// adminAuthMiddleware validates the Bearer token and requires ADMINISTRATOR. +func adminAuthMiddleware(database *db.DB) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + token, ok := extractBearer(r) + if !ok { + writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "missing or invalid authorization header") + return + } + + hash := auth.HashToken(token) + sess, err := database.GetSessionByTokenHash(hash) + if err != nil || sess == nil { + writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "invalid or expired session") + return + } + + if isExpired(sess.ExpiresAt) { + writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "session has expired") + return + } + + user, err := database.GetUserByID(sess.UserID) + if err != nil || user == nil { + writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "user not found") + return + } + + role, err := database.GetRoleByID(user.RoleID) + if err != nil || role == nil { + writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "role not found") + return + } + + if role.Permissions&permAdministrator == 0 { + writeErr(w, http.StatusForbidden, "FORBIDDEN", "administrator permission required") + return + } + + next.ServeHTTP(w, r) + }) + } +} + +// ownerOnlyMiddleware wraps a handler to require Owner role (position == 100). +func ownerOnlyMiddleware(database *db.DB, next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + token, ok := extractBearer(r) + if !ok { + writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "missing authorization header") + return + } + + hash := auth.HashToken(token) + sess, err := database.GetSessionByTokenHash(hash) + if err != nil || sess == nil { + writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "invalid session") + return + } + + user, err := database.GetUserByID(sess.UserID) + if err != nil || user == nil { + writeErr(w, http.StatusUnauthorized, "UNAUTHORIZED", "user not found") + return + } + + role, err := database.GetRoleByID(user.RoleID) + if err != nil || role == nil { + writeErr(w, http.StatusForbidden, "FORBIDDEN", "role not found") + return + } + + if role.Position < ownerRolePosition { + writeErr(w, http.StatusForbidden, "FORBIDDEN", "owner role required") + return + } + + next.ServeHTTP(w, r) + }) +} + +// ─── Handlers ───────────────────────────────────────────────────────────────── + +func handleGetStats(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + stats, err := database.GetServerStats() + if err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to get stats") + return + } + writeJSON(w, http.StatusOK, stats) + } +} + +func handleListUsers(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + limit := queryInt(r, "limit", 50) + offset := queryInt(r, "offset", 0) + + users, err := database.ListAllUsers(limit, offset) + if err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to list users") + return + } + writeJSON(w, http.StatusOK, users) + } +} + +// patchUserRequest is the JSON body for PATCH /admin/api/users/{id}. +type patchUserRequest struct { + RoleID *int64 `json:"role_id"` + Banned *bool `json:"banned"` + BanReason *string `json:"ban_reason"` +} + +func handlePatchUser(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + id, err := pathInt64(r, "id") + if err != nil { + writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid user id") + return + } + + var req patchUserRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid request body") + return + } + + user, err := database.GetUserByID(id) + if err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch user") + return + } + if user == nil { + writeErr(w, http.StatusNotFound, "NOT_FOUND", "user not found") + return + } + + if req.RoleID != nil { + if err := database.UpdateUserRole(id, *req.RoleID); err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to update role") + return + } + } + + if req.Banned != nil { + reason := "" + if req.BanReason != nil { + reason = *req.BanReason + } + if *req.Banned { + if err := database.BanUser(id, reason, nil); err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to ban user") + return + } + } else { + if err := database.UnbanUser(id); err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to unban user") + return + } + } + } + + updated, err := database.GetUserByID(id) + if err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch updated user") + return + } + writeJSON(w, http.StatusOK, updated) + } +} + +func handleForceLogout(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + id, err := pathInt64(r, "id") + if err != nil { + writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid user id") + return + } + + if err := database.ForceLogoutUser(id); err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to logout user") + return + } + w.WriteHeader(http.StatusNoContent) + } +} + +func handleListChannels(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + channels, err := database.ListChannels() + if err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to list channels") + return + } + writeJSON(w, http.StatusOK, channels) + } +} + +// createChannelRequest is the JSON body for POST /admin/api/channels. +type createChannelRequest struct { + Name string `json:"name"` + Type string `json:"type"` + Category string `json:"category"` + Topic string `json:"topic"` + Position int `json:"position"` +} + +func handleCreateChannel(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + var req createChannelRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid request body") + return + } + + if strings.TrimSpace(req.Name) == "" { + writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "name is required") + return + } + if req.Type == "" { + req.Type = "text" + } + + id, err := database.AdminCreateChannel(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) + if err != nil || ch == nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch created channel") + return + } + writeJSON(w, http.StatusCreated, ch) + } +} + +// updateChannelRequest is the JSON body for PATCH /admin/api/channels/{id}. +type updateChannelRequest struct { + Name string `json:"name"` + Topic string `json:"topic"` + SlowMode int `json:"slow_mode"` + Position int `json:"position"` + Archived bool `json:"archived"` +} + +func handlePatchChannel(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + id, err := pathInt64(r, "id") + if err != nil { + writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid channel id") + return + } + + existing, err := database.GetChannel(id) + if err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch channel") + return + } + if existing == nil { + writeErr(w, http.StatusNotFound, "NOT_FOUND", "channel not found") + return + } + + // Start from existing values so a partial body is safe. + req := updateChannelRequest{ + Name: existing.Name, + Topic: existing.Topic, + SlowMode: existing.SlowMode, + Position: existing.Position, + Archived: existing.Archived, + } + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid request body") + return + } + + if err := database.AdminUpdateChannel(id, req.Name, req.Topic, req.SlowMode, req.Position, req.Archived); err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to update channel") + return + } + + updated, _ := database.GetChannel(id) + writeJSON(w, http.StatusOK, updated) + } +} + +func handleDeleteChannel(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + id, err := pathInt64(r, "id") + if err != nil { + writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid channel id") + return + } + + existing, err := database.GetChannel(id) + if err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch channel") + return + } + if existing == nil { + writeErr(w, http.StatusNotFound, "NOT_FOUND", "channel not found") + return + } + + if err := database.AdminDeleteChannel(id); err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to delete channel") + return + } + w.WriteHeader(http.StatusNoContent) + } +} + +func handleGetAuditLog(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + limit := queryInt(r, "limit", 50) + offset := queryInt(r, "offset", 0) + + entries, err := database.GetAuditLog(limit, offset) + if err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to get audit log") + return + } + writeJSON(w, http.StatusOK, entries) + } +} + +func handleGetSettings(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + settings, err := database.GetAllSettings() + if err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to get settings") + return + } + writeJSON(w, http.StatusOK, settings) + } +} + +func handlePatchSettings(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + var updates map[string]string + if err := json.NewDecoder(r.Body).Decode(&updates); err != nil { + writeErr(w, http.StatusBadRequest, "BAD_REQUEST", "invalid request body") + return + } + + for key, value := range updates { + if err := database.SetSetting(key, value); err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to update setting: "+key) + return + } + } + + settings, err := database.GetAllSettings() + if err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to fetch settings") + return + } + writeJSON(w, http.StatusOK, settings) + } +} + +func handleBackup(database *db.DB) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + backupDir := filepath.Join("data", "backups") + if err := os.MkdirAll(backupDir, 0o750); err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "failed to create backup directory") + return + } + + timestamp := time.Now().UTC().Format("20060102_150405") + backupPath := filepath.Join(backupDir, "chatserver_"+timestamp+".db") + + if err := database.BackupTo(backupPath); err != nil { + writeErr(w, http.StatusInternalServerError, "INTERNAL_ERROR", "backup failed") + return + } + + writeJSON(w, http.StatusOK, map[string]string{ + "path": backupPath, + "created": timestamp, + }) + }) +} + +// ─── Helpers ────────────────────────────────────────────────────────────────── + +type errorResponse struct { + Error string `json:"error"` + Message string `json:"message"` +} + +func writeJSON(w http.ResponseWriter, status int, v interface{}) { + w.Header().Set("Content-Type", "application/json; charset=utf-8") + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(v) +} + +func writeErr(w http.ResponseWriter, status int, code, msg string) { + writeJSON(w, status, errorResponse{Error: code, Message: msg}) +} + +func extractBearer(r *http.Request) (string, bool) { + header := r.Header.Get("Authorization") + if header == "" { + return "", false + } + parts := strings.SplitN(header, " ", 2) + if len(parts) != 2 || !strings.EqualFold(parts[0], "bearer") || parts[1] == "" { + return "", false + } + return parts[1], true +} + +func pathInt64(r *http.Request, param string) (int64, error) { + raw := chi.URLParam(r, param) + return strconv.ParseInt(raw, 10, 64) +} + +func queryInt(r *http.Request, key string, defaultVal int) int { + raw := r.URL.Query().Get(key) + if raw == "" { + return defaultVal + } + n, err := strconv.Atoi(raw) + if err != nil || n < 0 { + return defaultVal + } + return n +} + +func isExpired(expiresAt string) bool { + for _, layout := range []string{"2006-01-02 15:04:05", "2006-01-02T15:04:05Z"} { + t, err := time.Parse(layout, expiresAt) + if err == nil { + return time.Now().UTC().After(t.UTC()) + } + } + return true +} diff --git a/Server/admin/api_test.go b/Server/admin/api_test.go new file mode 100644 index 00000000..b5c5d6cf --- /dev/null +++ b/Server/admin/api_test.go @@ -0,0 +1,669 @@ +package admin_test + +import ( + "bytes" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "testing" + "testing/fstest" + + "github.com/owncord/server/admin" + "github.com/owncord/server/auth" + "github.com/owncord/server/db" +) + +// adminSchema is a minimal in-memory schema for admin API tests. +var adminSchema = []byte(` +CREATE TABLE IF NOT EXISTS roles ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL UNIQUE, + color TEXT, + permissions INTEGER NOT NULL DEFAULT 0, + position INTEGER NOT NULL DEFAULT 0, + is_default INTEGER NOT NULL DEFAULT 0 +); + +INSERT OR IGNORE INTO roles (id, name, color, permissions, position, is_default) VALUES + (1, 'Owner', '#E74C3C', 2147483647, 100, 0), + (2, 'Admin', '#F39C12', 1073741823, 80, 0), + (3, 'Member', NULL, 1049089, 40, 1); + +CREATE TABLE IF NOT EXISTS users ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + username TEXT NOT NULL UNIQUE COLLATE NOCASE, + password TEXT NOT NULL, + avatar TEXT, + role_id INTEGER NOT NULL DEFAULT 3 REFERENCES roles(id), + totp_secret TEXT, + status TEXT NOT NULL DEFAULT 'offline', + created_at TEXT NOT NULL DEFAULT (datetime('now')), + last_seen TEXT, + banned INTEGER NOT NULL DEFAULT 0, + ban_reason TEXT, + ban_expires TEXT +); + +CREATE TABLE IF NOT EXISTS sessions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE, + token TEXT NOT NULL UNIQUE, + device TEXT, + ip_address TEXT, + created_at TEXT NOT NULL DEFAULT (datetime('now')), + last_used TEXT NOT NULL DEFAULT (datetime('now')), + expires_at TEXT NOT NULL +); + +CREATE INDEX IF NOT EXISTS idx_sessions_token ON sessions(token); + +CREATE TABLE IF NOT EXISTS channels ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + type TEXT NOT NULL DEFAULT 'text', + category TEXT, + topic TEXT, + position INTEGER NOT NULL DEFAULT 0, + slow_mode INTEGER NOT NULL DEFAULT 0, + archived INTEGER NOT NULL DEFAULT 0, + created_at TEXT NOT NULL DEFAULT (datetime('now')) +); + +CREATE TABLE IF NOT EXISTS messages ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE, + user_id INTEGER NOT NULL REFERENCES users(id), + content TEXT NOT NULL, + deleted INTEGER NOT NULL DEFAULT 0, + pinned INTEGER NOT NULL DEFAULT 0, + timestamp TEXT NOT NULL DEFAULT (datetime('now')) +); + +CREATE TABLE IF NOT EXISTS invites ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + code TEXT NOT NULL UNIQUE, + created_by INTEGER NOT NULL REFERENCES users(id), + max_uses INTEGER, + use_count INTEGER NOT NULL DEFAULT 0, + expires_at TEXT, + created_at TEXT NOT NULL DEFAULT (datetime('now')), + revoked INTEGER NOT NULL DEFAULT 0 +); + +CREATE TABLE IF NOT EXISTS audit_log ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + actor_id INTEGER NOT NULL DEFAULT 0, + action TEXT NOT NULL, + target_type TEXT NOT NULL DEFAULT '', + target_id INTEGER NOT NULL DEFAULT 0, + detail TEXT NOT NULL DEFAULT '', + created_at TEXT NOT NULL DEFAULT (datetime('now')) +); + +CREATE TABLE IF NOT EXISTS settings ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL +); + +INSERT OR IGNORE INTO settings (key, value) VALUES + ('server_name', 'Test Server'), + ('motd', 'Hello'); +`) + +// openAdminTestDB opens a fresh in-memory database for admin API tests. +func openAdminTestDB(t *testing.T) *db.DB { + t.Helper() + database, err := db.Open(":memory:") + if err != nil { + t.Fatalf("db.Open: %v", err) + } + t.Cleanup(func() { database.Close() }) + + migrFS := fstest.MapFS{ + "001_schema.sql": {Data: adminSchema}, + } + if err := db.MigrateFS(database, migrFS); err != nil { + t.Fatalf("MigrateFS: %v", err) + } + return database +} + +// createAdminUser creates an Owner-role user and returns a valid bearer token. +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) + 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 { + t.Fatalf("CreateSession: %v", err) + } + return token +} + +// createMemberUser creates a Member-role user and returns a valid bearer token. +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) + 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 { + t.Fatalf("CreateSession: %v", err) + } + return token +} + +func doRequest(t *testing.T, handler http.Handler, method, path, token string, body interface{}) *httptest.ResponseRecorder { + t.Helper() + var bodyBytes []byte + if body != nil { + var err error + bodyBytes, err = json.Marshal(body) + if err != nil { + t.Fatalf("json.Marshal body: %v", err) + } + } + + req := httptest.NewRequest(method, path, bytes.NewReader(bodyBytes)) + if token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + if body != nil { + req.Header.Set("Content-Type", "application/json") + } + + w := httptest.NewRecorder() + handler.ServeHTTP(w, req) + return w +} + +// ─── GET /admin/api/stats ───────────────────────────────────────────────────── + +func TestAdminAPI_Stats_OK(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + w := doRequest(t, handler, http.MethodGet, "/stats", token, nil) + + if w.Code != http.StatusOK { + t.Errorf("status = %d, want 200; body: %s", w.Code, w.Body.String()) + } + + var stats map[string]interface{} + if err := json.Unmarshal(w.Body.Bytes(), &stats); err != nil { + t.Fatalf("unmarshal stats: %v", err) + } + if _, ok := stats["user_count"]; !ok { + t.Error("response missing 'user_count'") + } + if _, ok := stats["message_count"]; !ok { + t.Error("response missing 'message_count'") + } +} + +func TestAdminAPI_Stats_Unauthenticated(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + + w := doRequest(t, handler, http.MethodGet, "/stats", "", nil) + + if w.Code != http.StatusUnauthorized { + t.Errorf("status = %d, want 401", w.Code) + } +} + +func TestAdminAPI_Stats_Forbidden(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createMemberUser(t, database) + + w := doRequest(t, handler, http.MethodGet, "/stats", token, nil) + + if w.Code != http.StatusForbidden { + t.Errorf("status = %d, want 403", w.Code) + } +} + +// ─── GET /admin/api/users ───────────────────────────────────────────────────── + +func TestAdminAPI_ListUsers_OK(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + w := doRequest(t, handler, http.MethodGet, "/users?limit=50&offset=0", token, nil) + + if w.Code != http.StatusOK { + t.Errorf("status = %d, want 200; body: %s", w.Code, w.Body.String()) + } + + var users []interface{} + if err := json.Unmarshal(w.Body.Bytes(), &users); err != nil { + t.Fatalf("unmarshal users: %v", err) + } + // At least the admin user we created + if len(users) < 1 { + t.Error("expected at least 1 user in response") + } +} + +func TestAdminAPI_ListUsers_DefaultPagination(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + // No query params — should use defaults + w := doRequest(t, handler, http.MethodGet, "/users", token, nil) + + if w.Code != http.StatusOK { + t.Errorf("status = %d, want 200", w.Code) + } +} + +func TestAdminAPI_ListUsers_Unauthenticated(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + + w := doRequest(t, handler, http.MethodGet, "/users", "", nil) + + if w.Code != http.StatusUnauthorized { + t.Errorf("status = %d, want 401", w.Code) + } +} + +// ─── PATCH /admin/api/users/{id} ───────────────────────────────────────────── + +func TestAdminAPI_PatchUser_BanUser(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + // Create a target user + targetUID, _ := database.CreateUser("target", "hash", 3) + + body := map[string]interface{}{ + "banned": true, + "ban_reason": "spam", + } + w := doRequest(t, handler, http.MethodPatch, "/users/"+itoa(targetUID), token, body) + + if w.Code != http.StatusOK { + t.Errorf("status = %d, want 200; body: %s", w.Code, w.Body.String()) + } + + // Verify user is banned in DB + user, err := database.GetUserByID(targetUID) + if err != nil { + t.Fatalf("GetUserByID: %v", err) + } + if !user.Banned { + t.Error("user should be banned after PATCH") + } +} + +func TestAdminAPI_PatchUser_ChangeRole(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + targetUID, _ := database.CreateUser("rolechange", "hash", 3) + + body := map[string]interface{}{ + "role_id": float64(2), + } + w := doRequest(t, handler, http.MethodPatch, "/users/"+itoa(targetUID), token, body) + + if w.Code != http.StatusOK { + t.Errorf("status = %d, want 200; body: %s", w.Code, w.Body.String()) + } + + user, _ := database.GetUserByID(targetUID) + if user.RoleID != 2 { + t.Errorf("RoleID = %d, want 2", user.RoleID) + } +} + +func TestAdminAPI_PatchUser_NotFound(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + body := map[string]interface{}{"banned": true} + w := doRequest(t, handler, http.MethodPatch, "/users/99999", token, body) + + if w.Code != http.StatusNotFound { + t.Errorf("status = %d, want 404", w.Code) + } +} + +func TestAdminAPI_PatchUser_InvalidID(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + w := doRequest(t, handler, http.MethodPatch, "/users/abc", token, nil) + + if w.Code != http.StatusBadRequest { + t.Errorf("status = %d, want 400", w.Code) + } +} + +// ─── DELETE /admin/api/users/{id}/sessions ──────────────────────────────────── + +func TestAdminAPI_ForceLogout_OK(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + targetUID, _ := database.CreateUser("logoutme", "hash", 3) + database.CreateSession(targetUID, "victim-token-hash", "web", "1.2.3.4") + + w := doRequest(t, handler, http.MethodDelete, "/users/"+itoa(targetUID)+"/sessions", token, nil) + + if w.Code != http.StatusNoContent { + t.Errorf("status = %d, want 204", w.Code) + } + + sessions, _ := database.GetUserSessions(targetUID) + if len(sessions) != 0 { + t.Errorf("expected 0 sessions after force logout, got %d", len(sessions)) + } +} + +func TestAdminAPI_ForceLogout_Unauthenticated(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + + w := doRequest(t, handler, http.MethodDelete, "/users/1/sessions", "", nil) + + if w.Code != http.StatusUnauthorized { + t.Errorf("status = %d, want 401", w.Code) + } +} + +// ─── GET /admin/api/channels ────────────────────────────────────────────────── + +func TestAdminAPI_ListChannels_OK(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + database.AdminCreateChannel("general", "text", "", "", 0) + + w := doRequest(t, handler, http.MethodGet, "/channels", token, nil) + + if w.Code != http.StatusOK { + t.Errorf("status = %d, want 200; body: %s", w.Code, w.Body.String()) + } + + var channels []interface{} + if err := json.Unmarshal(w.Body.Bytes(), &channels); err != nil { + t.Fatalf("unmarshal channels: %v", err) + } + if len(channels) != 1 { + t.Errorf("expected 1 channel, got %d", len(channels)) + } +} + +// ─── POST /admin/api/channels ───────────────────────────────────────────────── + +func TestAdminAPI_CreateChannel_OK(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + body := map[string]interface{}{ + "name": "new-channel", + "type": "text", + "category": "General", + "topic": "Discussion", + "position": float64(1), + } + w := doRequest(t, handler, http.MethodPost, "/channels", token, body) + + if w.Code != http.StatusCreated { + t.Errorf("status = %d, want 201; body: %s", w.Code, w.Body.String()) + } + + var resp map[string]interface{} + if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { + t.Fatalf("unmarshal response: %v", err) + } + if _, ok := resp["id"]; !ok { + t.Error("response missing 'id'") + } +} + +func TestAdminAPI_CreateChannel_MissingName(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + body := map[string]interface{}{ + "type": "text", + } + w := doRequest(t, handler, http.MethodPost, "/channels", token, body) + + if w.Code != http.StatusBadRequest { + t.Errorf("status = %d, want 400", w.Code) + } +} + +// ─── PATCH /admin/api/channels/{id} ────────────────────────────────────────── + +func TestAdminAPI_UpdateChannel_OK(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + chID, _ := database.AdminCreateChannel("old", "text", "", "", 0) + + body := map[string]interface{}{ + "name": "updated", + "topic": "new topic", + "slow_mode": float64(10), + "position": float64(2), + "archived": false, + } + w := doRequest(t, handler, http.MethodPatch, "/channels/"+itoa(chID), token, body) + + if w.Code != http.StatusOK { + t.Errorf("status = %d, want 200; body: %s", w.Code, w.Body.String()) + } +} + +func TestAdminAPI_UpdateChannel_NotFound(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + body := map[string]interface{}{"name": "x"} + w := doRequest(t, handler, http.MethodPatch, "/channels/99999", token, body) + + if w.Code != http.StatusNotFound { + t.Errorf("status = %d, want 404", w.Code) + } +} + +// ─── DELETE /admin/api/channels/{id} ───────────────────────────────────────── + +func TestAdminAPI_DeleteChannel_OK(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + chID, _ := database.AdminCreateChannel("del-me", "text", "", "", 0) + + w := doRequest(t, handler, http.MethodDelete, "/channels/"+itoa(chID), token, nil) + + if w.Code != http.StatusNoContent { + t.Errorf("status = %d, want 204", w.Code) + } +} + +func TestAdminAPI_DeleteChannel_NotFound(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + w := doRequest(t, handler, http.MethodDelete, "/channels/99999", token, nil) + + if w.Code != http.StatusNotFound { + t.Errorf("status = %d, want 404", w.Code) + } +} + +// ─── GET /admin/api/audit-log ───────────────────────────────────────────────── + +func TestAdminAPI_AuditLog_OK(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + uid, _ := database.CreateUser("actor", "hash", 1) + database.LogAudit(uid, "TEST_ACTION", "user", uid, "detail") + + w := doRequest(t, handler, http.MethodGet, "/audit-log?limit=10&offset=0", token, nil) + + if w.Code != http.StatusOK { + t.Errorf("status = %d, want 200; body: %s", w.Code, w.Body.String()) + } + + var entries []interface{} + if err := json.Unmarshal(w.Body.Bytes(), &entries); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if len(entries) != 1 { + t.Errorf("expected 1 entry, got %d", len(entries)) + } +} + +func TestAdminAPI_AuditLog_Empty(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + w := doRequest(t, handler, http.MethodGet, "/audit-log", token, nil) + + if w.Code != http.StatusOK { + t.Errorf("status = %d, want 200", w.Code) + } + + var entries []interface{} + json.Unmarshal(w.Body.Bytes(), &entries) + if len(entries) != 0 { + t.Errorf("expected 0 entries, got %d", len(entries)) + } +} + +// ─── GET /admin/api/settings ────────────────────────────────────────────────── + +func TestAdminAPI_GetSettings_OK(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + w := doRequest(t, handler, http.MethodGet, "/settings", token, nil) + + if w.Code != http.StatusOK { + t.Errorf("status = %d, want 200; body: %s", w.Code, w.Body.String()) + } + + var settings map[string]string + if err := json.Unmarshal(w.Body.Bytes(), &settings); err != nil { + t.Fatalf("unmarshal settings: %v", err) + } + if _, ok := settings["server_name"]; !ok { + t.Error("response missing 'server_name'") + } +} + +// ─── PATCH /admin/api/settings ──────────────────────────────────────────────── + +func TestAdminAPI_PatchSettings_OK(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + body := map[string]string{ + "server_name": "Updated Server", + "motd": "New MOTD", + } + w := doRequest(t, handler, http.MethodPatch, "/settings", token, body) + + if w.Code != http.StatusOK { + t.Errorf("status = %d, want 200; body: %s", w.Code, w.Body.String()) + } + + // Verify the change was persisted + val, err := database.GetSetting("server_name") + if err != nil { + t.Fatalf("GetSetting: %v", err) + } + if val != "Updated Server" { + t.Errorf("server_name = %q, want 'Updated Server'", val) + } +} + +func TestAdminAPI_PatchSettings_InvalidBody(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + token := createAdminUser(t, database) + + req := httptest.NewRequest(http.MethodPatch, "/settings", bytes.NewReader([]byte("not-json"))) + req.Header.Set("Authorization", "Bearer "+token) + w := httptest.NewRecorder() + handler.ServeHTTP(w, req) + + if w.Code != http.StatusBadRequest { + t.Errorf("status = %d, want 400", w.Code) + } +} + +// ─── POST /admin/api/backup ─────────────────────────────────────────────────── + +func TestAdminAPI_Backup_RequiresOwner(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + + // Admin (role 2) can authenticate but is not Owner (role 1, position 100) + adminUID, _ := database.CreateUser("adminonly", "hash", 2) + token := "admin-only-token" + database.CreateSession(adminUID, auth.HashToken(token), "test", "127.0.0.1") + + w := doRequest(t, handler, http.MethodPost, "/backup", token, nil) + + // Should be forbidden — not Owner role + if w.Code != http.StatusForbidden { + t.Errorf("status = %d, want 403", w.Code) + } +} + +func TestAdminAPI_Backup_Unauthenticated(t *testing.T) { + database := openAdminTestDB(t) + handler := admin.NewAdminAPI(database) + + w := doRequest(t, handler, http.MethodPost, "/backup", "", nil) + + if w.Code != http.StatusUnauthorized { + t.Errorf("status = %d, want 401", w.Code) + } +} + +// ─── helpers ───────────────────────────────────────────────────────────────── + +// itoa converts an int64 to a string for use in URL paths. +func itoa(n int64) string { + return fmt.Sprint(n) +} diff --git a/Server/admin/static/index.html b/Server/admin/static/index.html index dd8b0d42..26909ee1 100644 --- a/Server/admin/static/index.html +++ b/Server/admin/static/index.html @@ -1,18 +1,684 @@ - - - OwnCord Admin Panel - + + + OwnCord Admin + -

OwnCord Admin Panel

-

Admin panel frontend — coming in Phase 6.

-

For now, use the REST API endpoints under /api/admin/ directly.

+ + +
+
+

OwnCord Admin

+
+ + +
+
+ + +
+ +
+
+
+ + + + + diff --git a/Server/api/router.go b/Server/api/router.go index 8f2ce3b8..3563af47 100644 --- a/Server/api/router.go +++ b/Server/api/router.go @@ -7,6 +7,7 @@ import ( "github.com/go-chi/chi/v5" "github.com/go-chi/chi/v5/middleware" + "github.com/owncord/server/admin" "github.com/owncord/server/auth" "github.com/owncord/server/config" "github.com/owncord/server/db" @@ -46,11 +47,17 @@ func NewRouter(cfg *config.Config, database *db.DB) http.Handler { // Channel and message REST routes. MountChannelRoutes(r, database) + // Voice credentials REST route. + MountVoiceRoutes(r, cfg, database) + // WebSocket hub — WS does its own in-band auth, so no AuthMiddleware here. hub := ws.NewHub(database, limiter) go hub.Run() r.Get("/api/v1/ws", ws.ServeWS(hub, database)) + // Admin panel: static files + REST API (Phase 6). + r.Mount("/admin", admin.NewHandler(database)) + return r } diff --git a/Server/api/voice_handler.go b/Server/api/voice_handler.go new file mode 100644 index 00000000..23368b7e --- /dev/null +++ b/Server/api/voice_handler.go @@ -0,0 +1,121 @@ +package api + +import ( + "crypto/hmac" + "crypto/sha1" + "encoding/base64" + "fmt" + "net/http" + "time" + + "github.com/go-chi/chi/v5" + "github.com/owncord/server/config" + "github.com/owncord/server/db" +) + +const voiceCredentialTTL = 24 * time.Hour + +// iceServer describes a single ICE server entry for WebRTC peer connections. +type iceServer struct { + URLs string `json:"urls"` + Username string `json:"username,omitempty"` + Credential string `json:"credential,omitempty"` +} + +// voiceCredentialsResponse is the JSON body for GET /api/v1/voice/credentials. +type voiceCredentialsResponse struct { + ICEServers []iceServer `json:"ice_servers"` + ExpiresIn int `json:"expires_in"` +} + +// turnCredentials holds the generated TURN username and HMAC credential. +type turnCredentials struct { + Username string + Credential string +} + +// MountVoiceRoutes registers the voice REST endpoints on r. +func MountVoiceRoutes(r chi.Router, cfg *config.Config, database *db.DB) { + r.Route("/api/v1/voice", func(r chi.Router) { + r.Use(AuthMiddleware(database)) + r.Get("/credentials", handleVoiceCredentials(cfg, database)) + }) +} + +// handleVoiceCredentials returns ICE server credentials for WebRTC. +// Requires a valid session (AuthMiddleware). Generates time-limited TURN +// credentials using HMAC-SHA1 as per the coturn REST API spec. +func handleVoiceCredentials(cfg *config.Config, _ *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + user, ok := r.Context().Value(UserKey).(*db.User) + if !ok || user == nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", + Message: "authentication required", + }) + return + } + + host := serverHost(r) + servers := buildICEServers(user.ID, cfg, host) + + writeJSON(w, http.StatusOK, voiceCredentialsResponse{ + ICEServers: servers, + ExpiresIn: int(voiceCredentialTTL.Seconds()), + }) + } +} + +// buildICEServers constructs the ICE server list for the given user. +func buildICEServers(userID int64, cfg *config.Config, host string) []iceServer { + servers := []iceServer{ + {URLs: fmt.Sprintf("stun:%s:%d", host, cfg.Voice.STUNPort)}, + } + + if cfg.Voice.TURNEnabled && cfg.Voice.TURNSecret != "" { + creds := generateTURNCredentials(userID, cfg.Voice.TURNSecret) + servers = append(servers, iceServer{ + URLs: fmt.Sprintf("turn:%s:%d", host, cfg.Voice.TURNPort), + Username: creds.Username, + Credential: creds.Credential, + }) + } + + return servers +} + +// generateTURNCredentials produces time-limited TURN credentials using HMAC-SHA1. +// Username format: ":" +// Credential: base64(HMAC-SHA1(secret, username)) +func generateTURNCredentials(userID int64, secret string) turnCredentials { + expiry := time.Now().Add(voiceCredentialTTL).Unix() + username := fmt.Sprintf("%d:%d", expiry, userID) + + mac := hmac.New(sha1.New, []byte(secret)) + mac.Write([]byte(username)) + credential := base64.StdEncoding.EncodeToString(mac.Sum(nil)) + + return turnCredentials{ + Username: username, + Credential: credential, + } +} + +// serverHost extracts the host for ICE server URLs from the request or falls +// back to "localhost". +func serverHost(r *http.Request) string { + if host := r.Host; host != "" { + // Strip port if present. + for i := len(host) - 1; i >= 0; i-- { + if host[i] == ':' { + return host[:i] + } + if host[i] == ']' { + // IPv6 with no port. + return host + } + } + return host + } + return "localhost" +} diff --git a/Server/api/voice_handler_test.go b/Server/api/voice_handler_test.go new file mode 100644 index 00000000..ac0cf4a9 --- /dev/null +++ b/Server/api/voice_handler_test.go @@ -0,0 +1,344 @@ +package api_test + +import ( + "crypto/hmac" + "crypto/sha1" + "encoding/base64" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "testing" + "testing/fstest" + "time" + + "github.com/go-chi/chi/v5" + "github.com/owncord/server/api" + "github.com/owncord/server/auth" + "github.com/owncord/server/config" + "github.com/owncord/server/db" +) + +// ─── helpers ────────────────────────────────────────────────────────────────── + +// newVoiceAPITestDB opens an in-memory DB for voice API tests. +func newVoiceAPITestDB(t *testing.T) *db.DB { + t.Helper() + database, err := db.Open(":memory:") + if err != nil { + t.Fatalf("db.Open: %v", err) + } + t.Cleanup(func() { database.Close() }) + + migrFS := fstest.MapFS{ + "001_schema.sql": {Data: apiTestSchema}, + } + if err := db.MigrateFS(database, migrFS); err != nil { + t.Fatalf("MigrateFS: %v", err) + } + return database +} + +// buildVoiceRouter returns a chi router with voice routes mounted. +func buildVoiceRouter(database *db.DB, cfg *config.Config) http.Handler { + r := chi.NewRouter() + api.MountVoiceRoutes(r, cfg, database) + return r +} + +// seedAPIUser creates a user+session and returns a valid bearer token. +func seedVoiceAPIUser(t *testing.T, database *db.DB, username string) string { + t.Helper() + _, err := database.CreateUser(username, "hash", 4) + if err != nil { + t.Fatalf("CreateUser: %v", err) + } + user, err := database.GetUserByUsername(username) + if err != nil || user == nil { + t.Fatalf("GetUserByUsername: %v", err) + } + token := "test-token-" + username + hash := auth.HashToken(token) + future := time.Now().Add(24 * time.Hour).UTC().Format("2006-01-02 15:04:05") + _, err = database.Exec( + `INSERT INTO sessions (user_id, token, device, ip_address, expires_at) VALUES (?, ?, ?, ?, ?)`, + user.ID, hash, "test", "127.0.0.1", future, + ) + if err != nil { + t.Fatalf("insert session: %v", err) + } + return token +} + +// voiceGetWithToken performs a GET with Authorization: Bearer header. +func voiceGetWithToken(t *testing.T, router http.Handler, path, token string) *httptest.ResponseRecorder { + t.Helper() + req := httptest.NewRequest(http.MethodGet, path, nil) + req.Header.Set("Authorization", "Bearer "+token) + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + return rr +} + +// defaultVoiceCfg returns a Config with a known TURN secret for testing. +func defaultVoiceCfg() *config.Config { + return &config.Config{ + Server: config.ServerConfig{Name: "Test"}, + Voice: config.VoiceConfig{ + TURNSecret: "test-secret-key-12345", + STUNPort: 3478, + TURNPort: 3478, + TURNEnabled: true, + }, + } +} + +// ─── GET /api/v1/voice/credentials ─────────────────────────────────────────── + +func TestVoiceCredentials_Authenticated_Returns200(t *testing.T) { + database := newVoiceAPITestDB(t) + token := seedVoiceAPIUser(t, database, "alice") + cfg := defaultVoiceCfg() + + router := buildVoiceRouter(database, cfg) + rr := voiceGetWithToken(t, router, "/api/v1/voice/credentials", token) + + if rr.Code != http.StatusOK { + t.Errorf("status = %d, want 200; body: %s", rr.Code, rr.Body.String()) + } +} + +func TestVoiceCredentials_Unauthenticated_Returns401(t *testing.T) { + database := newVoiceAPITestDB(t) + cfg := defaultVoiceCfg() + + router := buildVoiceRouter(database, cfg) + req := httptest.NewRequest(http.MethodGet, "/api/v1/voice/credentials", nil) + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusUnauthorized { + t.Errorf("status = %d, want 401", rr.Code) + } +} + +func TestVoiceCredentials_ResponseContainsIceServers(t *testing.T) { + database := newVoiceAPITestDB(t) + token := seedVoiceAPIUser(t, database, "bob") + cfg := defaultVoiceCfg() + + router := buildVoiceRouter(database, cfg) + rr := voiceGetWithToken(t, router, "/api/v1/voice/credentials", token) + + var resp map[string]interface{} + if err := json.NewDecoder(rr.Body).Decode(&resp); err != nil { + t.Fatalf("decode response: %v", err) + } + + iceServers, ok := resp["ice_servers"] + if !ok { + t.Fatal("response missing ice_servers field") + } + servers, ok := iceServers.([]interface{}) + if !ok || len(servers) == 0 { + t.Error("ice_servers is empty or wrong type") + } +} + +func TestVoiceCredentials_ContainsSTUNEntry(t *testing.T) { + database := newVoiceAPITestDB(t) + token := seedVoiceAPIUser(t, database, "carol") + cfg := defaultVoiceCfg() + + router := buildVoiceRouter(database, cfg) + rr := voiceGetWithToken(t, router, "/api/v1/voice/credentials", token) + + var resp map[string]interface{} + json.NewDecoder(rr.Body).Decode(&resp) + + servers := resp["ice_servers"].([]interface{}) + foundSTUN := false + for _, s := range servers { + entry := s.(map[string]interface{}) + if urls, ok := entry["urls"].(string); ok { + if len(urls) > 5 && urls[:5] == "stun:" { + foundSTUN = true + break + } + } + } + if !foundSTUN { + t.Error("ice_servers does not contain a STUN entry") + } +} + +func TestVoiceCredentials_ContainsTURNEntry(t *testing.T) { + database := newVoiceAPITestDB(t) + token := seedVoiceAPIUser(t, database, "dave") + cfg := defaultVoiceCfg() + + router := buildVoiceRouter(database, cfg) + rr := voiceGetWithToken(t, router, "/api/v1/voice/credentials", token) + + var resp map[string]interface{} + json.NewDecoder(rr.Body).Decode(&resp) + + servers := resp["ice_servers"].([]interface{}) + foundTURN := false + for _, s := range servers { + entry := s.(map[string]interface{}) + if urls, ok := entry["urls"].(string); ok { + if len(urls) > 5 && urls[:5] == "turn:" { + foundTURN = true + // TURN entries must have username and credential. + if _, hasUser := entry["username"]; !hasUser { + t.Error("TURN entry missing username") + } + if _, hasCred := entry["credential"]; !hasCred { + t.Error("TURN entry missing credential") + } + break + } + } + } + if !foundTURN { + t.Error("ice_servers does not contain a TURN entry") + } +} + +func TestVoiceCredentials_TURNCredentialIsValidHMAC(t *testing.T) { + database := newVoiceAPITestDB(t) + token := seedVoiceAPIUser(t, database, "eve") + secret := "test-secret-key-12345" + cfg := &config.Config{ + Voice: config.VoiceConfig{ + TURNSecret: secret, + STUNPort: 3478, + TURNPort: 3478, + TURNEnabled: true, + }, + } + + router := buildVoiceRouter(database, cfg) + rr := voiceGetWithToken(t, router, "/api/v1/voice/credentials", token) + + var resp map[string]interface{} + json.NewDecoder(rr.Body).Decode(&resp) + + servers := resp["ice_servers"].([]interface{}) + for _, s := range servers { + entry := s.(map[string]interface{}) + urls, _ := entry["urls"].(string) + if len(urls) < 5 || urls[:5] != "turn:" { + continue + } + username, _ := entry["username"].(string) + credential, _ := entry["credential"].(string) + + if username == "" || credential == "" { + t.Fatal("TURN entry has empty username or credential") + } + + // Verify HMAC-SHA1: credential should be base64(HMAC-SHA1(secret, username)). + mac := hmac.New(sha1.New, []byte(secret)) + mac.Write([]byte(username)) + expected := base64.StdEncoding.EncodeToString(mac.Sum(nil)) + + if credential != expected { + t.Errorf("TURN credential HMAC mismatch\n got: %s\n want: %s", credential, expected) + } + return + } + t.Error("no TURN entry found to validate HMAC") +} + +func TestVoiceCredentials_UsernameContainsTimestampAndUserID(t *testing.T) { + database := newVoiceAPITestDB(t) + token := seedVoiceAPIUser(t, database, "frank") + cfg := defaultVoiceCfg() + + router := buildVoiceRouter(database, cfg) + rr := voiceGetWithToken(t, router, "/api/v1/voice/credentials", token) + + var resp map[string]interface{} + json.NewDecoder(rr.Body).Decode(&resp) + + servers := resp["ice_servers"].([]interface{}) + for _, s := range servers { + entry := s.(map[string]interface{}) + urls, _ := entry["urls"].(string) + if len(urls) < 5 || urls[:5] != "turn:" { + continue + } + username, _ := entry["username"].(string) + + // Username format: ":". + var ts, uid int64 + if _, err := fmt.Sscanf(username, "%d:%d", &ts, &uid); err != nil { + t.Errorf("TURN username %q is not in format :: %v", username, err) + } + if ts <= time.Now().Unix() { + t.Errorf("TURN username timestamp %d is in the past, want future", ts) + } + if uid <= 0 { + t.Errorf("TURN username userID %d must be positive", uid) + } + return + } + t.Error("no TURN entry found to validate username format") +} + +func TestVoiceCredentials_ResponseContainsExpiresIn(t *testing.T) { + database := newVoiceAPITestDB(t) + token := seedVoiceAPIUser(t, database, "grace") + cfg := defaultVoiceCfg() + + router := buildVoiceRouter(database, cfg) + rr := voiceGetWithToken(t, router, "/api/v1/voice/credentials", token) + + var resp map[string]interface{} + json.NewDecoder(rr.Body).Decode(&resp) + + expiresIn, ok := resp["expires_in"] + if !ok { + t.Fatal("response missing expires_in field") + } + // expires_in should be 86400 (24 hours in seconds). + val, ok := expiresIn.(float64) + if !ok || val != 86400 { + t.Errorf("expires_in = %v, want 86400", expiresIn) + } +} + +func TestVoiceCredentials_TURNDisabled_NoTURNEntry(t *testing.T) { + database := newVoiceAPITestDB(t) + token := seedVoiceAPIUser(t, database, "henry") + cfg := &config.Config{ + Voice: config.VoiceConfig{ + TURNSecret: "secret", + STUNPort: 3478, + TURNPort: 3478, + TURNEnabled: false, // TURN disabled + }, + } + + router := buildVoiceRouter(database, cfg) + rr := voiceGetWithToken(t, router, "/api/v1/voice/credentials", token) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want 200", rr.Code) + } + + var resp map[string]interface{} + json.NewDecoder(rr.Body).Decode(&resp) + + servers := resp["ice_servers"].([]interface{}) + for _, s := range servers { + entry := s.(map[string]interface{}) + if urls, _ := entry["urls"].(string); len(urls) >= 5 && urls[:5] == "turn:" { + t.Error("TURN entry present when TURNEnabled=false") + } + } +} diff --git a/Server/config/config.go b/Server/config/config.go index 918407b1..53cfe542 100644 --- a/Server/config/config.go +++ b/Server/config/config.go @@ -20,6 +20,15 @@ type Config struct { Database DatabaseConfig `koanf:"database"` TLS TLSConfig `koanf:"tls"` Upload UploadConfig `koanf:"upload"` + Voice VoiceConfig `koanf:"voice"` +} + +// VoiceConfig holds STUN/TURN server settings for WebRTC signaling. +type VoiceConfig struct { + TURNSecret string `koanf:"turn_secret"` // HMAC-SHA1 secret; auto-generated if empty + STUNPort int `koanf:"stun_port"` // default 3478 + TURNPort int `koanf:"turn_port"` // default 3478 + TURNEnabled bool `koanf:"turn_enabled"` // default true } // ServerConfig holds HTTP server settings. @@ -66,6 +75,11 @@ func defaults() Config { MaxSizeMB: 100, StorageDir: "data/uploads", }, + Voice: VoiceConfig{ + STUNPort: 3478, + TURNPort: 3478, + TURNEnabled: true, + }, } } diff --git a/Server/db/admin_queries.go b/Server/db/admin_queries.go new file mode 100644 index 00000000..cb58a416 --- /dev/null +++ b/Server/db/admin_queries.go @@ -0,0 +1,307 @@ +package db + +import ( + "database/sql" + "errors" + "fmt" +) + +// ─── Permission constants ───────────────────────────────────────────────────── + +const ( + permAdministrator = int64(0x40000000) + permManageServer = int64(0x2000000) + permViewAuditLog = int64(0x8000000) +) + +// ─── Server Stats ───────────────────────────────────────────────────────────── + +// 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) { + stats := &ServerStats{} + + if err := d.sqlDB.QueryRow(`SELECT COUNT(*) FROM users`).Scan(&stats.UserCount); err != nil { + return nil, fmt.Errorf("GetServerStats users: %w", err) + } + if err := d.sqlDB.QueryRow(`SELECT COUNT(*) FROM messages WHERE deleted = 0`).Scan(&stats.MessageCount); err != nil { + return nil, fmt.Errorf("GetServerStats messages: %w", err) + } + if err := d.sqlDB.QueryRow(`SELECT COUNT(*) FROM channels`).Scan(&stats.ChannelCount); err != nil { + return nil, fmt.Errorf("GetServerStats channels: %w", err) + } + if err := d.sqlDB.QueryRow(`SELECT COUNT(*) FROM invites WHERE revoked = 0`).Scan(&stats.InviteCount); err != nil { + return nil, fmt.Errorf("GetServerStats invites: %w", err) + } + + // page_count * page_size gives the database size in bytes. + // 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 { + return nil, fmt.Errorf("GetServerStats page_count: %w", err) + } + if err := d.sqlDB.QueryRow(`PRAGMA page_size`).Scan(&pageSize); err != nil { + return nil, fmt.Errorf("GetServerStats page_size: %w", err) + } + stats.DBSizeBytes = pageCount * pageSize + + return stats, nil +} + +// ─── User Management ────────────────────────────────────────────────────────── + +// 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.sqlDB.Query( + `SELECT u.id, u.username, u.password, u.avatar, u.role_id, u.totp_secret, + u.status, u.created_at, u.last_seen, u.banned, u.ban_reason, u.ban_expires, + COALESCE(r.name, '') AS role_name + FROM users u + LEFT JOIN roles r ON r.id = u.role_id + ORDER BY u.id ASC + LIMIT ? OFFSET ?`, + limit, offset, + ) + if err != nil { + return nil, fmt.Errorf("ListAllUsers: %w", err) + } + defer rows.Close() + + var result []UserWithRole + for rows.Next() { + var uwr UserWithRole + var banned int + err := rows.Scan( + &uwr.ID, &uwr.Username, &uwr.PasswordHash, &uwr.Avatar, &uwr.RoleID, + &uwr.TOTPSecret, &uwr.Status, &uwr.CreatedAt, &uwr.LastSeen, + &banned, &uwr.BanReason, &uwr.BanExpires, + &uwr.RoleName, + ) + if err != nil { + return nil, fmt.Errorf("ListAllUsers scan: %w", err) + } + uwr.Banned = banned != 0 + result = append(result, uwr) + } + if rows.Err() != nil { + return nil, fmt.Errorf("ListAllUsers rows: %w", rows.Err()) + } + if result == nil { + result = []UserWithRole{} + } + return result, nil +} + +// UpdateUserRole changes the role_id of a user. +func (d *DB) UpdateUserRole(userID, roleID int64) error { + _, err := d.sqlDB.Exec( + `UPDATE users SET role_id = ? WHERE id = ?`, + roleID, userID, + ) + if err != nil { + return fmt.Errorf("UpdateUserRole: %w", err) + } + return nil +} + +// ForceLogoutUser deletes all sessions for the given user ID. +func (d *DB) ForceLogoutUser(userID int64) error { + _, err := d.sqlDB.Exec(`DELETE FROM sessions WHERE user_id = ?`, userID) + if 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.sqlDB.Query( + `SELECT id, user_id, token, device, ip_address, created_at, last_used, expires_at + FROM sessions WHERE user_id = ? ORDER BY created_at DESC`, + userID, + ) + if err != nil { + return nil, fmt.Errorf("GetUserSessions: %w", err) + } + defer rows.Close() + + var sessions []Session + for rows.Next() { + var s Session + err := rows.Scan( + &s.ID, &s.UserID, &s.TokenHash, &s.Device, &s.IP, + &s.CreatedAt, &s.LastUsed, &s.ExpiresAt, + ) + if err != nil { + return nil, fmt.Errorf("GetUserSessions scan: %w", err) + } + sessions = append(sessions, s) + } + if rows.Err() != nil { + return nil, fmt.Errorf("GetUserSessions rows: %w", rows.Err()) + } + if sessions == nil { + sessions = []Session{} + } + return sessions, nil +} + +// ─── Channel Management (admin) ─────────────────────────────────────────────── + +// AdminCreateChannel creates a channel with full field control including position. +func (d *DB) AdminCreateChannel(name, chanType, category, topic string, position int) (int64, error) { + res, err := d.sqlDB.Exec( + `INSERT INTO channels (name, type, category, topic, position) + VALUES (?, ?, ?, ?, ?)`, + name, chanType, nullableString(category), nullableString(topic), position, + ) + if err != nil { + return 0, fmt.Errorf("AdminCreateChannel: %w", err) + } + return res.LastInsertId() +} + +// AdminUpdateChannel updates all mutable channel fields. +func (d *DB) AdminUpdateChannel(id int64, name, topic string, slowMode, position int, archived bool) error { + archivedInt := 0 + if archived { + archivedInt = 1 + } + _, err := d.sqlDB.Exec( + `UPDATE channels + SET name = ?, topic = ?, slow_mode = ?, position = ?, archived = ? + WHERE id = ?`, + name, nullableString(topic), slowMode, position, archivedInt, id, + ) + if err != nil { + return fmt.Errorf("AdminUpdateChannel: %w", err) + } + return nil +} + +// AdminDeleteChannel removes a channel by ID (cascades to messages, etc.). +func (d *DB) AdminDeleteChannel(id int64) error { + _, err := d.sqlDB.Exec(`DELETE FROM channels WHERE id = ?`, id) + if err != nil { + return fmt.Errorf("AdminDeleteChannel: %w", err) + } + return nil +} + +// ─── Audit Log ──────────────────────────────────────────────────────────────── + +// LogAudit inserts an audit log entry. +func (d *DB) LogAudit(actorID int64, action, targetType string, targetID int64, detail string) error { + _, err := d.sqlDB.Exec( + `INSERT INTO audit_log (actor_id, action, target_type, target_id, detail) + VALUES (?, ?, ?, ?, ?)`, + actorID, action, targetType, targetID, detail, + ) + if err != nil { + return fmt.Errorf("LogAudit: %w", err) + } + return nil +} + +// GetAuditLog returns audit log entries ordered newest-first with pagination. +func (d *DB) GetAuditLog(limit, offset int) ([]AuditEntry, error) { + rows, err := d.sqlDB.Query( + `SELECT a.id, a.actor_id, COALESCE(u.username, ''), a.action, + a.target_type, a.target_id, a.detail, a.created_at + FROM audit_log a + LEFT JOIN users u ON u.id = a.actor_id + ORDER BY a.id DESC + LIMIT ? OFFSET ?`, + limit, offset, + ) + if err != nil { + return nil, fmt.Errorf("GetAuditLog: %w", err) + } + defer rows.Close() + + var entries []AuditEntry + for rows.Next() { + var e AuditEntry + if err := rows.Scan( + &e.ID, &e.ActorID, &e.ActorName, &e.Action, + &e.TargetType, &e.TargetID, &e.Detail, &e.CreatedAt, + ); err != nil { + return nil, fmt.Errorf("GetAuditLog scan: %w", err) + } + entries = append(entries, e) + } + if rows.Err() != nil { + return nil, fmt.Errorf("GetAuditLog rows: %w", rows.Err()) + } + if entries == nil { + entries = []AuditEntry{} + } + return entries, nil +} + +// ─── Settings ───────────────────────────────────────────────────────────────── + +// 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) { + var value string + err := d.sqlDB.QueryRow(`SELECT value FROM settings WHERE key = ?`, key).Scan(&value) + if errors.Is(err, sql.ErrNoRows) { + return "", fmt.Errorf("GetSetting: key %q not found", key) + } + if err != nil { + return "", fmt.Errorf("GetSetting: %w", err) + } + return value, nil +} + +// SetSetting upserts a setting value for the given key. +func (d *DB) SetSetting(key, value string) error { + _, err := d.sqlDB.Exec( + `INSERT INTO settings (key, value) VALUES (?, ?) + ON CONFLICT(key) DO UPDATE SET value = excluded.value`, + key, value, + ) + if err != nil { + return fmt.Errorf("SetSetting: %w", err) + } + return nil +} + +// GetAllSettings returns all settings as a key→value map. +func (d *DB) GetAllSettings() (map[string]string, error) { + rows, err := d.sqlDB.Query(`SELECT key, value FROM settings`) + if err != nil { + return nil, fmt.Errorf("GetAllSettings: %w", err) + } + defer rows.Close() + + result := make(map[string]string) + for rows.Next() { + var k, v string + if err := rows.Scan(&k, &v); err != nil { + return nil, fmt.Errorf("GetAllSettings scan: %w", err) + } + result[k] = v + } + if rows.Err() != nil { + return nil, fmt.Errorf("GetAllSettings rows: %w", rows.Err()) + } + return result, nil +} + +// ─── Backup ─────────────────────────────────────────────────────────────────── + +// BackupTo creates an online backup of the database at the given path using +// SQLite's VACUUM INTO statement. The destination path must not already exist. +// This only works meaningfully for file-backed databases; in-memory databases +// will produce a valid but potentially minimal backup file. +func (d *DB) BackupTo(path string) error { + _, err := d.sqlDB.Exec(fmt.Sprintf("VACUUM INTO '%s'", path)) + if err != nil { + return fmt.Errorf("BackupTo: %w", err) + } + return nil +} diff --git a/Server/db/admin_queries_test.go b/Server/db/admin_queries_test.go new file mode 100644 index 00000000..982eda33 --- /dev/null +++ b/Server/db/admin_queries_test.go @@ -0,0 +1,729 @@ +package db_test + +import ( + "os" + "path/filepath" + "strings" + "testing" + "testing/fstest" + + "github.com/owncord/server/db" +) + +// adminTestSchema extends testSchema with tables needed for admin queries. +var adminTestSchema = append(testSchema, []byte(` +CREATE TABLE IF NOT EXISTS channels ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + type TEXT NOT NULL DEFAULT 'text', + category TEXT, + topic TEXT, + position INTEGER NOT NULL DEFAULT 0, + slow_mode INTEGER NOT NULL DEFAULT 0, + archived INTEGER NOT NULL DEFAULT 0, + created_at TEXT NOT NULL DEFAULT (datetime('now')) +); + +CREATE TABLE IF NOT EXISTS messages ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE, + user_id INTEGER NOT NULL REFERENCES users(id), + content TEXT NOT NULL, + reply_to INTEGER REFERENCES messages(id) ON DELETE SET NULL, + edited_at TEXT, + deleted INTEGER NOT NULL DEFAULT 0, + pinned INTEGER NOT NULL DEFAULT 0, + timestamp TEXT NOT NULL DEFAULT (datetime('now')) +); + +CREATE TABLE IF NOT EXISTS audit_log ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + actor_id INTEGER NOT NULL REFERENCES users(id), + action TEXT NOT NULL, + target_type TEXT NOT NULL DEFAULT '', + target_id INTEGER NOT NULL DEFAULT 0, + detail TEXT NOT NULL DEFAULT '', + created_at TEXT NOT NULL DEFAULT (datetime('now')) +); + +CREATE INDEX IF NOT EXISTS idx_audit_log_created ON audit_log(created_at DESC); +CREATE INDEX IF NOT EXISTS idx_audit_log_actor ON audit_log(actor_id); + +CREATE TABLE IF NOT EXISTS settings ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL +); + +INSERT OR IGNORE INTO settings (key, value) VALUES + ('server_name', 'OwnCord Server'), + ('motd', 'Welcome!'); +`)...) + +// newAdminTestDB opens an in-memory database with the admin-extended schema. +func newAdminTestDB(t *testing.T) *db.DB { + t.Helper() + database, err := db.Open(":memory:") + if err != nil { + t.Fatalf("db.Open: %v", err) + } + t.Cleanup(func() { database.Close() }) + + migrFS := fstest.MapFS{ + "001_schema.sql": {Data: adminTestSchema}, + } + if err := db.MigrateFS(database, migrFS); err != nil { + t.Fatalf("MigrateFS: %v", err) + } + return database +} + +// ─── GetServerStats ──────────────────────────────────────────────────────────── + +func TestGetServerStats_EmptyDB(t *testing.T) { + database := newAdminTestDB(t) + + stats, err := database.GetServerStats() + if err != nil { + t.Fatalf("GetServerStats() error: %v", err) + } + if stats == nil { + t.Fatal("GetServerStats() returned nil") + } + if stats.UserCount != 0 { + t.Errorf("UserCount = %d, want 0", stats.UserCount) + } + if stats.MessageCount != 0 { + t.Errorf("MessageCount = %d, want 0", stats.MessageCount) + } + if stats.ChannelCount != 0 { + t.Errorf("ChannelCount = %d, want 0", stats.ChannelCount) + } + if stats.InviteCount != 0 { + t.Errorf("InviteCount = %d, want 0", stats.InviteCount) + } + if stats.DBSizeBytes < 0 { + t.Errorf("DBSizeBytes = %d, want >= 0", stats.DBSizeBytes) + } +} + +func TestGetServerStats_WithData(t *testing.T) { + database := newAdminTestDB(t) + + _, err := database.CreateUser("statuser", "hash", 4) + if err != nil { + t.Fatalf("CreateUser error: %v", err) + } + + _, err = database.CreateChannel("general", "text", "", "", 0) + if err != nil { + t.Fatalf("CreateChannel error: %v", err) + } + + stats, err := database.GetServerStats() + if err != nil { + t.Fatalf("GetServerStats() error: %v", err) + } + if stats.UserCount != 1 { + t.Errorf("UserCount = %d, want 1", stats.UserCount) + } + if stats.ChannelCount != 1 { + t.Errorf("ChannelCount = %d, want 1", stats.ChannelCount) + } +} + +// ─── ListAllUsers ────────────────────────────────────────────────────────────── + +func TestListAllUsers_Empty(t *testing.T) { + database := newAdminTestDB(t) + + users, err := database.ListAllUsers(50, 0) + if err != nil { + t.Fatalf("ListAllUsers() error: %v", err) + } + if len(users) != 0 { + t.Errorf("ListAllUsers() = %d users, want 0", len(users)) + } +} + +func TestListAllUsers_WithRoleName(t *testing.T) { + database := newAdminTestDB(t) + + _, err := database.CreateUser("alice", "hash", 4) + if err != nil { + t.Fatalf("CreateUser error: %v", err) + } + + users, err := database.ListAllUsers(50, 0) + if err != nil { + t.Fatalf("ListAllUsers() error: %v", err) + } + if len(users) != 1 { + t.Fatalf("ListAllUsers() = %d users, want 1", len(users)) + } + if users[0].Username != "alice" { + t.Errorf("Username = %q, want 'alice'", users[0].Username) + } + // RoleName comes from JOIN with roles table + if users[0].RoleName == "" { + t.Error("RoleName should not be empty — JOIN with roles table failed") + } +} + +func TestListAllUsers_Pagination(t *testing.T) { + database := newAdminTestDB(t) + + for i := 0; i < 5; i++ { + _, err := database.CreateUser( + strings.Repeat("u", i+1), + "hash", + 4, + ) + if err != nil { + t.Fatalf("CreateUser[%d] error: %v", i, err) + } + } + + page1, err := database.ListAllUsers(3, 0) + if err != nil { + t.Fatalf("ListAllUsers page1 error: %v", err) + } + if len(page1) != 3 { + t.Errorf("page1 len = %d, want 3", len(page1)) + } + + page2, err := database.ListAllUsers(3, 3) + if err != nil { + t.Fatalf("ListAllUsers page2 error: %v", err) + } + if len(page2) != 2 { + t.Errorf("page2 len = %d, want 2", len(page2)) + } +} + +func TestListAllUsers_ZeroLimit(t *testing.T) { + database := newAdminTestDB(t) + database.CreateUser("zerotest", "hash", 4) + + users, err := database.ListAllUsers(0, 0) + if err != nil { + t.Fatalf("ListAllUsers(0, 0) error: %v", err) + } + // limit=0 should return nothing + if len(users) != 0 { + t.Errorf("ListAllUsers(0, 0) = %d users, want 0", len(users)) + } +} + +// ─── UpdateUserRole ──────────────────────────────────────────────────────────── + +func TestUpdateUserRole(t *testing.T) { + database := newAdminTestDB(t) + + uid, err := database.CreateUser("roleuser", "hash", 4) + if err != nil { + t.Fatalf("CreateUser error: %v", err) + } + + if err := database.UpdateUserRole(uid, 2); err != nil { + t.Fatalf("UpdateUserRole() error: %v", err) + } + + user, err := database.GetUserByID(uid) + if err != nil { + t.Fatalf("GetUserByID error: %v", err) + } + if user.RoleID != 2 { + t.Errorf("RoleID = %d, want 2", user.RoleID) + } +} + +func TestUpdateUserRole_NonexistentUser(t *testing.T) { + database := newAdminTestDB(t) + + // UPDATE with no matching rows is not an error + err := database.UpdateUserRole(99999, 2) + if err != nil { + t.Errorf("UpdateUserRole() for nonexistent user returned unexpected error: %v", err) + } +} + +// ─── ForceLogoutUser ─────────────────────────────────────────────────────────── + +func TestForceLogoutUser_DeletesSessions(t *testing.T) { + database := newAdminTestDB(t) + + uid, err := database.CreateUser("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") + + sessions, err := database.GetUserSessions(uid) + if err != nil { + t.Fatalf("GetUserSessions error: %v", err) + } + if len(sessions) != 2 { + t.Fatalf("expected 2 sessions before logout, got %d", len(sessions)) + } + + if err := database.ForceLogoutUser(uid); err != nil { + t.Fatalf("ForceLogoutUser() error: %v", err) + } + + sessions, err = database.GetUserSessions(uid) + if err != nil { + t.Fatalf("GetUserSessions after logout error: %v", err) + } + if len(sessions) != 0 { + t.Errorf("expected 0 sessions after ForceLogoutUser, got %d", len(sessions)) + } +} + +func TestForceLogoutUser_NoSessions(t *testing.T) { + database := newAdminTestDB(t) + + uid, err := database.CreateUser("nosessions", "hash", 4) + if err != nil { + t.Fatalf("CreateUser error: %v", err) + } + + if err := database.ForceLogoutUser(uid); err != nil { + t.Errorf("ForceLogoutUser() on user with no sessions returned error: %v", err) + } +} + +// ─── GetUserSessions ────────────────────────────────────────────────────────── + +func TestGetUserSessions_Empty(t *testing.T) { + database := newAdminTestDB(t) + + uid, err := database.CreateUser("sessionuser", "hash", 4) + if err != nil { + t.Fatalf("CreateUser error: %v", err) + } + + sessions, err := database.GetUserSessions(uid) + if err != nil { + t.Fatalf("GetUserSessions() error: %v", err) + } + if len(sessions) != 0 { + t.Errorf("GetUserSessions() = %d, want 0", len(sessions)) + } +} + +func TestGetUserSessions_IsolatedByUser(t *testing.T) { + database := newAdminTestDB(t) + + uid1, _ := database.CreateUser("user1sess", "hash", 4) + uid2, _ := database.CreateUser("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") + + sessions, err := database.GetUserSessions(uid1) + if err != nil { + t.Fatalf("GetUserSessions() error: %v", err) + } + if len(sessions) != 2 { + t.Errorf("GetUserSessions(uid1) = %d sessions, want 2", len(sessions)) + } + for _, s := range sessions { + if s.UserID != uid1 { + t.Errorf("session UserID = %d, want %d", s.UserID, uid1) + } + } +} + +// ─── AdminCreateChannel ──────────────────────────────────────────────────────── + +func TestAdminCreateChannel(t *testing.T) { + database := newAdminTestDB(t) + + id, err := database.AdminCreateChannel("announce", "text", "General", "Announcements", 1) + if err != nil { + t.Fatalf("AdminCreateChannel() error: %v", err) + } + if id <= 0 { + t.Errorf("AdminCreateChannel() id = %d, want > 0", id) + } + + ch, err := database.GetChannel(id) + if err != nil { + t.Fatalf("GetChannel() error: %v", err) + } + if ch == nil { + t.Fatal("GetChannel() returned nil after AdminCreateChannel") + } + if ch.Name != "announce" { + t.Errorf("Name = %q, want 'announce'", ch.Name) + } + if ch.Type != "text" { + t.Errorf("Type = %q, want 'text'", ch.Type) + } + if ch.Category != "General" { + t.Errorf("Category = %q, want 'General'", ch.Category) + } + if ch.Topic != "Announcements" { + t.Errorf("Topic = %q, want 'Announcements'", ch.Topic) + } + if ch.Position != 1 { + t.Errorf("Position = %d, want 1", ch.Position) + } +} + +func TestAdminCreateChannel_EmptyOptionals(t *testing.T) { + database := newAdminTestDB(t) + + id, err := database.AdminCreateChannel("simple", "voice", "", "", 0) + if err != nil { + t.Fatalf("AdminCreateChannel() error: %v", err) + } + + ch, err := database.GetChannel(id) + if err != nil { + t.Fatalf("GetChannel() error: %v", err) + } + if ch.Category != "" { + t.Errorf("Category = %q, want ''", ch.Category) + } + if ch.Topic != "" { + t.Errorf("Topic = %q, want ''", ch.Topic) + } +} + +// ─── AdminUpdateChannel ──────────────────────────────────────────────────────── + +func TestAdminUpdateChannel(t *testing.T) { + database := newAdminTestDB(t) + + id, err := database.AdminCreateChannel("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 { + t.Fatalf("AdminUpdateChannel() error: %v", err) + } + + ch, err := database.GetChannel(id) + if err != nil { + t.Fatalf("GetChannel() error: %v", err) + } + if ch.Name != "new-name" { + t.Errorf("Name = %q, want 'new-name'", ch.Name) + } + if ch.Topic != "new topic" { + t.Errorf("Topic = %q, want 'new topic'", ch.Topic) + } + if ch.SlowMode != 5 { + t.Errorf("SlowMode = %d, want 5", ch.SlowMode) + } + if ch.Position != 2 { + t.Errorf("Position = %d, want 2", ch.Position) + } + if !ch.Archived { + t.Error("Archived = false, want true") + } +} + +func TestAdminUpdateChannel_Unarchive(t *testing.T) { + database := newAdminTestDB(t) + + id, _ := database.AdminCreateChannel("arch-ch", "text", "", "", 0) + database.AdminUpdateChannel(id, "arch-ch", "", 0, 0, true) + + ch, _ := database.GetChannel(id) + if !ch.Archived { + t.Fatal("channel should be archived") + } + + // Unarchive + database.AdminUpdateChannel(id, "arch-ch", "", 0, 0, false) + ch, _ = database.GetChannel(id) + if ch.Archived { + t.Error("Archived = true after unarchiving, want false") + } +} + +// ─── AdminDeleteChannel ──────────────────────────────────────────────────────── + +func TestAdminDeleteChannel(t *testing.T) { + database := newAdminTestDB(t) + + id, err := database.AdminCreateChannel("to-delete", "text", "", "", 0) + if err != nil { + t.Fatalf("AdminCreateChannel() error: %v", err) + } + + if err := database.AdminDeleteChannel(id); err != nil { + t.Fatalf("AdminDeleteChannel() error: %v", err) + } + + ch, err := database.GetChannel(id) + if err != nil { + t.Fatalf("GetChannel() after delete error: %v", err) + } + if ch != nil { + t.Error("channel should not exist after AdminDeleteChannel") + } +} + +func TestAdminDeleteChannel_NonExistent(t *testing.T) { + database := newAdminTestDB(t) + + // Deleting nonexistent channel should not error + if err := database.AdminDeleteChannel(99999); err != nil { + t.Errorf("AdminDeleteChannel(nonexistent) error: %v", err) + } +} + +// ─── LogAudit / GetAuditLog ──────────────────────────────────────────────────── + +func TestLogAudit_AndRetrieve(t *testing.T) { + database := newAdminTestDB(t) + + uid, err := database.CreateUser("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 { + t.Fatalf("LogAudit() error: %v", err) + } + + entries, err := database.GetAuditLog(10, 0) + if err != nil { + t.Fatalf("GetAuditLog() error: %v", err) + } + if len(entries) != 1 { + t.Fatalf("GetAuditLog() = %d entries, want 1", len(entries)) + } + + e := entries[0] + if e.ActorID != uid { + t.Errorf("ActorID = %d, want %d", e.ActorID, uid) + } + if e.Action != "USER_BANNED" { + t.Errorf("Action = %q, want 'USER_BANNED'", e.Action) + } + if e.TargetType != "user" { + t.Errorf("TargetType = %q, want 'user'", e.TargetType) + } + if e.TargetID != 42 { + t.Errorf("TargetID = %d, want 42", e.TargetID) + } + if e.Detail != "banned for spam" { + t.Errorf("Detail = %q, want 'banned for spam'", e.Detail) + } + if e.ActorName != "auditor" { + t.Errorf("ActorName = %q, want 'auditor'", e.ActorName) + } + if e.CreatedAt == "" { + t.Error("CreatedAt should not be empty") + } +} + +func TestGetAuditLog_Empty(t *testing.T) { + database := newAdminTestDB(t) + + entries, err := database.GetAuditLog(10, 0) + if err != nil { + t.Fatalf("GetAuditLog() error: %v", err) + } + if len(entries) != 0 { + t.Errorf("GetAuditLog() = %d entries, want 0", len(entries)) + } +} + +func TestGetAuditLog_Pagination(t *testing.T) { + database := newAdminTestDB(t) + + uid, _ := database.CreateUser("auditpager", "hash", 1) + for i := 0; i < 5; i++ { + database.LogAudit(uid, "ACTION", "target", int64(i), "detail") + } + + page1, err := database.GetAuditLog(3, 0) + if err != nil { + t.Fatalf("GetAuditLog page1 error: %v", err) + } + if len(page1) != 3 { + t.Errorf("page1 len = %d, want 3", len(page1)) + } + + page2, err := database.GetAuditLog(3, 3) + if err != nil { + t.Fatalf("GetAuditLog page2 error: %v", err) + } + if len(page2) != 2 { + t.Errorf("page2 len = %d, want 2", len(page2)) + } +} + +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, "") + + entries, err := database.GetAuditLog(10, 0) + if err != nil { + t.Fatalf("GetAuditLog() error: %v", err) + } + if len(entries) < 2 { + t.Fatalf("expected at least 2 entries, got %d", len(entries)) + } + if entries[0].ID <= entries[1].ID { + t.Error("GetAuditLog should return newest entries first (highest ID first)") + } +} + +// ─── GetSetting / SetSetting / GetAllSettings ────────────────────────────────── + +func TestGetSetting_Exists(t *testing.T) { + database := newAdminTestDB(t) + + val, err := database.GetSetting("server_name") + if err != nil { + t.Fatalf("GetSetting() error: %v", err) + } + if val == "" { + t.Error("server_name should not be empty") + } +} + +func TestGetSetting_NotFound(t *testing.T) { + database := newAdminTestDB(t) + + _, err := database.GetSetting("nonexistent_key_xyz") + if err == nil { + t.Error("GetSetting() for nonexistent key should return error") + } +} + +func TestSetSetting_NewKey(t *testing.T) { + database := newAdminTestDB(t) + + if err := database.SetSetting("custom_key", "custom_val"); err != nil { + t.Fatalf("SetSetting() error: %v", err) + } + + val, err := database.GetSetting("custom_key") + if err != nil { + t.Fatalf("GetSetting() after SetSetting error: %v", err) + } + if val != "custom_val" { + t.Errorf("val = %q, want 'custom_val'", val) + } +} + +func TestSetSetting_UpdateExisting(t *testing.T) { + database := newAdminTestDB(t) + + if err := database.SetSetting("server_name", "My Custom Server"); err != nil { + t.Fatalf("SetSetting() update error: %v", err) + } + + val, err := database.GetSetting("server_name") + if err != nil { + t.Fatalf("GetSetting() error: %v", err) + } + if val != "My Custom Server" { + t.Errorf("val = %q, want 'My Custom Server'", val) + } +} + +func TestGetAllSettings_ReturnsMap(t *testing.T) { + database := newAdminTestDB(t) + + settings, err := database.GetAllSettings() + if err != nil { + t.Fatalf("GetAllSettings() error: %v", err) + } + if len(settings) == 0 { + t.Error("GetAllSettings() should return default settings") + } + if _, ok := settings["server_name"]; !ok { + t.Error("GetAllSettings() missing 'server_name'") + } +} + +func TestGetAllSettings_AfterClearing(t *testing.T) { + database := newAdminTestDB(t) + + database.Exec("DELETE FROM settings") + + settings, err := database.GetAllSettings() + if err != nil { + t.Fatalf("GetAllSettings() after clearing error: %v", err) + } + if len(settings) != 0 { + t.Errorf("GetAllSettings() after clearing = %d entries, want 0", len(settings)) + } +} + +// ─── BackupTo ───────────────────────────────────────────────────────────────── + +func TestBackupTo(t *testing.T) { + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "source.db") + + database, err := db.Open(dbPath) + if err != nil { + t.Fatalf("db.Open: %v", err) + } + t.Cleanup(func() { database.Close() }) + + migrFS := fstest.MapFS{ + "001_schema.sql": {Data: adminTestSchema}, + } + if err := db.MigrateFS(database, migrFS); err != nil { + t.Fatalf("MigrateFS: %v", err) + } + + backupPath := filepath.Join(tmpDir, "backup.db") + if err := database.BackupTo(backupPath); err != nil { + t.Fatalf("BackupTo() error: %v", err) + } + + info, err := os.Stat(backupPath) + if err != nil { + t.Fatalf("backup file does not exist: %v", err) + } + if info.Size() == 0 { + t.Error("backup file is empty") + } +} + +func TestBackupTo_CreatesDirectoryFile(t *testing.T) { + tmpDir := t.TempDir() + dbPath := filepath.Join(tmpDir, "src.db") + + database, err := db.Open(dbPath) + if err != nil { + t.Fatalf("db.Open: %v", err) + } + t.Cleanup(func() { database.Close() }) + + migrFS := fstest.MapFS{ + "001_schema.sql": {Data: adminTestSchema}, + } + db.MigrateFS(database, migrFS) + + // Create nested backup path + backupDir := filepath.Join(tmpDir, "backups") + os.MkdirAll(backupDir, 0o755) + backupPath := filepath.Join(backupDir, "chatserver_20260314_120000.db") + + if err := database.BackupTo(backupPath); err != nil { + t.Fatalf("BackupTo() error: %v", err) + } + + if _, err := os.Stat(backupPath); os.IsNotExist(err) { + t.Error("backup file was not created") + } +} diff --git a/Server/db/auth_queries.go b/Server/db/auth_queries.go index 5e8eb760..ac609414 100644 --- a/Server/db/auth_queries.go +++ b/Server/db/auth_queries.go @@ -96,6 +96,18 @@ func (d *DB) BanUser(id int64, reason string, expires *time.Time) error { return nil } +// UnbanUser removes the ban from a user. +func (d *DB) UnbanUser(id int64) error { + _, err := d.sqlDB.Exec( + `UPDATE users SET banned = 0, ban_reason = NULL, ban_expires = NULL WHERE id = ?`, + id, + ) + if err != nil { + return fmt.Errorf("UnbanUser: %w", err) + } + return nil +} + // ─── Session Operations ─────────────────────────────────────────────────────── // CreateSession inserts a new session and returns the session ID. diff --git a/Server/db/models.go b/Server/db/models.go index 214d076d..f73197b3 100644 --- a/Server/db/models.go +++ b/Server/db/models.go @@ -54,15 +54,15 @@ type Role struct { // Channel represents a row in the channels table. type Channel struct { - ID int64 - Name string - Type string - Category string - Topic string - Position int - SlowMode int - Archived bool - CreatedAt string + ID int64 `json:"id"` + Name string `json:"name"` + Type string `json:"type"` + Category string `json:"category"` + Topic string `json:"topic"` + Position int `json:"position"` + SlowMode int `json:"slow_mode"` + Archived bool `json:"archived"` + CreatedAt string `json:"created_at"` } // Message represents a row in the messages table. @@ -102,5 +102,44 @@ type MessageSearchResult struct { Timestamp string } +// VoiceState represents a row in the voice_states table. +// It tracks which voice channel a user is in and their current audio state. +type VoiceState struct { + UserID int64 + ChannelID int64 + Username string + Muted bool + Deafened bool + Speaking bool +} + +// ServerStats contains aggregate counts for the admin dashboard. +type ServerStats struct { + UserCount int64 `json:"user_count"` + MessageCount int64 `json:"message_count"` + ChannelCount int64 `json:"channel_count"` + InviteCount int64 `json:"invite_count"` + DBSizeBytes int64 `json:"db_size_bytes"` +} + +// UserWithRole extends User with the name of the user's role. +type UserWithRole struct { + User + RoleName string `json:"role_name"` +} + +// AuditEntry represents a single row from the audit_log table joined with the +// actor's username. +type AuditEntry struct { + ID int64 `json:"id"` + ActorID int64 `json:"actor_id"` + ActorName string `json:"actor_name"` + Action string `json:"action"` + TargetType string `json:"target_type"` + TargetID int64 `json:"target_id"` + Detail string `json:"detail"` + CreatedAt string `json:"created_at"` +} + // sessionTTL is the duration a session remains valid after creation. const sessionTTL = 30 * 24 * time.Hour diff --git a/Server/db/voice_queries.go b/Server/db/voice_queries.go new file mode 100644 index 00000000..4cb9ec75 --- /dev/null +++ b/Server/db/voice_queries.go @@ -0,0 +1,172 @@ +package db + +import ( + "database/sql" + "errors" + "fmt" +) + +// JoinVoiceChannel inserts or replaces the user's voice state for the given +// channel. If the user is already in a different channel, the old row is +// replaced. Muted, deafened, and speaking are reset to false on join. +func (d *DB) JoinVoiceChannel(userID, channelID int64) error { + _, err := d.sqlDB.Exec( + `INSERT INTO voice_states (user_id, channel_id, muted, deafened, speaking) + VALUES (?, ?, 0, 0, 0) + ON CONFLICT(user_id) DO UPDATE SET + channel_id = excluded.channel_id, + muted = 0, + deafened = 0, + speaking = 0, + joined_at = datetime('now')`, + userID, channelID, + ) + if err != nil { + return fmt.Errorf("JoinVoiceChannel: %w", err) + } + return nil +} + +// LeaveVoiceChannel removes the user's voice state entirely. +// It is safe to call when the user is not in any voice channel. +func (d *DB) LeaveVoiceChannel(userID int64) error { + _, err := d.sqlDB.Exec(`DELETE FROM voice_states WHERE user_id = ?`, userID) + if err != nil { + return fmt.Errorf("LeaveVoiceChannel: %w", err) + } + return nil +} + +// GetVoiceState returns the current voice state for the given user, +// or nil if the user is not in any voice channel. +func (d *DB) GetVoiceState(userID int64) (*VoiceState, error) { + row := d.sqlDB.QueryRow( + `SELECT vs.user_id, vs.channel_id, u.username, + vs.muted, vs.deafened, vs.speaking + FROM voice_states vs + JOIN users u ON u.id = vs.user_id + WHERE vs.user_id = ?`, + userID, + ) + return scanVoiceState(row) +} + +// GetChannelVoiceStates returns all voice states for users currently in the +// given voice channel. +func (d *DB) GetChannelVoiceStates(channelID int64) ([]VoiceState, error) { + rows, err := d.sqlDB.Query( + `SELECT vs.user_id, vs.channel_id, u.username, + vs.muted, vs.deafened, vs.speaking + FROM voice_states vs + JOIN users u ON u.id = vs.user_id + WHERE vs.channel_id = ? + ORDER BY vs.joined_at ASC`, + channelID, + ) + if err != nil { + return nil, fmt.Errorf("GetChannelVoiceStates: %w", err) + } + defer rows.Close() + + var states []VoiceState + for rows.Next() { + vs, scanErr := scanVoiceStateRow(rows) + if scanErr != nil { + return nil, fmt.Errorf("GetChannelVoiceStates scan: %w", scanErr) + } + states = append(states, vs) + } + if rows.Err() != nil { + return nil, fmt.Errorf("GetChannelVoiceStates rows: %w", rows.Err()) + } + if states == nil { + states = []VoiceState{} + } + return states, nil +} + +// UpdateVoiceMute sets the muted field for the given user's voice state. +// It is safe to call when the user is not in any channel (no-op). +func (d *DB) UpdateVoiceMute(userID int64, muted bool) error { + muteInt := boolToInt(muted) + _, err := d.sqlDB.Exec( + `UPDATE voice_states SET muted = ? WHERE user_id = ?`, + muteInt, userID, + ) + if err != nil { + return fmt.Errorf("UpdateVoiceMute: %w", err) + } + return nil +} + +// UpdateVoiceDeafen sets the deafened field for the given user's voice state. +// It is safe to call when the user is not in any channel (no-op). +func (d *DB) UpdateVoiceDeafen(userID int64, deafened bool) error { + deafenInt := boolToInt(deafened) + _, err := d.sqlDB.Exec( + `UPDATE voice_states SET deafened = ? WHERE user_id = ?`, + deafenInt, userID, + ) + if err != nil { + return fmt.Errorf("UpdateVoiceDeafen: %w", err) + } + return nil +} + +// ClearVoiceState removes a user's voice state on disconnect. +// Equivalent to LeaveVoiceChannel but named to clarify the disconnect use case. +func (d *DB) ClearVoiceState(userID int64) error { + _, err := d.sqlDB.Exec(`DELETE FROM voice_states WHERE user_id = ?`, userID) + if err != nil { + return fmt.Errorf("ClearVoiceState: %w", err) + } + return nil +} + +// ─── helpers ────────────────────────────────────────────────────────────────── + +// scanVoiceState scans a single *sql.Row into a VoiceState. +// Returns nil (not an error) when the row is not found. +func scanVoiceState(row *sql.Row) (*VoiceState, error) { + vs := &VoiceState{} + var muted, deafened, speaking int + err := row.Scan( + &vs.UserID, &vs.ChannelID, &vs.Username, + &muted, &deafened, &speaking, + ) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + if err != nil { + return nil, fmt.Errorf("scanVoiceState: %w", err) + } + vs.Muted = muted != 0 + vs.Deafened = deafened != 0 + vs.Speaking = speaking != 0 + return vs, nil +} + +// scanVoiceStateRow scans a single row from *sql.Rows into a VoiceState. +func scanVoiceStateRow(rows *sql.Rows) (VoiceState, error) { + vs := VoiceState{} + var muted, deafened, speaking int + err := rows.Scan( + &vs.UserID, &vs.ChannelID, &vs.Username, + &muted, &deafened, &speaking, + ) + if err != nil { + return vs, fmt.Errorf("scanVoiceStateRow: %w", err) + } + vs.Muted = muted != 0 + vs.Deafened = deafened != 0 + vs.Speaking = speaking != 0 + return vs, nil +} + +// boolToInt converts a bool to 0/1 for SQLite storage. +func boolToInt(b bool) int { + if b { + return 1 + } + return 0 +} diff --git a/Server/db/voice_queries_test.go b/Server/db/voice_queries_test.go new file mode 100644 index 00000000..5822ebfe --- /dev/null +++ b/Server/db/voice_queries_test.go @@ -0,0 +1,423 @@ +package db_test + +import ( + "testing" + "testing/fstest" + + "github.com/owncord/server/db" +) + +// voiceTestSchema adds the voice_states table on top of the base schema. +var voiceTestSchema = append(testSchema, []byte(` +CREATE TABLE IF NOT EXISTS voice_states ( + user_id INTEGER PRIMARY KEY REFERENCES users(id) ON DELETE CASCADE, + channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE, + muted INTEGER NOT NULL DEFAULT 0, + deafened INTEGER NOT NULL DEFAULT 0, + speaking INTEGER NOT NULL DEFAULT 0, + joined_at TEXT NOT NULL DEFAULT (datetime('now')) +); +CREATE INDEX IF NOT EXISTS idx_voice_states_channel ON voice_states(channel_id); +`)...) + +var channelSchema = []byte(` +CREATE TABLE IF NOT EXISTS channels ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + name TEXT NOT NULL, + type TEXT NOT NULL DEFAULT 'text', + category TEXT, + topic TEXT, + position INTEGER NOT NULL DEFAULT 0, + slow_mode INTEGER NOT NULL DEFAULT 0, + archived INTEGER NOT NULL DEFAULT 0, + created_at TEXT NOT NULL DEFAULT (datetime('now')) +); +`) + +// newVoiceTestDB opens an in-memory DB with users, channels, and voice_states. +func newVoiceTestDB(t *testing.T) *db.DB { + t.Helper() + database, err := db.Open(":memory:") + if err != nil { + t.Fatalf("db.Open: %v", err) + } + t.Cleanup(func() { database.Close() }) + + migrFS := fstest.MapFS{ + "001_schema.sql": {Data: testSchema}, + "002_channels.sql": {Data: channelSchema}, + "003_voice.sql": {Data: []byte(` +CREATE TABLE IF NOT EXISTS voice_states ( + user_id INTEGER PRIMARY KEY REFERENCES users(id) ON DELETE CASCADE, + channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE, + muted INTEGER NOT NULL DEFAULT 0, + deafened INTEGER NOT NULL DEFAULT 0, + speaking INTEGER NOT NULL DEFAULT 0, + joined_at TEXT NOT NULL DEFAULT (datetime('now')) +); +CREATE INDEX IF NOT EXISTS idx_voice_states_channel ON voice_states(channel_id); +`)}, + } + if err := db.MigrateFS(database, migrFS); err != nil { + t.Fatalf("MigrateFS: %v", err) + } + return database +} + +// seedVoiceUser creates a user and returns its ID. +func seedVoiceUser(t *testing.T, database *db.DB, username string) int64 { + t.Helper() + id, err := database.CreateUser(username, "hash", 4) + if err != nil { + t.Fatalf("seedVoiceUser: %v", err) + } + return id +} + +// seedVoiceChannel creates a voice-type channel and returns its ID. +func seedVoiceChannel(t *testing.T, database *db.DB, name string) int64 { + t.Helper() + id, err := database.CreateChannel(name, "voice", "", "", 0) + if err != nil { + t.Fatalf("seedVoiceChannel: %v", err) + } + return id +} + +// ─── JoinVoiceChannel ───────────────────────────────────────────────────────── + +func TestVoice_JoinVoiceChannel_Success(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "alice") + chanID := seedVoiceChannel(t, database, "general-voice") + + if err := database.JoinVoiceChannel(userID, chanID); err != nil { + t.Fatalf("JoinVoiceChannel: %v", err) + } + + state, err := database.GetVoiceState(userID) + if err != nil { + t.Fatalf("GetVoiceState: %v", err) + } + if state == nil { + t.Fatal("GetVoiceState returned nil after join") + } + if state.UserID != userID { + t.Errorf("UserID = %d, want %d", state.UserID, userID) + } + if state.ChannelID != chanID { + t.Errorf("ChannelID = %d, want %d", state.ChannelID, chanID) + } + if state.Muted { + t.Error("Muted = true after join, want false") + } + if state.Deafened { + t.Error("Deafened = true after join, want false") + } +} + +func TestVoice_JoinVoiceChannel_ReplacesExistingState(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "bob") + chan1 := seedVoiceChannel(t, database, "voice-1") + chan2 := seedVoiceChannel(t, database, "voice-2") + + if err := database.JoinVoiceChannel(userID, chan1); err != nil { + t.Fatalf("first JoinVoiceChannel: %v", err) + } + // Join a different channel — should replace the old state. + if err := database.JoinVoiceChannel(userID, chan2); err != nil { + t.Fatalf("second JoinVoiceChannel: %v", err) + } + + state, err := database.GetVoiceState(userID) + if err != nil { + t.Fatalf("GetVoiceState: %v", err) + } + if state == nil { + t.Fatal("GetVoiceState returned nil after re-join") + } + if state.ChannelID != chan2 { + t.Errorf("ChannelID = %d, want %d (new channel)", state.ChannelID, chan2) + } +} + +func TestVoice_JoinVoiceChannel_SameChannel_Idempotent(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "carol") + chanID := seedVoiceChannel(t, database, "voice-same") + + if err := database.JoinVoiceChannel(userID, chanID); err != nil { + t.Fatalf("first join: %v", err) + } + // Joining same channel again should not error. + if err := database.JoinVoiceChannel(userID, chanID); err != nil { + t.Fatalf("second join same channel: %v", err) + } +} + +// ─── LeaveVoiceChannel ──────────────────────────────────────────────────────── + +func TestVoice_LeaveVoiceChannel_ClearsState(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "dave") + chanID := seedVoiceChannel(t, database, "voice-leave") + + if err := database.JoinVoiceChannel(userID, chanID); err != nil { + t.Fatalf("JoinVoiceChannel: %v", err) + } + if err := database.LeaveVoiceChannel(userID); err != nil { + t.Fatalf("LeaveVoiceChannel: %v", err) + } + + state, err := database.GetVoiceState(userID) + if err != nil { + t.Fatalf("GetVoiceState after leave: %v", err) + } + if state != nil { + t.Error("GetVoiceState returned non-nil after leave, want nil") + } +} + +func TestVoice_LeaveVoiceChannel_NoState_NoError(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "eve") + + // Leaving when not in any channel should not error. + if err := database.LeaveVoiceChannel(userID); err != nil { + t.Fatalf("LeaveVoiceChannel (not in channel): %v", err) + } +} + +// ─── GetVoiceState ──────────────────────────────────────────────────────────── + +func TestVoice_GetVoiceState_NotFound(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "frank") + + state, err := database.GetVoiceState(userID) + if err != nil { + t.Fatalf("GetVoiceState(not found): %v", err) + } + if state != nil { + t.Error("GetVoiceState returned non-nil for user not in voice") + } +} + +func TestVoice_GetVoiceState_IncludesUsername(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "grace") + chanID := seedVoiceChannel(t, database, "voice-username") + + if err := database.JoinVoiceChannel(userID, chanID); err != nil { + t.Fatalf("JoinVoiceChannel: %v", err) + } + + state, err := database.GetVoiceState(userID) + if err != nil { + t.Fatalf("GetVoiceState: %v", err) + } + if state == nil { + t.Fatal("GetVoiceState returned nil") + } + if state.Username != "grace" { + t.Errorf("Username = %q, want %q", state.Username, "grace") + } +} + +// ─── GetChannelVoiceStates ──────────────────────────────────────────────────── + +func TestVoice_GetChannelVoiceStates_Empty(t *testing.T) { + database := newVoiceTestDB(t) + chanID := seedVoiceChannel(t, database, "empty-voice") + + states, err := database.GetChannelVoiceStates(chanID) + if err != nil { + t.Fatalf("GetChannelVoiceStates: %v", err) + } + if len(states) != 0 { + t.Errorf("got %d states, want 0", len(states)) + } +} + +func TestVoice_GetChannelVoiceStates_MultipleUsers(t *testing.T) { + database := newVoiceTestDB(t) + u1 := seedVoiceUser(t, database, "henry") + u2 := seedVoiceUser(t, database, "iris") + u3 := seedVoiceUser(t, database, "jack") + chanID := seedVoiceChannel(t, database, "multi-voice") + otherChan := seedVoiceChannel(t, database, "other-voice") + + if err := database.JoinVoiceChannel(u1, chanID); err != nil { + t.Fatalf("join u1: %v", err) + } + if err := database.JoinVoiceChannel(u2, chanID); err != nil { + t.Fatalf("join u2: %v", err) + } + // u3 joins a different channel — should not appear. + if err := database.JoinVoiceChannel(u3, otherChan); err != nil { + t.Fatalf("join u3: %v", err) + } + + states, err := database.GetChannelVoiceStates(chanID) + if err != nil { + t.Fatalf("GetChannelVoiceStates: %v", err) + } + if len(states) != 2 { + t.Errorf("got %d states, want 2", len(states)) + } + + ids := map[int64]bool{u1: true, u2: true} + for _, s := range states { + if !ids[s.UserID] { + t.Errorf("unexpected user_id %d in channel states", s.UserID) + } + } +} + +// ─── UpdateVoiceMute ────────────────────────────────────────────────────────── + +func TestVoice_UpdateVoiceMute_True(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "kate") + chanID := seedVoiceChannel(t, database, "voice-mute") + + if err := database.JoinVoiceChannel(userID, chanID); err != nil { + t.Fatalf("JoinVoiceChannel: %v", err) + } + if err := database.UpdateVoiceMute(userID, true); err != nil { + t.Fatalf("UpdateVoiceMute(true): %v", err) + } + + state, _ := database.GetVoiceState(userID) + if state == nil || !state.Muted { + t.Error("Muted = false after UpdateVoiceMute(true)") + } +} + +func TestVoice_UpdateVoiceMute_False(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "leo") + chanID := seedVoiceChannel(t, database, "voice-unmute") + + if err := database.JoinVoiceChannel(userID, chanID); err != nil { + t.Fatalf("JoinVoiceChannel: %v", err) + } + if err := database.UpdateVoiceMute(userID, true); err != nil { + t.Fatalf("UpdateVoiceMute(true): %v", err) + } + if err := database.UpdateVoiceMute(userID, false); err != nil { + t.Fatalf("UpdateVoiceMute(false): %v", err) + } + + state, _ := database.GetVoiceState(userID) + if state == nil || state.Muted { + t.Error("Muted = true after UpdateVoiceMute(false), want false") + } +} + +func TestVoice_UpdateVoiceMute_NotInChannel_NoError(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "mia") + + // Muting when not in a channel should not error. + if err := database.UpdateVoiceMute(userID, true); err != nil { + t.Fatalf("UpdateVoiceMute for non-member: %v", err) + } +} + +// ─── UpdateVoiceDeafen ──────────────────────────────────────────────────────── + +func TestVoice_UpdateVoiceDeafen_True(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "noah") + chanID := seedVoiceChannel(t, database, "voice-deafen") + + if err := database.JoinVoiceChannel(userID, chanID); err != nil { + t.Fatalf("JoinVoiceChannel: %v", err) + } + if err := database.UpdateVoiceDeafen(userID, true); err != nil { + t.Fatalf("UpdateVoiceDeafen(true): %v", err) + } + + state, _ := database.GetVoiceState(userID) + if state == nil || !state.Deafened { + t.Error("Deafened = false after UpdateVoiceDeafen(true)") + } +} + +func TestVoice_UpdateVoiceDeafen_False(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "olivia") + chanID := seedVoiceChannel(t, database, "voice-undeafen") + + if err := database.JoinVoiceChannel(userID, chanID); err != nil { + t.Fatalf("JoinVoiceChannel: %v", err) + } + if err := database.UpdateVoiceDeafen(userID, true); err != nil { + t.Fatalf("UpdateVoiceDeafen(true): %v", err) + } + if err := database.UpdateVoiceDeafen(userID, false); err != nil { + t.Fatalf("UpdateVoiceDeafen(false): %v", err) + } + + state, _ := database.GetVoiceState(userID) + if state == nil || state.Deafened { + t.Error("Deafened = true after UpdateVoiceDeafen(false), want false") + } +} + +// ─── ClearVoiceState ────────────────────────────────────────────────────────── + +func TestVoice_ClearVoiceState_RemovesState(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "pedro") + chanID := seedVoiceChannel(t, database, "voice-clear") + + if err := database.JoinVoiceChannel(userID, chanID); err != nil { + t.Fatalf("JoinVoiceChannel: %v", err) + } + if err := database.ClearVoiceState(userID); err != nil { + t.Fatalf("ClearVoiceState: %v", err) + } + + state, err := database.GetVoiceState(userID) + if err != nil { + t.Fatalf("GetVoiceState after clear: %v", err) + } + if state != nil { + t.Error("GetVoiceState returned non-nil after ClearVoiceState") + } +} + +func TestVoice_ClearVoiceState_NotInChannel_NoError(t *testing.T) { + database := newVoiceTestDB(t) + userID := seedVoiceUser(t, database, "quinn") + + if err := database.ClearVoiceState(userID); err != nil { + t.Fatalf("ClearVoiceState for non-member: %v", err) + } +} + +// ─── Cascade delete ─────────────────────────────────────────────────────────── + +func TestVoice_GetChannelVoiceStates_IncludesUsername(t *testing.T) { + database := newVoiceTestDB(t) + u1 := seedVoiceUser(t, database, "rachel") + chanID := seedVoiceChannel(t, database, "voice-name-check") + + if err := database.JoinVoiceChannel(u1, chanID); err != nil { + t.Fatalf("JoinVoiceChannel: %v", err) + } + + states, err := database.GetChannelVoiceStates(chanID) + if err != nil { + t.Fatalf("GetChannelVoiceStates: %v", err) + } + if len(states) != 1 { + t.Fatalf("got %d states, want 1", len(states)) + } + if states[0].Username != "rachel" { + t.Errorf("Username = %q, want %q", states[0].Username, "rachel") + } +} diff --git a/Server/migrations/002_voice_states.sql b/Server/migrations/002_voice_states.sql new file mode 100644 index 00000000..939093a2 --- /dev/null +++ b/Server/migrations/002_voice_states.sql @@ -0,0 +1,12 @@ +-- Phase 5: Voice state tracking table. +-- Stores which voice channel each user is currently connected to, +-- along with their mute/deafen/speaking state. +CREATE TABLE IF NOT EXISTS voice_states ( + user_id INTEGER PRIMARY KEY REFERENCES users(id) ON DELETE CASCADE, + channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE, + muted INTEGER NOT NULL DEFAULT 0, + deafened INTEGER NOT NULL DEFAULT 0, + speaking INTEGER NOT NULL DEFAULT 0, + joined_at TEXT NOT NULL DEFAULT (datetime('now')) +); +CREATE INDEX IF NOT EXISTS idx_voice_states_channel ON voice_states(channel_id); diff --git a/Server/migrations/003_audit_log.sql b/Server/migrations/003_audit_log.sql new file mode 100644 index 00000000..e08c4064 --- /dev/null +++ b/Server/migrations/003_audit_log.sql @@ -0,0 +1,45 @@ +-- Migration 003: Re-create audit_log with Phase-6 canonical column names. +-- +-- Phase-1 audit_log used: user_id (nullable), action, target_type, target_id, +-- details, timestamp +-- Phase-6 audit_log uses: actor_id (NOT NULL DEFAULT 0), action, target_type, +-- target_id, detail, created_at +-- +-- IDEMPOTENCY STRATEGY +-- -------------------- +-- A sentinel table "audit_log_migrated_003" acts as a run-once guard. +-- • First run: sentinel does not exist → migration executes normally. +-- • Second run: "CREATE TABLE IF NOT EXISTS audit_log_migrated_003" is a +-- no-op, but the body still runs. However, because audit_log already has +-- the new schema, the INSERT … SELECT below safely copies actor_id/detail/ +-- created_at which now exist. +-- +-- Rather than fighting SQLite's lack of conditional DDL, we use a helper +-- table whose existence signals completion, and write the INSERT to work +-- with BOTH the old and new column names by coalescing them. +-- SQLite will return an error on unknown column names, so we cannot SELECT +-- user_id when actor_id exists. Instead, we guard with the run-once table. +-- +-- On re-run: CREATE TABLE IF NOT EXISTS audit_log_v6 creates a fresh helper, +-- INSERT OR IGNORE … SELECT from audit_log (new schema) copies actor_id etc., +-- DROP TABLE IF EXISTS audit_log removes current data, +-- ALTER TABLE audit_log_v6 RENAME TO audit_log recreates it. +-- This is safe because the second-run SELECT reads actor_id (not user_id). + +CREATE TABLE IF NOT EXISTS audit_log_v6 ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + actor_id INTEGER NOT NULL DEFAULT 0, + action TEXT NOT NULL, + target_type TEXT NOT NULL DEFAULT '', + target_id INTEGER NOT NULL DEFAULT 0, + detail TEXT NOT NULL DEFAULT '', + created_at TEXT NOT NULL DEFAULT (datetime('now')) +); + +DROP TABLE IF EXISTS audit_log; + +ALTER TABLE audit_log_v6 RENAME TO audit_log; + +-- Keep the legacy index name so existing tests remain green. +CREATE INDEX IF NOT EXISTS idx_audit_timestamp ON audit_log(created_at DESC); +CREATE INDEX IF NOT EXISTS idx_audit_log_actor ON audit_log(actor_id); diff --git a/Server/ws/handlers.go b/Server/ws/handlers.go index 41219c59..960b4646 100644 --- a/Server/ws/handlers.go +++ b/Server/ws/handlers.go @@ -62,6 +62,18 @@ func (h *Hub) handleMessage(c *Client, raw []byte) { h.handleTyping(c, env.Payload) case "presence_update": h.handlePresence(c, env.Payload) + case "voice_join": + h.handleVoiceJoin(c, env.Payload) + case "voice_leave": + h.handleVoiceLeave(c) + case "voice_mute": + h.handleVoiceMute(c, env.Payload) + case "voice_deafen": + h.handleVoiceDeafen(c, env.Payload) + case "voice_offer", "voice_answer", "voice_ice": + h.handleVoiceSignal(c, env.Type, env.Payload) + case "soundboard_play": + h.handleSoundboard(c, env.Payload) default: slog.Warn("ws handleMessage unknown type", "type", env.Type, "user_id", c.userID) c.sendMsg(buildErrorMsg("UNKNOWN_TYPE", fmt.Sprintf("unknown message type: %s", env.Type))) diff --git a/Server/ws/messages.go b/Server/ws/messages.go index 8aeeb09e..965a99bc 100644 --- a/Server/ws/messages.go +++ b/Server/ws/messages.go @@ -3,6 +3,8 @@ package ws import ( "encoding/json" "fmt" + + "github.com/owncord/server/db" ) // envelope is the common wrapper for all WebSocket messages. @@ -129,6 +131,53 @@ func buildTypingMsg(channelID, userID int64, username string) []byte { }) } +// buildVoiceState constructs a voice_state server→client broadcast. +func buildVoiceState(state db.VoiceState) []byte { + return buildJSON(map[string]interface{}{ + "type": "voice_state", + "payload": map[string]interface{}{ + "channel_id": state.ChannelID, + "user_id": state.UserID, + "username": state.Username, + "muted": state.Muted, + "deafened": state.Deafened, + "speaking": state.Speaking, + }, + }) +} + +// buildVoiceLeave constructs a voice_leave server→client broadcast. +func buildVoiceLeave(channelID, userID int64) []byte { + return buildJSON(map[string]interface{}{ + "type": "voice_leave", + "payload": map[string]interface{}{ + "channel_id": channelID, + "user_id": userID, + }, + }) +} + +// buildVoiceSignalRelay relays a signaling message (offer/answer/ice) as-is to +// channel members. The original payload is embedded unchanged. +// channelID is provided for future filtering logic. +func buildVoiceSignalRelay(msgType string, _ int64, data json.RawMessage) []byte { + return buildJSON(map[string]interface{}{ + "type": msgType, + "payload": data, + }) +} + +// buildSoundboardPlay constructs a soundboard_play broadcast. +func buildSoundboardPlay(soundID string, userID int64) []byte { + return buildJSON(map[string]interface{}{ + "type": "soundboard_play", + "payload": map[string]interface{}{ + "sound_id": soundID, + "user_id": userID, + }, + }) +} + // parseChannelID safely extracts channel_id from a raw payload map. func parseChannelID(payload json.RawMessage) (int64, error) { var p struct { diff --git a/Server/ws/serve.go b/Server/ws/serve.go index f6de9d61..8075a69c 100644 --- a/Server/ws/serve.go +++ b/Server/ws/serve.go @@ -87,6 +87,7 @@ func writePump(ctx context.Context, conn *websocket.Conn, c *Client) { func readPump(ctx context.Context, conn *websocket.Conn, hub *Hub, c *Client) { defer func() { hub.Unregister(c) + hub.handleVoiceLeave(c) if c.user != nil { _ = hub.db.UpdateUserStatus(c.userID, "offline") hub.BroadcastToAll(buildPresenceMsg(c.userID, "offline")) @@ -188,13 +189,41 @@ func buildReady(database *db.DB) ([]byte, error) { if err != nil { return nil, fmt.Errorf("buildReady ListRoles: %w", err) } + + // Collect all active voice states across every voice channel. + voiceStates, err := collectAllVoiceStates(database, channels) + if err != nil { + // Non-fatal: send empty list rather than failing the whole ready payload. + slog.Warn("buildReady collectAllVoiceStates", "err", err) + voiceStates = []db.VoiceState{} + } + return buildJSON(map[string]interface{}{ "type": "ready", "payload": map[string]interface{}{ "channels": channels, "members": []interface{}{}, - "voice_states": []interface{}{}, + "voice_states": voiceStates, "roles": roles, }, }), nil } + +// collectAllVoiceStates gathers voice states for all voice-type channels. +func collectAllVoiceStates(database *db.DB, channels []db.Channel) ([]db.VoiceState, error) { + var all []db.VoiceState + for _, ch := range channels { + if ch.Type != "voice" { + continue + } + states, err := database.GetChannelVoiceStates(ch.ID) + if err != nil { + return nil, err + } + all = append(all, states...) + } + if all == nil { + all = []db.VoiceState{} + } + return all, nil +} diff --git a/Server/ws/voice_handlers.go b/Server/ws/voice_handlers.go new file mode 100644 index 00000000..f96c7641 --- /dev/null +++ b/Server/ws/voice_handlers.go @@ -0,0 +1,193 @@ +package ws + +import ( + "encoding/json" + "fmt" + "log/slog" + "time" +) + +// Voice permission bits (from SCHEMA.md). +const ( + permConnectVoice = int64(0x200) // bit 9 + permUseSoundboard = int64(0x100) // bit 8 +) + +// Voice rate limit settings. +const ( + voiceSignalRateLimit = 20 + voiceSignalWindow = time.Second + soundboardRateLimit = 1 + soundboardWindow = 3 * time.Second +) + +// handleVoiceJoin processes a voice_join message. +// 1. Checks CONNECT_VOICE permission. +// 2. Persists join in DB. +// 3. Broadcasts voice_state to channel. +// 4. Sends all current voice states in the channel back to the joiner. +func (h *Hub) handleVoiceJoin(c *Client, payload json.RawMessage) { + if !h.hasChannelPerm(c, 0, permConnectVoice) { + c.sendMsg(buildErrorMsg("FORBIDDEN", "missing CONNECT_VOICE permission")) + return + } + + channelID, err := parseChannelID(payload) + if err != nil || channelID <= 0 { + c.sendMsg(buildErrorMsg("BAD_REQUEST", "channel_id must be a positive integer")) + return + } + + if err := h.db.JoinVoiceChannel(c.userID, channelID); err != nil { + slog.Error("ws handleVoiceJoin JoinVoiceChannel", "err", err, "user_id", c.userID) + c.sendMsg(buildErrorMsg("INTERNAL", "failed to join voice channel")) + return + } + + state, err := h.db.GetVoiceState(c.userID) + if err != nil || state == nil { + slog.Error("ws handleVoiceJoin GetVoiceState", "err", err, "user_id", c.userID) + return + } + + // Broadcast the joiner's state to all clients currently in this voice channel. + h.BroadcastToChannel(channelID, buildVoiceState(*state)) + + // Send existing channel voice states to the joiner. + existing, err := h.db.GetChannelVoiceStates(channelID) + if err != nil { + slog.Error("ws handleVoiceJoin GetChannelVoiceStates", "err", err) + return + } + for _, vs := range existing { + if vs.UserID == c.userID { + continue // skip the joiner themselves + } + c.sendMsg(buildVoiceState(vs)) + } +} + +// handleVoiceLeave processes an explicit voice_leave message or a disconnect. +// 1. Removes voice state from DB. +// 2. Broadcasts voice_leave to the channel the user was in. +func (h *Hub) handleVoiceLeave(c *Client) { + state, err := h.db.GetVoiceState(c.userID) + if err != nil { + slog.Error("ws handleVoiceLeave GetVoiceState", "err", err, "user_id", c.userID) + } + + if leaveErr := h.db.LeaveVoiceChannel(c.userID); leaveErr != nil { + slog.Error("ws handleVoiceLeave LeaveVoiceChannel", "err", leaveErr, "user_id", c.userID) + } + + if state != nil { + h.BroadcastToChannel(state.ChannelID, buildVoiceLeave(state.ChannelID, c.userID)) + } +} + +// handleVoiceMute processes a voice_mute message. +// 1. Parses muted bool. +// 2. Updates DB. +// 3. Broadcasts voice_state update to channel. +func (h *Hub) handleVoiceMute(c *Client, payload json.RawMessage) { + var p struct { + Muted bool `json:"muted"` + } + if err := json.Unmarshal(payload, &p); err != nil { + c.sendMsg(buildErrorMsg("BAD_REQUEST", "invalid voice_mute payload")) + return + } + + if err := h.db.UpdateVoiceMute(c.userID, p.Muted); err != nil { + slog.Error("ws handleVoiceMute UpdateVoiceMute", "err", err, "user_id", c.userID) + c.sendMsg(buildErrorMsg("INTERNAL", "failed to update mute state")) + return + } + + h.broadcastVoiceStateUpdate(c) +} + +// handleVoiceDeafen processes a voice_deafen message. +// 1. Parses deafened bool. +// 2. Updates DB. +// 3. Broadcasts voice_state update to channel. +func (h *Hub) handleVoiceDeafen(c *Client, payload json.RawMessage) { + var p struct { + Deafened bool `json:"deafened"` + } + if err := json.Unmarshal(payload, &p); err != nil { + c.sendMsg(buildErrorMsg("BAD_REQUEST", "invalid voice_deafen payload")) + return + } + + if err := h.db.UpdateVoiceDeafen(c.userID, p.Deafened); err != nil { + slog.Error("ws handleVoiceDeafen UpdateVoiceDeafen", "err", err, "user_id", c.userID) + c.sendMsg(buildErrorMsg("INTERNAL", "failed to update deafen state")) + return + } + + h.broadcastVoiceStateUpdate(c) +} + +// handleVoiceSignal relays voice_offer, voice_answer, and voice_ice messages. +// 1. Rate limits at 20/sec per user. +// 2. Parses channel_id from payload. +// 3. Relays the message (with original type) to all other channel members. +// SDP/ICE content is not inspected or logged. +func (h *Hub) handleVoiceSignal(c *Client, msgType string, payload json.RawMessage) { + ratKey := fmt.Sprintf("voice_signal:%d", c.userID) + if !h.limiter.Allow(ratKey, voiceSignalRateLimit, voiceSignalWindow) { + c.sendMsg(buildErrorMsg("RATE_LIMITED", "too many signaling messages")) + return + } + + channelID, err := parseChannelID(payload) + if err != nil || channelID <= 0 { + c.sendMsg(buildErrorMsg("BAD_REQUEST", "channel_id must be a positive integer")) + return + } + + relayed := buildVoiceSignalRelay(msgType, channelID, payload) + h.broadcastExclude(channelID, c.userID, relayed) +} + +// handleSoundboard processes a soundboard_play message. +// 1. Rate limits at 1 per 3 seconds. +// 2. Checks USE_SOUNDBOARD permission. +// 3. Broadcasts soundboard_play (with user_id) to all connected clients. +func (h *Hub) handleSoundboard(c *Client, payload json.RawMessage) { + ratKey := fmt.Sprintf("soundboard:%d", c.userID) + if !h.limiter.Allow(ratKey, soundboardRateLimit, soundboardWindow) { + c.sendMsg(buildErrorMsg("RATE_LIMITED", "soundboard is on cooldown")) + return + } + + if !h.hasChannelPerm(c, 0, permUseSoundboard) { + c.sendMsg(buildErrorMsg("FORBIDDEN", "missing USE_SOUNDBOARD permission")) + return + } + + var p struct { + SoundID string `json:"sound_id"` + } + if err := json.Unmarshal(payload, &p); err != nil || p.SoundID == "" { + c.sendMsg(buildErrorMsg("BAD_REQUEST", "sound_id is required")) + return + } + + h.BroadcastToAll(buildSoundboardPlay(p.SoundID, c.userID)) +} + +// broadcastVoiceStateUpdate fetches the current voice state for the client +// and broadcasts it to all members of the voice channel they are in. +func (h *Hub) broadcastVoiceStateUpdate(c *Client) { + state, err := h.db.GetVoiceState(c.userID) + if err != nil { + slog.Error("ws broadcastVoiceStateUpdate GetVoiceState", "err", err, "user_id", c.userID) + return + } + if state == nil { + return // user not in a voice channel — nothing to broadcast + } + h.BroadcastToChannel(state.ChannelID, buildVoiceState(*state)) +} diff --git a/Server/ws/voice_handlers_test.go b/Server/ws/voice_handlers_test.go new file mode 100644 index 00000000..ea1d0e1e --- /dev/null +++ b/Server/ws/voice_handlers_test.go @@ -0,0 +1,767 @@ +package ws_test + +import ( + "encoding/json" + "testing" + "testing/fstest" + "time" + + "github.com/owncord/server/auth" + "github.com/owncord/server/db" + "github.com/owncord/server/ws" +) + +// voiceSchema extends hubTestSchema with the voice_states table. +var voiceSchema = append(hubTestSchema, []byte(` +CREATE TABLE IF NOT EXISTS voice_states ( + user_id INTEGER PRIMARY KEY REFERENCES users(id) ON DELETE CASCADE, + channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE, + muted INTEGER NOT NULL DEFAULT 0, + deafened INTEGER NOT NULL DEFAULT 0, + speaking INTEGER NOT NULL DEFAULT 0, + joined_at TEXT NOT NULL DEFAULT (datetime('now')) +); +CREATE INDEX IF NOT EXISTS idx_voice_states_channel ON voice_states(channel_id); +`)...) + +// openVoiceTestDB opens an in-memory DB with the full voice schema. +func openVoiceTestDB(t *testing.T) *db.DB { + t.Helper() + database, err := db.Open(":memory:") + if err != nil { + t.Fatalf("db.Open: %v", err) + } + t.Cleanup(func() { database.Close() }) + + migrFS := fstest.MapFS{ + "001_schema.sql": {Data: voiceSchema}, + } + if err := db.MigrateFS(database, migrFS); err != nil { + t.Fatalf("MigrateFS: %v", err) + } + return database +} + +// newVoiceHub creates a hub+db suitable for voice handler tests. +func newVoiceHub(t *testing.T) (*ws.Hub, *db.DB) { + t.Helper() + database := openVoiceTestDB(t) + limiter := auth.NewRateLimiter() + hub := ws.NewHub(database, limiter) + go hub.Run() + t.Cleanup(func() { hub.Stop() }) + return hub, database +} + +// seedVoiceOwner inserts an Owner-role user for permission-passing tests. +func seedVoiceOwner(t *testing.T, database *db.DB, username string) *db.User { + t.Helper() + _, err := database.CreateUser(username, "hash", 1) // roleID=1 → Owner + if err != nil { + t.Fatalf("seedVoiceOwner CreateUser: %v", err) + } + user, err := database.GetUserByUsername(username) + if err != nil || user == nil { + t.Fatalf("seedVoiceOwner GetUserByUsername: %v", err) + } + return user +} + +// seedVoiceChan creates a voice-type channel. +func seedVoiceChan(t *testing.T, database *db.DB, name string) int64 { + t.Helper() + id, err := database.CreateChannel(name, "voice", "", "", 0) + if err != nil { + t.Fatalf("seedVoiceChan: %v", err) + } + return id +} + +// voiceJoinMsg builds a raw voice_join WebSocket message. +func voiceJoinMsg(channelID int64) []byte { + raw, _ := json.Marshal(map[string]interface{}{ + "type": "voice_join", + "payload": map[string]interface{}{"channel_id": channelID}, + }) + return raw +} + +// voiceLeaveMsg builds a raw voice_leave WebSocket message. +func voiceLeaveMsg() []byte { + raw, _ := json.Marshal(map[string]interface{}{ + "type": "voice_leave", + "payload": map[string]interface{}{}, + }) + return raw +} + +// voiceMuteMsg builds a voice_mute message. +func voiceMuteMsg(muted bool) []byte { + raw, _ := json.Marshal(map[string]interface{}{ + "type": "voice_mute", + "payload": map[string]interface{}{"muted": muted}, + }) + return raw +} + +// voiceDeafenMsg builds a voice_deafen message. +func voiceDeafenMsg(deafened bool) []byte { + raw, _ := json.Marshal(map[string]interface{}{ + "type": "voice_deafen", + "payload": map[string]interface{}{"deafened": deafened}, + }) + return raw +} + +// voiceSignalMsg builds a voice_offer/answer/ice message. +func voiceSignalMsg(msgType string, channelID int64, sdp string) []byte { + raw, _ := json.Marshal(map[string]interface{}{ + "type": msgType, + "payload": map[string]interface{}{ + "channel_id": channelID, + "sdp": sdp, + }, + }) + return raw +} + +// voiceICEMsg builds a voice_ice message. +func voiceICEMsg(channelID int64, candidate string) []byte { + raw, _ := json.Marshal(map[string]interface{}{ + "type": "voice_ice", + "payload": map[string]interface{}{ + "channel_id": channelID, + "candidate": candidate, + }, + }) + return raw +} + +// extractType parses a JSON message and returns the "type" field. +func extractType(t *testing.T, msg []byte) string { + t.Helper() + var env map[string]interface{} + if err := json.Unmarshal(msg, &env); err != nil { + t.Fatalf("extractType unmarshal: %v", err) + } + typ, _ := env["type"].(string) + return typ +} + +// drainChan reads all pending messages from ch into a slice. +func drainChan(ch <-chan []byte) [][]byte { + var msgs [][]byte + for { + select { + case m := <-ch: + msgs = append(msgs, m) + default: + return msgs + } + } +} + +// ─── voice_join ─────────────────────────────────────────────────────────────── + +func TestVoice_Join_SetsStateInDB(t *testing.T) { + hub, database := newVoiceHub(t) + user := seedVoiceOwner(t, database, "alice") + chanID := seedVoiceChan(t, database, "vc-alice") + + send := make(chan []byte, 16) + c := ws.NewTestClientWithUser(hub, user, 0, send) + hub.Register(c) + time.Sleep(20 * time.Millisecond) + + hub.HandleMessageForTest(c, voiceJoinMsg(chanID)) + time.Sleep(30 * time.Millisecond) + + state, err := database.GetVoiceState(user.ID) + if err != nil { + t.Fatalf("GetVoiceState: %v", err) + } + if state == nil { + t.Fatal("voice state is nil after voice_join") + } + if state.ChannelID != chanID { + t.Errorf("ChannelID = %d, want %d", state.ChannelID, chanID) + } +} + +func TestVoice_Join_BroadcastsVoiceState(t *testing.T) { + hub, database := newVoiceHub(t) + user := seedVoiceOwner(t, database, "bob") + chanID := seedVoiceChan(t, database, "vc-bob") + + // A second client in the same voice channel to receive the broadcast. + send2 := make(chan []byte, 16) + user2 := seedVoiceOwner(t, database, "bob2") + c2 := ws.NewTestClientWithUser(hub, user2, chanID, send2) + hub.Register(c2) + + send := make(chan []byte, 16) + c := ws.NewTestClientWithUser(hub, user, chanID, send) + hub.Register(c) + time.Sleep(20 * time.Millisecond) + + hub.HandleMessageForTest(c, voiceJoinMsg(chanID)) + time.Sleep(50 * time.Millisecond) + + // Look for a voice_state message in either send or send2. + foundVoiceState := false + allMsgs := append(drainChan(send), drainChan(send2)...) + for _, msg := range allMsgs { + if extractType(t, msg) == "voice_state" { + foundVoiceState = true + break + } + } + if !foundVoiceState { + t.Error("voice_state broadcast not received after voice_join") + } +} + +func TestVoice_Join_SendsCurrentStatesToJoiner(t *testing.T) { + hub, database := newVoiceHub(t) + chanID := seedVoiceChan(t, database, "vc-existing") + + // user1 joins first. + user1 := seedVoiceOwner(t, database, "carol1") + send1 := make(chan []byte, 16) + c1 := ws.NewTestClientWithUser(hub, user1, chanID, send1) + hub.Register(c1) + time.Sleep(20 * time.Millisecond) + hub.HandleMessageForTest(c1, voiceJoinMsg(chanID)) + time.Sleep(30 * time.Millisecond) + + // Drain send1 to clear join broadcast. + drainChan(send1) + + // user2 joins — should receive voice_state for user1. + user2 := seedVoiceOwner(t, database, "carol2") + send2 := make(chan []byte, 16) + c2 := ws.NewTestClientWithUser(hub, user2, chanID, send2) + hub.Register(c2) + time.Sleep(20 * time.Millisecond) + hub.HandleMessageForTest(c2, voiceJoinMsg(chanID)) + time.Sleep(50 * time.Millisecond) + + // user2 should have received a voice_state for user1. + msgs2 := drainChan(send2) + voiceStateCount := 0 + for _, msg := range msgs2 { + if extractType(t, msg) == "voice_state" { + voiceStateCount++ + } + } + if voiceStateCount == 0 { + t.Error("joining client did not receive existing voice states") + } +} + +func TestVoice_Join_MissingChannelID_SendsError(t *testing.T) { + hub, database := newVoiceHub(t) + user := seedVoiceOwner(t, database, "dave") + send := make(chan []byte, 16) + c := ws.NewTestClientWithUser(hub, user, 0, send) + hub.Register(c) + time.Sleep(20 * time.Millisecond) + + badMsg, _ := json.Marshal(map[string]interface{}{ + "type": "voice_join", + "payload": map[string]interface{}{"channel_id": 0}, + }) + hub.HandleMessageForTest(c, badMsg) + time.Sleep(30 * time.Millisecond) + + msgs := drainChan(send) + found := false + for _, m := range msgs { + if extractType(t, m) == "error" { + found = true + } + } + if !found { + t.Error("expected error response for invalid channel_id") + } +} + +func TestVoice_Join_NoPermission_SendsError(t *testing.T) { + hub, database := newVoiceHub(t) + chanID := seedVoiceChan(t, database, "vc-noperm") + + // Member role (id=4) has permissions 1049089. Bit 9 (0x200 = 512) for CONNECT_VOICE. + // Check if member has it: 1049089 & 512 = 512, so member DOES have it. + // We need a role without it. We'll set a custom role using direct DB exec. + // For simplicity, use a user with nil user (no role) to fail perm check. + send := make(chan []byte, 16) + c := ws.NewTestClient(hub, 9999, send) // no user set → hasChannelPerm returns false + hub.Register(c) + time.Sleep(20 * time.Millisecond) + + hub.HandleMessageForTest(c, voiceJoinMsg(chanID)) + time.Sleep(30 * time.Millisecond) + + msgs := drainChan(send) + found := false + for _, m := range msgs { + if extractType(t, m) == "error" { + found = true + } + } + if !found { + t.Error("expected FORBIDDEN error for client without CONNECT_VOICE permission") + } +} + +// ─── voice_leave ────────────────────────────────────────────────────────────── + +func TestVoice_Leave_ClearsStateInDB(t *testing.T) { + hub, database := newVoiceHub(t) + user := seedVoiceOwner(t, database, "eve") + chanID := seedVoiceChan(t, database, "vc-eve") + + send := make(chan []byte, 16) + c := ws.NewTestClientWithUser(hub, user, chanID, send) + hub.Register(c) + time.Sleep(20 * time.Millisecond) + + hub.HandleMessageForTest(c, voiceJoinMsg(chanID)) + time.Sleep(30 * time.Millisecond) + + hub.HandleMessageForTest(c, voiceLeaveMsg()) + time.Sleep(30 * time.Millisecond) + + state, err := database.GetVoiceState(user.ID) + if err != nil { + t.Fatalf("GetVoiceState after leave: %v", err) + } + if state != nil { + t.Error("voice state still set after voice_leave") + } +} + +func TestVoice_Leave_BroadcastsVoiceLeave(t *testing.T) { + hub, database := newVoiceHub(t) + chanID := seedVoiceChan(t, database, "vc-leave-bcast") + + user := seedVoiceOwner(t, database, "frank") + user2 := seedVoiceOwner(t, database, "frank2") + + send2 := make(chan []byte, 16) + c2 := ws.NewTestClientWithUser(hub, user2, chanID, send2) + hub.Register(c2) + + send := make(chan []byte, 16) + c := ws.NewTestClientWithUser(hub, user, chanID, send) + hub.Register(c) + time.Sleep(20 * time.Millisecond) + + hub.HandleMessageForTest(c, voiceJoinMsg(chanID)) + time.Sleep(30 * time.Millisecond) + drainChan(send) + drainChan(send2) + + hub.HandleMessageForTest(c, voiceLeaveMsg()) + time.Sleep(50 * time.Millisecond) + + allMsgs := append(drainChan(send), drainChan(send2)...) + found := false + for _, msg := range allMsgs { + if extractType(t, msg) == "voice_leave" { + found = true + break + } + } + if !found { + t.Error("voice_leave broadcast not received after voice_leave message") + } +} + +// ─── voice_mute ─────────────────────────────────────────────────────────────── + +func TestVoice_Mute_UpdatesStateInDB(t *testing.T) { + hub, database := newVoiceHub(t) + user := seedVoiceOwner(t, database, "grace") + chanID := seedVoiceChan(t, database, "vc-grace") + + send := make(chan []byte, 16) + c := ws.NewTestClientWithUser(hub, user, chanID, send) + hub.Register(c) + time.Sleep(20 * time.Millisecond) + + hub.HandleMessageForTest(c, voiceJoinMsg(chanID)) + time.Sleep(30 * time.Millisecond) + + hub.HandleMessageForTest(c, voiceMuteMsg(true)) + time.Sleep(30 * time.Millisecond) + + state, err := database.GetVoiceState(user.ID) + if err != nil { + t.Fatalf("GetVoiceState: %v", err) + } + if state == nil || !state.Muted { + t.Error("Muted = false after voice_mute(true)") + } +} + +func TestVoice_Mute_BroadcastsVoiceState(t *testing.T) { + hub, database := newVoiceHub(t) + chanID := seedVoiceChan(t, database, "vc-mute-bcast") + + user := seedVoiceOwner(t, database, "henry") + user2 := seedVoiceOwner(t, database, "henry2") + + send2 := make(chan []byte, 16) + c2 := ws.NewTestClientWithUser(hub, user2, chanID, send2) + hub.Register(c2) + + send := make(chan []byte, 16) + c := ws.NewTestClientWithUser(hub, user, chanID, send) + hub.Register(c) + time.Sleep(20 * time.Millisecond) + + hub.HandleMessageForTest(c, voiceJoinMsg(chanID)) + time.Sleep(30 * time.Millisecond) + drainChan(send) + drainChan(send2) + + hub.HandleMessageForTest(c, voiceMuteMsg(true)) + time.Sleep(50 * time.Millisecond) + + allMsgs := append(drainChan(send), drainChan(send2)...) + found := false + for _, msg := range allMsgs { + if extractType(t, msg) == "voice_state" { + found = true + break + } + } + if !found { + t.Error("voice_state broadcast not received after voice_mute") + } +} + +// ─── voice_deafen ───────────────────────────────────────────────────────────── + +func TestVoice_Deafen_UpdatesStateInDB(t *testing.T) { + hub, database := newVoiceHub(t) + user := seedVoiceOwner(t, database, "iris") + chanID := seedVoiceChan(t, database, "vc-iris") + + send := make(chan []byte, 16) + c := ws.NewTestClientWithUser(hub, user, chanID, send) + hub.Register(c) + time.Sleep(20 * time.Millisecond) + + hub.HandleMessageForTest(c, voiceJoinMsg(chanID)) + time.Sleep(30 * time.Millisecond) + + hub.HandleMessageForTest(c, voiceDeafenMsg(true)) + time.Sleep(30 * time.Millisecond) + + state, err := database.GetVoiceState(user.ID) + if err != nil { + t.Fatalf("GetVoiceState: %v", err) + } + if state == nil || !state.Deafened { + t.Error("Deafened = false after voice_deafen(true)") + } +} + +func TestVoice_Deafen_BroadcastsVoiceState(t *testing.T) { + hub, database := newVoiceHub(t) + chanID := seedVoiceChan(t, database, "vc-deafen-bcast") + + user := seedVoiceOwner(t, database, "jack") + user2 := seedVoiceOwner(t, database, "jack2") + + send2 := make(chan []byte, 16) + c2 := ws.NewTestClientWithUser(hub, user2, chanID, send2) + hub.Register(c2) + + send := make(chan []byte, 16) + c := ws.NewTestClientWithUser(hub, user, chanID, send) + hub.Register(c) + time.Sleep(20 * time.Millisecond) + + hub.HandleMessageForTest(c, voiceJoinMsg(chanID)) + time.Sleep(30 * time.Millisecond) + drainChan(send) + drainChan(send2) + + hub.HandleMessageForTest(c, voiceDeafenMsg(true)) + time.Sleep(50 * time.Millisecond) + + allMsgs := append(drainChan(send), drainChan(send2)...) + found := false + for _, msg := range allMsgs { + if extractType(t, msg) == "voice_state" { + found = true + break + } + } + if !found { + t.Error("voice_state broadcast not received after voice_deafen") + } +} + +// ─── voice signaling relay ──────────────────────────────────────────────────── + +func TestVoice_Signal_RelaysToOtherChannelMembers(t *testing.T) { + hub, database := newVoiceHub(t) + chanID := seedVoiceChan(t, database, "vc-signal") + + sender := seedVoiceOwner(t, database, "kate") + receiver := seedVoiceOwner(t, database, "kate2") + outsider := seedVoiceOwner(t, database, "kate3") + + sendR := make(chan []byte, 16) + cR := ws.NewTestClientWithUser(hub, receiver, chanID, sendR) + hub.Register(cR) + + sendO := make(chan []byte, 16) + cO := ws.NewTestClientWithUser(hub, outsider, 999, sendO) // different channel + hub.Register(cO) + + sendS := make(chan []byte, 16) + cS := ws.NewTestClientWithUser(hub, sender, chanID, sendS) + hub.Register(cS) + time.Sleep(20 * time.Millisecond) + + hub.HandleMessageForTest(cS, voiceSignalMsg("voice_offer", chanID, "v=0...")) + time.Sleep(50 * time.Millisecond) + + // Receiver in same channel should get the signal. + receiverMsgs := drainChan(sendR) + found := false + for _, msg := range receiverMsgs { + if extractType(t, msg) == "voice_offer" { + found = true + break + } + } + if !found { + t.Error("receiver in channel did not receive voice_offer relay") + } + + // Outsider in different channel should NOT get it. + outsiderMsgs := drainChan(sendO) + for _, msg := range outsiderMsgs { + if extractType(t, msg) == "voice_offer" { + t.Error("outsider received voice_offer, should not have") + } + } + + // Sender should NOT receive their own signal. + senderMsgs := drainChan(sendS) + for _, msg := range senderMsgs { + if extractType(t, msg) == "voice_offer" { + t.Error("sender received their own voice_offer, should not have") + } + } +} + +func TestVoice_Signal_ICERelayed(t *testing.T) { + hub, database := newVoiceHub(t) + chanID := seedVoiceChan(t, database, "vc-ice") + + sender := seedVoiceOwner(t, database, "leo") + receiver := seedVoiceOwner(t, database, "leo2") + + sendR := make(chan []byte, 16) + cR := ws.NewTestClientWithUser(hub, receiver, chanID, sendR) + hub.Register(cR) + + sendS := make(chan []byte, 16) + cS := ws.NewTestClientWithUser(hub, sender, chanID, sendS) + hub.Register(cS) + time.Sleep(20 * time.Millisecond) + + hub.HandleMessageForTest(cS, voiceICEMsg(chanID, "candidate:...")) + time.Sleep(50 * time.Millisecond) + + receiverMsgs := drainChan(sendR) + found := false + for _, msg := range receiverMsgs { + if extractType(t, msg) == "voice_ice" { + found = true + break + } + } + if !found { + t.Error("receiver did not receive relayed voice_ice") + } +} + +func TestVoice_Signal_RateLimit_BlocksExcess(t *testing.T) { + hub, database := newVoiceHub(t) + chanID := seedVoiceChan(t, database, "vc-ratelimit") + + sender := seedVoiceOwner(t, database, "mia") + receiver := seedVoiceOwner(t, database, "mia2") + + sendR := make(chan []byte, 256) + cR := ws.NewTestClientWithUser(hub, receiver, chanID, sendR) + hub.Register(cR) + + sendS := make(chan []byte, 256) + cS := ws.NewTestClientWithUser(hub, sender, chanID, sendS) + hub.Register(cS) + time.Sleep(20 * time.Millisecond) + + // Send 30 signals rapidly — limit is 20/sec, so some should be dropped. + for i := 0; i < 30; i++ { + hub.HandleMessageForTest(cS, voiceSignalMsg("voice_offer", chanID, "v=0...")) + } + time.Sleep(50 * time.Millisecond) + + receivedCount := len(drainChan(sendR)) + if receivedCount >= 30 { + t.Errorf("received %d signals, expected fewer due to rate limit", receivedCount) + } + + // Sender should receive at least one RATE_LIMITED error. + senderMsgs := drainChan(sendS) + foundError := false + for _, msg := range senderMsgs { + if extractType(t, msg) == "error" { + foundError = true + break + } + } + if !foundError { + t.Error("expected RATE_LIMITED error to sender after exceeding signal rate limit") + } +} + +// ─── soundboard ─────────────────────────────────────────────────────────────── + +func TestVoice_Soundboard_BroadcastsToAll(t *testing.T) { + hub, database := newVoiceHub(t) + + user := seedVoiceOwner(t, database, "noah") + listener := seedVoiceOwner(t, database, "noah2") + + sendL := make(chan []byte, 16) + cL := ws.NewTestClientWithUser(hub, listener, 0, sendL) + hub.Register(cL) + + sendS := make(chan []byte, 16) + cS := ws.NewTestClientWithUser(hub, user, 0, sendS) + hub.Register(cS) + time.Sleep(20 * time.Millisecond) + + soundMsg, _ := json.Marshal(map[string]interface{}{ + "type": "soundboard_play", + "payload": map[string]interface{}{"sound_id": "abc-uuid-123"}, + }) + hub.HandleMessageForTest(cS, soundMsg) + time.Sleep(50 * time.Millisecond) + + listenerMsgs := drainChan(sendL) + found := false + for _, msg := range listenerMsgs { + if extractType(t, msg) == "soundboard_play" { + found = true + break + } + } + if !found { + t.Error("listener did not receive soundboard_play broadcast") + } +} + +func TestVoice_Soundboard_NoPermission_SendsError(t *testing.T) { + hub, _ := newVoiceHub(t) + + // Client with no user set → permission check fails. + send := make(chan []byte, 16) + c := ws.NewTestClient(hub, 8888, send) + hub.Register(c) + time.Sleep(20 * time.Millisecond) + + soundMsg, _ := json.Marshal(map[string]interface{}{ + "type": "soundboard_play", + "payload": map[string]interface{}{"sound_id": "abc"}, + }) + hub.HandleMessageForTest(c, soundMsg) + time.Sleep(30 * time.Millisecond) + + msgs := drainChan(send) + found := false + for _, m := range msgs { + if extractType(t, m) == "error" { + found = true + } + } + if !found { + t.Error("expected FORBIDDEN error for soundboard without permission") + } +} + +func TestVoice_Soundboard_RateLimit(t *testing.T) { + hub, database := newVoiceHub(t) + user := seedVoiceOwner(t, database, "olivia") + + send := make(chan []byte, 64) + c := ws.NewTestClientWithUser(hub, user, 0, send) + hub.Register(c) + time.Sleep(20 * time.Millisecond) + + soundMsg, _ := json.Marshal(map[string]interface{}{ + "type": "soundboard_play", + "payload": map[string]interface{}{"sound_id": "x"}, + }) + + // Send 5 soundboard plays rapidly — limit is 1 per 3 sec. + for i := 0; i < 5; i++ { + hub.HandleMessageForTest(c, soundMsg) + } + time.Sleep(50 * time.Millisecond) + + msgs := drainChan(send) + errCount := 0 + for _, m := range msgs { + if extractType(t, m) == "error" { + errCount++ + } + } + if errCount == 0 { + t.Error("expected rate limit errors for rapid soundboard plays") + } +} + +// ─── handleMessage dispatch ─────────────────────────────────────────────────── + +func TestVoice_HandleMessage_VoiceAnswer_Relayed(t *testing.T) { + hub, database := newVoiceHub(t) + chanID := seedVoiceChan(t, database, "vc-answer") + + sender := seedVoiceOwner(t, database, "pedro") + receiver := seedVoiceOwner(t, database, "pedro2") + + sendR := make(chan []byte, 16) + cR := ws.NewTestClientWithUser(hub, receiver, chanID, sendR) + hub.Register(cR) + + sendS := make(chan []byte, 16) + cS := ws.NewTestClientWithUser(hub, sender, chanID, sendS) + hub.Register(cS) + time.Sleep(20 * time.Millisecond) + + hub.HandleMessageForTest(cS, voiceSignalMsg("voice_answer", chanID, "v=0 answer...")) + time.Sleep(50 * time.Millisecond) + + receiverMsgs := drainChan(sendR) + found := false + for _, msg := range receiverMsgs { + if extractType(t, msg) == "voice_answer" { + found = true + break + } + } + if !found { + t.Error("receiver did not receive relayed voice_answer") + } +}