160 lines
4.6 KiB
Go
160 lines
4.6 KiB
Go
package membership
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"log/slog"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"testing"
|
|
|
|
"proxui/backend/internal/auth"
|
|
"proxui/backend/internal/rbac"
|
|
)
|
|
|
|
func TestRequireTenantMembershipStoresMembership(t *testing.T) {
|
|
resolver := &stubResolver{
|
|
membership: Membership{
|
|
TenantID: "tenant-1",
|
|
Role: rbac.RoleAdmin,
|
|
},
|
|
found: true,
|
|
}
|
|
middleware := NewMiddleware(resolver, slog.Default())
|
|
handler := middleware.RequireTenantMembership(staticTenantID("tenant-1"), http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
membership, ok := FromRequest(r)
|
|
if !ok {
|
|
t.Fatal("membership missing from request")
|
|
}
|
|
if membership.TenantID != "tenant-1" {
|
|
t.Fatalf("tenant id = %q, want tenant-1", membership.TenantID)
|
|
}
|
|
if membership.Role != rbac.RoleAdmin {
|
|
t.Fatalf("role = %q, want %q", membership.Role, rbac.RoleAdmin)
|
|
}
|
|
w.WriteHeader(http.StatusNoContent)
|
|
}))
|
|
|
|
req := requestWithPrincipal()
|
|
rec := httptest.NewRecorder()
|
|
|
|
handler.ServeHTTP(rec, req)
|
|
|
|
if rec.Code != http.StatusNoContent {
|
|
t.Fatalf("status = %d, want %d", rec.Code, http.StatusNoContent)
|
|
}
|
|
if resolver.profileID != "profile-1" {
|
|
t.Fatalf("profile id = %q, want profile-1", resolver.profileID)
|
|
}
|
|
if resolver.tenantID != "tenant-1" {
|
|
t.Fatalf("tenant id = %q, want tenant-1", resolver.tenantID)
|
|
}
|
|
}
|
|
|
|
func TestRequireTenantMembershipRejectsMissingPrincipal(t *testing.T) {
|
|
resolver := &stubResolver{found: true}
|
|
middleware := NewMiddleware(resolver, slog.Default())
|
|
handler := middleware.RequireTenantMembership(staticTenantID("tenant-1"), http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.WriteHeader(http.StatusNoContent)
|
|
}))
|
|
|
|
req := httptest.NewRequest(http.MethodGet, "/tenants/tenant-1/membership", nil)
|
|
rec := httptest.NewRecorder()
|
|
|
|
handler.ServeHTTP(rec, req)
|
|
|
|
if rec.Code != http.StatusUnauthorized {
|
|
t.Fatalf("status = %d, want %d", rec.Code, http.StatusUnauthorized)
|
|
}
|
|
if resolver.calls != 0 {
|
|
t.Fatalf("resolver calls = %d, want 0", resolver.calls)
|
|
}
|
|
}
|
|
|
|
func TestRequireTenantMembershipRejectsMissingTenantID(t *testing.T) {
|
|
resolver := &stubResolver{found: true}
|
|
middleware := NewMiddleware(resolver, slog.Default())
|
|
handler := middleware.RequireTenantMembership(func(*http.Request) (string, bool) {
|
|
return "", false
|
|
}, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.WriteHeader(http.StatusNoContent)
|
|
}))
|
|
|
|
req := requestWithPrincipal()
|
|
rec := httptest.NewRecorder()
|
|
|
|
handler.ServeHTTP(rec, req)
|
|
|
|
if rec.Code != http.StatusBadRequest {
|
|
t.Fatalf("status = %d, want %d", rec.Code, http.StatusBadRequest)
|
|
}
|
|
if resolver.calls != 0 {
|
|
t.Fatalf("resolver calls = %d, want 0", resolver.calls)
|
|
}
|
|
}
|
|
|
|
func TestRequireTenantMembershipRejectsNonMember(t *testing.T) {
|
|
resolver := &stubResolver{found: false}
|
|
middleware := NewMiddleware(resolver, slog.Default())
|
|
handler := middleware.RequireTenantMembership(staticTenantID("tenant-1"), http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.WriteHeader(http.StatusNoContent)
|
|
}))
|
|
|
|
req := requestWithPrincipal()
|
|
rec := httptest.NewRecorder()
|
|
|
|
handler.ServeHTTP(rec, req)
|
|
|
|
if rec.Code != http.StatusForbidden {
|
|
t.Fatalf("status = %d, want %d", rec.Code, http.StatusForbidden)
|
|
}
|
|
}
|
|
|
|
func TestRequireTenantMembershipReturnsServerErrorOnResolverFailure(t *testing.T) {
|
|
resolver := &stubResolver{err: errors.New("db failed")}
|
|
middleware := NewMiddleware(resolver, slog.Default())
|
|
handler := middleware.RequireTenantMembership(staticTenantID("tenant-1"), http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.WriteHeader(http.StatusNoContent)
|
|
}))
|
|
|
|
req := requestWithPrincipal()
|
|
rec := httptest.NewRecorder()
|
|
|
|
handler.ServeHTTP(rec, req)
|
|
|
|
if rec.Code != http.StatusInternalServerError {
|
|
t.Fatalf("status = %d, want %d", rec.Code, http.StatusInternalServerError)
|
|
}
|
|
}
|
|
|
|
func requestWithPrincipal() *http.Request {
|
|
req := httptest.NewRequest(http.MethodGet, "/tenants/tenant-1/membership", nil)
|
|
return req.WithContext(auth.ContextWithPrincipal(req.Context(), auth.Principal{
|
|
Subject: "profile-1",
|
|
Email: "user@example.test",
|
|
Role: "authenticated",
|
|
}))
|
|
}
|
|
|
|
func staticTenantID(tenantID string) TenantIDExtractor {
|
|
return func(*http.Request) (string, bool) {
|
|
return tenantID, true
|
|
}
|
|
}
|
|
|
|
type stubResolver struct {
|
|
calls int
|
|
profileID string
|
|
tenantID string
|
|
membership Membership
|
|
found bool
|
|
err error
|
|
}
|
|
|
|
func (s *stubResolver) ResolveTenantMembership(_ context.Context, profileID string, tenantID string) (Membership, bool, error) {
|
|
s.calls++
|
|
s.profileID = profileID
|
|
s.tenantID = tenantID
|
|
return s.membership, s.found, s.err
|
|
}
|