feat: add tenant authorization middleware
This commit is contained in:
@@ -0,0 +1,28 @@
|
||||
package membership
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"proxui/backend/internal/rbac"
|
||||
)
|
||||
|
||||
type Membership struct {
|
||||
TenantID string `json:"tenant_id"`
|
||||
Role rbac.Role `json:"role"`
|
||||
}
|
||||
|
||||
type contextKey struct{}
|
||||
|
||||
func ContextWithMembership(ctx context.Context, membership Membership) context.Context {
|
||||
return context.WithValue(ctx, contextKey{}, membership)
|
||||
}
|
||||
|
||||
func FromContext(ctx context.Context) (Membership, bool) {
|
||||
membership, ok := ctx.Value(contextKey{}).(Membership)
|
||||
return membership, ok
|
||||
}
|
||||
|
||||
func FromRequest(r *http.Request) (Membership, bool) {
|
||||
return FromContext(r.Context())
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package membership
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
|
||||
"proxui/backend/internal/auth"
|
||||
)
|
||||
|
||||
type Resolver interface {
|
||||
ResolveTenantMembership(ctx context.Context, profileID string, tenantID string) (Membership, bool, error)
|
||||
}
|
||||
|
||||
type TenantIDExtractor func(*http.Request) (string, bool)
|
||||
|
||||
type Middleware struct {
|
||||
resolver Resolver
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
func NewMiddleware(resolver Resolver, logger *slog.Logger) Middleware {
|
||||
return Middleware{
|
||||
resolver: resolver,
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
func (m Middleware) RequireTenantMembership(extractTenantID TenantIDExtractor, next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
principal, ok := auth.PrincipalFromRequest(r)
|
||||
if !ok {
|
||||
writeError(w, http.StatusUnauthorized, "unauthorized")
|
||||
return
|
||||
}
|
||||
|
||||
tenantID, ok := extractTenantID(r)
|
||||
if !ok || tenantID == "" {
|
||||
writeError(w, http.StatusBadRequest, "tenant_id_required")
|
||||
return
|
||||
}
|
||||
|
||||
membership, found, err := m.resolver.ResolveTenantMembership(r.Context(), principal.Subject, tenantID)
|
||||
if err != nil {
|
||||
m.logger.Error("membership resolve failed", "error", err)
|
||||
writeError(w, http.StatusInternalServerError, "membership_resolve_failed")
|
||||
return
|
||||
}
|
||||
if !found {
|
||||
writeError(w, http.StatusForbidden, "tenant_forbidden")
|
||||
return
|
||||
}
|
||||
|
||||
next.ServeHTTP(w, r.WithContext(ContextWithMembership(r.Context(), membership)))
|
||||
})
|
||||
}
|
||||
|
||||
func writeError(w http.ResponseWriter, status int, message string) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(status)
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{"error": message})
|
||||
}
|
||||
@@ -0,0 +1,159 @@
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package membership
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
|
||||
"proxui/backend/internal/rbac"
|
||||
)
|
||||
|
||||
type Repository struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
func NewRepository(db *sql.DB) Repository {
|
||||
return Repository{db: db}
|
||||
}
|
||||
|
||||
func (r Repository) ResolveTenantMembership(ctx context.Context, profileID string, tenantID string) (Membership, bool, error) {
|
||||
var membership Membership
|
||||
var role string
|
||||
err := r.db.QueryRowContext(ctx, `
|
||||
select m.tenant_id::text, m.role::text
|
||||
from public.memberships m
|
||||
join public.tenants t on t.id = m.tenant_id
|
||||
where m.profile_id = $1
|
||||
and m.tenant_id = $2
|
||||
and t.status = 'active'
|
||||
`, profileID, tenantID).Scan(&membership.TenantID, &role)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return Membership{}, false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return Membership{}, false, err
|
||||
}
|
||||
|
||||
membership.Role = rbac.Role(role)
|
||||
return membership, true, nil
|
||||
}
|
||||
Reference in New Issue
Block a user