Implement organization assumption check in authorization layer

Signed-off-by: Émile Ré <emile@getprobo.com>
This commit is contained in:
Émile Ré
2026-01-08 17:45:25 +01:00
committed by Bryan Frimin
parent 5e602d744b
commit c6094ff572
25 changed files with 1614 additions and 968 deletions

View File

@@ -1,8 +1,8 @@
import { graphql, usePreloadedQuery, type PreloadedQuery } from "react-relay";
import { Link, Outlet } from "react-router";
import { Badge, Button, IconPeopleAdd, Layout, Skeleton } from "@probo/ui";
import { Sidebar } from "../memberships/_components/Sidebar";
import { MembershipsDropdown } from "../memberships/MembershipsDropdown";
import { Sidebar } from "./_components/Sidebar";
import { MembershipsDropdown } from "./_components/MembershipsDropdown";
import type { ViewerMembershipLayoutQuery } from "/__generated__/iam/ViewerMembershipLayoutQuery.graphql";
import { ViewerMembershipDropdown } from "./_components/ViewerMembershipDropdown";
import { Suspense } from "react";

View File

@@ -14,14 +14,14 @@ import {
} from "@probo/ui";
import { Suspense, useCallback, useState } from "react";
import { useTranslate } from "@probo/i18n";
import {
MembershipsDropdownMenu,
membershipsDropdownMenuQuery,
} from "./MembershipsDropdownMenu";
import type { MembershipsDropdownMenuQuery } from "/__generated__/iam/MembershipsDropdownMenuQuery.graphql";
import { Link } from "react-router";
import type { MembershipsDropdown_organizationFragment$key } from "/__generated__/iam/MembershipsDropdown_organizationFragment.graphql";
import type { MembershipsDropdown_viewerFragment$key } from "/__generated__/iam/MembershipsDropdown_viewerFragment.graphql";
import {
MembershipsDropdownMenu,
membershipsDropdownMenuQuery,
} from "./MembershipsDropdownMenu";
const organizationFragment = graphql`
fragment MembershipsDropdown_organizationFragment on Organization {

View File

@@ -341,7 +341,12 @@ WHERE
return result.RowsAffected(), nil
}
func (s *Session) LoadByRootSessionIDAndMembershipID(ctx context.Context, conn pg.Conn, rootSessionID gid.GID, membershipID gid.GID) error {
func (s *Session) LoadByRootSessionIDAndMembershipID(
ctx context.Context,
conn pg.Conn,
rootSessionID gid.GID,
membershipID gid.GID,
) error {
q := `
SELECT
id,
@@ -390,3 +395,58 @@ LIMIT 1
return nil
}
func (s *Session) LoadByRootSessionIDAndOrganizationID(
ctx context.Context,
conn pg.Conn,
rootSessionID gid.GID,
organizationID gid.GID,
) error {
q := `
SELECT
id,
identity_id,
tenant_id,
membership_id,
data,
parent_session_id,
auth_method,
authenticated_at,
expire_reason,
user_agent,
ip_address,
expired_at,
created_at,
updated_at
FROM
iam_sessions
WHERE
parent_session_id = @root_session_id
AND organization_id = @organization_id
ORDER BY created_at DESC
LIMIT 1
`
args := pgx.StrictNamedArgs{
"root_session_id": rootSessionID,
"organization_id": organizationID,
}
rows, err := conn.Query(ctx, q, args)
if err != nil {
return fmt.Errorf("cannot query session: %w", err)
}
session, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Session])
if err != nil {
if err == pgx.ErrNoRows {
return ErrResourceNotFound
}
return fmt.Errorf("cannot collect session: %w", err)
}
*s = session
return nil
}

View File

@@ -16,8 +16,10 @@ package iam
import (
"context"
"errors"
"fmt"
"maps"
"time"
"go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/coredata"
@@ -35,6 +37,7 @@ type AuthorizationAttributer interface {
type AuthorizeParams struct {
Principal gid.GID
Resource gid.GID
Session *gid.GID
Action string
ResourceAttributes map[string]string
}
@@ -70,29 +73,52 @@ func (a *Authorizer) Authorize(ctx context.Context, params AuthorizeParams) erro
}
func (a *Authorizer) authorize(ctx context.Context, conn pg.Conn, params AuthorizeParams) error {
memberships, err := a.loadMemberships(ctx, conn, params.Principal)
if err != nil {
return err
}
resourceAttrs, err := a.buildResourceAttributes(ctx, conn, params)
if err != nil {
return err
return fmt.Errorf("cannot build resource attributes: %w", err)
}
// Find role for resource's organization
resourceOrgID := resourceAttrs["organization_id"]
role := findRoleForOrg(memberships, resourceOrgID)
// Find role for resource's organization
memberships, err := a.loadMemberships(ctx, conn, params.Principal)
if err != nil {
return fmt.Errorf("cannot load memberships for principal: %w", err)
}
membership := findMembershipForOrg(memberships, resourceOrgID)
if membership != nil && params.Session != nil {
if _, err := a.getActiveChildSessionForMembership(
ctx,
conn,
*params.Session,
membership.ID,
); err != nil {
var errSessionNotFound *ErrSessionNotFound
var errSessionExpired *ErrSessionExpired
if errors.As(err, &errSessionNotFound) || errors.As(err, &errSessionExpired) {
return NewInsufficientPermissionsError(params.Principal, params.Resource, params.Action)
}
return fmt.Errorf("cannot get active child session for membership: %w", err)
}
}
var role string
if membership != nil {
role = membership.Role.String()
}
// Only set principal.organization_id if they have a role in this org
var principalOrgID string
if role != "" {
principalOrgID = resourceOrgID
if membership != nil && role != "" {
principalOrgID = membership.OrganizationID.String()
}
principalAttrs, err := a.buildPrincipalAttributes(ctx, conn, params.Principal, principalOrgID)
if err != nil {
return err
return fmt.Errorf("cannot build principal attributes: %w", err)
}
policies := a.buildPoliciesForRole(role)
@@ -122,6 +148,29 @@ func (a *Authorizer) loadMemberships(ctx context.Context, conn pg.Conn, principa
return memberships, nil
}
func (a *Authorizer) getActiveChildSessionForMembership(
ctx context.Context,
conn pg.Conn,
rootSessionID gid.GID,
membershipID gid.GID,
) (*coredata.Session, error) {
childSession := &coredata.Session{}
if err := childSession.LoadByRootSessionIDAndMembershipID(ctx, conn, rootSessionID, membershipID); err != nil {
if err == coredata.ErrResourceNotFound {
return nil, NewSessionNotFoundError(gid.Nil)
}
return nil, fmt.Errorf("cannot load child session: %w", err)
}
if childSession.ExpireReason != nil || time.Now().After(childSession.ExpiredAt) {
return nil, NewSessionExpiredError(childSession.ID)
}
return childSession, nil
}
func (a *Authorizer) buildPrincipalAttributes(
ctx context.Context,
conn pg.Conn,
@@ -188,11 +237,12 @@ func (a *Authorizer) buildPoliciesForRole(role string) []*policy.Policy {
return policies
}
func findRoleForOrg(memberships coredata.Memberships, orgID string) string {
func findMembershipForOrg(memberships coredata.Memberships, orgID string) *coredata.Membership {
for _, m := range memberships {
if m.OrganizationID.String() == orgID && m.State == coredata.MembershipStateActive {
return string(m.Role)
return m
}
}
return ""
return nil
}

View File

@@ -190,6 +190,10 @@ func NewSessionNotFoundError(sessionID gid.GID) error {
}
func (e ErrSessionNotFound) Error() string {
if e.SessionID == gid.Nil {
return "session not found"
}
return fmt.Sprintf("session %q not found", e.SessionID)
}
@@ -353,18 +357,3 @@ func NewNoSCIMConfigurationFoundError(organizationID gid.GID) error {
func (e ErrNoSCIMConfigurationFound) Error() string {
return fmt.Sprintf("SCIM configuration not found for organization %q", e.OrganizationID)
}
// TenantAccessError is used by API recovery middleware to translate authorization/tenant failures
// into a consistent client-facing error response.
//
// NOTE: This is intentionally generic to avoid leaking resource existence.
type TenantAccessError struct {
Message string
}
func (e *TenantAccessError) Error() string {
if e == nil || e.Message == "" {
return "tenant access denied"
}
return e.Message
}

View File

@@ -297,9 +297,6 @@ const (
// TrustCenterDocumentAccess actions
ActionTrustCenterDocumentAccessList = "core:trust-center-document-access:list"
// DataProtectionOfficer actions
ActionDataProtectionOfficerList = "core:data-protection-officer:list"
// RightsRequest actions
ActionRightsRequestList = "core:rights-request:list"
ActionRightsRequestGet = "core:rights-request:get"

View File

@@ -20,11 +20,14 @@ import (
"fmt"
"net/http"
"github.com/99designs/gqlgen/graphql"
"github.com/vektah/gqlparser/v2/gqlerror"
"go.gearno.de/kit/httpserver"
"go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/iam"
"go.probo.inc/probo/pkg/securetoken"
"go.probo.inc/probo/pkg/server/gqlutils"
)
var (
@@ -56,7 +59,15 @@ func NewAPIKeyMiddleware(svc *iam.Service, tokenSecret string) func(next http.Ha
session := SessionFromContext(ctx)
if keyID != gid.Nil && session != nil {
httpserver.RenderError(w, http.StatusBadRequest, errors.New("api key authentication cannot be used with session authentication"))
httpserver.RenderJSON(
w,
http.StatusUnauthorized,
&graphql.Response{
Errors: gqlerror.List{
gqlutils.Conflictf(ctx, "API key authentication cannot be used with session authentication"),
},
},
)
return
}

View File

@@ -0,0 +1,83 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package connect_v1
import (
"context"
"errors"
"go.gearno.de/kit/log"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/iam"
"go.probo.inc/probo/pkg/server/gqlutils"
)
type (
AuthorizeFuncOption func(*iam.AuthorizeParams)
AuthorizeFunc func(context.Context, gid.GID, string, ...AuthorizeFuncOption) error
)
func WithAttr(key, value string) AuthorizeFuncOption {
return func(params *iam.AuthorizeParams) {
params.ResourceAttributes[key] = value
}
}
func WithSession(sessionID *gid.GID) AuthorizeFuncOption {
return func(params *iam.AuthorizeParams) {
params.Session = sessionID
}
}
func NewAuthorizeFunc(
svc *iam.Service,
logger *log.Logger,
) AuthorizeFunc {
return func(
ctx context.Context,
objectID gid.GID,
action string,
options ...AuthorizeFuncOption,
) error {
identity := IdentityFromContext(ctx)
session := SessionFromContext(ctx)
params := iam.AuthorizeParams{
Principal: identity.ID,
Resource: objectID,
Action: action,
ResourceAttributes: make(map[string]string),
}
if session != nil {
params.Session = &session.ID
}
for _, option := range options {
option(&params)
}
if err := svc.Authorizer.Authorize(ctx, params); err != nil {
var errInsufficientPermissions *iam.ErrInsufficientPermissions
if errors.As(err, &errInsufficientPermissions) {
return gqlutils.Forbidden(ctx, err)
}
logger.ErrorCtx(ctx, "cannot authorize", log.Error(err))
return gqlutils.Internal(ctx)
}
return nil
}
}

View File

@@ -19,7 +19,6 @@ import (
"net/http"
"github.com/99designs/gqlgen/graphql"
"github.com/vektah/gqlparser/v2/gqlerror"
"go.gearno.de/kit/log"
"go.probo.inc/probo/pkg/baseurl"
"go.probo.inc/probo/pkg/iam"
@@ -29,29 +28,6 @@ import (
"go.probo.inc/probo/pkg/server/gqlutils"
)
var (
ErrForbidden = &gqlerror.Error{
Message: "You are not authorized to access this resource",
Extensions: map[string]any{
"code": "FORBIDDEN",
},
}
ErrUnauthenticated = &gqlerror.Error{
Message: "You must be authenticated to access this resouce",
Extensions: map[string]any{
"code": "UNAUTHENTICATED",
},
}
ErrAlreadyAuthenticated = &gqlerror.Error{
Message: "authentication not allowed for this resource/action",
Extensions: map[string]any{
"code": "ALREADY_AUTHENTICATED",
},
}
)
func SessionDirective(ctx context.Context, obj any, next graphql.Resolver, required types.SessionRequirement) (any, error) {
session := SessionFromContext(ctx)
apiKey := APIKeyFromContext(ctx)
@@ -60,42 +36,27 @@ func SessionDirective(ctx context.Context, obj any, next graphql.Resolver, requi
case types.SessionRequirementOptional:
case types.SessionRequirementPresent:
if session == nil && apiKey == nil {
return nil, ErrUnauthenticated
return nil, gqlutils.Unauthenticatedf(
ctx,
"authentication is required to access this resouce",
)
}
case types.SessionRequirementNone:
if session != nil && apiKey != nil {
return nil, ErrAlreadyAuthenticated
return nil, gqlutils.Invalidf(
ctx,
"authentication not allowed for this resource/action",
)
}
}
return next(ctx)
}
func IsViewerDirective(ctx context.Context, obj any, next graphql.Resolver) (any, error) {
identity := IdentityFromContext(ctx)
switch node := obj.(type) {
case *types.Identity:
if identity.ID != node.ID {
return nil, ErrForbidden
}
case *types.Membership:
if identity.ID != node.Identity.ID {
return nil, ErrForbidden
}
case *types.Session:
if identity.ID != node.Identity.ID {
return nil, ErrForbidden
}
default:
}
return next(ctx)
}
func NewGraphQLHandler(svc *iam.Service, logger *log.Logger, baseURL *baseurl.BaseURL, cookieConfig securecookie.Config) http.Handler {
config := schema.Config{
Resolvers: &Resolver{
authorize: NewAuthorizeFunc(svc, logger),
logger: logger,
iam: svc,
baseURL: baseURL,
@@ -103,7 +64,6 @@ func NewGraphQLHandler(svc *iam.Service, logger *log.Logger, baseURL *baseurl.Ba
},
Directives: schema.DirectiveRoot{
Session: SessionDirective,
IsViewer: IsViewerDirective,
},
}

View File

@@ -0,0 +1,52 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package connect_v1
import (
"net/http"
"github.com/99designs/gqlgen/graphql"
"github.com/vektah/gqlparser/v2/gqlerror"
"go.gearno.de/kit/httpserver"
"go.probo.inc/probo/pkg/server/gqlutils"
)
func NewIdentityPresenceMiddleware() func(next http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(
func(w http.ResponseWriter, r *http.Request) {
identity := IdentityFromContext(r.Context())
if identity == nil {
httpserver.RenderJSON(
w,
http.StatusUnauthorized,
&graphql.Response{
Errors: gqlerror.List{
gqlutils.Unauthenticatedf(
r.Context(),
"authentication is required to access this resouce",
),
},
},
)
return
}
next.ServeHTTP(w, r)
},
)
}
}

View File

@@ -18,23 +18,20 @@ package connect_v1
import (
"context"
"errors"
"net/http"
"time"
"github.com/99designs/gqlgen/graphql"
"github.com/go-chi/chi/v5"
"go.gearno.de/kit/log"
"go.probo.inc/probo/pkg/baseurl"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/iam"
"go.probo.inc/probo/pkg/securecookie"
"go.probo.inc/probo/pkg/server/api/connect/v1/types"
"go.probo.inc/probo/pkg/server/gqlutils"
)
type (
Resolver struct {
authorize AuthorizeFunc
logger *log.Logger
iam *iam.Service
baseURL *baseurl.BaseURL
@@ -81,54 +78,5 @@ func NewMux(logger *log.Logger, svc *iam.Service, cookieConfig securecookie.Conf
}
func (r *Resolver) Permission(ctx context.Context, obj types.Node, action string) (bool, error) {
identity := IdentityFromContext(ctx)
err := r.iam.Authorizer.Authorize(
ctx,
iam.AuthorizeParams{
Principal: identity.ID,
Resource: obj.GetID(),
Action: action,
},
)
if err != nil {
var errInsufficientPermissions *iam.ErrInsufficientPermissions
if errors.As(err, &errInsufficientPermissions) {
return false, nil
}
r.logger.ErrorCtx(ctx, "cannot authorize", log.Error(err))
return false, gqlutils.InternalServerError(ctx)
}
return true, nil
}
func (r *Resolver) Authorize(ctx context.Context, objectID gid.GID, action string, attrs map[string]string) bool {
identity := IdentityFromContext(ctx)
err := r.iam.Authorizer.Authorize(
ctx,
iam.AuthorizeParams{
Principal: identity.ID,
Resource: objectID,
Action: action,
ResourceAttributes: attrs,
},
)
if err != nil {
var errInsufficientPermissions *iam.ErrInsufficientPermissions
if errors.As(err, &errInsufficientPermissions) {
graphql.AddError(ctx, err)
return false
}
r.logger.ErrorCtx(ctx, "cannot authorize", log.Error(err))
graphql.AddError(ctx, gqlutils.InternalServerError(ctx))
return false
}
return true
return r.authorize(ctx, obj.GetID(), action) == nil, nil
}

View File

@@ -37,7 +37,7 @@ func (h *SAMLHandler) MetadataHandler(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/samlmetadata+xml")
w.WriteHeader(http.StatusOK)
w.Write(metadataXML)
_, _ = w.Write(metadataXML)
}
func (h *SAMLHandler) ConsumeHandler(w http.ResponseWriter, r *http.Request) {
@@ -97,7 +97,11 @@ func (h *SAMLHandler) ConsumeHandler(w http.ResponseWriter, r *http.Request) {
return
}
securecookie.Set(w, h.cookieConfig, rootSession.ID.String())
if err := securecookie.Set(w, h.cookieConfig, rootSession.ID.String()); err != nil {
h.logger.ErrorCtx(ctx, "cannot set cookie", log.Error(err))
h.renderInternalServerError(w, r)
return
}
redirectURL := h.baseURL.WithPath("/organizations/" + membership.OrganizationID.String()).MustString()
http.Redirect(w, r, redirectURL, http.StatusFound)
}

View File

@@ -13,8 +13,6 @@ directive @goEnum(value: String) on ENUM_VALUE
directive @session(required: SessionRequirement!) on FIELD_DEFINITION
directive @isViewer on FIELD_DEFINITION
scalar CursorKey
scalar Datetime
scalar Upload
@@ -150,7 +148,7 @@ type Identity implements Node {
last: Int
before: CursorKey
orderBy: MembershipOrder
): MembershipConnection @goField(forceResolver: true) @isViewer
): MembershipConnection @goField(forceResolver: true)
pendingInvitations(
first: Int
@@ -158,7 +156,7 @@ type Identity implements Node {
last: Int
before: CursorKey
orderBy: InvitationOrder
): InvitationConnection @goField(forceResolver: true) @isViewer
): InvitationConnection @goField(forceResolver: true)
sessions(
first: Int
@@ -166,14 +164,14 @@ type Identity implements Node {
last: Int
before: CursorKey
orderBy: SessionOrder
): SessionConnection @goField(forceResolver: true) @isViewer
): SessionConnection @goField(forceResolver: true)
personalAPIKeys(
first: Int
after: CursorKey
last: Int
before: CursorKey
): PersonalAPIKeyConnection @goField(forceResolver: true) @isViewer
): PersonalAPIKeyConnection @goField(forceResolver: true)
permission(action: String!): Boolean!
@goField(forceResolver: true)
@@ -272,7 +270,7 @@ type Membership implements Node {
source: MembershipSource!
state: MembershipState!
lastSession: Session @goField(forceResolver: true) @isViewer
lastSession: Session @goField(forceResolver: true)
permission(action: String!): Boolean!
@goField(forceResolver: true)
@@ -297,7 +295,7 @@ type Invitation implements Node {
type Session implements Node {
id: ID!
identity: Identity @goField(forceResolver: true) @isViewer
identity: Identity @goField(forceResolver: true)
ipAddress: String!
userAgent: String!
updatedAt: Datetime!

View File

@@ -67,7 +67,6 @@ type ResolverRoot interface {
}
type DirectiveRoot struct {
IsViewer func(ctx context.Context, obj any, next graphql.Resolver) (res any, err error)
Session func(ctx context.Context, obj any, next graphql.Resolver, required types.SessionRequirement) (res any, err error)
}
@@ -2281,8 +2280,6 @@ directive @goEnum(value: String) on ENUM_VALUE
directive @session(required: SessionRequirement!) on FIELD_DEFINITION
directive @isViewer on FIELD_DEFINITION
scalar CursorKey
scalar Datetime
scalar Upload
@@ -2418,7 +2415,7 @@ type Identity implements Node {
last: Int
before: CursorKey
orderBy: MembershipOrder
): MembershipConnection @goField(forceResolver: true) @isViewer
): MembershipConnection @goField(forceResolver: true)
pendingInvitations(
first: Int
@@ -2426,7 +2423,7 @@ type Identity implements Node {
last: Int
before: CursorKey
orderBy: InvitationOrder
): InvitationConnection @goField(forceResolver: true) @isViewer
): InvitationConnection @goField(forceResolver: true)
sessions(
first: Int
@@ -2434,14 +2431,14 @@ type Identity implements Node {
last: Int
before: CursorKey
orderBy: SessionOrder
): SessionConnection @goField(forceResolver: true) @isViewer
): SessionConnection @goField(forceResolver: true)
personalAPIKeys(
first: Int
after: CursorKey
last: Int
before: CursorKey
): PersonalAPIKeyConnection @goField(forceResolver: true) @isViewer
): PersonalAPIKeyConnection @goField(forceResolver: true)
permission(action: String!): Boolean!
@goField(forceResolver: true)
@@ -2540,7 +2537,7 @@ type Membership implements Node {
source: MembershipSource!
state: MembershipState!
lastSession: Session @goField(forceResolver: true) @isViewer
lastSession: Session @goField(forceResolver: true)
permission(action: String!): Boolean!
@goField(forceResolver: true)
@@ -2565,7 +2562,7 @@ type Invitation implements Node {
type Session implements Node {
id: ID!
identity: Identity @goField(forceResolver: true) @isViewer
identity: Identity @goField(forceResolver: true)
ipAddress: String!
userAgent: String!
updatedAt: Datetime!
@@ -4683,20 +4680,7 @@ func (ec *executionContext) _Identity_memberships(ctx context.Context, field gra
fc := graphql.GetFieldContext(ctx)
return ec.resolvers.Identity().Memberships(ctx, obj, fc.Args["first"].(*int), fc.Args["after"].(*page.CursorKey), fc.Args["last"].(*int), fc.Args["before"].(*page.CursorKey), fc.Args["orderBy"].(*types.MembershipOrderBy))
},
func(ctx context.Context, next graphql.Resolver) graphql.Resolver {
directive0 := next
directive1 := func(ctx context.Context) (any, error) {
if ec.directives.IsViewer == nil {
var zeroVal *types.MembershipConnection
return zeroVal, errors.New("directive isViewer is not implemented")
}
return ec.directives.IsViewer(ctx, obj, directive0)
}
next = directive1
return next
},
nil,
ec.marshalOMembershipConnection2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋserverᚋapiᚋconnectᚋv1ᚋtypesᚐMembershipConnection,
true,
false,
@@ -4745,20 +4729,7 @@ func (ec *executionContext) _Identity_pendingInvitations(ctx context.Context, fi
fc := graphql.GetFieldContext(ctx)
return ec.resolvers.Identity().PendingInvitations(ctx, obj, fc.Args["first"].(*int), fc.Args["after"].(*page.CursorKey), fc.Args["last"].(*int), fc.Args["before"].(*page.CursorKey), fc.Args["orderBy"].(*types.InvitationOrderBy))
},
func(ctx context.Context, next graphql.Resolver) graphql.Resolver {
directive0 := next
directive1 := func(ctx context.Context) (any, error) {
if ec.directives.IsViewer == nil {
var zeroVal *types.InvitationConnection
return zeroVal, errors.New("directive isViewer is not implemented")
}
return ec.directives.IsViewer(ctx, obj, directive0)
}
next = directive1
return next
},
nil,
ec.marshalOInvitationConnection2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋserverᚋapiᚋconnectᚋv1ᚋtypesᚐInvitationConnection,
true,
false,
@@ -4807,20 +4778,7 @@ func (ec *executionContext) _Identity_sessions(ctx context.Context, field graphq
fc := graphql.GetFieldContext(ctx)
return ec.resolvers.Identity().Sessions(ctx, obj, fc.Args["first"].(*int), fc.Args["after"].(*page.CursorKey), fc.Args["last"].(*int), fc.Args["before"].(*page.CursorKey), fc.Args["orderBy"].(*types.SessionOrder))
},
func(ctx context.Context, next graphql.Resolver) graphql.Resolver {
directive0 := next
directive1 := func(ctx context.Context) (any, error) {
if ec.directives.IsViewer == nil {
var zeroVal *types.SessionConnection
return zeroVal, errors.New("directive isViewer is not implemented")
}
return ec.directives.IsViewer(ctx, obj, directive0)
}
next = directive1
return next
},
nil,
ec.marshalOSessionConnection2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋserverᚋapiᚋconnectᚋv1ᚋtypesᚐSessionConnection,
true,
false,
@@ -4869,20 +4827,7 @@ func (ec *executionContext) _Identity_personalAPIKeys(ctx context.Context, field
fc := graphql.GetFieldContext(ctx)
return ec.resolvers.Identity().PersonalAPIKeys(ctx, obj, fc.Args["first"].(*int), fc.Args["after"].(*page.CursorKey), fc.Args["last"].(*int), fc.Args["before"].(*page.CursorKey))
},
func(ctx context.Context, next graphql.Resolver) graphql.Resolver {
directive0 := next
directive1 := func(ctx context.Context) (any, error) {
if ec.directives.IsViewer == nil {
var zeroVal *types.PersonalAPIKeyConnection
return zeroVal, errors.New("directive isViewer is not implemented")
}
return ec.directives.IsViewer(ctx, obj, directive0)
}
next = directive1
return next
},
nil,
ec.marshalOPersonalAPIKeyConnection2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋserverᚋapiᚋconnectᚋv1ᚋtypesᚐPersonalAPIKeyConnection,
true,
false,
@@ -5863,20 +5808,7 @@ func (ec *executionContext) _Membership_lastSession(ctx context.Context, field g
func(ctx context.Context) (any, error) {
return ec.resolvers.Membership().LastSession(ctx, obj)
},
func(ctx context.Context, next graphql.Resolver) graphql.Resolver {
directive0 := next
directive1 := func(ctx context.Context) (any, error) {
if ec.directives.IsViewer == nil {
var zeroVal *types.Session
return zeroVal, errors.New("directive isViewer is not implemented")
}
return ec.directives.IsViewer(ctx, obj, directive0)
}
next = directive1
return next
},
nil,
ec.marshalOSession2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋserverᚋapiᚋconnectᚋv1ᚋtypesᚐSession,
true,
false,
@@ -11579,20 +11511,7 @@ func (ec *executionContext) _Session_identity(ctx context.Context, field graphql
func(ctx context.Context) (any, error) {
return ec.resolvers.Session().Identity(ctx, obj)
},
func(ctx context.Context, next graphql.Resolver) graphql.Resolver {
directive0 := next
directive1 := func(ctx context.Context) (any, error) {
if ec.directives.IsViewer == nil {
var zeroVal *types.Identity
return zeroVal, errors.New("directive isViewer is not implemented")
}
return ec.directives.IsViewer(ctx, obj, directive0)
}
next = directive1
return next
},
nil,
ec.marshalOIdentity2ᚖgoᚗproboᚗincᚋproboᚋpkgᚋserverᚋapiᚋconnectᚋv1ᚋtypesᚐIdentity,
true,
false,

View File

@@ -21,11 +21,14 @@ import (
"net"
"net/http"
"github.com/99designs/gqlgen/graphql"
"github.com/vektah/gqlparser/v2/gqlerror"
"go.gearno.de/kit/httpserver"
"go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/iam"
"go.probo.inc/probo/pkg/securecookie"
"go.probo.inc/probo/pkg/server/gqlutils"
)
var (
@@ -64,7 +67,15 @@ func NewSessionMiddleware(svc *iam.Service, cookieConfig securecookie.Config) fu
apiKey := APIKeyFromContext(ctx)
if sessionID != gid.Nil && apiKey != nil {
httpserver.RenderError(w, http.StatusBadRequest, errors.New("session authentication cannot be used with API key authentication"))
httpserver.RenderJSON(
w,
http.StatusUnauthorized,
&graphql.Response{
Errors: gqlerror.List{
gqlutils.Conflictf(ctx, "session authentication cannot be used with API key authentication"),
},
},
)
return
}

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,41 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package console_v1
import (
"net/http"
"go.gearno.de/kit/log"
"go.probo.inc/probo/pkg/iam"
"go.probo.inc/probo/pkg/probo"
connect_v1 "go.probo.inc/probo/pkg/server/api/connect/v1"
"go.probo.inc/probo/pkg/server/api/console/v1/schema"
"go.probo.inc/probo/pkg/server/gqlutils"
)
func NewGraphQLHandler(iamSvc *iam.Service, proboSvc *probo.Service, customDomainCname string, logger *log.Logger) http.Handler {
config := schema.Config{
Resolvers: &Resolver{
authorize: connect_v1.NewAuthorizeFunc(iamSvc, logger),
probo: proboSvc,
iam: iamSvc,
customDomainCname: customDomainCname,
},
}
es := schema.NewExecutableSchema(config)
gqlh := gqlutils.NewHandler(es, logger)
return gqlh
}

View File

@@ -19,14 +19,11 @@ package console_v1
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"strings"
"github.com/99designs/gqlgen/graphql"
"github.com/go-chi/chi/v5"
"github.com/vektah/gqlparser/v2/gqlerror"
"go.gearno.de/crypto/uuid"
"go.gearno.de/kit/httpserver"
"go.gearno.de/kit/log"
@@ -39,36 +36,19 @@ import (
"go.probo.inc/probo/pkg/saferedirect"
"go.probo.inc/probo/pkg/securecookie"
connect_v1 "go.probo.inc/probo/pkg/server/api/connect/v1"
"go.probo.inc/probo/pkg/server/api/console/v1/schema"
"go.probo.inc/probo/pkg/server/api/console/v1/types"
"go.probo.inc/probo/pkg/server/gqlutils"
"go.probo.inc/probo/pkg/statelesstoken"
)
type (
Resolver struct {
authorize connect_v1.AuthorizeFunc
probo *probo.Service
iam *iam.Service
customDomainCname string
}
)
func ensureAuthenticated(ctx context.Context, next graphql.OperationHandler) graphql.ResponseHandler {
identity := connect_v1.IdentityFromContext(ctx)
if identity == nil {
return func(ctx context.Context) *graphql.Response {
return &graphql.Response{
Errors: gqlerror.List{
gqlutils.Unauthorized(),
},
}
}
}
return next(ctx)
}
func NewMux(
logger *log.Logger,
proboSvc *probo.Service,
@@ -85,19 +65,11 @@ func NewMux(
r.Use(connect_v1.NewSessionMiddleware(iamSvc, cookieConfig))
r.Use(connect_v1.NewAPIKeyMiddleware(iamSvc, tokenSecret))
r.Use(connect_v1.NewIdentityPresenceMiddleware())
config := schema.Config{
Resolvers: &Resolver{
probo: proboSvc,
iam: iamSvc,
customDomainCname: customDomainCname,
},
}
es := schema.NewExecutableSchema(config)
h := gqlutils.NewHandler(es, logger)
h.AroundOperations(ensureAuthenticated)
graphqlHandler := NewGraphQLHandler(iamSvc, proboSvc, customDomainCname, logger)
r.Handle("/graphql", h)
r.Handle("/graphql", graphqlHandler)
r.Get(
"/documents/signing-requests",
@@ -238,10 +210,16 @@ func NewMux(
httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("authentication required"))
return
}
session := connect_v1.SessionFromContext(r.Context())
if session == nil {
httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("authentication required"))
return
}
if err := iamSvc.Authorizer.Authorize(r.Context(), iam.AuthorizeParams{
Principal: identity.ID,
Resource: organizationID,
Session: &session.ID,
Action: probo.ActionConnectorInitiate,
}); err != nil {
httpserver.RenderError(w, http.StatusForbidden, err)
@@ -317,42 +295,6 @@ func (r *Resolver) ProboService(ctx context.Context, tenantID gid.TenantID) *pro
return r.probo.WithTenant(tenantID)
}
func (r *Resolver) MustAuthorize(ctx context.Context, entityID gid.GID, action iam.Action) {
identity := connect_v1.IdentityFromContext(ctx)
err := r.iam.Authorizer.Authorize(
ctx,
iam.AuthorizeParams{
Principal: identity.ID,
Resource: entityID,
Action: action,
},
)
if err != nil {
panic(err)
}
}
func (r *Resolver) Permission(ctx context.Context, obj types.Node, action string) (bool, error) {
identity := connect_v1.IdentityFromContext(ctx)
err := r.iam.Authorizer.Authorize(
ctx,
iam.AuthorizeParams{
Principal: identity.ID,
Resource: obj.GetID(),
Action: action,
},
)
if err != nil {
var errInsufficientPermissions *iam.ErrInsufficientPermissions
if errors.As(err, &errInsufficientPermissions) {
return false, nil
}
panic(fmt.Errorf("cannot authorize: %w", err))
}
return true, nil
return r.authorize(ctx, obj.GetID(), action) == nil, nil
}

File diff suppressed because it is too large Load Diff

View File

@@ -47,11 +47,6 @@ func convertPanicToError(ctx context.Context, logger *log.Logger, panicValue any
return fmt.Errorf("internal server error")
}
var tenantAccessErr *iam.TenantAccessError
if errTyped, ok := panicValue.(error); ok && errors.As(errTyped, &tenantAccessErr) {
return fmt.Errorf("not authorized: %s", tenantAccessErr.Message)
}
var permissionDeniedErr *iam.ErrInsufficientPermissions
if errTyped, ok := panicValue.(error); ok && errors.As(errTyped, &permissionDeniedErr) {
return fmt.Errorf("permission denied: %s", permissionDeniedErr.Error())

View File

@@ -16,76 +16,106 @@ package gqlutils
import (
"context"
"errors"
"fmt"
"maps"
"github.com/99designs/gqlgen/graphql"
"github.com/vektah/gqlparser/v2/gqlerror"
"go.probo.inc/probo/pkg/validator"
)
func Unauthorized() *gqlerror.Error {
func Unauthenticated(ctx context.Context, err error) *gqlerror.Error {
return &gqlerror.Error{
Message: "not authorized",
Message: err.Error(),
Path: graphql.GetPath(ctx),
Extensions: map[string]any{
"code": "UNAUTHORIZED",
"code": "UNAUTHENTICATED",
},
}
}
func Forbidden(err error) *gqlerror.Error {
func Unauthenticatedf(ctx context.Context, format string, a ...any) *gqlerror.Error {
return Unauthenticated(ctx, fmt.Errorf(format, a...))
}
func AlreadyUnauthenticated(ctx context.Context, err error) *gqlerror.Error {
return &gqlerror.Error{
Message: "Authentication not allowed for this resource/action",
Extensions: map[string]any{
"code": "ALREADY_AUTHENTICATED",
},
}
}
func Forbidden(ctx context.Context, err error) *gqlerror.Error {
return &gqlerror.Error{
Message: err.Error(),
Path: graphql.GetPath(ctx),
Extensions: map[string]any{
"code": "FORBIDDEN",
},
}
}
func AuthenticationRequired(details map[string]any) *gqlerror.Error {
extensions := map[string]any{"code": "AUTHENTICATION_REQUIRED"}
maps.Copy(extensions, details)
return &gqlerror.Error{
Message: "Additional authentication required to access this organization",
Extensions: extensions,
}
}
func NotFound(err error) *gqlerror.Error {
func NotFound(ctx context.Context, err error) *gqlerror.Error {
return &gqlerror.Error{
Message: err.Error(),
Path: graphql.GetPath(ctx),
Extensions: map[string]any{
"code": "NOT_FOUND",
},
}
}
func Conflict(err error) *gqlerror.Error {
func Conflict(ctx context.Context, err error) *gqlerror.Error {
return &gqlerror.Error{
Message: err.Error(),
Path: graphql.GetPath(ctx),
Extensions: map[string]any{
"code": "CONFLICT",
},
}
}
func Invalid(err error, details map[string]any) *gqlerror.Error {
extensions := map[string]any{"code": "INVALID_REQUEST"}
func Conflictf(ctx context.Context, format string, a ...any) *gqlerror.Error {
return Conflict(ctx, fmt.Errorf(format, a...))
}
func Invalid(ctx context.Context, err error) *gqlerror.Error {
var errValidation *validator.ValidationError
var details map[string]any
if errors.As(err, &errValidation) {
details = map[string]any{
"cause": errValidation.Code,
"field": errValidation.Field,
"value": errValidation.Value,
}
}
extensions := map[string]any{"code": "INVALID"}
if details != nil {
maps.Copy(extensions, details)
}
return &gqlerror.Error{
Message: err.Error(),
Path: graphql.GetPath(ctx),
Extensions: extensions,
}
}
func InternalServerError(ctx context.Context) *gqlerror.Error {
func Invalidf(ctx context.Context, format string, a ...any) *gqlerror.Error {
return Invalid(ctx, fmt.Errorf(format, a...))
}
func Internal(ctx context.Context) *gqlerror.Error {
return &gqlerror.Error{
Message: "An internal server error occurred. Please try again later.",
Path: graphql.GetPath(ctx),
Extensions: map[string]any{
"code": "INTERNAL_SERVER_ERROR",
"code": "INTERNAL",
},
}
}

View File

@@ -31,26 +31,6 @@ func RecoverFunc(ctx context.Context, err any) error {
return gqlErr
}
// TODO: multi session here
// var errSAMLRequired iam.ErrSAMLAuthRequired
// if errors.As(asError(err), &errSAMLRequired) {
// return AuthenticationRequired(map[string]any{
// "requiresSaml": true,
// "redirectUrl": errSAMLRequired.RedirectURL,
// "samlConfigId": errSAMLRequired.ConfigID.String(),
// "organizationId": errSAMLRequired.OrganizationID.String(),
// })
// }
// var errPasswordRequired iam.ErrPasswordAuthRequired
// if errors.As(asError(err), &errPasswordRequired) {
// return AuthenticationRequired(map[string]any{
// "requiresSaml": false,
// "redirectUrl": errPasswordRequired.RedirectURL,
// "organizationId": errPasswordRequired.OrganizationID.String(),
// })
// }
var errValidations validator.ValidationErrors
if errors.As(asError(err), &errValidations) {
gqlErrors := gqlerror.List{}
@@ -59,12 +39,8 @@ func RecoverFunc(ctx context.Context, err any) error {
gqlErrors = append(
gqlErrors,
Invalid(
ctx,
err,
map[string]any{
"cause": err.Code,
"field": err.Field,
"value": err.Value,
},
),
)
}
@@ -72,14 +48,9 @@ func RecoverFunc(ctx context.Context, err any) error {
return gqlErrors
}
var tenantAccessErr *iam.TenantAccessError
if errTyped, ok := err.(error); ok && errors.As(errTyped, &tenantAccessErr) {
return Unauthorized()
}
var permissionDeniedErr *iam.ErrInsufficientPermissions
if errTyped, ok := err.(error); ok && errors.As(errTyped, &permissionDeniedErr) {
return Forbidden(permissionDeniedErr)
return Forbidden(ctx, permissionDeniedErr)
}
logger := httpserver.LoggerFromContext(ctx)