diff --git a/apps/console/src/components/PageError.tsx b/apps/console/src/components/PageError.tsx
index d69438b24..1c0675c34 100644
--- a/apps/console/src/components/PageError.tsx
+++ b/apps/console/src/components/PageError.tsx
@@ -1,5 +1,4 @@
import { useTranslate } from "@probo/i18n";
-import { AuthenticationRequiredError } from "@probo/relay";
import { IconPageCross } from "@probo/ui";
import { useEffect, useRef } from "react";
import { useLocation, useRouteError } from "react-router";
@@ -33,27 +32,6 @@ export function PageError({ resetErrorBoundary, error: propsError }: Props) {
}
}, [location, resetErrorBoundary]);
- useEffect(() => {
- if (error instanceof AuthenticationRequiredError) {
- window.location.href = error.redirectUrl;
- }
- }, [error]);
-
- if (error instanceof AuthenticationRequiredError) {
- return (
-
diff --git a/apps/console/src/routes.tsx b/apps/console/src/routes.tsx
index cfa9825d9..3fb931782 100644
--- a/apps/console/src/routes.tsx
+++ b/apps/console/src/routes.tsx
@@ -1,6 +1,7 @@
import { Role } from "@probo/helpers";
import { lazy } from "@probo/react-lazy";
import {
+ NotAssumingError,
UnAuthenticatedError,
} from "@probo/relay";
import { type AppRoute, routeFromAppRoute } from "@probo/routes";
@@ -47,6 +48,11 @@ function ErrorBoundary() {
return
;
}
+ if (error instanceof NotAssumingError) {
+ // TODO redirect to right URL
+ return
;
+ }
+
return
;
}
diff --git a/packages/relay/src/errors.ts b/packages/relay/src/errors.ts
index dce1adca8..84dcb9630 100644
--- a/packages/relay/src/errors.ts
+++ b/packages/relay/src/errors.ts
@@ -14,25 +14,11 @@ export class InternalServerError extends Error {
}
}
-export class AuthenticationRequiredError extends Error {
- public redirectUrl: string;
- public requiresSaml: boolean;
- public organizationId: string;
- public samlConfigId?: string;
-
- constructor(extensions: {
- redirectUrl: string;
- requiresSaml: boolean;
- organizationId: string;
- samlConfigId?: string;
- }) {
- super("AUTHENTICATION_REQUIRED");
- this.name = "AuthenticationRequiredError";
- Object.setPrototypeOf(this, AuthenticationRequiredError.prototype);
- this.redirectUrl = extensions.redirectUrl;
- this.requiresSaml = extensions.requiresSaml;
- this.organizationId = extensions.organizationId;
- this.samlConfigId = extensions.samlConfigId;
+export class NotAssumingError extends Error {
+ constructor(message?: string) {
+ super(message ?? "NOT_ASSUMING");
+ this.name = "NotAssumingError";
+ Object.setPrototypeOf(this, NotAssumingError.prototype)
}
}
diff --git a/packages/relay/src/fetch.ts b/packages/relay/src/fetch.ts
index 3ea14078a..e5f31dd58 100644
--- a/packages/relay/src/fetch.ts
+++ b/packages/relay/src/fetch.ts
@@ -2,17 +2,17 @@ import { type FetchFunction } from "relay-runtime";
import {
InternalServerError,
UnAuthenticatedError,
- AuthenticationRequiredError,
UnauthorizedError,
ForbiddenError,
+ NotAssumingError,
} from "./errors";
import { GraphQLError } from "graphql";
const hasUnauthenticatedError = (error: GraphQLError) =>
error.extensions?.code == "UNAUTHENTICATED";
-const hasAuthenticationRequiredError = (error: GraphQLError) =>
- error.extensions?.code == "AUTHENTICATION_REQUIRED";
+const hasNotAssumingError = (error: GraphQLError) =>
+ error.extensions?.code == "NOT_ASSUMING";
const hasUnauthorizedError = (error: GraphQLError) =>
error.extensions?.code == "UNAUTHORIZED";
@@ -83,17 +83,9 @@ export const makeFetchQuery = (endpoint: string): FetchFunction => {
throw new UnAuthenticatedError(unauthenticatedError.message);
}
- const authRequiredError = errors.find(hasAuthenticationRequiredError);
- if (authRequiredError?.extensions) {
- const { redirectUrl, requiresSaml, organizationId, samlConfigId } =
- authRequiredError.extensions;
-
- throw new AuthenticationRequiredError({
- redirectUrl: redirectUrl as string,
- requiresSaml: requiresSaml as boolean,
- organizationId: organizationId as string,
- samlConfigId: samlConfigId as string | undefined,
- });
+ const notAssumingError = errors.find(hasNotAssumingError);
+ if (notAssumingError) {
+ throw new NotAssumingError(notAssumingError.message)
}
const unauthorizedError = errors.find(hasUnauthorizedError);
diff --git a/pkg/iam/authorizer.go b/pkg/iam/authorizer.go
index 625e9bafb..67c8a91f7 100644
--- a/pkg/iam/authorizer.go
+++ b/pkg/iam/authorizer.go
@@ -99,7 +99,7 @@ func (a *Authorizer) authorize(ctx context.Context, conn pg.Conn, params Authori
var errSessionExpired *ErrSessionExpired
if errors.As(err, &errSessionNotFound) || errors.As(err, &errSessionExpired) {
- return NewInsufficientPermissionsError(params.Principal, params.Resource, params.Action)
+ return NewAssumptionNeededError(params.Principal, membership.ID)
}
return fmt.Errorf("cannot get active child session for membership: %w", err)
diff --git a/pkg/iam/errors.go b/pkg/iam/errors.go
index 9c0d606ba..314f1c998 100644
--- a/pkg/iam/errors.go
+++ b/pkg/iam/errors.go
@@ -183,6 +183,19 @@ func (e ErrInsufficientPermissions) Error() string {
return fmt.Sprintf("identity %q does not have sufficient permissions to perform action %s on entity %q", e.IdentityID, e.Action, e.EntityID)
}
+type ErrAssumptionNeeded struct {
+ IdentityID gid.GID
+ MembershipID gid.GID
+}
+
+func NewAssumptionNeededError(identityID gid.GID, membershipID gid.GID) error {
+ return &ErrAssumptionNeeded{IdentityID: identityID, MembershipID: membershipID}
+}
+
+func (e ErrAssumptionNeeded) Error() string {
+ return fmt.Sprintf("assumption for identity %q needed for membership %q", e.IdentityID, e.MembershipID)
+}
+
type ErrSessionNotFound struct{ SessionID gid.GID }
func NewSessionNotFoundError(sessionID gid.GID) error {
diff --git a/pkg/server/api/authz/authorization.go b/pkg/server/api/authz/authorization.go
index eef15ff67..a2f270e3a 100644
--- a/pkg/server/api/authz/authorization.go
+++ b/pkg/server/api/authz/authorization.go
@@ -72,6 +72,11 @@ func NewAuthorizeFunc(
}
if err := svc.Authorizer.Authorize(ctx, params); err != nil {
+ var errAssumptionNeeded *iam.ErrAssumptionNeeded
+ if errors.As(err, &errAssumptionNeeded) {
+ return gqlutils.NotAssuming(ctx, err)
+ }
+
var errInsufficientPermissions *iam.ErrInsufficientPermissions
if errors.As(err, &errInsufficientPermissions) {
return gqlutils.Forbidden(ctx, err)
diff --git a/pkg/server/gqlutils/errors.go b/pkg/server/gqlutils/errors.go
index 825638001..24e5f3c6c 100644
--- a/pkg/server/gqlutils/errors.go
+++ b/pkg/server/gqlutils/errors.go
@@ -52,6 +52,20 @@ func Unauthenticatedf(ctx context.Context, format string, a ...any) *gqlerror.Er
return Unauthenticated(ctx, fmt.Errorf(format, a...))
}
+func NotAssuming(ctx context.Context, err error) *gqlerror.Error {
+ return &gqlerror.Error{
+ Message: err.Error(),
+ Path: graphql.GetPath(ctx),
+ Extensions: map[string]any{
+ "code": "NOT_ASSUMING",
+ },
+ }
+}
+
+func NotAssumingf(ctx context.Context, format string, a ...any) *gqlerror.Error {
+ return NotAssuming(ctx, fmt.Errorf(format, a...))
+}
+
func Forbidden(ctx context.Context, err error) *gqlerror.Error {
return &gqlerror.Error{
Message: err.Error(),