mirror of
https://github.com/warmbly/warmbly.git
synced 2026-10-03 16:02:02 +00:00
156 lines
5.2 KiB
Go
156 lines
5.2 KiB
Go
package middleware
|
|
|
|
import (
|
|
"context"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"testing"
|
|
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/google/uuid"
|
|
"github.com/warmbly/warmbly/internal/app/organization"
|
|
"github.com/warmbly/warmbly/internal/app/token"
|
|
"github.com/warmbly/warmbly/internal/errx"
|
|
"github.com/warmbly/warmbly/internal/models"
|
|
)
|
|
|
|
type sessionTokens struct {
|
|
token.TokenService
|
|
session *models.Session
|
|
left *int
|
|
}
|
|
|
|
func (f sessionTokens) ValidateAccessToken(context.Context, string) (*models.Session, *errx.Error) {
|
|
return f.session, nil
|
|
}
|
|
|
|
func (f sessionTokens) LeaveOrganization(context.Context, uuid.UUID, uuid.UUID) *errx.Error {
|
|
if f.left != nil {
|
|
*f.left++
|
|
}
|
|
return nil
|
|
}
|
|
|
|
type membershipOrgs struct {
|
|
organization.OrganizationService
|
|
members map[uuid.UUID]*models.OrganizationMember
|
|
permissionHits int
|
|
}
|
|
|
|
func (f *membershipOrgs) GetMembership(_ context.Context, orgID, userID uuid.UUID) (*models.OrganizationMember, *errx.Error) {
|
|
m := f.members[orgID]
|
|
if m == nil || m.UserID != userID {
|
|
return nil, nil
|
|
}
|
|
return m, nil
|
|
}
|
|
|
|
func (f *membershipOrgs) HasPermission(_ context.Context, orgID, userID uuid.UUID, perm models.OrganizationPermission) (bool, *errx.Error) {
|
|
f.permissionHits++
|
|
m, _ := f.GetMembership(context.Background(), orgID, userID)
|
|
return m != nil && m.HasPermission(perm), nil
|
|
}
|
|
|
|
func (f *membershipOrgs) Get(_ context.Context, orgID uuid.UUID) (*models.Organization, *errx.Error) {
|
|
return &models.Organization{ID: orgID}, nil
|
|
}
|
|
|
|
func serveSession(t *testing.T, h *Handler, auth gin.HandlerFunc, gates ...gin.HandlerFunc) (int, *uuid.UUID) {
|
|
t.Helper()
|
|
var seen *uuid.UUID
|
|
r := gin.New()
|
|
r.Use(RequestIDMiddleware())
|
|
chain := append([]gin.HandlerFunc{auth}, gates...)
|
|
chain = append(chain, func(c *gin.Context) {
|
|
seen = GetOrganizationID(c)
|
|
c.Status(http.StatusOK)
|
|
})
|
|
r.GET("/x", chain...)
|
|
|
|
req := httptest.NewRequestWithContext(context.Background(), http.MethodGet, "/x", nil)
|
|
req.Header.Set("Authorization", "Bearer t")
|
|
rec := httptest.NewRecorder()
|
|
r.ServeHTTP(rec, req)
|
|
return rec.Code, seen
|
|
}
|
|
|
|
func TestSessionOrganizationRequiresMembership(t *testing.T) {
|
|
userID, orgID := uuid.New(), uuid.New()
|
|
session := &models.Session{ID: uuid.New(), UserID: userID, CurrentOrganizationID: &orgID}
|
|
|
|
member := &models.OrganizationMember{OrganizationID: orgID, UserID: userID, Permissions: models.PermViewContacts}
|
|
for _, tt := range []struct {
|
|
name string
|
|
members map[uuid.UUID]*models.OrganizationMember
|
|
wantCode int
|
|
wantOrg bool
|
|
}{
|
|
{"member keeps the workspace", map[uuid.UUID]*models.OrganizationMember{orgID: member}, http.StatusOK, true},
|
|
{"former member has none", map[uuid.UUID]*models.OrganizationMember{}, http.StatusBadRequest, false},
|
|
} {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
orgs := &membershipOrgs{members: tt.members}
|
|
h := &Handler{TokenService: sessionTokens{session: session}, OrganizationService: orgs}
|
|
|
|
for name, auth := range map[string]gin.HandlerFunc{
|
|
"AuthMiddleware": h.AuthMiddleware(),
|
|
"CombinedAuthMiddleware": h.CombinedAuthMiddleware(),
|
|
} {
|
|
code, seen := serveSession(t, h, auth, h.RequireOrganization())
|
|
if code != tt.wantCode {
|
|
t.Fatalf("%s: status = %d, want %d", name, code, tt.wantCode)
|
|
}
|
|
if (seen != nil) != tt.wantOrg {
|
|
t.Fatalf("%s: organization in request = %v, want present=%v", name, seen, tt.wantOrg)
|
|
}
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestSessionMemberAnswersPermissionGates(t *testing.T) {
|
|
userID, orgID := uuid.New(), uuid.New()
|
|
session := &models.Session{ID: uuid.New(), UserID: userID, CurrentOrganizationID: &orgID}
|
|
orgs := &membershipOrgs{members: map[uuid.UUID]*models.OrganizationMember{
|
|
orgID: {OrganizationID: orgID, UserID: userID, Permissions: models.PermViewContacts},
|
|
}}
|
|
h := &Handler{TokenService: sessionTokens{session: session}, OrganizationService: orgs}
|
|
|
|
if code, _ := serveSession(t, h, h.AuthMiddleware(), h.RequirePermission(models.PermViewContacts)); code != http.StatusOK {
|
|
t.Fatalf("granted permission: status = %d, want 200", code)
|
|
}
|
|
if code, _ := serveSession(t, h, h.AuthMiddleware(), h.RequirePermission(models.PermManageTeam)); code != http.StatusForbidden {
|
|
t.Fatalf("missing permission: status = %d, want 403", code)
|
|
}
|
|
if orgs.permissionHits != 0 {
|
|
t.Fatalf("permission lookups = %d, want 0 (answered from the session's membership)", orgs.permissionHits)
|
|
}
|
|
}
|
|
|
|
func TestFormerMemberSelectionIsClearedAndDetached(t *testing.T) {
|
|
userID, orgID := uuid.New(), uuid.New()
|
|
session := &models.Session{ID: uuid.New(), UserID: userID, CurrentOrganizationID: &orgID}
|
|
left := 0
|
|
h := &Handler{TokenService: sessionTokens{session: session, left: &left}, OrganizationService: &membershipOrgs{}}
|
|
|
|
var seen *models.Session
|
|
r := gin.New()
|
|
r.GET("/x", h.AuthMiddleware(), func(c *gin.Context) {
|
|
seen = GetSession(c)
|
|
c.Status(http.StatusOK)
|
|
})
|
|
req := httptest.NewRequestWithContext(context.Background(), http.MethodGet, "/x", nil)
|
|
req.Header.Set("Authorization", "Bearer t")
|
|
r.ServeHTTP(httptest.NewRecorder(), req)
|
|
|
|
if left != 1 {
|
|
t.Fatalf("selection cleared %d times, want 1", left)
|
|
}
|
|
if seen == nil || seen.CurrentOrganizationID != nil {
|
|
t.Fatalf("the request's session still names the workspace")
|
|
}
|
|
if session.CurrentOrganizationID == nil {
|
|
t.Fatalf("the shared session was changed instead of copied")
|
|
}
|
|
}
|