feat: implement Phase 2 auth & security with TDD

- 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%
This commit is contained in:
jevb
2026-03-14 20:52:11 +01:00
parent a1434ad07f
commit b7dd6eabe9
21 changed files with 3329 additions and 0 deletions
+311
View File
@@ -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,
}
}
+456
View File
@@ -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)
}
}
+160
View File
@@ -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,
}
}
+294
View File
@@ -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
}
+194
View File
@@ -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 <token>" 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 <token>" 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"`
}
+361
View File
@@ -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);
`)
+10
View File
@@ -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
}
+50
View File
@@ -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
}
+106
View File
@@ -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?)")
}
}
+104
View File
@@ -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)
}
+126
View File
@@ -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
}
+24
View File
@@ -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[:])
}
+77
View File
@@ -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")
}
}
+285
View File
@@ -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
}
+468
View File
@@ -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)
}
}
+30
View File
@@ -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()
}
+56
View File
@@ -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
+156
View File
@@ -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)
}
}
+49
View File
@@ -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()
}
+4
View File
@@ -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
+8
View File
@@ -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=