mirror of
https://github.com/J3vb/OwnCord.git
synced 2026-09-03 03:50:00 +03:00
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:
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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);
|
||||
`)
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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?)")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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[:])
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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=
|
||||
|
||||
Reference in New Issue
Block a user