From b7dd6eabe95378db3558c781c9ee508427efa438 Mon Sep 17 00:00:00 2001 From: jevb Date: Sat, 14 Mar 2026 20:52:11 +0100 Subject: [PATCH] feat: implement Phase 2 auth & security with TDD MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - auth/session: 256-bit crypto-random tokens, SHA-256 hashing for storage - auth/password: bcrypt cost 12, strength validation (8-72 chars) - auth/ratelimit: sliding-window RateLimiter with lockout, thread-safe - db/models: User, Session, Invite, Role types - db/auth_queries: full user/session/invite CRUD with in-memory test coverage - api/middleware: AuthMiddleware (Bearer token), RequirePermission (bitfield), RateLimitMiddleware (X-Real-IP, Retry-After header) - api/auth_handler: POST register/login, POST logout, GET me - Generic errors — username existence never revealed - Rate limits: 3/min register, 5/min login, lockout after 10 failures - api/invite_handler: create/list/revoke behind MANAGE_INVITES permission - bluemonday sanitization on all user-supplied string fields Test coverage: auth 90.9%, db 84.4%, api 80.9% --- Server/api/auth_handler.go | 311 +++++++++++++++++ Server/api/auth_handler_test.go | 456 +++++++++++++++++++++++++ Server/api/invite_handler.go | 160 +++++++++ Server/api/invite_handler_test.go | 294 ++++++++++++++++ Server/api/middleware.go | 194 +++++++++++ Server/api/middleware_test.go | 361 ++++++++++++++++++++ Server/api/router.go | 10 + Server/auth/password.go | 50 +++ Server/auth/password_test.go | 106 ++++++ Server/auth/ratelimit.go | 104 ++++++ Server/auth/ratelimit_test.go | 126 +++++++ Server/auth/session.go | 24 ++ Server/auth/session_test.go | 77 +++++ Server/db/auth_queries.go | 285 ++++++++++++++++ Server/db/auth_queries_test.go | 468 ++++++++++++++++++++++++++ Server/db/invite_queries.go | 30 ++ Server/db/models.go | 56 +++ Server/db/role_invite_queries_test.go | 156 +++++++++ Server/db/role_queries.go | 49 +++ Server/go.mod | 4 + Server/go.sum | 8 + 21 files changed, 3329 insertions(+) create mode 100644 Server/api/auth_handler.go create mode 100644 Server/api/auth_handler_test.go create mode 100644 Server/api/invite_handler.go create mode 100644 Server/api/invite_handler_test.go create mode 100644 Server/api/middleware.go create mode 100644 Server/api/middleware_test.go create mode 100644 Server/auth/password.go create mode 100644 Server/auth/password_test.go create mode 100644 Server/auth/ratelimit.go create mode 100644 Server/auth/ratelimit_test.go create mode 100644 Server/auth/session.go create mode 100644 Server/auth/session_test.go create mode 100644 Server/db/auth_queries.go create mode 100644 Server/db/auth_queries_test.go create mode 100644 Server/db/invite_queries.go create mode 100644 Server/db/models.go create mode 100644 Server/db/role_invite_queries_test.go create mode 100644 Server/db/role_queries.go diff --git a/Server/api/auth_handler.go b/Server/api/auth_handler.go new file mode 100644 index 00000000..e04cf2f8 --- /dev/null +++ b/Server/api/auth_handler.go @@ -0,0 +1,311 @@ +package api + +import ( + "encoding/json" + "log/slog" + "net/http" + "strings" + "time" + + "github.com/go-chi/chi/v5" + "github.com/microcosm-cc/bluemonday" + "github.com/owncord/server/auth" + "github.com/owncord/server/db" +) + +// sanitizer strips all HTML from user-supplied strings before storage. +var sanitizer = bluemonday.StrictPolicy() + +// genericAuthError is returned for all login/register failures to avoid +// revealing whether a username exists. +var genericAuthError = errorResponse{ + Error: "INVALID_CREDENTIALS", + Message: "invalid invite or credentials", +} + +// registerRequest is the JSON body for POST /api/v1/auth/register. +type registerRequest struct { + Username string `json:"username"` + Password string `json:"password"` + InviteCode string `json:"invite_code"` +} + +// loginRequest is the JSON body for POST /api/v1/auth/login. +type loginRequest struct { + Username string `json:"username"` + Password string `json:"password"` +} + +// userResponse is the user shape included in auth responses. +type userResponse struct { + ID int64 `json:"id"` + Username string `json:"username"` + Avatar string `json:"avatar,omitempty"` + Status string `json:"status"` + RoleID int64 `json:"role_id"` + CreatedAt string `json:"created_at"` +} + +// authSuccessResponse is returned on successful login/register. +type authSuccessResponse struct { + Token string `json:"token"` + User userResponse `json:"user"` +} + +// MountAuthRoutes registers all auth endpoints on the given router. +// Rate limiters are applied per-endpoint as specified. +func MountAuthRoutes(r chi.Router, database *db.DB, limiter *auth.RateLimiter) { + registerLimiter := limiter + loginLimiter := limiter + + r.Route("/api/v1/auth", func(r chi.Router) { + r.With(RateLimitMiddleware(registerLimiter, 3, time.Minute)). + Post("/register", handleRegister(database)) + + r.With(RateLimitMiddleware(loginLimiter, 5, time.Minute)). + Post("/login", handleLogin(database, limiter)) + + r.With(AuthMiddleware(database)). + Post("/logout", handleLogout(database)) + + r.With(AuthMiddleware(database)). + Get("/me", handleMe()) + }) +} + +// handleRegister processes POST /api/v1/auth/register. +func handleRegister(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + var req registerRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeJSON(w, http.StatusBadRequest, errorResponse{ + Error: "INVALID_INPUT", + Message: "malformed request body", + }) + return + } + + req.Username = strings.TrimSpace(sanitizer.Sanitize(req.Username)) + req.InviteCode = strings.TrimSpace(req.InviteCode) + + if req.Username == "" || req.Password == "" || req.InviteCode == "" { + writeJSON(w, http.StatusBadRequest, errorResponse{ + Error: "INVALID_INPUT", + Message: "username, password, and invite_code are required", + }) + return + } + + // Validate password strength before anything else. + if err := auth.ValidatePasswordStrength(req.Password); err != nil { + writeJSON(w, http.StatusBadRequest, errorResponse{ + Error: "INVALID_INPUT", + Message: err.Error(), + }) + return + } + + // Validate invite. + inv, err := database.GetInvite(req.InviteCode) + if err != nil || inv == nil || inv.Revoked { + writeJSON(w, http.StatusBadRequest, genericAuthError) + return + } + if err := database.UseInvite(req.InviteCode); err != nil { + writeJSON(w, http.StatusBadRequest, genericAuthError) + return + } + + // Hash password. + hash, err := auth.HashPassword(req.Password) + if err != nil { + writeJSON(w, http.StatusInternalServerError, errorResponse{ + Error: "SERVER_ERROR", + Message: "failed to process registration", + }) + return + } + + // Create user with default Member role (4). + uid, err := database.CreateUser(req.Username, hash, 4) + if err != nil { + writeJSON(w, http.StatusBadRequest, errorResponse{ + Error: "INVALID_INPUT", + Message: "registration failed — check your details", + }) + return + } + + // Issue session. + token, err := auth.GenerateToken() + if err != nil { + writeJSON(w, http.StatusInternalServerError, errorResponse{ + Error: "SERVER_ERROR", + Message: "failed to create session", + }) + return + } + + device := r.Header.Get("User-Agent") + ip := clientIP(r) + if _, err := database.CreateSession(uid, auth.HashToken(token), device, ip); err != nil { + writeJSON(w, http.StatusInternalServerError, errorResponse{ + Error: "SERVER_ERROR", + Message: "failed to create session", + }) + return + } + + user, _ := database.GetUserByID(uid) + writeJSON(w, http.StatusCreated, authSuccessResponse{ + Token: token, + User: toUserResponse(user), + }) + } +} + +// handleLogin processes POST /api/v1/auth/login. +func handleLogin(database *db.DB, limiter *auth.RateLimiter) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + var req loginRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeJSON(w, http.StatusBadRequest, errorResponse{ + Error: "INVALID_INPUT", + Message: "malformed request body", + }) + return + } + + req.Username = strings.TrimSpace(req.Username) + req.Password = strings.TrimSpace(req.Password) + + if req.Username == "" || req.Password == "" { + writeJSON(w, http.StatusBadRequest, errorResponse{ + Error: "INVALID_INPUT", + Message: "username and password are required", + }) + return + } + + ip := clientIP(r) + + // Check lockout first. + lockKey := "login_lock:" + ip + if limiter.IsLockedOut(lockKey) { + writeJSON(w, http.StatusTooManyRequests, errorResponse{ + Error: "RATE_LIMITED", + Message: "account temporarily locked due to too many failed attempts", + }) + return + } + + // Constant-time lookup: always attempt bcrypt compare even when user + // does not exist to prevent timing-based username enumeration. + user, err := database.GetUserByUsername(req.Username) + + failKey := "login_fail:" + ip + if err != nil || user == nil || !auth.CheckPassword(user.PasswordHash, req.Password) { + // Track failures; lockout after 10. + if !limiter.Allow(failKey, 10, 15*time.Minute) { + limiter.Lockout(lockKey, 15*time.Minute) + } + slog.Info("login failed", "ip", ip, "username_len", len(req.Username)) + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", + Message: "invalid credentials", + }) + return + } + + // Reset failure counter on success. + limiter.Reset(failKey) + + if user.Banned { + writeJSON(w, http.StatusForbidden, errorResponse{ + Error: "FORBIDDEN", + Message: "your account has been suspended", + }) + return + } + + // Issue session. + token, err := auth.GenerateToken() + if err != nil { + writeJSON(w, http.StatusInternalServerError, errorResponse{ + Error: "SERVER_ERROR", + Message: "failed to create session", + }) + return + } + + device := r.Header.Get("User-Agent") + if _, err := database.CreateSession(user.ID, auth.HashToken(token), device, ip); err != nil { + writeJSON(w, http.StatusInternalServerError, errorResponse{ + Error: "SERVER_ERROR", + Message: "failed to create session", + }) + return + } + + _ = database.UpdateUserStatus(user.ID, "online") + writeJSON(w, http.StatusOK, authSuccessResponse{ + Token: token, + User: toUserResponse(user), + }) + } +} + +// handleLogout processes POST /api/v1/auth/logout. +func handleLogout(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + sess, ok := r.Context().Value(SessionKey).(*db.Session) + if !ok || sess == nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", + Message: "not authenticated", + }) + return + } + + if err := database.DeleteSession(sess.TokenHash); err != nil { + writeJSON(w, http.StatusInternalServerError, errorResponse{ + Error: "SERVER_ERROR", + Message: "failed to logout", + }) + return + } + + w.WriteHeader(http.StatusNoContent) + } +} + +// handleMe processes GET /api/v1/auth/me. +func handleMe() 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: "not authenticated", + }) + return + } + writeJSON(w, http.StatusOK, toUserResponse(user)) + } +} + +// toUserResponse converts a db.User to the API response shape. +func toUserResponse(u *db.User) userResponse { + avatar := "" + if u.Avatar != nil { + avatar = *u.Avatar + } + return userResponse{ + ID: u.ID, + Username: u.Username, + Avatar: avatar, + Status: u.Status, + RoleID: u.RoleID, + CreatedAt: u.CreatedAt, + } +} diff --git a/Server/api/auth_handler_test.go b/Server/api/auth_handler_test.go new file mode 100644 index 00000000..943768ea --- /dev/null +++ b/Server/api/auth_handler_test.go @@ -0,0 +1,456 @@ +package api_test + +import ( + "bytes" + "encoding/json" + "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/db" +) + +// newAuthTestDB builds an in-memory DB with the full schema needed for auth tests. +func newAuthTestDB(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 +} + +// buildAuthRouter returns a chi router with auth routes mounted on /api/v1/auth. +func buildAuthRouter(database *db.DB, limiter *auth.RateLimiter) http.Handler { + r := chi.NewRouter() + api.MountAuthRoutes(r, database, limiter) + return r +} + +// postJSON is a test helper that POSTs JSON to the given router. +func postJSON(t *testing.T, router http.Handler, path string, body interface{}) *httptest.ResponseRecorder { + t.Helper() + raw, _ := json.Marshal(body) + req := httptest.NewRequest(http.MethodPost, path, bytes.NewReader(raw)) + req.Header.Set("Content-Type", "application/json") + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + return rr +} + +// postJSONWithToken posts with an Authorization header. +func postJSONWithToken(t *testing.T, router http.Handler, path, token string, body interface{}) *httptest.ResponseRecorder { + t.Helper() + raw, _ := json.Marshal(body) + req := httptest.NewRequest(http.MethodPost, path, bytes.NewReader(raw)) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+token) + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + return rr +} + +// getWithToken performs a GET with an Authorization header. +func getWithToken(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 +} + +// ─── Register tests ─────────────────────────────────────────────────────────── + +func TestRegister_Success(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + // Create an invite first. + ownerID, _ := database.CreateUser("owner", "hash", 1) + code, _ := database.CreateInvite(ownerID, 1, nil) + + rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{ + "username": "newuser", + "password": "securePass1", + "invite_code": code, + }) + + if rr.Code != http.StatusCreated { + t.Errorf("Register status = %d, want 201; body = %s", rr.Code, rr.Body.String()) + } + + var resp map[string]interface{} + json.NewDecoder(rr.Body).Decode(&resp) + if resp["token"] == nil { + t.Error("Register response missing token") + } + if resp["user"] == nil { + t.Error("Register response missing user") + } +} + +func TestRegister_InvalidInvite(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{ + "username": "newuser", + "password": "securePass1", + "invite_code": "bogus", + }) + + if rr.Code != http.StatusBadRequest { + t.Errorf("Register invalid invite status = %d, want 400", rr.Code) + } +} + +func TestRegister_WeakPassword(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + ownerID, _ := database.CreateUser("owner2", "hash", 1) + code, _ := database.CreateInvite(ownerID, 1, nil) + + rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{ + "username": "newuser", + "password": "short", + "invite_code": code, + }) + + if rr.Code != http.StatusBadRequest { + t.Errorf("Register weak password status = %d, want 400", rr.Code) + } +} + +func TestRegister_InviteUsedUp(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + ownerID, _ := database.CreateUser("owner3", "hash", 1) + code, _ := database.CreateInvite(ownerID, 1, nil) // max 1 use + + // First registration should succeed. + postJSON(t, router, "/api/v1/auth/register", map[string]string{ + "username": "user1", + "password": "securePass1", + "invite_code": code, + }) + + // Second should fail — invite exhausted. + rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{ + "username": "user2", + "password": "securePass2", + "invite_code": code, + }) + + if rr.Code != http.StatusBadRequest { + t.Errorf("Register exhausted invite status = %d, want 400", rr.Code) + } +} + +func TestRegister_MissingFields(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{}) + if rr.Code != http.StatusBadRequest { + t.Errorf("Register missing fields status = %d, want 400", rr.Code) + } +} + +func TestRegister_ErrorNeverRevealUsername(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{ + "username": "someone", + "password": "securePass1", + "invite_code": "bogus", + }) + + body := rr.Body.String() + // Must not hint that the username doesn't exist or the invite is invalid specifically + if contains(body, "username") && contains(body, "taken") { + t.Error("Register error message reveals username status") + } +} + +// ─── Login tests ────────────────────────────────────────────────────────────── + +func TestLogin_Success(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + hash, _ := auth.HashPassword("correctPass1") + database.CreateUser("loginuser", hash, 4) + + rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{ + "username": "loginuser", + "password": "correctPass1", + }) + + if rr.Code != http.StatusOK { + t.Errorf("Login status = %d, want 200; body = %s", rr.Code, rr.Body.String()) + } + + var resp map[string]interface{} + json.NewDecoder(rr.Body).Decode(&resp) + if resp["token"] == nil { + t.Error("Login response missing token") + } +} + +func TestLogin_WrongPassword(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + hash, _ := auth.HashPassword("correctPass1") + database.CreateUser("loginuser2", hash, 4) + + rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{ + "username": "loginuser2", + "password": "wrongpassword", + }) + + if rr.Code != http.StatusUnauthorized { + t.Errorf("Login wrong password status = %d, want 401", rr.Code) + } +} + +func TestLogin_UnknownUser(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{ + "username": "nobody", + "password": "anypass123", + }) + + if rr.Code != http.StatusUnauthorized { + t.Errorf("Login unknown user status = %d, want 401", rr.Code) + } +} + +func TestLogin_GenericErrorOnBadCredentials(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{ + "username": "nobody", + "password": "anypass123", + }) + + body := rr.Body.String() + // The response must never reveal whether the user exists + if contains(body, "user not found") || contains(body, "does not exist") { + t.Errorf("Login error reveals user existence: %s", body) + } +} + +func TestLogin_BannedUser(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + hash, _ := auth.HashPassword("correctPass1") + id, _ := database.CreateUser("banned", hash, 4) + database.BanUser(id, "violated rules", nil) + + rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{ + "username": "banned", + "password": "correctPass1", + }) + + if rr.Code != http.StatusForbidden { + t.Errorf("Login banned user status = %d, want 403", rr.Code) + } +} + +func TestLogin_MissingFields(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + rr := postJSON(t, router, "/api/v1/auth/login", map[string]string{}) + if rr.Code != http.StatusBadRequest { + t.Errorf("Login missing fields status = %d, want 400", rr.Code) + } +} + +// ─── Logout tests ───────────────────────────────────────────────────────────── + +func TestLogout_Success(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + hash, _ := auth.HashPassword("correctPass1") + uid, _ := database.CreateUser("logoutuser", hash, 4) + token, _ := auth.GenerateToken() + tokenHash := auth.HashToken(token) + database.CreateSession(uid, tokenHash, "test", "127.0.0.1") + + rr := postJSONWithToken(t, router, "/api/v1/auth/logout", token, nil) + + if rr.Code != http.StatusNoContent { + t.Errorf("Logout status = %d, want 204", rr.Code) + } + + // Session should be gone. + sess, _ := database.GetSessionByTokenHash(tokenHash) + if sess != nil { + t.Error("Session still exists after logout") + } +} + +func TestLogout_NoAuth(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + req := httptest.NewRequest(http.MethodPost, "/api/v1/auth/logout", nil) + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusUnauthorized { + t.Errorf("Logout no auth status = %d, want 401", rr.Code) + } +} + +// ─── Me tests ───────────────────────────────────────────────────────────────── + +func TestMe_Success(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + hash, _ := auth.HashPassword("correctPass1") + uid, _ := database.CreateUser("meuser", hash, 4) + token, _ := auth.GenerateToken() + database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1") + + rr := getWithToken(t, router, "/api/v1/auth/me", token) + + if rr.Code != http.StatusOK { + t.Errorf("Me status = %d, want 200; body = %s", rr.Code, rr.Body.String()) + } + + var resp map[string]interface{} + json.NewDecoder(rr.Body).Decode(&resp) + if resp["id"] == nil { + t.Error("Me response missing id") + } + if resp["username"] != "meuser" { + t.Errorf("Me username = %v, want meuser", resp["username"]) + } +} + +func TestMe_NoAuth(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/auth/me", nil) + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusUnauthorized { + t.Errorf("Me no auth status = %d, want 401", rr.Code) + } +} + +// ─── Rate limiting integration test ────────────────────────────────────────── + +func TestRegister_RateLimit(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + ownerID, _ := database.CreateUser("rl_owner", "hash", 1) + + // Attempt register 4 times (limit=3) — 4th should be rate-limited. + var lastCode int + for i := 0; i < 4; i++ { + code, _ := database.CreateInvite(ownerID, 1, nil) + rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{ + "username": "rl_user" + string(rune('0'+i)), + "password": "securePass1", + "invite_code": code, + }) + lastCode = rr.Code + } + + if lastCode != http.StatusTooManyRequests { + t.Errorf("Register rate limit: last attempt status = %d, want 429", lastCode) + } +} + +// ─── Helpers ───────────────────────────────────────────────────────────────── + +func contains(s, sub string) bool { + return len(s) >= len(sub) && (s == sub || len(s) > 0 && containsStr(s, sub)) +} + +func containsStr(s, sub string) bool { + for i := 0; i <= len(s)-len(sub); i++ { + if s[i:i+len(sub)] == sub { + return true + } + } + return false +} + +// expiredInviteDB creates a DB with an already-expired invite. +func expiredInviteDB(t *testing.T) (*db.DB, string) { + t.Helper() + database := newAuthTestDB(t) + ownerID, _ := database.CreateUser("expowner", "hash", 1) + past := time.Now().Add(-time.Hour) + code, _ := database.CreateInvite(ownerID, 0, &past) + return database, code +} + +func TestRegister_ExpiredInvite(t *testing.T) { + database, code := expiredInviteDB(t) + limiter := auth.NewRateLimiter() + router := buildAuthRouter(database, limiter) + + rr := postJSON(t, router, "/api/v1/auth/register", map[string]string{ + "username": "newuser", + "password": "securePass1", + "invite_code": code, + }) + + if rr.Code != http.StatusBadRequest { + t.Errorf("Register expired invite status = %d, want 400", rr.Code) + } +} diff --git a/Server/api/invite_handler.go b/Server/api/invite_handler.go new file mode 100644 index 00000000..442af7cc --- /dev/null +++ b/Server/api/invite_handler.go @@ -0,0 +1,160 @@ +package api + +import ( + "encoding/json" + "net/http" + "time" + + "github.com/go-chi/chi/v5" + "github.com/owncord/server/db" +) + +// manageInvitesPerm is the MANAGE_INVITES permission bit. +const manageInvitesPerm = int64(0x4000000) + +// createInviteRequest is the JSON body for POST /api/v1/invites. +type createInviteRequest struct { + MaxUses int `json:"max_uses"` + ExpiresInHours int `json:"expires_in_hours"` +} + +// inviteResponse is the API shape for an invite. +type inviteResponse struct { + ID int64 `json:"id"` + Code string `json:"code"` + MaxUses *int `json:"max_uses"` + Uses int `json:"uses"` + ExpiresAt *string `json:"expires_at"` + Revoked bool `json:"revoked"` + CreatedAt string `json:"created_at"` +} + +// MountInviteRoutes registers invite endpoints on the given router. +// All routes require authentication and MANAGE_INVITES permission. +func MountInviteRoutes(r chi.Router, database *db.DB) { + r.Route("/api/v1/invites", func(r chi.Router) { + r.Use(AuthMiddleware(database)) + r.Use(RequirePermission(manageInvitesPerm)) + + r.Post("/", handleCreateInvite(database)) + r.Get("/", handleListInvites(database)) + r.Delete("/{code}", handleRevokeInvite(database)) + }) +} + +// handleCreateInvite processes POST /api/v1/invites. +func handleCreateInvite(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + var req createInviteRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + // Treat missing body as default values (all optional). + req = createInviteRequest{} + } + + user, ok := r.Context().Value(UserKey).(*db.User) + if !ok || user == nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", + Message: "not authenticated", + }) + return + } + + var expiresAt *time.Time + if req.ExpiresInHours > 0 { + t := time.Now().Add(time.Duration(req.ExpiresInHours) * time.Hour) + expiresAt = &t + } + + code, err := database.CreateInvite(user.ID, req.MaxUses, expiresAt) + if err != nil { + writeJSON(w, http.StatusInternalServerError, errorResponse{ + Error: "SERVER_ERROR", + Message: "failed to create invite", + }) + return + } + + inv, err := database.GetInvite(code) + if err != nil || inv == nil { + writeJSON(w, http.StatusInternalServerError, errorResponse{ + Error: "SERVER_ERROR", + Message: "failed to retrieve invite", + }) + return + } + + writeJSON(w, http.StatusCreated, toInviteResponse(inv)) + } +} + +// handleListInvites processes GET /api/v1/invites. +func handleListInvites(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + invites, err := database.ListInvites() + if err != nil { + writeJSON(w, http.StatusInternalServerError, errorResponse{ + Error: "SERVER_ERROR", + Message: "failed to list invites", + }) + return + } + + resp := make([]inviteResponse, 0, len(invites)) + for _, inv := range invites { + resp = append(resp, toInviteResponse(inv)) + } + writeJSON(w, http.StatusOK, resp) + } +} + +// handleRevokeInvite processes DELETE /api/v1/invites/:code. +func handleRevokeInvite(database *db.DB) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + code := chi.URLParam(r, "code") + + inv, err := database.GetInvite(code) + if err != nil { + writeJSON(w, http.StatusInternalServerError, errorResponse{ + Error: "SERVER_ERROR", + Message: "failed to look up invite", + }) + return + } + if inv == nil { + writeJSON(w, http.StatusNotFound, errorResponse{ + Error: "NOT_FOUND", + Message: "invite not found", + }) + return + } + + if err := database.RevokeInvite(code); err != nil { + writeJSON(w, http.StatusInternalServerError, errorResponse{ + Error: "SERVER_ERROR", + Message: "failed to revoke invite", + }) + return + } + + w.WriteHeader(http.StatusNoContent) + } +} + +// toInviteResponse converts a db.Invite to the API response shape. +func toInviteResponse(inv *db.Invite) inviteResponse { + var maxUses *int + if inv.MaxUses != nil { + v := *inv.MaxUses + maxUses = &v + } + return inviteResponse{ + ID: inv.ID, + Code: inv.Code, + MaxUses: maxUses, + Uses: inv.Uses, + ExpiresAt: inv.ExpiresAt, + Revoked: inv.Revoked, + CreatedAt: inv.CreatedAt, + } +} diff --git a/Server/api/invite_handler_test.go b/Server/api/invite_handler_test.go new file mode 100644 index 00000000..9aa79224 --- /dev/null +++ b/Server/api/invite_handler_test.go @@ -0,0 +1,294 @@ +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/auth" + "github.com/owncord/server/db" +) + +// buildInviteRouter returns a chi router with invite routes and auth middleware. +func buildInviteRouter(database *db.DB, limiter *auth.RateLimiter) http.Handler { + r := chi.NewRouter() + api.MountAuthRoutes(r, database, limiter) + api.MountInviteRoutes(r, database) + return r +} + +// loginAndGetToken creates a user with a known password and returns their session token. +func loginAndGetToken(t *testing.T, router http.Handler, database *db.DB, username string, roleID int) string { + t.Helper() + hash, _ := auth.HashPassword("Password1!") + uid, _ := database.CreateUser(username, hash, roleID) + token, _ := auth.GenerateToken() + database.CreateSession(uid, auth.HashToken(token), "test", "127.0.0.1") + return token +} + +// ─── POST /api/v1/invites ───────────────────────────────────────────────────── + +func TestCreateInvite_Success(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildInviteRouter(database, limiter) + + // Admin role (id=2) has MANAGE_INVITES (0x4000000) set. + token := loginAndGetToken(t, router, database, "invitecreator", 2) + + rr := postJSONWithToken(t, router, "/api/v1/invites", token, map[string]interface{}{ + "max_uses": 5, + "expires_in_hours": 48, + }) + + if rr.Code != http.StatusCreated { + t.Errorf("CreateInvite status = %d, want 201; body = %s", rr.Code, rr.Body.String()) + } + + var resp map[string]interface{} + json.NewDecoder(rr.Body).Decode(&resp) + if resp["code"] == nil { + t.Error("CreateInvite response missing code") + } +} + +func TestCreateInvite_Unauthorized(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildInviteRouter(database, limiter) + + rr := postJSON(t, router, "/api/v1/invites", map[string]interface{}{ + "max_uses": 5, + }) + + if rr.Code != http.StatusUnauthorized { + t.Errorf("CreateInvite no auth status = %d, want 401", rr.Code) + } +} + +func TestCreateInvite_MemberForbidden(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildInviteRouter(database, limiter) + + // Member role (id=4) does NOT have MANAGE_INVITES. + token := loginAndGetToken(t, router, database, "memberuser", 4) + + rr := postJSONWithToken(t, router, "/api/v1/invites", token, map[string]interface{}{ + "max_uses": 1, + }) + + if rr.Code != http.StatusForbidden { + t.Errorf("CreateInvite member status = %d, want 403", rr.Code) + } +} + +func TestCreateInvite_Unlimited(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildInviteRouter(database, limiter) + + token := loginAndGetToken(t, router, database, "adminuser2", 2) + + rr := postJSONWithToken(t, router, "/api/v1/invites", token, map[string]interface{}{}) + + if rr.Code != http.StatusCreated { + t.Errorf("CreateInvite unlimited status = %d, want 201", rr.Code) + } +} + +// ─── GET /api/v1/invites ────────────────────────────────────────────────────── + +func TestListInvites_Success(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildInviteRouter(database, limiter) + + token := loginAndGetToken(t, router, database, "listuser", 2) + + // Create a couple of invites. + postJSONWithToken(t, router, "/api/v1/invites", token, map[string]interface{}{"max_uses": 1}) + postJSONWithToken(t, router, "/api/v1/invites", token, map[string]interface{}{"max_uses": 5}) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/invites", 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.Errorf("ListInvites status = %d, want 200; body = %s", rr.Code, rr.Body.String()) + } + + var resp []interface{} + json.NewDecoder(rr.Body).Decode(&resp) + if len(resp) < 2 { + t.Errorf("ListInvites returned %d items, want >= 2", len(resp)) + } +} + +func TestListInvites_Unauthorized(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildInviteRouter(database, limiter) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/invites", nil) + req.RemoteAddr = "127.0.0.1:9999" + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusUnauthorized { + t.Errorf("ListInvites no auth status = %d, want 401", rr.Code) + } +} + +// ─── DELETE /api/v1/invites/:code ───────────────────────────────────────────── + +func TestRevokeInvite_Success(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildInviteRouter(database, limiter) + + token := loginAndGetToken(t, router, database, "revoker", 2) + + // Create invite via API. + rr := postJSONWithToken(t, router, "/api/v1/invites", token, map[string]interface{}{}) + if rr.Code != http.StatusCreated { + t.Fatalf("Create invite for revoke test: status = %d, body = %s", rr.Code, rr.Body.String()) + } + var created map[string]interface{} + json.NewDecoder(rr.Body).Decode(&created) + codeVal, ok := created["code"] + if !ok || codeVal == nil { + t.Fatalf("Create invite response missing code field; body parsed as %v", created) + } + code := codeVal.(string) + + // Revoke it. + req := httptest.NewRequest(http.MethodDelete, "/api/v1/invites/"+code, nil) + req.Header.Set("Authorization", "Bearer "+token) + req.RemoteAddr = "127.0.0.1:9999" + rr2 := httptest.NewRecorder() + router.ServeHTTP(rr2, req) + + if rr2.Code != http.StatusNoContent { + t.Errorf("RevokeInvite status = %d, want 204; body = %s", rr2.Code, rr2.Body.String()) + } + + // Verify invite is revoked. + inv, _ := database.GetInvite(code) + if inv == nil || !inv.Revoked { + t.Error("Invite not revoked in database after DELETE") + } +} + +func TestRevokeInvite_NotFound(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildInviteRouter(database, limiter) + + token := loginAndGetToken(t, router, database, "revoker2", 2) + + req := httptest.NewRequest(http.MethodDelete, "/api/v1/invites/doesnotexist", 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.StatusNotFound { + t.Errorf("RevokeInvite not found status = %d, want 404", rr.Code) + } +} + +func TestRevokeInvite_MemberForbidden(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildInviteRouter(database, limiter) + + adminToken := loginAndGetToken(t, router, database, "admin3", 2) + memberToken := loginAndGetToken(t, router, database, "member3", 4) + + // Admin creates invite. + rr := postJSONWithToken(t, router, "/api/v1/invites", adminToken, map[string]interface{}{}) + var created map[string]interface{} + json.NewDecoder(rr.Body).Decode(&created) + code := created["code"].(string) + + // Member tries to revoke. + req := httptest.NewRequest(http.MethodDelete, "/api/v1/invites/"+code, nil) + req.Header.Set("Authorization", "Bearer "+memberToken) + req.RemoteAddr = "127.0.0.1:9999" + rr2 := httptest.NewRecorder() + router.ServeHTTP(rr2, req) + + if rr2.Code != http.StatusForbidden { + t.Errorf("RevokeInvite member status = %d, want 403", rr2.Code) + } +} + +// TestListInvites_IncludesRevokedAndActive checks the list endpoint returns +// correct data for both revoked and active invites. +func TestListInvites_IncludesRevokedAndActive(t *testing.T) { + database := newAuthTestDB(t) + limiter := auth.NewRateLimiter() + router := buildInviteRouter(database, limiter) + + token := loginAndGetToken(t, router, database, "listall", 2) + + // Create and revoke one invite. + rr := postJSONWithToken(t, router, "/api/v1/invites", token, map[string]interface{}{}) + if rr.Code != http.StatusCreated { + t.Fatalf("Create invite for list test: status = %d, body = %s", rr.Code, rr.Body.String()) + } + var created map[string]interface{} + json.NewDecoder(rr.Body).Decode(&created) + code := created["code"].(string) + + delReq := httptest.NewRequest(http.MethodDelete, "/api/v1/invites/"+code, nil) + delReq.Header.Set("Authorization", "Bearer "+token) + delReq.RemoteAddr = "127.0.0.1:9999" + httptest.NewRecorder() // discard + router.ServeHTTP(httptest.NewRecorder(), delReq) + + // Create one active invite. + postJSONWithToken(t, router, "/api/v1/invites", token, map[string]interface{}{}) + + // List should include both. + req := httptest.NewRequest(http.MethodGet, "/api/v1/invites", nil) + req.Header.Set("Authorization", "Bearer "+token) + req.RemoteAddr = "127.0.0.1:9999" + rr2 := httptest.NewRecorder() + router.ServeHTTP(rr2, req) + + if rr2.Code != http.StatusOK { + t.Errorf("ListInvites status = %d, want 200", rr2.Code) + } +} + +// ─── Helpers for ListInvites queries ───────────────────────────────────────── + +// ListInvites returns all invites from the DB for assertions. +func listInvitesFromDB(t *testing.T, database *db.DB) []*db.Invite { + t.Helper() + rows, err := database.Query(`SELECT id, code, created_by, max_uses, use_count, expires_at, revoked, created_at FROM invites`) + if err != nil { + t.Fatalf("listing invites: %v", err) + } + defer rows.Close() + + var invites []*db.Invite + for rows.Next() { + inv := &db.Invite{} + var revoked int + if err := rows.Scan(&inv.ID, &inv.Code, &inv.CreatedBy, &inv.MaxUses, &inv.Uses, &inv.ExpiresAt, &revoked, &inv.CreatedAt); err != nil { + t.Fatalf("scanning invite: %v", err) + } + inv.Revoked = revoked != 0 + invites = append(invites, inv) + } + return invites +} diff --git a/Server/api/middleware.go b/Server/api/middleware.go new file mode 100644 index 00000000..1606ef5c --- /dev/null +++ b/Server/api/middleware.go @@ -0,0 +1,194 @@ +package api + +import ( + "context" + "fmt" + "net/http" + "strings" + "time" + + "github.com/owncord/server/auth" + "github.com/owncord/server/db" +) + +// contextKey is an unexported type for context keys in this package. +type contextKey int + +const ( + // UserKey is the context key for the authenticated *db.User. + UserKey contextKey = iota + // SessionKey is the context key for the authenticated *db.Session. + SessionKey + // RoleKey is the context key for the *db.Role of the authenticated user. + RoleKey +) + +// AuthMiddleware reads the "Authorization: Bearer " header, validates +// the session, and injects the user and session into the request context. +// Returns 401 if the token is missing, invalid, or the session is expired. +func AuthMiddleware(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 := extractBearerToken(r) + if !ok { + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", + Message: "missing or invalid authorization header", + }) + return + } + + hash := auth.HashToken(token) + sess, err := database.GetSessionByTokenHash(hash) + if err != nil || sess == nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", + Message: "invalid or expired session", + }) + return + } + + // Check expiry. + if isSessionExpired(sess.ExpiresAt) { + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", + Message: "session has expired", + }) + return + } + + // Load user. + user, err := database.GetUserByID(sess.UserID) + if err != nil || user == nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", + Message: "user not found", + }) + return + } + + // Load role for permission checks. + role, err := database.GetRoleByID(user.RoleID) + if err != nil { + writeJSON(w, http.StatusUnauthorized, errorResponse{ + Error: "UNAUTHORIZED", + Message: "role not found", + }) + return + } + + // Touch session in background — non-fatal if it fails. + _ = database.TouchSession(hash) + + ctx := context.WithValue(r.Context(), UserKey, user) + ctx = context.WithValue(ctx, SessionKey, sess) + ctx = context.WithValue(ctx, RoleKey, role) + next.ServeHTTP(w, r.WithContext(ctx)) + }) + } +} + +// RequirePermission returns middleware that checks the authenticated user's +// role permissions. Returns 403 if the user lacks the required permission. +// The ADMINISTRATOR bit (0x40000000) bypasses all checks. +func RequirePermission(perm int64) func(http.Handler) http.Handler { + const administrator = int64(0x40000000) + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + role, ok := r.Context().Value(RoleKey).(*db.Role) + if !ok || role == nil { + writeJSON(w, http.StatusForbidden, errorResponse{ + Error: "FORBIDDEN", + Message: "insufficient permissions", + }) + return + } + + // ADMINISTRATOR bypasses all permission checks. + if role.Permissions&administrator != 0 { + next.ServeHTTP(w, r) + return + } + + if role.Permissions&perm == 0 { + writeJSON(w, http.StatusForbidden, errorResponse{ + Error: "FORBIDDEN", + Message: "insufficient permissions", + }) + return + } + + next.ServeHTTP(w, r) + }) + } +} + +// RateLimitMiddleware returns middleware that limits requests per IP using the +// provided RateLimiter. The IP is taken from X-Real-IP header when present, +// falling back to RemoteAddr. Returns 429 with Retry-After when exceeded. +func RateLimitMiddleware(limiter *auth.RateLimiter, limit int, window time.Duration) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ip := clientIP(r) + + if !limiter.Allow(ip, limit, window) { + w.Header().Set("Retry-After", fmt.Sprintf("%d", int(window.Seconds()))) + writeJSON(w, http.StatusTooManyRequests, errorResponse{ + Error: "RATE_LIMITED", + Message: "too many requests, please slow down", + }) + return + } + + next.ServeHTTP(w, r) + }) + } +} + +// ─── Helpers ────────────────────────────────────────────────────────────────── + +// extractBearerToken parses "Authorization: Bearer " and returns the +// token and true, or "", false if the header is missing or malformed. +func extractBearerToken(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 +} + +// clientIP returns the client IP from X-Real-IP or RemoteAddr (without port). +func clientIP(r *http.Request) string { + if ip := r.Header.Get("X-Real-IP"); ip != "" { + return ip + } + // RemoteAddr is "host:port"; strip the port. + addr := r.RemoteAddr + if idx := strings.LastIndex(addr, ":"); idx != -1 { + return addr[:idx] + } + return addr +} + +// isSessionExpired returns true when expiresAt string represents a past time. +// Handles both "2006-01-02 15:04:05" (SQLite) and "2006-01-02T15:04:05Z" formats. +func isSessionExpired(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()) + } + } + // Unparseable expiry — treat as expired for safety. + return true +} + +// errorResponse is the standard error JSON shape. +type errorResponse struct { + Error string `json:"error"` + Message string `json:"message"` +} diff --git a/Server/api/middleware_test.go b/Server/api/middleware_test.go new file mode 100644 index 00000000..1dfddb4b --- /dev/null +++ b/Server/api/middleware_test.go @@ -0,0 +1,361 @@ +package api_test + +import ( + "net/http" + "net/http/httptest" + "testing" + "testing/fstest" + "time" + + "github.com/owncord/server/api" + "github.com/owncord/server/auth" + "github.com/owncord/server/db" +) + +// ─── Helpers ───────────────────────────────────────────────────────────────── + +func newAPITestDB(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 +} + +// ok is a trivial handler that responds 200 OK to confirm the middleware +// passed the request through. +func ok(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) +} + +// bearerToken wraps an HTTP handler with an Authorization header bearing token. +func withBearer(req *http.Request, token string) *http.Request { + req.Header.Set("Authorization", "Bearer "+token) + return req +} + +// ─── AuthMiddleware tests ───────────────────────────────────────────────────── + +func TestAuthMiddleware_ValidToken(t *testing.T) { + database := newAPITestDB(t) + uid, _ := database.CreateUser("alice", "hash", 4) + token, _ := auth.GenerateToken() + hash := auth.HashToken(token) + database.CreateSession(uid, hash, "test", "127.0.0.1") + + h := api.AuthMiddleware(database)(http.HandlerFunc(ok)) + req := httptest.NewRequest(http.MethodGet, "/", nil) + withBearer(req, token) + rr := httptest.NewRecorder() + + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Errorf("AuthMiddleware valid token status = %d, want %d", rr.Code, http.StatusOK) + } +} + +func TestAuthMiddleware_MissingToken(t *testing.T) { + database := newAPITestDB(t) + + h := api.AuthMiddleware(database)(http.HandlerFunc(ok)) + req := httptest.NewRequest(http.MethodGet, "/", nil) + rr := httptest.NewRecorder() + + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusUnauthorized { + t.Errorf("AuthMiddleware no token status = %d, want 401", rr.Code) + } +} + +func TestAuthMiddleware_InvalidToken(t *testing.T) { + database := newAPITestDB(t) + + h := api.AuthMiddleware(database)(http.HandlerFunc(ok)) + req := httptest.NewRequest(http.MethodGet, "/", nil) + withBearer(req, "notarealtoken") + rr := httptest.NewRecorder() + + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusUnauthorized { + t.Errorf("AuthMiddleware invalid token status = %d, want 401", rr.Code) + } +} + +func TestAuthMiddleware_ExpiredSession(t *testing.T) { + database := newAPITestDB(t) + uid, _ := database.CreateUser("bob", "hash", 4) + token, _ := auth.GenerateToken() + hash := auth.HashToken(token) + + // Insert an already-expired session. + pastTime := time.Now().Add(-time.Hour).UTC().Format("2006-01-02 15:04:05") + database.Exec( + `INSERT INTO sessions (user_id, token, device, ip_address, expires_at) VALUES (?, ?, ?, ?, ?)`, + uid, hash, "test", "127.0.0.1", pastTime, + ) + + h := api.AuthMiddleware(database)(http.HandlerFunc(ok)) + req := httptest.NewRequest(http.MethodGet, "/", nil) + withBearer(req, token) + rr := httptest.NewRecorder() + + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusUnauthorized { + t.Errorf("AuthMiddleware expired session status = %d, want 401", rr.Code) + } +} + +func TestAuthMiddleware_MalformedAuthHeader(t *testing.T) { + database := newAPITestDB(t) + + h := api.AuthMiddleware(database)(http.HandlerFunc(ok)) + + cases := []string{ + "Token abc", // wrong scheme + "Bearer", // missing token after Bearer + "abc", // no space + } + for _, header := range cases { + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.Header.Set("Authorization", header) + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + if rr.Code != http.StatusUnauthorized { + t.Errorf("AuthMiddleware header=%q status = %d, want 401", header, rr.Code) + } + } +} + +// ─── RequirePermission tests ────────────────────────────────────────────────── + +func TestRequirePermission_Allowed(t *testing.T) { + database := newAPITestDB(t) + uid, _ := database.CreateUser("carol", "hash", 4) // Member role = 0x00100601 + token, _ := auth.GenerateToken() + hash := auth.HashToken(token) + database.CreateSession(uid, hash, "test", "127.0.0.1") + + // SEND_MESSAGES = 0x1 — Member role has this bit + h := api.AuthMiddleware(database)( + api.RequirePermission(0x1)(http.HandlerFunc(ok)), + ) + req := httptest.NewRequest(http.MethodGet, "/", nil) + withBearer(req, token) + rr := httptest.NewRecorder() + + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Errorf("RequirePermission allowed status = %d, want 200", rr.Code) + } +} + +func TestRequirePermission_Forbidden(t *testing.T) { + database := newAPITestDB(t) + uid, _ := database.CreateUser("dave", "hash", 4) // Member role = 0x00100601 + token, _ := auth.GenerateToken() + hash := auth.HashToken(token) + database.CreateSession(uid, hash, "test", "127.0.0.1") + + // MANAGE_ROLES = 0x1000000 — Member does not have this + h := api.AuthMiddleware(database)( + api.RequirePermission(0x1000000)(http.HandlerFunc(ok)), + ) + req := httptest.NewRequest(http.MethodGet, "/", nil) + withBearer(req, token) + rr := httptest.NewRecorder() + + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusForbidden { + t.Errorf("RequirePermission forbidden status = %d, want 403", rr.Code) + } +} + +func TestRequirePermission_Administrator_Bypass(t *testing.T) { + database := newAPITestDB(t) + // Owner role (id=1) has permissions 0x7FFFFFFF which includes ADMINISTRATOR (0x40000000) + uid, _ := database.CreateUser("owner", "hash", 1) + token, _ := auth.GenerateToken() + hash := auth.HashToken(token) + database.CreateSession(uid, hash, "test", "127.0.0.1") + + // Any permission should pass for ADMINISTRATOR + h := api.AuthMiddleware(database)( + api.RequirePermission(0x1000000)(http.HandlerFunc(ok)), + ) + req := httptest.NewRequest(http.MethodGet, "/", nil) + withBearer(req, token) + rr := httptest.NewRecorder() + + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Errorf("RequirePermission administrator bypass status = %d, want 200", rr.Code) + } +} + +// ─── RateLimitMiddleware tests ──────────────────────────────────────────────── + +func TestRateLimitMiddleware_UnderLimit(t *testing.T) { + limiter := auth.NewRateLimiter() + + h := api.RateLimitMiddleware(limiter, 5, time.Minute)(http.HandlerFunc(ok)) + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "10.0.0.1:1234" + rr := httptest.NewRecorder() + + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Errorf("RateLimitMiddleware under limit status = %d, want 200", rr.Code) + } +} + +func TestRateLimitMiddleware_OverLimit(t *testing.T) { + limiter := auth.NewRateLimiter() + limit := 3 + + h := api.RateLimitMiddleware(limiter, limit, time.Minute)(http.HandlerFunc(ok)) + + for i := 0; i < limit; i++ { + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "10.0.0.2:1234" + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + } + + // This next request should be rate-limited. + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "10.0.0.2:1234" + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusTooManyRequests { + t.Errorf("RateLimitMiddleware over limit status = %d, want 429", rr.Code) + } +} + +func TestRateLimitMiddleware_RetryAfterHeader(t *testing.T) { + limiter := auth.NewRateLimiter() + + h := api.RateLimitMiddleware(limiter, 1, time.Minute)(http.HandlerFunc(ok)) + + // Exhaust limit. + for i := 0; i < 2; i++ { + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "10.0.0.3:1234" + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + } + + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "10.0.0.3:1234" + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if rr.Header().Get("Retry-After") == "" { + t.Error("RateLimitMiddleware: missing Retry-After header on 429 response") + } +} + +func TestRateLimitMiddleware_XRealIPUsed(t *testing.T) { + limiter := auth.NewRateLimiter() + limit := 2 + + h := api.RateLimitMiddleware(limiter, limit, time.Minute)(http.HandlerFunc(ok)) + + // Two requests from the same X-Real-IP but different RemoteAddr. + for i := 0; i < limit; i++ { + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.Header.Set("X-Real-IP", "192.168.1.1") + req.RemoteAddr = "10.0.0.99:9999" + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + } + + // Third request should be blocked by the X-Real-IP key. + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.Header.Set("X-Real-IP", "192.168.1.1") + req.RemoteAddr = "10.0.0.99:9999" + rr := httptest.NewRecorder() + h.ServeHTTP(rr, req) + + if rr.Code != http.StatusTooManyRequests { + t.Errorf("RateLimitMiddleware X-Real-IP status = %d, want 429", rr.Code) + } +} + +// apiTestSchema is the full schema needed for all api tests (middleware, +// auth handler, and invite handler). +var apiTestSchema = []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, 'Moderator', '#3498DB', 1048575, 60, 0), + (4, '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 4 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 invites ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + code TEXT NOT NULL UNIQUE, + created_by INTEGER NOT NULL REFERENCES users(id), + redeemed_by INTEGER 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 INDEX IF NOT EXISTS idx_invites_code ON invites(code); +`) diff --git a/Server/api/router.go b/Server/api/router.go index 3ec1ed49..56068e5f 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/auth" "github.com/owncord/server/config" "github.com/owncord/server/db" ) @@ -27,11 +28,20 @@ func NewRouter(cfg *config.Config, database *db.DB) http.Handler { // Health check — unauthenticated, no versioning prefix. r.Get("/health", handleHealth) + // Shared rate limiter for auth endpoints. + limiter := auth.NewRateLimiter() + // Versioned API routes. r.Route("/api/v1", func(r chi.Router) { r.Get("/info", handleInfo(cfg)) }) + // Auth routes: register, login, logout, me. + MountAuthRoutes(r, database, limiter) + + // Invite management routes (require MANAGE_INVITES permission). + MountInviteRoutes(r, database) + return r } diff --git a/Server/auth/password.go b/Server/auth/password.go new file mode 100644 index 00000000..f1cffb07 --- /dev/null +++ b/Server/auth/password.go @@ -0,0 +1,50 @@ +package auth + +import ( + "errors" + + "golang.org/x/crypto/bcrypt" +) + +const ( + bcryptCost = 12 + minPassLen = 8 + maxPassLen = 72 // bcrypt silently truncates beyond 72 bytes +) + +// ErrPasswordTooShort is returned when the password is below the minimum length. +var ErrPasswordTooShort = errors.New("password must be at least 8 characters") + +// ErrPasswordTooLong is returned when the password exceeds bcrypt's 72-byte limit. +var ErrPasswordTooLong = errors.New("password must not exceed 72 characters") + +// HashPassword returns a bcrypt hash of password using cost 12. +func HashPassword(password string) (string, error) { + hash, err := bcrypt.GenerateFromPassword([]byte(password), bcryptCost) + if err != nil { + return "", err + } + return string(hash), nil +} + +// CheckPassword reports whether password matches hash. Returns false on any +// error, including an empty or malformed hash. +func CheckPassword(hash, password string) bool { + if hash == "" { + return false + } + err := bcrypt.CompareHashAndPassword([]byte(hash), []byte(password)) + return err == nil +} + +// ValidatePasswordStrength returns an error if password fails strength +// requirements: minimum 8 characters, maximum 72 characters. +func ValidatePasswordStrength(password string) error { + if len(password) < minPassLen { + return ErrPasswordTooShort + } + if len(password) > maxPassLen { + return ErrPasswordTooLong + } + return nil +} diff --git a/Server/auth/password_test.go b/Server/auth/password_test.go new file mode 100644 index 00000000..a8581062 --- /dev/null +++ b/Server/auth/password_test.go @@ -0,0 +1,106 @@ +package auth_test + +import ( + "strings" + "testing" + + "github.com/owncord/server/auth" +) + +func TestHashPassword_DiffersFromPlaintext(t *testing.T) { + hash, err := auth.HashPassword("mypassword") + if err != nil { + t.Fatalf("HashPassword() error = %v", err) + } + if hash == "mypassword" { + t.Error("HashPassword() hash equals plaintext") + } +} + +func TestHashPassword_BcryptPrefix(t *testing.T) { + hash, err := auth.HashPassword("mypassword") + if err != nil { + t.Fatalf("HashPassword() error = %v", err) + } + if !strings.HasPrefix(hash, "$2") { + t.Errorf("HashPassword() = %q, want bcrypt prefix $2*", hash) + } +} + +func TestCheckPassword_CorrectPassword(t *testing.T) { + hash, err := auth.HashPassword("correctpassword") + if err != nil { + t.Fatalf("HashPassword() error = %v", err) + } + if !auth.CheckPassword(hash, "correctpassword") { + t.Error("CheckPassword() returned false for correct password") + } +} + +func TestCheckPassword_WrongPassword(t *testing.T) { + hash, err := auth.HashPassword("correctpassword") + if err != nil { + t.Fatalf("HashPassword() error = %v", err) + } + if auth.CheckPassword(hash, "wrongpassword") { + t.Error("CheckPassword() returned true for wrong password") + } +} + +func TestCheckPassword_EmptyPassword(t *testing.T) { + hash, err := auth.HashPassword("somepassword") + if err != nil { + t.Fatalf("HashPassword() error = %v", err) + } + if auth.CheckPassword(hash, "") { + t.Error("CheckPassword() returned true for empty password") + } +} + +func TestCheckPassword_EmptyHash(t *testing.T) { + if auth.CheckPassword("", "somepassword") { + t.Error("CheckPassword() returned true with empty hash") + } +} + +func TestValidatePasswordStrength_Valid(t *testing.T) { + cases := []string{ + "12345678", // exactly 8 chars + "abcdefghij", // 10 chars + strings.Repeat("a", 72), // exactly 72 chars (bcrypt max) + } + for _, pw := range cases { + if err := auth.ValidatePasswordStrength(pw); err != nil { + t.Errorf("ValidatePasswordStrength(%q) error = %v, want nil", pw, err) + } + } +} + +func TestValidatePasswordStrength_TooShort(t *testing.T) { + cases := []string{ + "", // empty + "1234567", // 7 chars + "abc", // 3 chars + } + for _, pw := range cases { + if err := auth.ValidatePasswordStrength(pw); err == nil { + t.Errorf("ValidatePasswordStrength(%q) error = nil, want error", pw) + } + } +} + +func TestValidatePasswordStrength_TooLong(t *testing.T) { + pw := strings.Repeat("a", 73) // 73 chars — over bcrypt 72 byte limit + if err := auth.ValidatePasswordStrength(pw); err == nil { + t.Errorf("ValidatePasswordStrength(%q) error = nil, want error for >72 chars", pw) + } +} + +func TestHashPassword_TwoCallsDifferentHashes(t *testing.T) { + // bcrypt includes a random salt + h1, _ := auth.HashPassword("password") + h2, _ := auth.HashPassword("password") + if h1 == h2 { + t.Error("HashPassword() produced identical hashes for the same password (salt missing?)") + } +} diff --git a/Server/auth/ratelimit.go b/Server/auth/ratelimit.go new file mode 100644 index 00000000..5af5a7ec --- /dev/null +++ b/Server/auth/ratelimit.go @@ -0,0 +1,104 @@ +package auth + +import ( + "sync" + "time" +) + +// entry records individual request timestamps for sliding-window limiting. +type entry struct { + timestamps []time.Time +} + +// lockoutEntry records when a lockout expires. +type lockoutEntry struct { + expiresAt time.Time +} + +// RateLimiter is an in-memory, thread-safe sliding-window rate limiter with +// optional IP lockout support. +type RateLimiter struct { + mu sync.Mutex + windows map[string]*entry + lockouts map[string]*lockoutEntry +} + +// NewRateLimiter returns an initialised RateLimiter. +func NewRateLimiter() *RateLimiter { + return &RateLimiter{ + windows: make(map[string]*entry), + lockouts: make(map[string]*lockoutEntry), + } +} + +// Allow reports whether a request from key is permitted given the limit and +// window. It records the current request timestamp regardless of the outcome. +// Returns false when key is locked out or has exceeded limit within window. +func (r *RateLimiter) Allow(key string, limit int, window time.Duration) bool { + r.mu.Lock() + defer r.mu.Unlock() + + // Lockout takes priority. + if lo, ok := r.lockouts[key]; ok { + if time.Now().Before(lo.expiresAt) { + return false + } + delete(r.lockouts, key) + } + + now := time.Now() + cutoff := now.Add(-window) + + e, ok := r.windows[key] + if !ok { + e = &entry{} + r.windows[key] = e + } + + // Prune timestamps outside the current window. + valid := e.timestamps[:0] + for _, ts := range e.timestamps { + if ts.After(cutoff) { + valid = append(valid, ts) + } + } + e.timestamps = valid + + if len(e.timestamps) >= limit { + return false + } + + e.timestamps = append(e.timestamps, now) + return true +} + +// Lockout prevents any requests from key for duration regardless of the +// sliding-window counter. +func (r *RateLimiter) Lockout(key string, duration time.Duration) { + r.mu.Lock() + defer r.mu.Unlock() + r.lockouts[key] = &lockoutEntry{expiresAt: time.Now().Add(duration)} +} + +// IsLockedOut reports whether key is currently under a lockout. +func (r *RateLimiter) IsLockedOut(key string) bool { + r.mu.Lock() + defer r.mu.Unlock() + lo, ok := r.lockouts[key] + if !ok { + return false + } + if time.Now().Before(lo.expiresAt) { + return true + } + delete(r.lockouts, key) + return false +} + +// Reset clears all rate-limit state (timestamps and lockout) for key. +func (r *RateLimiter) Reset(key string) { + r.mu.Lock() + defer r.mu.Unlock() + delete(r.windows, key) + delete(r.lockouts, key) +} diff --git a/Server/auth/ratelimit_test.go b/Server/auth/ratelimit_test.go new file mode 100644 index 00000000..ca7b7e31 --- /dev/null +++ b/Server/auth/ratelimit_test.go @@ -0,0 +1,126 @@ +package auth_test + +import ( + "testing" + "time" + + "github.com/owncord/server/auth" +) + +func TestRateLimiter_UnderLimitAllowed(t *testing.T) { + rl := auth.NewRateLimiter() + for i := 0; i < 5; i++ { + if !rl.Allow("key1", 5, time.Second) { + t.Errorf("Allow() = false at iteration %d, want true", i) + } + } +} + +func TestRateLimiter_AtLimitAllowed(t *testing.T) { + rl := auth.NewRateLimiter() + // Allow up to exactly the limit + for i := 0; i < 3; i++ { + rl.Allow("keyA", 3, time.Second) + } + // The 4th call should be blocked + if rl.Allow("keyA", 3, time.Second) { + t.Error("Allow() = true after limit exceeded, want false") + } +} + +func TestRateLimiter_OverLimitBlocked(t *testing.T) { + rl := auth.NewRateLimiter() + limit := 3 + for i := 0; i < limit; i++ { + rl.Allow("key2", limit, time.Second) + } + if rl.Allow("key2", limit, time.Second) { + t.Error("Allow() = true when over limit, want false") + } +} + +func TestRateLimiter_WindowExpiryResets(t *testing.T) { + rl := auth.NewRateLimiter() + window := 50 * time.Millisecond + limit := 2 + // Exhaust limit + rl.Allow("key3", limit, window) + rl.Allow("key3", limit, window) + if rl.Allow("key3", limit, window) { + t.Error("Allow() should be blocked after exhausting limit") + } + // Wait for window to expire + time.Sleep(window + 10*time.Millisecond) + if !rl.Allow("key3", limit, window) { + t.Error("Allow() should be permitted after window expires") + } +} + +func TestRateLimiter_DifferentKeysIndependent(t *testing.T) { + rl := auth.NewRateLimiter() + for i := 0; i < 5; i++ { + rl.Allow("keyX", 3, time.Second) + } + // keyY should still be allowed + if !rl.Allow("keyY", 3, time.Second) { + t.Error("Allow() blocked keyY even though only keyX exceeded limit") + } +} + +func TestRateLimiter_LockoutEnforced(t *testing.T) { + rl := auth.NewRateLimiter() + rl.Lockout("keyLock", time.Hour) + if !rl.IsLockedOut("keyLock") { + t.Error("IsLockedOut() = false after Lockout(), want true") + } +} + +func TestRateLimiter_LockoutExpires(t *testing.T) { + rl := auth.NewRateLimiter() + rl.Lockout("keyExp", 30*time.Millisecond) + time.Sleep(50 * time.Millisecond) + if rl.IsLockedOut("keyExp") { + t.Error("IsLockedOut() = true after lockout expired, want false") + } +} + +func TestRateLimiter_IsLockedOut_UnknownKey(t *testing.T) { + rl := auth.NewRateLimiter() + if rl.IsLockedOut("unknown") { + t.Error("IsLockedOut() = true for unknown key, want false") + } +} + +func TestRateLimiter_Reset(t *testing.T) { + rl := auth.NewRateLimiter() + rl.Allow("keyR", 1, time.Second) + rl.Allow("keyR", 1, time.Second) // now blocked + rl.Reset("keyR") + if !rl.Allow("keyR", 1, time.Second) { + t.Error("Allow() = false after Reset(), want true") + } +} + +func TestRateLimiter_LockoutBlocksAllow(t *testing.T) { + rl := auth.NewRateLimiter() + rl.Lockout("keyLB", time.Hour) + // Even under normal limit, lockout should block + if rl.Allow("keyLB", 100, time.Second) { + t.Error("Allow() = true for locked-out key, want false") + } +} + +func TestRateLimiter_ThreadSafe(t *testing.T) { + rl := auth.NewRateLimiter() + done := make(chan struct{}, 100) + for i := 0; i < 100; i++ { + go func() { + rl.Allow("concurrent", 50, time.Second) + done <- struct{}{} + }() + } + for i := 0; i < 100; i++ { + <-done + } + // If we get here without a race condition data race, we pass +} diff --git a/Server/auth/session.go b/Server/auth/session.go new file mode 100644 index 00000000..19e3cedb --- /dev/null +++ b/Server/auth/session.go @@ -0,0 +1,24 @@ +package auth + +import ( + "crypto/rand" + "crypto/sha256" + "encoding/hex" +) + +// GenerateToken returns a cryptographically random 256-bit token encoded as a +// 64-character lowercase hex string. +func GenerateToken() (string, error) { + raw := make([]byte, 32) // 256 bits + if _, err := rand.Read(raw); err != nil { + return "", err + } + return hex.EncodeToString(raw), nil +} + +// HashToken returns the SHA-256 hex digest of token. Store this hash in the +// database; never store the plaintext token. +func HashToken(token string) string { + sum := sha256.Sum256([]byte(token)) + return hex.EncodeToString(sum[:]) +} diff --git a/Server/auth/session_test.go b/Server/auth/session_test.go new file mode 100644 index 00000000..0b5865af --- /dev/null +++ b/Server/auth/session_test.go @@ -0,0 +1,77 @@ +package auth_test + +import ( + "testing" + + "github.com/owncord/server/auth" +) + +func TestGenerateToken_Length(t *testing.T) { + token, err := auth.GenerateToken() + if err != nil { + t.Fatalf("GenerateToken() error = %v", err) + } + if len(token) != 64 { + t.Errorf("GenerateToken() len = %d, want 64", len(token)) + } +} + +func TestGenerateToken_HexCharacters(t *testing.T) { + token, err := auth.GenerateToken() + if err != nil { + t.Fatalf("GenerateToken() error = %v", err) + } + for i, c := range token { + if !((c >= '0' && c <= '9') || (c >= 'a' && c <= 'f')) { + t.Errorf("GenerateToken() char[%d] = %q, not lowercase hex", i, c) + } + } +} + +func TestGenerateToken_Uniqueness(t *testing.T) { + const n = 1000 + seen := make(map[string]struct{}, n) + for i := 0; i < n; i++ { + tok, err := auth.GenerateToken() + if err != nil { + t.Fatalf("GenerateToken() iteration %d error = %v", i, err) + } + if _, dup := seen[tok]; dup { + t.Fatalf("GenerateToken() produced duplicate token at iteration %d", i) + } + seen[tok] = struct{}{} + } +} + +func TestHashToken_Deterministic(t *testing.T) { + token := "abc123" + h1 := auth.HashToken(token) + h2 := auth.HashToken(token) + if h1 != h2 { + t.Errorf("HashToken() not deterministic: %q != %q", h1, h2) + } +} + +func TestHashToken_DiffersFromPlaintext(t *testing.T) { + token := "abc123" + hash := auth.HashToken(token) + if hash == token { + t.Errorf("HashToken() hash equals plaintext token") + } +} + +func TestHashToken_Length(t *testing.T) { + // SHA-256 hex = 64 chars + hash := auth.HashToken("any-token") + if len(hash) != 64 { + t.Errorf("HashToken() len = %d, want 64", len(hash)) + } +} + +func TestHashToken_DifferentInputsDifferentHashes(t *testing.T) { + h1 := auth.HashToken("token-one") + h2 := auth.HashToken("token-two") + if h1 == h2 { + t.Errorf("HashToken() same hash for different inputs") + } +} diff --git a/Server/db/auth_queries.go b/Server/db/auth_queries.go new file mode 100644 index 00000000..5e8eb760 --- /dev/null +++ b/Server/db/auth_queries.go @@ -0,0 +1,285 @@ +package db + +import ( + "crypto/rand" + "database/sql" + "encoding/hex" + "errors" + "fmt" + "time" +) + +// ─── User Operations ────────────────────────────────────────────────────────── + +// CreateUser inserts a new user record and returns the assigned ID. +func (d *DB) CreateUser(username, passwordHash string, roleID int) (int64, error) { + res, err := d.sqlDB.Exec( + `INSERT INTO users (username, password, role_id) VALUES (?, ?, ?)`, + username, passwordHash, roleID, + ) + if err != nil { + return 0, fmt.Errorf("CreateUser: %w", err) + } + return res.LastInsertId() +} + +// GetUserByUsername returns the user with the given username (case-insensitive), +// or nil if not found. +func (d *DB) GetUserByUsername(username string) (*User, error) { + row := d.sqlDB.QueryRow( + `SELECT id, username, password, avatar, role_id, totp_secret, status, + created_at, last_seen, banned, ban_reason, ban_expires + FROM users WHERE username = ? COLLATE NOCASE`, + username, + ) + return scanUser(row) +} + +// GetUserByID returns the user with the given ID, or nil if not found. +func (d *DB) GetUserByID(id int64) (*User, error) { + row := d.sqlDB.QueryRow( + `SELECT id, username, password, avatar, role_id, totp_secret, status, + created_at, last_seen, banned, ban_reason, ban_expires + FROM users WHERE id = ?`, + id, + ) + return scanUser(row) +} + +// scanUser reads a User from a *sql.Row, returning nil (not an error) when the +// row is not found. +func scanUser(row *sql.Row) (*User, error) { + u := &User{} + var banned int + err := row.Scan( + &u.ID, &u.Username, &u.PasswordHash, &u.Avatar, &u.RoleID, + &u.TOTPSecret, &u.Status, &u.CreatedAt, &u.LastSeen, + &banned, &u.BanReason, &u.BanExpires, + ) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + if err != nil { + return nil, fmt.Errorf("scanUser: %w", err) + } + u.Banned = banned != 0 + return u, nil +} + +// UpdateUserStatus sets the status column for the given user ID. +func (d *DB) UpdateUserStatus(id int64, status string) error { + _, err := d.sqlDB.Exec( + `UPDATE users SET status = ?, last_seen = datetime('now') WHERE id = ?`, + status, id, + ) + if err != nil { + return fmt.Errorf("UpdateUserStatus: %w", err) + } + return nil +} + +// BanUser marks a user as banned with an optional expiry. Pass nil for a +// permanent ban. +func (d *DB) BanUser(id int64, reason string, expires *time.Time) error { + var expiresStr *string + if expires != nil { + s := expires.UTC().Format("2006-01-02T15:04:05Z") + expiresStr = &s + } + _, err := d.sqlDB.Exec( + `UPDATE users SET banned = 1, ban_reason = ?, ban_expires = ? WHERE id = ?`, + reason, expiresStr, id, + ) + if err != nil { + return fmt.Errorf("BanUser: %w", err) + } + return nil +} + +// ─── Session Operations ─────────────────────────────────────────────────────── + +// CreateSession inserts a new session and returns the session ID. +// tokenHash must already be hashed (never store plaintext tokens). +func (d *DB) CreateSession(userID int64, tokenHash, device, ip string) (int64, error) { + expiresAt := time.Now().Add(sessionTTL).UTC().Format("2006-01-02T15:04:05Z") + res, err := d.sqlDB.Exec( + `INSERT INTO sessions (user_id, token, device, ip_address, expires_at) + VALUES (?, ?, ?, ?, ?)`, + userID, tokenHash, device, ip, expiresAt, + ) + if err != nil { + return 0, fmt.Errorf("CreateSession: %w", err) + } + return res.LastInsertId() +} + +// GetSessionByTokenHash retrieves a session by its hashed token, or nil if +// not found. +func (d *DB) GetSessionByTokenHash(tokenHash string) (*Session, error) { + row := d.sqlDB.QueryRow( + `SELECT id, user_id, token, device, ip_address, created_at, last_used, expires_at + FROM sessions WHERE token = ?`, + tokenHash, + ) + s := &Session{} + err := row.Scan( + &s.ID, &s.UserID, &s.TokenHash, &s.Device, &s.IP, + &s.CreatedAt, &s.LastUsed, &s.ExpiresAt, + ) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + if err != nil { + return nil, fmt.Errorf("GetSessionByTokenHash: %w", err) + } + return s, nil +} + +// DeleteSession removes the session with the given token hash. +func (d *DB) DeleteSession(tokenHash string) error { + _, err := d.sqlDB.Exec(`DELETE FROM sessions WHERE token = ?`, tokenHash) + if err != nil { + return fmt.Errorf("DeleteSession: %w", err) + } + return nil +} + +// DeleteExpiredSessions removes all sessions whose expires_at is in the past. +// Compares using strftime to handle both ISO-8601 and SQLite datetime formats. +func (d *DB) DeleteExpiredSessions() error { + _, err := d.sqlDB.Exec( + `DELETE FROM sessions WHERE strftime('%s', expires_at) < strftime('%s', 'now')`, + ) + if err != nil { + return fmt.Errorf("DeleteExpiredSessions: %w", err) + } + return nil +} + +// TouchSession updates last_used for the session with the given token hash. +func (d *DB) TouchSession(tokenHash string) error { + _, err := d.sqlDB.Exec( + `UPDATE sessions SET last_used = datetime('now') WHERE token = ?`, + tokenHash, + ) + if err != nil { + return fmt.Errorf("TouchSession: %w", err) + } + return nil +} + +// ─── Invite Operations ──────────────────────────────────────────────────────── + +// CreateInvite generates a random invite code, persists it, and returns the +// code. maxUses=0 means unlimited. expiresAt=nil means never expires. +func (d *DB) CreateInvite(createdBy int64, maxUses int, expiresAt *time.Time) (string, error) { + code, err := generateInviteCode() + if err != nil { + return "", fmt.Errorf("CreateInvite generate code: %w", err) + } + + var maxUsesVal *int + if maxUses > 0 { + maxUsesVal = &maxUses + } + var expiresStr *string + if expiresAt != nil { + s := expiresAt.UTC().Format("2006-01-02T15:04:05Z") + expiresStr = &s + } + + _, err = d.sqlDB.Exec( + `INSERT INTO invites (code, created_by, max_uses, expires_at) VALUES (?, ?, ?, ?)`, + code, createdBy, maxUsesVal, expiresStr, + ) + if err != nil { + return "", fmt.Errorf("CreateInvite insert: %w", err) + } + return code, nil +} + +// GetInvite returns the invite for the given code, or nil if not found. +func (d *DB) GetInvite(code string) (*Invite, error) { + row := d.sqlDB.QueryRow( + `SELECT id, code, created_by, max_uses, use_count, expires_at, revoked, created_at + FROM invites WHERE code = ?`, + code, + ) + inv := &Invite{} + var revoked int + err := row.Scan( + &inv.ID, &inv.Code, &inv.CreatedBy, &inv.MaxUses, + &inv.Uses, &inv.ExpiresAt, &revoked, &inv.CreatedAt, + ) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + if err != nil { + return nil, fmt.Errorf("GetInvite: %w", err) + } + inv.Revoked = revoked != 0 + return inv, nil +} + +// UseInvite increments use_count after validating the invite is usable. +// Returns an error if the invite is revoked, expired, or has reached max uses. +func (d *DB) UseInvite(code string) error { + inv, err := d.GetInvite(code) + if err != nil { + return err + } + if inv == nil { + return errors.New("invite not found") + } + if inv.Revoked { + return errors.New("invite has been revoked") + } + if inv.ExpiresAt != nil { + // Try both SQLite datetime format and ISO-8601 format. + var expires time.Time + var parseErr error + for _, layout := range []string{"2006-01-02 15:04:05", "2006-01-02T15:04:05Z"} { + expires, parseErr = time.Parse(layout, *inv.ExpiresAt) + if parseErr == nil { + break + } + } + if parseErr != nil { + return fmt.Errorf("parsing invite expiry: %w", parseErr) + } + if time.Now().UTC().After(expires) { + return errors.New("invite has expired") + } + } + if inv.MaxUses != nil && inv.Uses >= *inv.MaxUses { + return errors.New("invite has reached its maximum uses") + } + _, err = d.sqlDB.Exec( + `UPDATE invites SET use_count = use_count + 1 WHERE code = ?`, + code, + ) + if err != nil { + return fmt.Errorf("UseInvite update: %w", err) + } + return nil +} + +// RevokeInvite marks an invite as revoked. +func (d *DB) RevokeInvite(code string) error { + _, err := d.sqlDB.Exec(`UPDATE invites SET revoked = 1 WHERE code = ?`, code) + if err != nil { + return fmt.Errorf("RevokeInvite: %w", err) + } + return nil +} + +// ─── Helpers ────────────────────────────────────────────────────────────────── + +// generateInviteCode produces a random 8-byte (16-char hex) code. +func generateInviteCode() (string, error) { + b := make([]byte, 8) + if _, err := rand.Read(b); err != nil { + return "", err + } + return hex.EncodeToString(b), nil +} diff --git a/Server/db/auth_queries_test.go b/Server/db/auth_queries_test.go new file mode 100644 index 00000000..a0225d26 --- /dev/null +++ b/Server/db/auth_queries_test.go @@ -0,0 +1,468 @@ +package db_test + +import ( + "testing" + "testing/fstest" + "time" + + "github.com/owncord/server/db" +) + +// newTestDB opens an in-memory SQLite database and runs migrations from the +// embedded FS so tests are fully self-contained. +func newTestDB(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() }) + + // Build a minimal migration FS with the initial schema. + migrFS := fstest.MapFS{ + "001_schema.sql": {Data: testSchema}, + } + if err := db.MigrateFS(database, migrFS); err != nil { + t.Fatalf("MigrateFS: %v", err) + } + return database +} + +// testSchema mirrors the production migration but kept inline so tests are +// portable and don't depend on the real migrations embed. +var testSchema = []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, 'Moderator', '#3498DB', 1048575, 60, 0), + (4, '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 4 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 invites ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + code TEXT NOT NULL UNIQUE, + created_by INTEGER NOT NULL REFERENCES users(id), + redeemed_by INTEGER 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 INDEX IF NOT EXISTS idx_invites_code ON invites(code); +`) + +// ─── User tests ────────────────────────────────────────────────────────────── + +func TestCreateUser_Success(t *testing.T) { + database := newTestDB(t) + id, err := database.CreateUser("alice", "hash123", 4) + if err != nil { + t.Fatalf("CreateUser: %v", err) + } + if id <= 0 { + t.Errorf("CreateUser returned id = %d, want > 0", id) + } +} + +func TestCreateUser_DuplicateUsername(t *testing.T) { + database := newTestDB(t) + if _, err := database.CreateUser("bob", "hash1", 4); err != nil { + t.Fatalf("first CreateUser: %v", err) + } + _, err := database.CreateUser("bob", "hash2", 4) + if err == nil { + t.Error("CreateUser() with duplicate username returned nil error, want error") + } +} + +func TestCreateUser_CaseInsensitiveDuplicate(t *testing.T) { + database := newTestDB(t) + if _, err := database.CreateUser("Charlie", "hash1", 4); err != nil { + t.Fatalf("first CreateUser: %v", err) + } + _, err := database.CreateUser("charlie", "hash2", 4) + if err == nil { + t.Error("CreateUser() with case-insensitive duplicate returned nil error, want error") + } +} + +func TestGetUserByUsername_Found(t *testing.T) { + database := newTestDB(t) + database.CreateUser("dave", "hashDave", 4) + + user, err := database.GetUserByUsername("dave") + if err != nil { + t.Fatalf("GetUserByUsername: %v", err) + } + if user.Username != "dave" { + t.Errorf("Username = %q, want %q", user.Username, "dave") + } + if user.PasswordHash != "hashDave" { + t.Errorf("PasswordHash = %q, want %q", user.PasswordHash, "hashDave") + } +} + +func TestGetUserByUsername_CaseInsensitive(t *testing.T) { + database := newTestDB(t) + database.CreateUser("Eve", "hashEve", 4) + + user, err := database.GetUserByUsername("EVE") + if err != nil { + t.Fatalf("GetUserByUsername case-insensitive: %v", err) + } + if user == nil { + t.Fatal("GetUserByUsername returned nil for case-insensitive match") + } +} + +func TestGetUserByUsername_NotFound(t *testing.T) { + database := newTestDB(t) + user, err := database.GetUserByUsername("nobody") + if err != nil { + t.Fatalf("GetUserByUsername(not found): %v", err) + } + if user != nil { + t.Error("GetUserByUsername returned non-nil for missing user") + } +} + +func TestGetUserByID_Found(t *testing.T) { + database := newTestDB(t) + id, _ := database.CreateUser("frank", "hashFrank", 4) + + user, err := database.GetUserByID(id) + if err != nil { + t.Fatalf("GetUserByID: %v", err) + } + if user.ID != id { + t.Errorf("ID = %d, want %d", user.ID, id) + } +} + +func TestGetUserByID_NotFound(t *testing.T) { + database := newTestDB(t) + user, err := database.GetUserByID(999) + if err != nil { + t.Fatalf("GetUserByID(not found): %v", err) + } + if user != nil { + t.Error("GetUserByID returned non-nil for missing user") + } +} + +func TestUpdateUserStatus(t *testing.T) { + database := newTestDB(t) + id, _ := database.CreateUser("grace", "hash", 4) + + if err := database.UpdateUserStatus(id, "online"); err != nil { + t.Fatalf("UpdateUserStatus: %v", err) + } + user, _ := database.GetUserByID(id) + if user.Status != "online" { + t.Errorf("Status = %q, want %q", user.Status, "online") + } +} + +func TestBanUser_Permanent(t *testing.T) { + database := newTestDB(t) + id, _ := database.CreateUser("hank", "hash", 4) + + if err := database.BanUser(id, "spam", nil); err != nil { + t.Fatalf("BanUser: %v", err) + } + user, _ := database.GetUserByID(id) + if !user.Banned { + t.Error("Banned = false after BanUser, want true") + } + if user.BanExpires != nil { + t.Errorf("BanExpires = %v, want nil for permanent ban", user.BanExpires) + } +} + +func TestBanUser_Temporary(t *testing.T) { + database := newTestDB(t) + id, _ := database.CreateUser("ivan", "hash", 4) + expires := time.Now().Add(24 * time.Hour) + + if err := database.BanUser(id, "temp ban", &expires); err != nil { + t.Fatalf("BanUser (temp): %v", err) + } + user, _ := database.GetUserByID(id) + if !user.Banned { + t.Error("Banned = false after temp ban") + } + if user.BanExpires == nil { + t.Error("BanExpires = nil for temp ban, want non-nil") + } +} + +// ─── Session tests ──────────────────────────────────────────────────────────── + +func TestCreateSession_Success(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("jack", "hash", 4) + + id, err := database.CreateSession(uid, "tokenHash1", "GoTest/1.0", "127.0.0.1") + if err != nil { + t.Fatalf("CreateSession: %v", err) + } + if id <= 0 { + t.Errorf("CreateSession id = %d, want > 0", id) + } +} + +func TestGetSessionByTokenHash_Found(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("kate", "hash", 4) + database.CreateSession(uid, "myTokenHash", "GoTest/1.0", "127.0.0.1") + + sess, err := database.GetSessionByTokenHash("myTokenHash") + if err != nil { + t.Fatalf("GetSessionByTokenHash: %v", err) + } + if sess == nil { + t.Fatal("GetSessionByTokenHash returned nil for existing session") + } + if sess.UserID != uid { + t.Errorf("UserID = %d, want %d", sess.UserID, uid) + } +} + +func TestGetSessionByTokenHash_NotFound(t *testing.T) { + database := newTestDB(t) + sess, err := database.GetSessionByTokenHash("nonexistent") + if err != nil { + t.Fatalf("GetSessionByTokenHash(not found): %v", err) + } + if sess != nil { + t.Error("GetSessionByTokenHash returned non-nil for missing session") + } +} + +func TestDeleteSession(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("leo", "hash", 4) + database.CreateSession(uid, "delToken", "GoTest/1.0", "127.0.0.1") + + if err := database.DeleteSession("delToken"); err != nil { + t.Fatalf("DeleteSession: %v", err) + } + sess, _ := database.GetSessionByTokenHash("delToken") + if sess != nil { + t.Error("Session still exists after DeleteSession") + } +} + +func TestDeleteExpiredSessions(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("mia", "hash", 4) + + // Insert an already-expired session directly via Exec. + // Use SQLite datetime format (space separator) to match what datetime('now') produces. + pastTime := time.Now().Add(-time.Hour).UTC().Format("2006-01-02 15:04:05") + _, err := database.Exec( + `INSERT INTO sessions (user_id, token, device, ip_address, expires_at) VALUES (?, ?, ?, ?, ?)`, + uid, "expiredToken", "test", "127.0.0.1", pastTime, + ) + if err != nil { + t.Fatalf("inserting expired session: %v", err) + } + + // Insert a valid session through the normal path. + database.CreateSession(uid, "validToken", "GoTest/1.0", "127.0.0.1") + + if err := database.DeleteExpiredSessions(); err != nil { + t.Fatalf("DeleteExpiredSessions: %v", err) + } + + expired, _ := database.GetSessionByTokenHash("expiredToken") + if expired != nil { + t.Error("Expired session still exists after DeleteExpiredSessions") + } + valid, _ := database.GetSessionByTokenHash("validToken") + if valid == nil { + t.Error("Valid session was deleted by DeleteExpiredSessions") + } +} + +func TestTouchSession(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("noah", "hash", 4) + database.CreateSession(uid, "touchToken", "GoTest/1.0", "127.0.0.1") + + sess1, _ := database.GetSessionByTokenHash("touchToken") + time.Sleep(2 * time.Millisecond) + + if err := database.TouchSession("touchToken"); err != nil { + t.Fatalf("TouchSession: %v", err) + } + + sess2, _ := database.GetSessionByTokenHash("touchToken") + if sess1.LastUsed == sess2.LastUsed { + // last_used should have advanced; if they're equal the touch had no effect + // (This can be flaky at millisecond resolution, but is a reasonable sanity check.) + t.Log("TouchSession: last_used unchanged (may be a timing issue on fast machines)") + } +} + +// ─── Invite tests ───────────────────────────────────────────────────────────── + +func TestCreateInvite_Success(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("olivia", "hash", 4) + + code, err := database.CreateInvite(uid, 0, nil) + if err != nil { + t.Fatalf("CreateInvite: %v", err) + } + if len(code) == 0 { + t.Error("CreateInvite returned empty code") + } +} + +func TestGetInvite_Found(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("pedro", "hash", 4) + code, _ := database.CreateInvite(uid, 5, nil) + + inv, err := database.GetInvite(code) + if err != nil { + t.Fatalf("GetInvite: %v", err) + } + if inv == nil { + t.Fatal("GetInvite returned nil for existing code") + } + if inv.Code != code { + t.Errorf("Code = %q, want %q", inv.Code, code) + } + if inv.MaxUses == nil || *inv.MaxUses != 5 { + t.Errorf("MaxUses = %v, want 5", inv.MaxUses) + } +} + +func TestGetInvite_NotFound(t *testing.T) { + database := newTestDB(t) + inv, err := database.GetInvite("bogus") + if err != nil { + t.Fatalf("GetInvite(not found): %v", err) + } + if inv != nil { + t.Error("GetInvite returned non-nil for missing code") + } +} + +func TestUseInvite_IncrementsUses(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("quinn", "hash", 4) + code, _ := database.CreateInvite(uid, 5, nil) + + if err := database.UseInvite(code); err != nil { + t.Fatalf("UseInvite: %v", err) + } + + inv, _ := database.GetInvite(code) + if inv.Uses != 1 { + t.Errorf("Uses = %d, want 1", inv.Uses) + } +} + +func TestUseInvite_ExceedsMaxUses(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("rachel", "hash", 4) + code, _ := database.CreateInvite(uid, 1, nil) + + if err := database.UseInvite(code); err != nil { + t.Fatalf("first UseInvite: %v", err) + } + // Second use should fail + if err := database.UseInvite(code); err == nil { + t.Error("UseInvite() returned nil error after exceeding max_uses") + } +} + +func TestUseInvite_Revoked(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("sam", "hash", 4) + code, _ := database.CreateInvite(uid, 0, nil) + + database.RevokeInvite(code) + if err := database.UseInvite(code); err == nil { + t.Error("UseInvite() returned nil error for revoked invite") + } +} + +func TestUseInvite_Expired(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("tina", "hash", 4) + + past := time.Now().Add(-time.Hour) + code, _ := database.CreateInvite(uid, 0, &past) + + if err := database.UseInvite(code); err == nil { + t.Error("UseInvite() returned nil error for expired invite") + } +} + +func TestRevokeInvite(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("uma", "hash", 4) + code, _ := database.CreateInvite(uid, 0, nil) + + if err := database.RevokeInvite(code); err != nil { + t.Fatalf("RevokeInvite: %v", err) + } + + inv, _ := database.GetInvite(code) + if !inv.Revoked { + t.Error("Revoked = false after RevokeInvite, want true") + } +} + +func TestCreateInvite_UnlimitedUses(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("vera", "hash", 4) + code, _ := database.CreateInvite(uid, 0, nil) // 0 = unlimited + + inv, _ := database.GetInvite(code) + if inv.MaxUses != nil { + t.Errorf("MaxUses = %v, want nil for unlimited", inv.MaxUses) + } +} diff --git a/Server/db/invite_queries.go b/Server/db/invite_queries.go new file mode 100644 index 00000000..1d4afdde --- /dev/null +++ b/Server/db/invite_queries.go @@ -0,0 +1,30 @@ +package db + +import "fmt" + +// ListInvites returns all invites ordered by creation time descending. +func (d *DB) ListInvites() ([]*Invite, error) { + rows, err := d.sqlDB.Query( + `SELECT id, code, created_by, max_uses, use_count, expires_at, revoked, created_at + FROM invites ORDER BY created_at DESC`, + ) + if err != nil { + return nil, fmt.Errorf("ListInvites: %w", err) + } + defer rows.Close() + + var invites []*Invite + for rows.Next() { + inv := &Invite{} + var revoked int + if err := rows.Scan( + &inv.ID, &inv.Code, &inv.CreatedBy, &inv.MaxUses, + &inv.Uses, &inv.ExpiresAt, &revoked, &inv.CreatedAt, + ); err != nil { + return nil, fmt.Errorf("ListInvites scan: %w", err) + } + inv.Revoked = revoked != 0 + invites = append(invites, inv) + } + return invites, rows.Err() +} diff --git a/Server/db/models.go b/Server/db/models.go new file mode 100644 index 00000000..6429317c --- /dev/null +++ b/Server/db/models.go @@ -0,0 +1,56 @@ +package db + +import "time" + +// User represents a row in the users table. +type User struct { + ID int64 + Username string + PasswordHash string + Avatar *string + RoleID int64 + TOTPSecret *string + Status string + CreatedAt string + LastSeen *string + Banned bool + BanReason *string + BanExpires *string +} + +// Session represents a row in the sessions table. +type Session struct { + ID int64 + UserID int64 + TokenHash string + Device string + IP string + CreatedAt string + LastUsed string + ExpiresAt string +} + +// Invite represents a row in the invites table. +type Invite struct { + ID int64 + Code string + CreatedBy int64 + Uses int + MaxUses *int + ExpiresAt *string + Revoked bool + CreatedAt string +} + +// Role represents a row in the roles table. +type Role struct { + ID int64 + Name string + Color *string + Permissions int64 + Position int + IsDefault bool +} + +// sessionTTL is the duration a session remains valid after creation. +const sessionTTL = 30 * 24 * time.Hour diff --git a/Server/db/role_invite_queries_test.go b/Server/db/role_invite_queries_test.go new file mode 100644 index 00000000..f93e5053 --- /dev/null +++ b/Server/db/role_invite_queries_test.go @@ -0,0 +1,156 @@ +package db_test + +import ( + "testing" +) + +// ─── GetRoleByID tests ──────────────────────────────────────────────────────── + +func TestGetRoleByID_Found(t *testing.T) { + database := newTestDB(t) + + role, err := database.GetRoleByID(4) // Member — inserted by migration + if err != nil { + t.Fatalf("GetRoleByID: %v", err) + } + if role == nil { + t.Fatal("GetRoleByID returned nil for Member role") + } + if role.Name != "Member" { + t.Errorf("Name = %q, want %q", role.Name, "Member") + } + if role.Permissions == 0 { + t.Error("Member permissions = 0, want non-zero") + } +} + +func TestGetRoleByID_NotFound(t *testing.T) { + database := newTestDB(t) + + role, err := database.GetRoleByID(9999) + if err != nil { + t.Fatalf("GetRoleByID(not found): %v", err) + } + if role != nil { + t.Error("GetRoleByID returned non-nil for missing role") + } +} + +func TestGetRoleByID_OwnerHasAllPermissions(t *testing.T) { + database := newTestDB(t) + + role, err := database.GetRoleByID(1) // Owner + if err != nil { + t.Fatalf("GetRoleByID Owner: %v", err) + } + if role == nil { + t.Fatal("GetRoleByID returned nil for Owner role") + } + // Owner has permissions = 0x7FFFFFFF = 2147483647 + if role.Permissions != 2147483647 { + t.Errorf("Owner Permissions = %d, want 2147483647", role.Permissions) + } +} + +func TestGetRoleByID_IsDefaultField(t *testing.T) { + database := newTestDB(t) + + owner, _ := database.GetRoleByID(1) + member, _ := database.GetRoleByID(4) + + if owner.IsDefault { + t.Error("Owner.IsDefault = true, want false") + } + // Member is the default role (is_default=1 in the migration). + if !member.IsDefault { + t.Error("Member.IsDefault = false, want true (Member is the default role for new users)") + } +} + +// ─── ListRoles tests ────────────────────────────────────────────────────────── + +func TestListRoles_ReturnsFourDefaultRoles(t *testing.T) { + database := newTestDB(t) + + roles, err := database.ListRoles() + if err != nil { + t.Fatalf("ListRoles: %v", err) + } + if len(roles) != 4 { + t.Errorf("ListRoles count = %d, want 4", len(roles)) + } +} + +func TestListRoles_OrderedByPositionDesc(t *testing.T) { + database := newTestDB(t) + + roles, err := database.ListRoles() + if err != nil { + t.Fatalf("ListRoles: %v", err) + } + + for i := 1; i < len(roles); i++ { + if roles[i].Position > roles[i-1].Position { + t.Errorf("ListRoles not ordered by position DESC: index %d (%d) > index %d (%d)", + i, roles[i].Position, i-1, roles[i-1].Position) + } + } +} + +// ─── ListInvites tests ──────────────────────────────────────────────────────── + +func TestListInvites_Empty(t *testing.T) { + database := newTestDB(t) + + invites, err := database.ListInvites() + if err != nil { + t.Fatalf("ListInvites empty: %v", err) + } + if len(invites) != 0 { + t.Errorf("ListInvites empty = %d items, want 0", len(invites)) + } +} + +func TestListInvites_Multiple(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("listowner", "hash", 4) + + database.CreateInvite(uid, 1, nil) + database.CreateInvite(uid, 5, nil) + database.CreateInvite(uid, 0, nil) + + invites, err := database.ListInvites() + if err != nil { + t.Fatalf("ListInvites multiple: %v", err) + } + if len(invites) != 3 { + t.Errorf("ListInvites count = %d, want 3", len(invites)) + } +} + +func TestListInvites_IncludesRevokedInvites(t *testing.T) { + database := newTestDB(t) + uid, _ := database.CreateUser("revokelistowner", "hash", 4) + + code, _ := database.CreateInvite(uid, 1, nil) + database.RevokeInvite(code) + database.CreateInvite(uid, 0, nil) // active + + invites, err := database.ListInvites() + if err != nil { + t.Fatalf("ListInvites with revoked: %v", err) + } + if len(invites) != 2 { + t.Errorf("ListInvites count = %d, want 2", len(invites)) + } + + var revokedCount int + for _, inv := range invites { + if inv.Revoked { + revokedCount++ + } + } + if revokedCount != 1 { + t.Errorf("ListInvites revoked count = %d, want 1", revokedCount) + } +} diff --git a/Server/db/role_queries.go b/Server/db/role_queries.go new file mode 100644 index 00000000..fad35d40 --- /dev/null +++ b/Server/db/role_queries.go @@ -0,0 +1,49 @@ +package db + +import ( + "database/sql" + "errors" + "fmt" +) + +// GetRoleByID returns the role with the given ID, or nil if not found. +func (d *DB) GetRoleByID(id int64) (*Role, error) { + row := d.sqlDB.QueryRow( + `SELECT id, name, color, permissions, position, is_default FROM roles WHERE id = ?`, + id, + ) + r := &Role{} + var isDefault int + err := row.Scan(&r.ID, &r.Name, &r.Color, &r.Permissions, &r.Position, &isDefault) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + if err != nil { + return nil, fmt.Errorf("GetRoleByID: %w", err) + } + r.IsDefault = isDefault != 0 + return r, nil +} + +// ListRoles returns all roles ordered by position descending. +func (d *DB) ListRoles() ([]*Role, error) { + rows, err := d.sqlDB.Query( + `SELECT id, name, color, permissions, position, is_default FROM roles ORDER BY position DESC`, + ) + if err != nil { + return nil, fmt.Errorf("ListRoles: %w", err) + } + defer rows.Close() + + var roles []*Role + for rows.Next() { + r := &Role{} + var isDefault int + if err := rows.Scan(&r.ID, &r.Name, &r.Color, &r.Permissions, &r.Position, &isDefault); err != nil { + return nil, fmt.Errorf("ListRoles scan: %w", err) + } + r.IsDefault = isDefault != 0 + roles = append(roles, r) + } + return roles, rows.Err() +} diff --git a/Server/go.mod b/Server/go.mod index a2240c41..67604841 100644 --- a/Server/go.mod +++ b/Server/go.mod @@ -3,12 +3,14 @@ module github.com/owncord/server go 1.25.0 require ( + github.com/aymerick/douceur v0.2.0 // indirect github.com/dustin/go-humanize v1.0.1 // indirect github.com/fatih/structs v1.1.0 // indirect github.com/fsnotify/fsnotify v1.9.0 // indirect github.com/go-chi/chi/v5 v5.2.5 // indirect github.com/go-viper/mapstructure/v2 v2.4.0 // indirect github.com/google/uuid v1.6.0 // indirect + github.com/gorilla/css v1.0.1 // indirect github.com/knadh/koanf/maps v0.1.2 // indirect github.com/knadh/koanf/parsers/yaml v1.1.0 // indirect github.com/knadh/koanf/providers/env v1.1.0 // indirect @@ -16,6 +18,7 @@ require ( github.com/knadh/koanf/providers/structs v1.0.0 // indirect github.com/knadh/koanf/v2 v2.3.3 // indirect github.com/mattn/go-isatty v0.0.20 // indirect + github.com/microcosm-cc/bluemonday v1.0.27 // indirect github.com/mitchellh/copystructure v1.2.0 // indirect github.com/mitchellh/reflectwalk v1.0.2 // indirect github.com/ncruces/go-strftime v1.0.0 // indirect @@ -23,6 +26,7 @@ require ( go.yaml.in/yaml/v3 v3.0.3 // indirect golang.org/x/crypto v0.49.0 // indirect golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect + golang.org/x/net v0.51.0 // indirect golang.org/x/sys v0.42.0 // indirect modernc.org/libc v1.67.6 // indirect modernc.org/mathutil v1.7.1 // indirect diff --git a/Server/go.sum b/Server/go.sum index 73a412ba..3f83a338 100644 --- a/Server/go.sum +++ b/Server/go.sum @@ -1,3 +1,5 @@ +github.com/aymerick/douceur v0.2.0 h1:Mv+mAeH1Q+n9Fr+oyamOlAkUNPWPlA8PPGR0QAaYuPk= +github.com/aymerick/douceur v0.2.0/go.mod h1:wlT5vV2O3h55X9m7iVYN0TBM0NH/MmbLnd30/FjWUq4= github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY= github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= github.com/fatih/structs v1.1.0 h1:Q7juDM0QtcnhCpeyLGQKyg4TOIghuNXrkL32pHAUMxo= @@ -10,6 +12,8 @@ github.com/go-viper/mapstructure/v2 v2.4.0 h1:EBsztssimR/CONLSZZ04E8qAkxNYq4Qp9L github.com/go-viper/mapstructure/v2 v2.4.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/gorilla/css v1.0.1 h1:ntNaBIghp6JmvWnxbZKANoLyuXTPZ4cAMlo6RyhlbO8= +github.com/gorilla/css v1.0.1/go.mod h1:BvnYkspnSzMmwRK+b8/xgNPLiIuNZr6vbZBTPQ2A3b0= github.com/knadh/koanf/maps v0.1.2 h1:RBfmAW5CnZT+PJ1CVc1QSJKf4Xu9kxfQgYVQSu8hpbo= github.com/knadh/koanf/maps v0.1.2/go.mod h1:npD/QZY3V6ghQDdcQzl1W4ICNVTkohC8E73eI2xW4yI= github.com/knadh/koanf/parsers/yaml v1.1.0 h1:3ltfm9ljprAHt4jxgeYLlFPmUaunuCgu1yILuTXRdM4= @@ -24,6 +28,8 @@ github.com/knadh/koanf/v2 v2.3.3 h1:jLJC8XCRfLC7n4F+ZKKdBsbq1bfXTpuFhf4L7t94D94= github.com/knadh/koanf/v2 v2.3.3/go.mod h1:gRb40VRAbd4iJMYYD5IxZ6hfuopFcXBpc9bbQpZwo28= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= +github.com/microcosm-cc/bluemonday v1.0.27 h1:MpEUotklkwCSLeH+Qdx1VJgNqLlpY2KXwXFM08ygZfk= +github.com/microcosm-cc/bluemonday v1.0.27/go.mod h1:jFi9vgW+H7c3V0lb6nR74Ib/DIB5OBs92Dimizgw2cA= github.com/mitchellh/copystructure v1.2.0 h1:vpKXTN4ewci03Vljg/q9QvCGUDttBOGBIa15WveJJGw= github.com/mitchellh/copystructure v1.2.0/go.mod h1:qLl+cE2AmVv+CoeAwDPye/v+N2HKCj9FbZEVFJRxO9s= github.com/mitchellh/reflectwalk v1.0.2 h1:G2LzWKi524PWgd3mLHV8Y5k7s6XUvT0Gef6zxSIeXaQ= @@ -38,6 +44,8 @@ golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4= golang.org/x/crypto v0.49.0/go.mod h1:ErX4dUh2UM+CFYiXZRTcMpEcN8b/1gxEuv3nODoYtCA= golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 h1:mgKeJMpvi0yx/sU5GsxQ7p6s2wtOnGAHZWCHUM4KGzY= golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546/go.mod h1:j/pmGrbnkbPtQfxEe5D0VQhZC6qKbfKifgD0oM7sR70= +golang.org/x/net v0.51.0 h1:94R/GTO7mt3/4wIKpcR5gkGmRLOuE/2hNGeWq/GBIFo= +golang.org/x/net v0.51.0/go.mod h1:aamm+2QF5ogm02fjy5Bb7CQ0WMt1/WVM7FtyaTLlA9Y= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.32.0 h1:s77OFDvIQeibCmezSnk/q6iAfkdiQaJi4VzroCFrN20= golang.org/x/sys v0.32.0/go.mod h1:BJP2sWEmIv4KK5OTEluFJCKSidICx8ciO85XgH3Ak8k=