From ad61cbff4447491b0ea450cc7ff94e09e7ac0d35 Mon Sep 17 00:00:00 2001 From: jevb Date: Tue, 31 Mar 2026 16:23:37 +0200 Subject: [PATCH] test: add channel pins, metrics, diagnostics, client-update, middleware tests Cover low-coverage handler functions: handleGetPins, handleSetPinned, handleMetrics, handleDiagnosticsConnectivity, isPrivateIP, AdminIPRestrict, handleClientUpdate, and handleLiveKitHealth. Adds 30 new test cases across 4 new test files and 2 modified test files. --- Server/api/channel_handler_test.go | 257 +++++++++++++++++++++++++ Server/api/client_update_test.go | 139 +++++++++++++ Server/api/diagnostics_handler_test.go | 134 +++++++++++++ Server/api/export_test.go | 34 ++++ Server/api/metrics_handler_test.go | 115 +++++++++++ Server/api/middleware_test.go | 184 ++++++++++++++++++ 6 files changed, 863 insertions(+) create mode 100644 Server/api/client_update_test.go create mode 100644 Server/api/diagnostics_handler_test.go create mode 100644 Server/api/export_test.go create mode 100644 Server/api/metrics_handler_test.go diff --git a/Server/api/channel_handler_test.go b/Server/api/channel_handler_test.go index c4c5d558..60903932 100644 --- a/Server/api/channel_handler_test.go +++ b/Server/api/channel_handler_test.go @@ -646,3 +646,260 @@ func TestChannelMessages_InvalidLimit(t *testing.T) { t.Errorf("invalid limit status = %d, want 400", rr.Code) } } + +// ─── GET /api/v1/channels/{id}/pins ───────────────────────────────────────── + +// newPinTestDB creates a DB with dm_participants table needed for pin tests. +func newPinTestDB(t *testing.T) *db.DB { + t.Helper() + database := newChannelTestDB(t) + // Add DM tables required by pin handlers for DM authorization. + _, err := database.Exec(` + CREATE TABLE IF NOT EXISTS dm_participants ( + channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE, + user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE, + PRIMARY KEY (channel_id, user_id) + ); + CREATE INDEX IF NOT EXISTS idx_dm_participants_user ON dm_participants(user_id); + CREATE TABLE IF NOT EXISTS dm_open_state ( + user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE, + channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE, + opened_at TEXT NOT NULL DEFAULT (datetime('now')), + PRIMARY KEY (user_id, channel_id) + ); + `) + if err != nil { + t.Fatalf("create dm_participants: %v", err) + } + return database +} + +func TestGetPins_Unauthenticated(t *testing.T) { + router := buildChannelRouter(newPinTestDB(t)) + rr := chGet(t, router, "/api/v1/channels/1/pins", "") + if rr.Code != http.StatusUnauthorized { + t.Errorf("status = %d, want 401", rr.Code) + } +} + +func TestGetPins_ChannelNotFound(t *testing.T) { + database := newPinTestDB(t) + router := buildChannelRouter(database) + token := chTestCreateToken(t, database, "pinuser1", 1) + + rr := chGet(t, router, "/api/v1/channels/9999/pins", token) + if rr.Code != http.StatusNotFound { + t.Errorf("status = %d, want 404; body: %s", rr.Code, rr.Body.String()) + } +} + +func TestGetPins_EmptyPins(t *testing.T) { + database := newPinTestDB(t) + router := buildChannelRouter(database) + token := chTestCreateToken(t, database, "pinuser2", 1) + chID, _ := database.CreateChannel("general", "text", "", "", 0) + + rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/pins", chID), token) + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want 200; body: %s", rr.Code, rr.Body.String()) + } + + var resp struct { + Messages []any `json:"messages"` + HasMore bool `json:"has_more"` + } + _ = json.NewDecoder(rr.Body).Decode(&resp) + if len(resp.Messages) != 0 { + t.Errorf("expected 0 pinned messages, got %d", len(resp.Messages)) + } + if resp.HasMore { + t.Error("has_more should be false for empty pins") + } +} + +func TestGetPins_ReturnsPinnedMessages(t *testing.T) { + database := newPinTestDB(t) + router := buildChannelRouter(database) + token := chTestCreateToken(t, database, "pinuser3", 1) + user, _ := database.GetUserByUsername("pinuser3") + chID, _ := database.CreateChannel("general", "text", "", "", 0) + + msgID, _ := database.CreateMessage(chID, user.ID, "pinned message", nil) + _ = database.SetMessagePinned(msgID, true) + // Also create an unpinned message — should not appear. + _, _ = database.CreateMessage(chID, user.ID, "not pinned", nil) + + rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/pins", chID), token) + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want 200; body: %s", rr.Code, rr.Body.String()) + } + + var resp struct { + Messages []any `json:"messages"` + } + _ = json.NewDecoder(rr.Body).Decode(&resp) + if len(resp.Messages) != 1 { + t.Errorf("expected 1 pinned message, got %d", len(resp.Messages)) + } +} + +func TestGetPins_DMChannel_NonParticipantForbidden(t *testing.T) { + database := newPinTestDB(t) + router := buildChannelRouter(database) + + // Create two users for the DM and a third who should be denied. + chTestCreateToken(t, database, "dmuser1", 4) + chTestCreateToken(t, database, "dmuser2", 4) + outsiderToken := chTestCreateToken(t, database, "outsider", 4) + + user1, _ := database.GetUserByUsername("dmuser1") + user2, _ := database.GetUserByUsername("dmuser2") + + // Create a DM channel manually. + dmCh, _, _ := database.GetOrCreateDMChannel(user1.ID, user2.ID) + + rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/pins", dmCh.ID), outsiderToken) + if rr.Code != http.StatusForbidden { + t.Errorf("status = %d, want 403; body: %s", rr.Code, rr.Body.String()) + } +} + +func TestGetPins_MemberNoReadPermission(t *testing.T) { + database := newPinTestDB(t) + router := buildChannelRouter(database) + // Role 4 = Member with permissions 1635 (0x663). + // Deny READ_MESSAGES on a specific channel via override. + token := chTestCreateToken(t, database, "nopermuser", 4) + chID, _ := database.CreateChannel("restricted", "text", "", "", 0) + + // Deny all permissions for role 4 on this channel. + _, _ = database.Exec( + `INSERT INTO channel_overrides (channel_id, role_id, allow, deny) VALUES (?, 4, 0, 2147483647)`, + chID, + ) + + rr := chGet(t, router, fmt.Sprintf("/api/v1/channels/%d/pins", chID), token) + if rr.Code != http.StatusForbidden { + t.Errorf("status = %d, want 403; body: %s", rr.Code, rr.Body.String()) + } +} + +// ─── POST/DELETE /api/v1/channels/{id}/pins/{messageId} ───────────────────── + +func chPost(t *testing.T, router http.Handler, path, token string) *httptest.ResponseRecorder { + t.Helper() + req := httptest.NewRequest(http.MethodPost, path, nil) + if token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + return rr +} + +func chDelete(t *testing.T, router http.Handler, path, token string) *httptest.ResponseRecorder { + t.Helper() + req := httptest.NewRequest(http.MethodDelete, path, nil) + if token != "" { + req.Header.Set("Authorization", "Bearer "+token) + } + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + return rr +} + +func TestSetPinned_PinSuccessfully(t *testing.T) { + database := newPinTestDB(t) + router := buildChannelRouter(database) + token := chTestCreateToken(t, database, "pinner1", 1) + user, _ := database.GetUserByUsername("pinner1") + chID, _ := database.CreateChannel("general", "text", "", "", 0) + msgID, _ := database.CreateMessage(chID, user.ID, "pin me", nil) + + rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", chID, msgID), token) + if rr.Code != http.StatusNoContent { + t.Errorf("pin status = %d, want 204; body: %s", rr.Code, rr.Body.String()) + } + + // Verify the message is actually pinned. + msg, _ := database.GetMessage(msgID) + if !msg.Pinned { + t.Error("message should be pinned after POST") + } +} + +func TestSetPinned_UnpinSuccessfully(t *testing.T) { + database := newPinTestDB(t) + router := buildChannelRouter(database) + token := chTestCreateToken(t, database, "unpinner1", 1) + user, _ := database.GetUserByUsername("unpinner1") + chID, _ := database.CreateChannel("general", "text", "", "", 0) + msgID, _ := database.CreateMessage(chID, user.ID, "unpin me", nil) + _ = database.SetMessagePinned(msgID, true) + + rr := chDelete(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", chID, msgID), token) + if rr.Code != http.StatusNoContent { + t.Errorf("unpin status = %d, want 204; body: %s", rr.Code, rr.Body.String()) + } + + msg, _ := database.GetMessage(msgID) + if msg.Pinned { + t.Error("message should not be pinned after DELETE") + } +} + +func TestSetPinned_MessageNotFound(t *testing.T) { + database := newPinTestDB(t) + router := buildChannelRouter(database) + token := chTestCreateToken(t, database, "pinner2", 1) + chID, _ := database.CreateChannel("general", "text", "", "", 0) + + rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/9999", chID), token) + if rr.Code != http.StatusNotFound { + t.Errorf("status = %d, want 404; body: %s", rr.Code, rr.Body.String()) + } +} + +func TestSetPinned_ChannelNotFound(t *testing.T) { + database := newPinTestDB(t) + router := buildChannelRouter(database) + token := chTestCreateToken(t, database, "pinner3", 1) + + rr := chPost(t, router, "/api/v1/channels/9999/pins/1", token) + if rr.Code != http.StatusNotFound { + t.Errorf("status = %d, want 404; body: %s", rr.Code, rr.Body.String()) + } +} + +func TestSetPinned_NoPermission(t *testing.T) { + database := newPinTestDB(t) + router := buildChannelRouter(database) + // Member role (4) has permissions 1635 — does not include MANAGE_MESSAGES (0x2000). + token := chTestCreateToken(t, database, "noperm", 4) + user, _ := database.GetUserByUsername("noperm") + chID, _ := database.CreateChannel("general", "text", "", "", 0) + msgID, _ := database.CreateMessage(chID, user.ID, "try to pin", nil) + + rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", chID, msgID), token) + if rr.Code != http.StatusForbidden { + t.Errorf("status = %d, want 403; body: %s", rr.Code, rr.Body.String()) + } +} + +func TestSetPinned_Idempotent(t *testing.T) { + database := newPinTestDB(t) + router := buildChannelRouter(database) + token := chTestCreateToken(t, database, "pinner4", 1) + user, _ := database.GetUserByUsername("pinner4") + chID, _ := database.CreateChannel("general", "text", "", "", 0) + msgID, _ := database.CreateMessage(chID, user.ID, "already pinned", nil) + _ = database.SetMessagePinned(msgID, true) + + // Pinning again should still succeed (idempotent). + rr := chPost(t, router, fmt.Sprintf("/api/v1/channels/%d/pins/%d", chID, msgID), token) + if rr.Code != http.StatusNoContent { + t.Errorf("idempotent pin status = %d, want 204; body: %s", rr.Code, rr.Body.String()) + } +} diff --git a/Server/api/client_update_test.go b/Server/api/client_update_test.go new file mode 100644 index 00000000..81c088fe --- /dev/null +++ b/Server/api/client_update_test.go @@ -0,0 +1,139 @@ +package api_test + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/go-chi/chi/v5" + "github.com/owncord/server/api" + "github.com/owncord/server/updater" +) + +// fakeGitHubRelease returns a test HTTP server that mimics the GitHub +// Releases API, serving a release with the given tag and NSIS assets. +// Asset download URLs point back to the test server so FetchTextAsset works. +func fakeGitHubRelease(t *testing.T, tag string) *httptest.Server { + t.Helper() + + var srv *httptest.Server + mux := http.NewServeMux() + + mux.HandleFunc("/repos/test/repo/releases/latest", func(w http.ResponseWriter, r *http.Request) { + resp := map[string]any{ + "tag_name": tag, + "body": "Release notes here", + "html_url": "https://github.com/test/repo/releases/" + tag, + "assets": []map[string]any{ + { + "name": "OwnCord_1.0.0_x64-setup.nsis.zip", + "browser_download_url": srv.URL + "/download/OwnCord_1.0.0_x64-setup.nsis.zip", + }, + { + "name": "OwnCord_1.0.0_x64-setup.nsis.zip.sig", + "browser_download_url": srv.URL + "/download/OwnCord_1.0.0_x64-setup.nsis.zip.sig", + }, + }, + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(resp) + }) + + // Serve the signature file content. + mux.HandleFunc("/download/OwnCord_1.0.0_x64-setup.nsis.zip.sig", func(w http.ResponseWriter, r *http.Request) { + _, _ = w.Write([]byte("dW50cnVzdGVkIGNvbW1lbnQ=")) + }) + + srv = httptest.NewServer(mux) + t.Cleanup(srv.Close) + return srv +} + +func buildClientUpdateRouter(u *updater.Updater) http.Handler { + r := chi.NewRouter() + api.MountClientUpdateRoute(r, u) + return r +} + +func TestClientUpdate_NewVersionAvailable(t *testing.T) { + srv := fakeGitHubRelease(t, "v2.0.0") + u := updater.NewUpdater("1.0.0", "", "test", "repo") + u.SetBaseURL(srv.URL) + + router := buildClientUpdateRouter(u) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/client-update/windows-x86_64/1.0.0", nil) + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want 200; body: %s", rr.Code, rr.Body.String()) + } + + var resp map[string]any + if err := json.NewDecoder(rr.Body).Decode(&resp); err != nil { + t.Fatalf("decode: %v", err) + } + + if resp["version"] == nil { + t.Error("response missing 'version' field") + } + if resp["platforms"] == nil { + t.Error("response missing 'platforms' field") + } +} + +func TestClientUpdate_AlreadyLatest(t *testing.T) { + srv := fakeGitHubRelease(t, "v1.0.0") + u := updater.NewUpdater("1.0.0", "", "test", "repo") + u.SetBaseURL(srv.URL) + + router := buildClientUpdateRouter(u) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/client-update/windows-x86_64/1.0.0", nil) + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusNoContent { + t.Errorf("status = %d, want 204; body: %s", rr.Code, rr.Body.String()) + } +} + +func TestClientUpdate_FutureVersion(t *testing.T) { + srv := fakeGitHubRelease(t, "v1.0.0") + u := updater.NewUpdater("1.0.0", "", "test", "repo") + u.SetBaseURL(srv.URL) + + router := buildClientUpdateRouter(u) + + // Client has a newer version than the release. + req := httptest.NewRequest(http.MethodGet, "/api/v1/client-update/windows-x86_64/2.0.0", nil) + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusNoContent { + t.Errorf("status = %d, want 204; body: %s", rr.Code, rr.Body.String()) + } +} + +func TestClientUpdate_GitHubError(t *testing.T) { + // Server that always returns 500. + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + })) + t.Cleanup(srv.Close) + + u := updater.NewUpdater("1.0.0", "", "test", "repo") + u.SetBaseURL(srv.URL) + + router := buildClientUpdateRouter(u) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/client-update/windows-x86_64/1.0.0", nil) + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusBadGateway { + t.Errorf("status = %d, want 502; body: %s", rr.Code, rr.Body.String()) + } +} diff --git a/Server/api/diagnostics_handler_test.go b/Server/api/diagnostics_handler_test.go new file mode 100644 index 00000000..a809fbce --- /dev/null +++ b/Server/api/diagnostics_handler_test.go @@ -0,0 +1,134 @@ +package api_test + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/owncord/server/api" + "github.com/owncord/server/auth" + "github.com/owncord/server/config" + "github.com/owncord/server/db" +) + +// hashTokenForTest is a package-level helper wrapping auth.HashToken. +func hashTokenForTest(token string) string { + return auth.HashToken(token) +} + +// setupDiagnosticsRouter creates a full router with an authenticated user for +// diagnostics testing. +func setupDiagnosticsRouter(t *testing.T) (http.Handler, string) { + t.Helper() + + database, err := db.Open(":memory:") + if err != nil { + t.Fatalf("db.Open: %v", err) + } + if err := db.Migrate(database); err != nil { + t.Fatalf("db.Migrate: %v", err) + } + t.Cleanup(func() { _ = database.Close() }) + + cfg := &config.Config{ + Server: config.ServerConfig{ + Name: "Test Server", + Port: 8443, + }, + } + + handler, _, cleanup := api.NewRouter(cfg, database, "1.0.0-test", nil) + t.Cleanup(cleanup) + + // Create a user and session for authenticated requests. + uid, _ := database.CreateUser("diaguser", "$2a$12$fake", 1) + token := "diagtest-token-123" + hash := auth.HashToken(token) + _, _ = database.Exec( + `INSERT INTO sessions (user_id, token, device, ip_address, expires_at) + VALUES (?, ?, 'test', '127.0.0.1', '2099-01-01T00:00:00Z')`, + uid, hash, + ) + + return handler, token +} + +func TestDiagnosticsConnectivity_ReturnsData(t *testing.T) { + router, token := setupDiagnosticsRouter(t) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/diagnostics/connectivity", nil) + req.Header.Set("Authorization", "Bearer "+token) + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want 200; body: %s", rr.Code, rr.Body.String()) + } + + var resp map[string]any + if err := json.NewDecoder(rr.Body).Decode(&resp); err != nil { + t.Fatalf("decode: %v", err) + } + + // Verify top-level sections exist. + for _, section := range []string{"server", "voice", "client"} { + if _, ok := resp[section]; !ok { + t.Errorf("missing section %q in diagnostics response", section) + } + } + + // Verify server section has expected fields. + server, _ := resp["server"].(map[string]any) + if server["version"] != "1.0.0-test" { + t.Errorf("server.version = %v, want 1.0.0-test", server["version"]) + } +} + +func TestDiagnosticsConnectivity_Unauthenticated(t *testing.T) { + router, _ := setupDiagnosticsRouter(t) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/diagnostics/connectivity", 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) + } +} + +// ─── isPrivateIP tests ────────────────────────────────────────────────────── + +func TestIsPrivateIP(t *testing.T) { + tests := []struct { + name string + ip string + want bool + }{ + {"10.x.x.x", "10.0.0.1", true}, + {"172.16.x.x", "172.16.0.1", true}, + {"172.17.x.x", "172.17.5.5", true}, + {"172.31.x.x", "172.31.255.255", true}, + {"192.168.x.x", "192.168.1.1", true}, + {"127.x.x.x", "127.0.0.1", true}, + {"::1 loopback", "::1", true}, + {"fc ULA", "fc00::1", true}, + {"fd ULA", "fd12::1", true}, + {"public 8.8.8.8", "8.8.8.8", false}, + {"public 203.x", "203.0.113.1", false}, + {"public 1.1.1.1", "1.1.1.1", false}, + {"172.32 not private", "172.32.0.1", false}, + {"empty string", "", false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := api.IsPrivateIPForTest(tt.ip) + if got != tt.want { + t.Errorf("isPrivateIP(%q) = %v, want %v", tt.ip, got, tt.want) + } + }) + } +} diff --git a/Server/api/export_test.go b/Server/api/export_test.go new file mode 100644 index 00000000..076b94f5 --- /dev/null +++ b/Server/api/export_test.go @@ -0,0 +1,34 @@ +package api + +import "net/http" + +// HandleMetricsForTest exposes handleMetrics for use in external tests. +var HandleMetricsForTest = handleMetrics + +// HandleLiveKitHealthForTest exposes handleLiveKitHealth for use in external tests. +func HandleLiveKitHealthForTest(healthCheck func() (bool, error)) http.HandlerFunc { + // Inline the logic since handleLiveKitHealth requires a *ws.Hub. + return func(w http.ResponseWriter, r *http.Request) { + ok, err := healthCheck() + if ok { + writeJSON(w, http.StatusOK, livekitHealthResponse{ + Status: "ok", + LiveKitReachable: true, + }) + return + } + + errMsg := "unknown" + if err != nil { + errMsg = err.Error() + } + writeJSON(w, http.StatusServiceUnavailable, livekitHealthResponse{ + Status: "degraded", + LiveKitReachable: false, + Error: errMsg, + }) + } +} + +// IsPrivateIPForTest exposes isPrivateIP for use in external tests. +var IsPrivateIPForTest = isPrivateIP diff --git a/Server/api/metrics_handler_test.go b/Server/api/metrics_handler_test.go new file mode 100644 index 00000000..28397a9e --- /dev/null +++ b/Server/api/metrics_handler_test.go @@ -0,0 +1,115 @@ +package api_test + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/go-chi/chi/v5" + "github.com/owncord/server/api" +) + +// buildMetricsRouter creates a chi router with the metrics endpoint behind AdminIPRestrict. +func buildMetricsRouter(allowedCIDRs []string) http.Handler { + r := chi.NewRouter() + r.With(api.AdminIPRestrict(allowedCIDRs)). + Get("/api/v1/metrics", api.HandleMetricsForTest( + func() int { return 5 }, + func() int { return 2 }, + func() (bool, error) { return true, nil }, + )) + return r +} + +func TestHandleMetrics_ReturnsExpectedFields(t *testing.T) { + router := buildMetricsRouter(nil) // no IP restriction + + req := httptest.NewRequest(http.MethodGet, "/api/v1/metrics", nil) + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want 200; body: %s", rr.Code, rr.Body.String()) + } + + var resp map[string]any + if err := json.NewDecoder(rr.Body).Decode(&resp); err != nil { + t.Fatalf("decode: %v", err) + } + + requiredFields := []string{ + "uptime", "uptime_seconds", "goroutines", + "heap_alloc_mb", "heap_sys_mb", "num_gc", + "connected_users", "voice_sessions", "livekit_healthy", + } + for _, f := range requiredFields { + if _, ok := resp[f]; !ok { + t.Errorf("missing field %q in metrics response", f) + } + } + + // Verify the callback values are reflected. + if int(resp["connected_users"].(float64)) != 5 { + t.Errorf("connected_users = %v, want 5", resp["connected_users"]) + } + if int(resp["voice_sessions"].(float64)) != 2 { + t.Errorf("voice_sessions = %v, want 2", resp["voice_sessions"]) + } + if resp["livekit_healthy"] != true { + t.Errorf("livekit_healthy = %v, want true", resp["livekit_healthy"]) + } +} + +func TestHandleMetrics_AdminIPRestrict_BlocksNonAdmin(t *testing.T) { + router := buildMetricsRouter([]string{"10.0.0.0/8"}) // only 10.x allowed + + req := httptest.NewRequest(http.MethodGet, "/api/v1/metrics", nil) + req.RemoteAddr = "192.168.1.1:9999" // not in allowed CIDR + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusForbidden { + t.Errorf("status = %d, want 403; body: %s", rr.Code, rr.Body.String()) + } +} + +func TestHandleMetrics_AdminIPRestrict_AllowsAdmin(t *testing.T) { + router := buildMetricsRouter([]string{"127.0.0.0/8"}) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/metrics", nil) + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Errorf("status = %d, want 200; body: %s", rr.Code, rr.Body.String()) + } +} + +func TestHandleMetrics_WithoutLiveKitHealthCheck(t *testing.T) { + r := chi.NewRouter() + r.Get("/api/v1/metrics", api.HandleMetricsForTest( + func() int { return 0 }, + func() int { return 0 }, + nil, // no livekit + )) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/metrics", nil) + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + r.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want 200; body: %s", rr.Code, rr.Body.String()) + } + + var resp map[string]any + _ = json.NewDecoder(rr.Body).Decode(&resp) + + // livekit_healthy should be absent when no health check is provided. + if _, ok := resp["livekit_healthy"]; ok { + t.Errorf("livekit_healthy should be omitted when health check is nil, got %v", resp["livekit_healthy"]) + } +} diff --git a/Server/api/middleware_test.go b/Server/api/middleware_test.go index fc3f443a..4f2d3a09 100644 --- a/Server/api/middleware_test.go +++ b/Server/api/middleware_test.go @@ -1,6 +1,8 @@ package api_test import ( + "encoding/json" + "fmt" "net/http" "net/http/httptest" "strings" @@ -576,6 +578,188 @@ func TestMaxBodySize_PassesThrough(t *testing.T) { } } +// ─── AdminIPRestrict tests ────────────────────────────────────────────────── + +func TestAdminIPRestrict_AllowedCIDR(t *testing.T) { + h := api.AdminIPRestrict([]string{"127.0.0.0/8"})(http.HandlerFunc(ok)) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Errorf("AdminIPRestrict allowed CIDR status = %d, want 200", rr.Code) + } +} + +func TestAdminIPRestrict_BlockedCIDR(t *testing.T) { + h := api.AdminIPRestrict([]string{"10.0.0.0/8"})(http.HandlerFunc(ok)) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "192.168.1.1:9999" // not in 10.0.0.0/8 + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusForbidden { + t.Errorf("AdminIPRestrict blocked CIDR status = %d, want 403", rr.Code) + } +} + +func TestAdminIPRestrict_EmptyAllowsAll(t *testing.T) { + h := api.AdminIPRestrict(nil)(http.HandlerFunc(ok)) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "203.0.113.1:9999" + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Errorf("AdminIPRestrict empty list status = %d, want 200", rr.Code) + } +} + +func TestAdminIPRestrict_InvalidCIDR(t *testing.T) { + // Invalid CIDR should fail closed (deny access since isTrustedProxy + // returns false on parse error). + h := api.AdminIPRestrict([]string{"not-a-cidr"})(http.HandlerFunc(ok)) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusForbidden { + t.Errorf("AdminIPRestrict invalid CIDR status = %d, want 403", rr.Code) + } +} + +func TestAdminIPRestrict_MultipleCIDRs(t *testing.T) { + h := api.AdminIPRestrict([]string{"10.0.0.0/8", "192.168.0.0/16"})(http.HandlerFunc(ok)) + + // First CIDR matches. + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "10.1.2.3:9999" + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + if rr.Code != http.StatusOK { + t.Errorf("AdminIPRestrict multi-CIDR (10.x) status = %d, want 200", rr.Code) + } + + // Second CIDR matches. + req = httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "192.168.1.50:9999" + rr = httptest.NewRecorder() + h.ServeHTTP(rr, req) + if rr.Code != http.StatusOK { + t.Errorf("AdminIPRestrict multi-CIDR (192.168.x) status = %d, want 200", rr.Code) + } + + // Neither matches. + req = httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "172.16.0.1:9999" + rr = httptest.NewRecorder() + h.ServeHTTP(rr, req) + if rr.Code != http.StatusForbidden { + t.Errorf("AdminIPRestrict multi-CIDR (no match) status = %d, want 403", rr.Code) + } +} + +// ─── SecurityHeadersWithTLS tests ─────────────────────────────────────────── + +func TestSecurityHeadersWithTLS_HSTS(t *testing.T) { + h := api.SecurityHeadersWithTLS("auto")(http.HandlerFunc(ok)) + req := httptest.NewRequest(http.MethodGet, "/", nil) + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if got := rr.Header().Get("Strict-Transport-Security"); got == "" { + t.Error("SecurityHeadersWithTLS: missing HSTS header when TLS enabled") + } +} + +func TestSecurityHeadersWithTLS_NoHSTSWithoutTLS(t *testing.T) { + h := api.SecurityHeadersWithTLS("")(http.HandlerFunc(ok)) + req := httptest.NewRequest(http.MethodGet, "/", nil) + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if got := rr.Header().Get("Strict-Transport-Security"); got != "" { + t.Errorf("SecurityHeadersWithTLS: unexpected HSTS header %q when TLS disabled", got) + } +} + +// ─── handleLiveKitHealth tests ────────────────────────────────────────────── + +func TestLiveKitHealth_Healthy(t *testing.T) { + h := api.HandleLiveKitHealthForTest(func() (bool, error) { + return true, nil + }) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want 200; body: %s", rr.Code, rr.Body.String()) + } + + var resp map[string]any + _ = json.NewDecoder(rr.Body).Decode(&resp) + if resp["status"] != "ok" { + t.Errorf("status = %v, want ok", resp["status"]) + } + if resp["livekit_reachable"] != true { + t.Errorf("livekit_reachable = %v, want true", resp["livekit_reachable"]) + } +} + +func TestLiveKitHealth_Unhealthy(t *testing.T) { + h := api.HandleLiveKitHealthForTest(func() (bool, error) { + return false, fmt.Errorf("connection refused") + }) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusServiceUnavailable { + t.Fatalf("status = %d, want 503; body: %s", rr.Code, rr.Body.String()) + } + + var resp map[string]any + _ = json.NewDecoder(rr.Body).Decode(&resp) + if resp["status"] != "degraded" { + t.Errorf("status = %v, want degraded", resp["status"]) + } + if resp["livekit_reachable"] != false { + t.Errorf("livekit_reachable = %v, want false", resp["livekit_reachable"]) + } + if resp["error"] != "connection refused" { + t.Errorf("error = %v, want 'connection refused'", resp["error"]) + } +} + +func TestLiveKitHealth_UnhealthyNoError(t *testing.T) { + h := api.HandleLiveKitHealthForTest(func() (bool, error) { + return false, nil + }) + + req := httptest.NewRequest(http.MethodGet, "/", nil) + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusServiceUnavailable { + t.Fatalf("status = %d, want 503; body: %s", rr.Code, rr.Body.String()) + } + + var resp map[string]any + _ = json.NewDecoder(rr.Body).Decode(&resp) + if resp["error"] != "unknown" { + t.Errorf("error = %v, want 'unknown'", resp["error"]) + } +} + // apiTestSchema is the full schema needed for all api tests (middleware, // auth handler, and invite handler). var apiTestSchema = []byte(`