From 238c19d50973ad15a0cf77d06e6739cd4e657273 Mon Sep 17 00:00:00 2001 From: Ludovic Vielle Date: Fri, 17 Jul 2026 16:54:57 +0200 Subject: [PATCH] Add pre-assume enrolled device status query The /enroll wait UI polled device state via node(), which requires an assumed org session, so confirmation never succeeded for unassumed viewers. Expose viewer.enrolledDevice behind itam:employee-device:get (own-device, skip assumption) and point the poller at it. Signed-off-by: Ludovic Vielle --- .../enroll/_components/EnrollDeviceButton.tsx | 9 +-- e2e/console/device_enrollment_test.go | 79 +++++++++++++++++++ pkg/itam/actions.go | 1 + pkg/itam/policies.go | 5 +- pkg/server/api/console/v1/device_resolvers.go | 2 + .../api/console/v1/graphql/viewer.graphql | 3 + pkg/server/api/console/v1/viewer_resolvers.go | 31 ++++++++ pkg/server/gqlutils/errors.go | 16 ++++ 8 files changed, 140 insertions(+), 6 deletions(-) diff --git a/apps/console/src/pages/organizations/enroll/_components/EnrollDeviceButton.tsx b/apps/console/src/pages/organizations/enroll/_components/EnrollDeviceButton.tsx index 58ca63a21..830ff0112 100644 --- a/apps/console/src/pages/organizations/enroll/_components/EnrollDeviceButton.tsx +++ b/apps/console/src/pages/organizations/enroll/_components/EnrollDeviceButton.tsx @@ -49,9 +49,8 @@ const enrollDeviceButtonMutation = graphql` const enrollDeviceButtonStatusQuery = graphql` query EnrollDeviceButtonStatusQuery($deviceId: ID!) { - device: node(id: $deviceId) { - __typename - ... on Device { + viewer @required(action: THROW) { + enrolledDevice(id: $deviceId) { id state hostname @@ -138,8 +137,8 @@ function EnrollDeviceButtonContent( return; } - const device = data?.device; - if (device?.__typename !== "Device") { + const device = data?.viewer.enrolledDevice; + if (device == null) { if (Date.now() > deadline) { setIsWaitingForActivity(false); setHasTimedOut(true); diff --git a/e2e/console/device_enrollment_test.go b/e2e/console/device_enrollment_test.go index 734207894..513dedf47 100644 --- a/e2e/console/device_enrollment_test.go +++ b/e2e/console/device_enrollment_test.go @@ -31,6 +31,8 @@ import ( "github.com/stretchr/testify/require" "go.probo.inc/probo/e2e/internal/testutil" + "go.probo.inc/probo/pkg/coredata" + "go.probo.inc/probo/pkg/gid" ) const ( @@ -108,6 +110,17 @@ const ( } } }` + + getEnrolledDeviceQuery = ` + query GetEnrolledDevice($id: ID!) { + viewer { + enrolledDevice(id: $id) { + id + state + hostname + } + } + }` ) type enrollDeviceResult struct { @@ -708,6 +721,72 @@ func TestDeviceEnrollment(t *testing.T) { testutil.RequireForbiddenError(t, err, "viewer should not enroll devices") }) + t.Run("unassumed session can poll own enrolledDevice", func(t *testing.T) { + t.Parallel() + + _, _, employee, _, orgID, _ := setupDeviceEnrollmentClients(t) + + enrolled := enrollAndActivateDevice(t, employee, orgID) + deviceID := enrolled.EnrollDevice.Device.ID + + unassumed := testutil.NewClientWithNewSession(t, employee) + + _, err := unassumed.Do(getDeviceQuery, map[string]any{"id": deviceID}) + testutil.RequireErrorCode(t, err, "ASSUMPTION_REQUIRED", "node get requires assumption") + + var result struct { + Viewer struct { + EnrolledDevice struct { + ID string `json:"id"` + State string `json:"state"` + } `json:"enrolledDevice"` + } `json:"viewer"` + } + unassumed.MustExecute(getEnrolledDeviceQuery, map[string]any{"id": deviceID}, &result) + require.Equal(t, deviceID, result.Viewer.EnrolledDevice.ID) + require.Equal(t, "ACTIVE", result.Viewer.EnrolledDevice.State) + }) + + t.Run("unassumed session cannot read another users enrolledDevice", func(t *testing.T) { + t.Parallel() + + owner, _, _, _, orgID, _ := setupDeviceEnrollmentClients(t) + + employeeA := testutil.NewClientInOrg(t, testutil.RoleEmployee, owner) + employeeB := testutil.NewClientInOrg(t, testutil.RoleEmployee, owner) + + enrolledB := enrollDevice(t, employeeB, orgID) + + unassumedA := testutil.NewClientWithNewSession(t, employeeA) + + _, err := unassumedA.Do(getEnrolledDeviceQuery, map[string]any{ + "id": enrolledB.EnrollDevice.Device.ID, + }) + testutil.RequireErrorCode(t, err, "NOT_FOUND", "employee cannot read another users enrolledDevice") + }) + + t.Run("enrolledDevice does not disclose foreign org device existence", func(t *testing.T) { + t.Parallel() + + _, _, employeeA, _, _, _ := setupDeviceEnrollmentClients(t) + _, _, employeeB, _, orgBID, _ := setupDeviceEnrollmentClients(t) + + enrolledB := enrollDevice(t, employeeB, orgBID) + + unassumedA := testutil.NewClientWithNewSession(t, employeeA) + + _, err := unassumedA.Do(getEnrolledDeviceQuery, map[string]any{ + "id": enrolledB.EnrollDevice.Device.ID, + }) + testutil.RequireErrorCode(t, err, "NOT_FOUND", "foreign org enrolledDevice must look like not found") + + unknownID := gid.New(employeeA.GetOrganizationID().TenantID(), coredata.DeviceEntityType).String() + _, err = unassumedA.Do(getEnrolledDeviceQuery, map[string]any{ + "id": unknownID, + }) + testutil.RequireErrorCode(t, err, "NOT_FOUND", "unknown enrolledDevice must look like not found") + }) + t.Run("owner retains admin access", func(t *testing.T) { t.Parallel() diff --git a/pkg/itam/actions.go b/pkg/itam/actions.go index 6f8ad8236..fcd5ec804 100644 --- a/pkg/itam/actions.go +++ b/pkg/itam/actions.go @@ -26,6 +26,7 @@ const ( // Device actions ActionDeviceList = "itam:device:list" ActionEmployeeDeviceList = "itam:employee-device:list" + ActionEmployeeDeviceGet = "itam:employee-device:get" ActionDeviceGet = "itam:device:get" ActionDeviceCreate = "itam:device:create" ActionDeviceEnroll = "itam:device:enroll" diff --git a/pkg/itam/policies.go b/pkg/itam/policies.go index 3820226d3..f9202d5d1 100644 --- a/pkg/itam/policies.go +++ b/pkg/itam/policies.go @@ -40,6 +40,9 @@ var FullAccessPolicy = policy.NewPolicy( ActionDeviceEnroll, ActionDeviceRevoke, ActionDeviceAssignOwner, ActionDevicePostureList, ).WithSID("itam-full-access").When(organizationCondition), + policy.Allow(ActionEmployeeDeviceGet). + WithSID("itam-full-access-get-own-device"). + When(organizationCondition, ownerCondition), ).WithDescription("Full ITAM access for organization owners and admins") // ViewerPolicy grants read-only access to ITAM entities for organization @@ -60,7 +63,7 @@ var EmployeePolicy = policy.NewPolicy( policy.Allow(ActionDeviceEnroll). WithSID("itam-employee-enroll-device"). When(organizationCondition), - policy.Allow(ActionDeviceGet). + policy.Allow(ActionDeviceGet, ActionEmployeeDeviceGet). WithSID("itam-employee-get-own-device"). When(organizationCondition, ownerCondition), policy.Allow(ActionEmployeeDeviceList). diff --git a/pkg/server/api/console/v1/device_resolvers.go b/pkg/server/api/console/v1/device_resolvers.go index 6a8f1ff5d..2ce20b75a 100644 --- a/pkg/server/api/console/v1/device_resolvers.go +++ b/pkg/server/api/console/v1/device_resolvers.go @@ -101,6 +101,8 @@ func (r *deviceConnectionResolver) TotalCount(ctx context.Context, obj *types.De } // EnrollDevice is the resolver for the enrollDevice field. +// SkipAssumptionCheck: self-enrollment from /enroll runs before the viewer +// assumes the target organization. func (r *mutationResolver) EnrollDevice(ctx context.Context, input types.EnrollDeviceInput) (*types.CreateDevicePayload, error) { identity := authn.IdentityFromContext(ctx) diff --git a/pkg/server/api/console/v1/graphql/viewer.graphql b/pkg/server/api/console/v1/graphql/viewer.graphql index a2cace8fd..6f34ddfd7 100644 --- a/pkg/server/api/console/v1/graphql/viewer.graphql +++ b/pkg/server/api/console/v1/graphql/viewer.graphql @@ -31,4 +31,7 @@ type Viewer { before: CursorKey orderBy: DeviceOrder ): DeviceConnection! @goField(forceResolver: true) + + # Own-device read for self-enrollment status polling before org assumption. + enrolledDevice(id: ID!): Device @goField(forceResolver: true) } diff --git a/pkg/server/api/console/v1/viewer_resolvers.go b/pkg/server/api/console/v1/viewer_resolvers.go index 5716e8299..188465bce 100644 --- a/pkg/server/api/console/v1/viewer_resolvers.go +++ b/pkg/server/api/console/v1/viewer_resolvers.go @@ -17,6 +17,7 @@ import ( "go.probo.inc/probo/pkg/page" "go.probo.inc/probo/pkg/probo" "go.probo.inc/probo/pkg/server/api/authn" + "go.probo.inc/probo/pkg/server/api/authz" "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" @@ -230,6 +231,36 @@ func (r *viewerResolver) EnrolledDevices(ctx context.Context, obj *types.Viewer, return types.NewOwnedDeviceConnection(devicesPage, r, organizationID, profile.ID), nil } +// EnrolledDevice is the resolver for the enrolledDevice field. +func (r *viewerResolver) EnrolledDevice(ctx context.Context, obj *types.Viewer, id gid.GID) (*types.Device, error) { + scope, err := r.authorize( + ctx, + id, + itam.ActionEmployeeDeviceGet, + authz.WithSkipAssumptionCheck(), + ) + if err != nil { + if gqlutils.IsForbidden(err) { + return nil, gqlutils.NotFoundf(ctx, "resource not found") + } + + return nil, err + } + + device, err := r.itam.GetDevice(ctx, scope, id) + if err != nil { + if errors.Is(err, coredata.ErrResourceNotFound) { + return nil, gqlutils.NotFound(ctx, err) + } + + r.logger.ErrorCtx(ctx, "cannot get enrolled device", log.Error(err)) + + return nil, gqlutils.Internal(ctx) + } + + return types.NewDevice(device), nil +} + // Viewer returns schema.ViewerResolver implementation. func (r *Resolver) Viewer() schema.ViewerResolver { return &viewerResolver{r} } diff --git a/pkg/server/gqlutils/errors.go b/pkg/server/gqlutils/errors.go index 80d454566..5bde7cebe 100644 --- a/pkg/server/gqlutils/errors.go +++ b/pkg/server/gqlutils/errors.go @@ -120,6 +120,22 @@ func Forbiddenf(ctx context.Context, format string, a ...any) *gqlerror.Error { return Forbidden(ctx, fmt.Errorf(format, a...)) } +// IsForbidden reports whether err is a GraphQL error with code FORBIDDEN. +func IsForbidden(err error) bool { + return hasCode(err, "FORBIDDEN") +} + +func hasCode(err error, code string) bool { + gqlErr, ok := errors.AsType[*gqlerror.Error](err) + if !ok { + return false + } + + got, _ := gqlErr.Extensions["code"].(string) + + return got == code +} + func NotFound(ctx context.Context, err error) *gqlerror.Error { return &gqlerror.Error{ Message: err.Error(),