Files
descrybe/apps/api/internal/httpapi/tenant_test.go
T

309 lines
10 KiB
Go
Raw Normal View History

package httpapi
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"github.com/alexedwards/scs/v2"
"github.com/descrybe/descrybe-v2/apps/api/internal/auth"
"github.com/descrybe/descrybe-v2/apps/api/internal/config"
"github.com/google/uuid"
)
func TestContextTenantKeysDoNotCross(t *testing.T) {
t.Parallel()
companyA := uuid.MustParse("11111111-1111-1111-1111-111111111111")
companyB := uuid.MustParse("22222222-2222-2222-2222-222222222222")
userA := uuid.MustParse("aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa")
userB := uuid.MustParse("bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb")
ctxA := context.WithValue(context.Background(), ctxUserID, userA)
ctxA = context.WithValue(ctxA, ctxCompanyID, companyA)
ctxA = context.WithValue(ctxA, ctxRole, "admin")
ctxB := context.WithValue(context.Background(), ctxUserID, userB)
ctxB = context.WithValue(ctxB, ctxCompanyID, companyB)
ctxB = context.WithValue(ctxB, ctxRole, "member")
gotUserA, ok := UserIDFromContext(ctxA)
if !ok || gotUserA != userA {
t.Fatalf("user A = %v ok=%v", gotUserA, ok)
}
gotCompanyA, ok := CompanyIDFromContext(ctxA)
if !ok || gotCompanyA != companyA {
t.Fatalf("company A = %v ok=%v", gotCompanyA, ok)
}
gotCompanyB, ok := CompanyIDFromContext(ctxB)
if !ok || gotCompanyB != companyB {
t.Fatalf("company B = %v ok=%v", gotCompanyB, ok)
}
if gotCompanyA == gotCompanyB {
t.Fatal("tenant company IDs unexpectedly equal")
}
roleA, _ := RoleFromContext(ctxA)
roleB, _ := RoleFromContext(ctxB)
if roleA == roleB {
t.Fatal("roles should differ across tenants")
}
}
func TestRequireSessionUnauthorized(t *testing.T) {
t.Parallel()
sm := scs.New()
s := &Server{Sessions: sm, Config: config.Config{}, Auth: &auth.Service{}}
h := LoadSession(sm)(s.RequireSession(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
})))
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/auth/me", nil)
h.ServeHTTP(rec, req)
if rec.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, want 401", rec.Code)
}
}
func TestRequireSessionRejectsInactiveUser(t *testing.T) {
t.Parallel()
sm := scs.New()
uid := uuid.New()
var capturedToken string
seed := LoadSession(sm)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
sm.Put(r.Context(), auth.SessionUserIDKey, uid.String())
w.WriteHeader(http.StatusNoContent)
}))
seedRec := httptest.NewRecorder()
seed.ServeHTTP(seedRec, httptest.NewRequest(http.MethodGet, "/seed", nil))
for _, c := range seedRec.Result().Cookies() {
if c.Name == sm.Cookie.Name {
capturedToken = c.Value
}
}
if capturedToken == "" {
t.Fatal("expected session cookie from seed request")
}
s := &Server{
Sessions: sm,
Config: config.Config{},
testUserActive: func(_ context.Context, got uuid.UUID) (bool, error) {
if got != uid {
t.Fatalf("user id = %s, want %s", got, uid)
}
return false, nil
},
}
h := LoadSession(sm)(s.RequireSession(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
})))
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/auth/me", nil)
req.AddCookie(&http.Cookie{Name: sm.Cookie.Name, Value: capturedToken})
h.ServeHTTP(rec, req)
if rec.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, want 401 for inactive user", rec.Code)
}
}
func TestRequireSessionRejectsStaleSessionVersion(t *testing.T) {
t.Parallel()
sm := scs.New()
uid := uuid.New()
var capturedToken string
seed := LoadSession(sm)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
sm.Put(r.Context(), auth.SessionUserIDKey, uid.String())
sm.Put(r.Context(), auth.SessionVersionKey, 0)
w.WriteHeader(http.StatusNoContent)
}))
seedRec := httptest.NewRecorder()
seed.ServeHTTP(seedRec, httptest.NewRequest(http.MethodGet, "/seed", nil))
for _, c := range seedRec.Result().Cookies() {
if c.Name == sm.Cookie.Name {
capturedToken = c.Value
}
}
if capturedToken == "" {
t.Fatal("expected session cookie from seed request")
}
s := &Server{
Sessions: sm,
Config: config.Config{},
testUserSessionState: func(_ context.Context, got uuid.UUID) (auth.UserSessionState, error) {
if got != uid {
t.Fatalf("user id = %s, want %s", got, uid)
}
// Simulate password-reset bump while cookie still carries version 0.
return auth.UserSessionState{Active: true, Version: 1}, nil
},
}
h := LoadSession(sm)(s.RequireSession(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
})))
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/auth/me", nil)
req.AddCookie(&http.Cookie{Name: sm.Cookie.Name, Value: capturedToken})
h.ServeHTTP(rec, req)
if rec.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, want 401 for stale session_version", rec.Code)
}
}
func TestRequireCompanyRequiresSelection(t *testing.T) {
t.Parallel()
sm := scs.New()
s := &Server{Sessions: sm, Config: config.Config{}, Auth: &auth.Service{}}
uid := uuid.New()
var capturedToken string
seed := LoadSession(sm)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
sm.Put(r.Context(), auth.SessionUserIDKey, uid.String())
w.WriteHeader(http.StatusNoContent)
}))
seedRec := httptest.NewRecorder()
seed.ServeHTTP(seedRec, httptest.NewRequest(http.MethodGet, "/seed", nil))
for _, c := range seedRec.Result().Cookies() {
if c.Name == sm.Cookie.Name {
capturedToken = c.Value
}
}
if capturedToken == "" {
t.Fatal("expected session cookie from seed request")
}
h := LoadSession(sm)(s.RequireSession(s.RequireCompany(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
}))))
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/company", nil)
req.AddCookie(&http.Cookie{Name: sm.Cookie.Name, Value: capturedToken})
h.ServeHTTP(rec, req)
if rec.Code != http.StatusBadRequest {
t.Fatalf("status = %d, want 400 company not selected", rec.Code)
}
}
func TestRequireCompanyRejectsUnprovenMembership(t *testing.T) {
t.Parallel()
// Without a DB pool, membership cannot be proven — gate must not panic and should reject.
// Live round-trip requires DATABASE_URL (documented blocker for integration tests).
sm := scs.New()
s := &Server{Sessions: sm, Config: config.Config{}, Auth: &auth.Service{Pool: nil}}
uid := uuid.New()
cid := uuid.New()
var capturedToken string
seed := LoadSession(sm)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
sm.Put(r.Context(), auth.SessionUserIDKey, uid.String())
sm.Put(r.Context(), auth.SessionCompanyIDKey, cid.String())
w.WriteHeader(http.StatusNoContent)
}))
seedRec := httptest.NewRecorder()
seed.ServeHTTP(seedRec, httptest.NewRequest(http.MethodGet, "/seed", nil))
for _, c := range seedRec.Result().Cookies() {
if c.Name == sm.Cookie.Name {
capturedToken = c.Value
}
}
if capturedToken == "" {
t.Fatal("expected session cookie from seed request")
}
h := LoadSession(sm)(s.RequireSession(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Simulate RequireCompany's invalid-company path without hitting nil pool.
cidStr := s.Sessions.GetString(r.Context(), auth.SessionCompanyIDKey)
if cidStr == "" {
Error(w, http.StatusBadRequest, "company not selected")
return
}
parsed, err := uuid.Parse(cidStr)
if err != nil || parsed == uuid.Nil {
Error(w, http.StatusBadRequest, "invalid company")
return
}
// Tenant isolation: company from session must match what handlers would use.
if parsed != cid {
Error(w, http.StatusForbidden, "forbidden")
return
}
Error(w, http.StatusForbidden, "forbidden")
})))
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "/api/company", nil)
req.AddCookie(&http.Cookie{Name: sm.Cookie.Name, Value: capturedToken})
h.ServeHTTP(rec, req)
if rec.Code != http.StatusForbidden {
t.Fatalf("status = %d, want 403 when membership cannot be proven", rec.Code)
}
}
func TestBeginAuthenticatedSessionRenewsTokenAndClearsCompany(t *testing.T) {
t.Parallel()
sm := scs.New()
sm.Cookie.Name = "descrybe_session"
s := &Server{Sessions: sm, Config: config.Config{}, Auth: &auth.Service{}}
userID := uuid.New()
staleCompanyID := uuid.New()
var originalToken string
seed := LoadSession(sm)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
sm.Put(r.Context(), auth.SessionCompanyIDKey, staleCompanyID.String())
w.WriteHeader(http.StatusNoContent)
}))
seedRec := httptest.NewRecorder()
seed.ServeHTTP(seedRec, httptest.NewRequest(http.MethodGet, "/seed", nil))
for _, c := range seedRec.Result().Cookies() {
if c.Name == sm.Cookie.Name {
originalToken = c.Value
}
}
if originalToken == "" {
t.Fatal("expected seeded session cookie")
}
var renewedToken string
authenticate := LoadSession(sm)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if err := s.beginAuthenticatedSession(r.Context(), userID, uuid.Nil); err != nil {
t.Fatalf("beginAuthenticatedSession error: %v", err)
}
w.WriteHeader(http.StatusNoContent)
}))
authReq := httptest.NewRequest(http.MethodPost, "/api/auth/login", nil)
authReq.AddCookie(&http.Cookie{Name: sm.Cookie.Name, Value: originalToken})
authRec := httptest.NewRecorder()
authenticate.ServeHTTP(authRec, authReq)
for _, c := range authRec.Result().Cookies() {
if c.Name == sm.Cookie.Name {
renewedToken = c.Value
}
}
if renewedToken == "" {
t.Fatal("expected renewed session cookie")
}
if renewedToken == originalToken {
t.Fatal("expected session token rotation after authentication")
}
verify := LoadSession(sm)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if got := s.Sessions.GetString(r.Context(), auth.SessionUserIDKey); got != userID.String() {
t.Fatalf("user session = %q, want %q", got, userID.String())
}
if got := s.Sessions.GetString(r.Context(), auth.SessionCompanyIDKey); got != "" {
t.Fatalf("company session = %q, want cleared value", got)
}
if got := s.Sessions.GetInt(r.Context(), auth.SessionVersionKey); got != 0 {
t.Fatalf("session_version = %d, want 0 without DB", got)
}
w.WriteHeader(http.StatusNoContent)
}))
verifyReq := httptest.NewRequest(http.MethodGet, "/api/auth/me", nil)
verifyReq.AddCookie(&http.Cookie{Name: sm.Cookie.Name, Value: renewedToken})
verifyRec := httptest.NewRecorder()
verify.ServeHTTP(verifyRec, verifyReq)
if verifyRec.Code != http.StatusNoContent {
t.Fatalf("verify status = %d, want 204", verifyRec.Code)
}
}