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 }