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 { graphql, usePreloadedQuery, type PreloadedQuery } from "react-relay";
import { Link, Outlet } from "react-router"; import { Link, Outlet } from "react-router";
import { Badge, Button, IconPeopleAdd, Layout, Skeleton } from "@probo/ui"; import { Badge, Button, IconPeopleAdd, Layout, Skeleton } from "@probo/ui";
import { Sidebar } from "../memberships/_components/Sidebar"; import { Sidebar } from "./_components/Sidebar";
import { MembershipsDropdown } from "../memberships/MembershipsDropdown"; import { MembershipsDropdown } from "./_components/MembershipsDropdown";
import type { ViewerMembershipLayoutQuery } from "/__generated__/iam/ViewerMembershipLayoutQuery.graphql"; import type { ViewerMembershipLayoutQuery } from "/__generated__/iam/ViewerMembershipLayoutQuery.graphql";
import { ViewerMembershipDropdown } from "./_components/ViewerMembershipDropdown"; import { ViewerMembershipDropdown } from "./_components/ViewerMembershipDropdown";
import { Suspense } from "react"; import { Suspense } from "react";

View File

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

View File

@@ -341,7 +341,12 @@ WHERE
return result.RowsAffected(), nil 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 := ` q := `
SELECT SELECT
id, id,
@@ -390,3 +395,58 @@ LIMIT 1
return nil 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 ( import (
"context" "context"
"errors"
"fmt" "fmt"
"maps" "maps"
"time"
"go.gearno.de/kit/pg" "go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/coredata"
@@ -35,6 +37,7 @@ type AuthorizationAttributer interface {
type AuthorizeParams struct { type AuthorizeParams struct {
Principal gid.GID Principal gid.GID
Resource gid.GID Resource gid.GID
Session *gid.GID
Action string Action string
ResourceAttributes map[string]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 { 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) resourceAttrs, err := a.buildResourceAttributes(ctx, conn, params)
if err != nil { if err != nil {
return err return fmt.Errorf("cannot build resource attributes: %w", err)
} }
// Find role for resource's organization
resourceOrgID := resourceAttrs["organization_id"] 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 // Only set principal.organization_id if they have a role in this org
var principalOrgID string var principalOrgID string
if role != "" { if membership != nil && role != "" {
principalOrgID = resourceOrgID principalOrgID = membership.OrganizationID.String()
} }
principalAttrs, err := a.buildPrincipalAttributes(ctx, conn, params.Principal, principalOrgID) principalAttrs, err := a.buildPrincipalAttributes(ctx, conn, params.Principal, principalOrgID)
if err != nil { if err != nil {
return err return fmt.Errorf("cannot build principal attributes: %w", err)
} }
policies := a.buildPoliciesForRole(role) policies := a.buildPoliciesForRole(role)
@@ -122,6 +148,29 @@ func (a *Authorizer) loadMemberships(ctx context.Context, conn pg.Conn, principa
return memberships, nil 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( func (a *Authorizer) buildPrincipalAttributes(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
@@ -188,11 +237,12 @@ func (a *Authorizer) buildPoliciesForRole(role string) []*policy.Policy {
return policies return policies
} }
func findRoleForOrg(memberships coredata.Memberships, orgID string) string { func findMembershipForOrg(memberships coredata.Memberships, orgID string) *coredata.Membership {
for _, m := range memberships { for _, m := range memberships {
if m.OrganizationID.String() == orgID && m.State == coredata.MembershipStateActive { 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 { func (e ErrSessionNotFound) Error() string {
if e.SessionID == gid.Nil {
return "session not found"
}
return fmt.Sprintf("session %q not found", e.SessionID) return fmt.Sprintf("session %q not found", e.SessionID)
} }
@@ -353,18 +357,3 @@ func NewNoSCIMConfigurationFoundError(organizationID gid.GID) error {
func (e ErrNoSCIMConfigurationFound) Error() string { func (e ErrNoSCIMConfigurationFound) Error() string {
return fmt.Sprintf("SCIM configuration not found for organization %q", e.OrganizationID) 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 // TrustCenterDocumentAccess actions
ActionTrustCenterDocumentAccessList = "core:trust-center-document-access:list" ActionTrustCenterDocumentAccessList = "core:trust-center-document-access:list"
// DataProtectionOfficer actions
ActionDataProtectionOfficerList = "core:data-protection-officer:list"
// RightsRequest actions // RightsRequest actions
ActionRightsRequestList = "core:rights-request:list" ActionRightsRequestList = "core:rights-request:list"
ActionRightsRequestGet = "core:rights-request:get" ActionRightsRequestGet = "core:rights-request:get"

View File

@@ -20,11 +20,14 @@ import (
"fmt" "fmt"
"net/http" "net/http"
"github.com/99designs/gqlgen/graphql"
"github.com/vektah/gqlparser/v2/gqlerror"
"go.gearno.de/kit/httpserver" "go.gearno.de/kit/httpserver"
"go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/iam" "go.probo.inc/probo/pkg/iam"
"go.probo.inc/probo/pkg/securetoken" "go.probo.inc/probo/pkg/securetoken"
"go.probo.inc/probo/pkg/server/gqlutils"
) )
var ( var (
@@ -56,7 +59,15 @@ func NewAPIKeyMiddleware(svc *iam.Service, tokenSecret string) func(next http.Ha
session := SessionFromContext(ctx) session := SessionFromContext(ctx)
if keyID != gid.Nil && session != nil { 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 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" "net/http"
"github.com/99designs/gqlgen/graphql" "github.com/99designs/gqlgen/graphql"
"github.com/vektah/gqlparser/v2/gqlerror"
"go.gearno.de/kit/log" "go.gearno.de/kit/log"
"go.probo.inc/probo/pkg/baseurl" "go.probo.inc/probo/pkg/baseurl"
"go.probo.inc/probo/pkg/iam" "go.probo.inc/probo/pkg/iam"
@@ -29,29 +28,6 @@ import (
"go.probo.inc/probo/pkg/server/gqlutils" "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) { func SessionDirective(ctx context.Context, obj any, next graphql.Resolver, required types.SessionRequirement) (any, error) {
session := SessionFromContext(ctx) session := SessionFromContext(ctx)
apiKey := APIKeyFromContext(ctx) apiKey := APIKeyFromContext(ctx)
@@ -60,50 +36,34 @@ func SessionDirective(ctx context.Context, obj any, next graphql.Resolver, requi
case types.SessionRequirementOptional: case types.SessionRequirementOptional:
case types.SessionRequirementPresent: case types.SessionRequirementPresent:
if session == nil && apiKey == nil { if session == nil && apiKey == nil {
return nil, ErrUnauthenticated return nil, gqlutils.Unauthenticatedf(
ctx,
"authentication is required to access this resouce",
)
} }
case types.SessionRequirementNone: case types.SessionRequirementNone:
if session != nil && apiKey != nil { if session != nil && apiKey != nil {
return nil, ErrAlreadyAuthenticated return nil, gqlutils.Invalidf(
ctx,
"authentication not allowed for this resource/action",
)
} }
} }
return next(ctx) 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 { func NewGraphQLHandler(svc *iam.Service, logger *log.Logger, baseURL *baseurl.BaseURL, cookieConfig securecookie.Config) http.Handler {
config := schema.Config{ config := schema.Config{
Resolvers: &Resolver{ Resolvers: &Resolver{
authorize: NewAuthorizeFunc(svc, logger),
logger: logger, logger: logger,
iam: svc, iam: svc,
baseURL: baseURL, baseURL: baseURL,
cookieConfig: cookieConfig, cookieConfig: cookieConfig,
}, },
Directives: schema.DirectiveRoot{ Directives: schema.DirectiveRoot{
Session: SessionDirective, 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 ( import (
"context" "context"
"errors"
"net/http" "net/http"
"time" "time"
"github.com/99designs/gqlgen/graphql"
"github.com/go-chi/chi/v5" "github.com/go-chi/chi/v5"
"go.gearno.de/kit/log" "go.gearno.de/kit/log"
"go.probo.inc/probo/pkg/baseurl" "go.probo.inc/probo/pkg/baseurl"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/iam" "go.probo.inc/probo/pkg/iam"
"go.probo.inc/probo/pkg/securecookie" "go.probo.inc/probo/pkg/securecookie"
"go.probo.inc/probo/pkg/server/api/connect/v1/types" "go.probo.inc/probo/pkg/server/api/connect/v1/types"
"go.probo.inc/probo/pkg/server/gqlutils"
) )
type ( type (
Resolver struct { Resolver struct {
authorize AuthorizeFunc
logger *log.Logger logger *log.Logger
iam *iam.Service iam *iam.Service
baseURL *baseurl.BaseURL 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) { func (r *Resolver) Permission(ctx context.Context, obj types.Node, action string) (bool, error) {
identity := IdentityFromContext(ctx) return r.authorize(ctx, obj.GetID(), action) == nil, nil
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
} }

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.Header().Set("Content-Type", "application/samlmetadata+xml")
w.WriteHeader(http.StatusOK) w.WriteHeader(http.StatusOK)
w.Write(metadataXML) _, _ = w.Write(metadataXML)
} }
func (h *SAMLHandler) ConsumeHandler(w http.ResponseWriter, r *http.Request) { 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 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() redirectURL := h.baseURL.WithPath("/organizations/" + membership.OrganizationID.String()).MustString()
http.Redirect(w, r, redirectURL, http.StatusFound) 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 @session(required: SessionRequirement!) on FIELD_DEFINITION
directive @isViewer on FIELD_DEFINITION
scalar CursorKey scalar CursorKey
scalar Datetime scalar Datetime
scalar Upload scalar Upload
@@ -150,7 +148,7 @@ type Identity implements Node {
last: Int last: Int
before: CursorKey before: CursorKey
orderBy: MembershipOrder orderBy: MembershipOrder
): MembershipConnection @goField(forceResolver: true) @isViewer ): MembershipConnection @goField(forceResolver: true)
pendingInvitations( pendingInvitations(
first: Int first: Int
@@ -158,7 +156,7 @@ type Identity implements Node {
last: Int last: Int
before: CursorKey before: CursorKey
orderBy: InvitationOrder orderBy: InvitationOrder
): InvitationConnection @goField(forceResolver: true) @isViewer ): InvitationConnection @goField(forceResolver: true)
sessions( sessions(
first: Int first: Int
@@ -166,14 +164,14 @@ type Identity implements Node {
last: Int last: Int
before: CursorKey before: CursorKey
orderBy: SessionOrder orderBy: SessionOrder
): SessionConnection @goField(forceResolver: true) @isViewer ): SessionConnection @goField(forceResolver: true)
personalAPIKeys( personalAPIKeys(
first: Int first: Int
after: CursorKey after: CursorKey
last: Int last: Int
before: CursorKey before: CursorKey
): PersonalAPIKeyConnection @goField(forceResolver: true) @isViewer ): PersonalAPIKeyConnection @goField(forceResolver: true)
permission(action: String!): Boolean! permission(action: String!): Boolean!
@goField(forceResolver: true) @goField(forceResolver: true)
@@ -272,7 +270,7 @@ type Membership implements Node {
source: MembershipSource! source: MembershipSource!
state: MembershipState! state: MembershipState!
lastSession: Session @goField(forceResolver: true) @isViewer lastSession: Session @goField(forceResolver: true)
permission(action: String!): Boolean! permission(action: String!): Boolean!
@goField(forceResolver: true) @goField(forceResolver: true)
@@ -297,7 +295,7 @@ type Invitation implements Node {
type Session implements Node { type Session implements Node {
id: ID! id: ID!
identity: Identity @goField(forceResolver: true) @isViewer identity: Identity @goField(forceResolver: true)
ipAddress: String! ipAddress: String!
userAgent: String! userAgent: String!
updatedAt: Datetime! updatedAt: Datetime!

View File

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

View File

@@ -21,11 +21,14 @@ import (
"net" "net"
"net/http" "net/http"
"github.com/99designs/gqlgen/graphql"
"github.com/vektah/gqlparser/v2/gqlerror"
"go.gearno.de/kit/httpserver" "go.gearno.de/kit/httpserver"
"go.probo.inc/probo/pkg/coredata" "go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/iam" "go.probo.inc/probo/pkg/iam"
"go.probo.inc/probo/pkg/securecookie" "go.probo.inc/probo/pkg/securecookie"
"go.probo.inc/probo/pkg/server/gqlutils"
) )
var ( var (
@@ -64,7 +67,15 @@ func NewSessionMiddleware(svc *iam.Service, cookieConfig securecookie.Config) fu
apiKey := APIKeyFromContext(ctx) apiKey := APIKeyFromContext(ctx)
if sessionID != gid.Nil && apiKey != nil { 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 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 ( import (
"context" "context"
"encoding/json" "encoding/json"
"errors"
"fmt" "fmt"
"net/http" "net/http"
"strings" "strings"
"github.com/99designs/gqlgen/graphql"
"github.com/go-chi/chi/v5" "github.com/go-chi/chi/v5"
"github.com/vektah/gqlparser/v2/gqlerror"
"go.gearno.de/crypto/uuid" "go.gearno.de/crypto/uuid"
"go.gearno.de/kit/httpserver" "go.gearno.de/kit/httpserver"
"go.gearno.de/kit/log" "go.gearno.de/kit/log"
@@ -39,36 +36,19 @@ import (
"go.probo.inc/probo/pkg/saferedirect" "go.probo.inc/probo/pkg/saferedirect"
"go.probo.inc/probo/pkg/securecookie" "go.probo.inc/probo/pkg/securecookie"
connect_v1 "go.probo.inc/probo/pkg/server/api/connect/v1" 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/api/console/v1/types"
"go.probo.inc/probo/pkg/server/gqlutils"
"go.probo.inc/probo/pkg/statelesstoken" "go.probo.inc/probo/pkg/statelesstoken"
) )
type ( type (
Resolver struct { Resolver struct {
authorize connect_v1.AuthorizeFunc
probo *probo.Service probo *probo.Service
iam *iam.Service iam *iam.Service
customDomainCname string 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( func NewMux(
logger *log.Logger, logger *log.Logger,
proboSvc *probo.Service, proboSvc *probo.Service,
@@ -85,19 +65,11 @@ func NewMux(
r.Use(connect_v1.NewSessionMiddleware(iamSvc, cookieConfig)) r.Use(connect_v1.NewSessionMiddleware(iamSvc, cookieConfig))
r.Use(connect_v1.NewAPIKeyMiddleware(iamSvc, tokenSecret)) r.Use(connect_v1.NewAPIKeyMiddleware(iamSvc, tokenSecret))
r.Use(connect_v1.NewIdentityPresenceMiddleware())
config := schema.Config{ graphqlHandler := NewGraphQLHandler(iamSvc, proboSvc, customDomainCname, logger)
Resolvers: &Resolver{
probo: proboSvc,
iam: iamSvc,
customDomainCname: customDomainCname,
},
}
es := schema.NewExecutableSchema(config)
h := gqlutils.NewHandler(es, logger)
h.AroundOperations(ensureAuthenticated)
r.Handle("/graphql", h) r.Handle("/graphql", graphqlHandler)
r.Get( r.Get(
"/documents/signing-requests", "/documents/signing-requests",
@@ -238,10 +210,16 @@ func NewMux(
httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("authentication required")) httpserver.RenderError(w, http.StatusUnauthorized, fmt.Errorf("authentication required"))
return 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{ if err := iamSvc.Authorizer.Authorize(r.Context(), iam.AuthorizeParams{
Principal: identity.ID, Principal: identity.ID,
Resource: organizationID, Resource: organizationID,
Session: &session.ID,
Action: probo.ActionConnectorInitiate, Action: probo.ActionConnectorInitiate,
}); err != nil { }); err != nil {
httpserver.RenderError(w, http.StatusForbidden, err) 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) 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) { func (r *Resolver) Permission(ctx context.Context, obj types.Node, action string) (bool, error) {
identity := connect_v1.IdentityFromContext(ctx) return r.authorize(ctx, obj.GetID(), action) == nil, nil
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
} }

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") 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 var permissionDeniedErr *iam.ErrInsufficientPermissions
if errTyped, ok := panicValue.(error); ok && errors.As(errTyped, &permissionDeniedErr) { if errTyped, ok := panicValue.(error); ok && errors.As(errTyped, &permissionDeniedErr) {
return fmt.Errorf("permission denied: %s", permissionDeniedErr.Error()) return fmt.Errorf("permission denied: %s", permissionDeniedErr.Error())

View File

@@ -16,76 +16,106 @@ package gqlutils
import ( import (
"context" "context"
"errors"
"fmt"
"maps" "maps"
"github.com/99designs/gqlgen/graphql" "github.com/99designs/gqlgen/graphql"
"github.com/vektah/gqlparser/v2/gqlerror" "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{ return &gqlerror.Error{
Message: "not authorized", Message: err.Error(),
Path: graphql.GetPath(ctx),
Extensions: map[string]any{ 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{ return &gqlerror.Error{
Message: err.Error(), Message: err.Error(),
Path: graphql.GetPath(ctx),
Extensions: map[string]any{ Extensions: map[string]any{
"code": "FORBIDDEN", "code": "FORBIDDEN",
}, },
} }
} }
func AuthenticationRequired(details map[string]any) *gqlerror.Error { func NotFound(ctx context.Context, err error) *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 {
return &gqlerror.Error{ return &gqlerror.Error{
Message: err.Error(), Message: err.Error(),
Path: graphql.GetPath(ctx),
Extensions: map[string]any{ Extensions: map[string]any{
"code": "NOT_FOUND", "code": "NOT_FOUND",
}, },
} }
} }
func Conflict(err error) *gqlerror.Error { func Conflict(ctx context.Context, err error) *gqlerror.Error {
return &gqlerror.Error{ return &gqlerror.Error{
Message: err.Error(), Message: err.Error(),
Path: graphql.GetPath(ctx),
Extensions: map[string]any{ Extensions: map[string]any{
"code": "CONFLICT", "code": "CONFLICT",
}, },
} }
} }
func Invalid(err error, details map[string]any) *gqlerror.Error { func Conflictf(ctx context.Context, format string, a ...any) *gqlerror.Error {
extensions := map[string]any{"code": "INVALID_REQUEST"} 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 { if details != nil {
maps.Copy(extensions, details) maps.Copy(extensions, details)
} }
return &gqlerror.Error{ return &gqlerror.Error{
Message: err.Error(), Message: err.Error(),
Path: graphql.GetPath(ctx),
Extensions: extensions, 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{ return &gqlerror.Error{
Message: "An internal server error occurred. Please try again later.", Message: "An internal server error occurred. Please try again later.",
Path: graphql.GetPath(ctx), Path: graphql.GetPath(ctx),
Extensions: map[string]any{ 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 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 var errValidations validator.ValidationErrors
if errors.As(asError(err), &errValidations) { if errors.As(asError(err), &errValidations) {
gqlErrors := gqlerror.List{} gqlErrors := gqlerror.List{}
@@ -59,12 +39,8 @@ func RecoverFunc(ctx context.Context, err any) error {
gqlErrors = append( gqlErrors = append(
gqlErrors, gqlErrors,
Invalid( Invalid(
ctx,
err, 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 return gqlErrors
} }
var tenantAccessErr *iam.TenantAccessError
if errTyped, ok := err.(error); ok && errors.As(errTyped, &tenantAccessErr) {
return Unauthorized()
}
var permissionDeniedErr *iam.ErrInsufficientPermissions var permissionDeniedErr *iam.ErrInsufficientPermissions
if errTyped, ok := err.(error); ok && errors.As(errTyped, &permissionDeniedErr) { if errTyped, ok := err.(error); ok && errors.As(errTyped, &permissionDeniedErr) {
return Forbidden(permissionDeniedErr) return Forbidden(ctx, permissionDeniedErr)
} }
logger := httpserver.LoggerFromContext(ctx) logger := httpserver.LoggerFromContext(ctx)