From 126b1bb3cd732ff628b88fc2ecc71185dab44fe5 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 16:41:52 -0600 Subject: [PATCH 01/88] feat(scim): add SCIM Users, Groups, and admin endpoints --- README.md | 8 + go.mod | 9 +- go.sum | 16 +- internal/api/admin.go | 10 + internal/api/admin_test.go | 84 ++ internal/api/api.go | 53 +- internal/api/apierrors/errorcode.go | 1 + internal/api/apilimiter/apilimiter.go | 26 + internal/api/apilimiter/apilimiter_test.go | 8 + internal/api/context.go | 33 +- internal/api/external.go | 26 + internal/api/external_test.go | 62 ++ internal/api/helpers.go | 1 + internal/api/identity.go | 16 +- internal/api/middleware.go | 2 +- internal/api/router.go | 4 + internal/api/scim.go | 377 ++++++++ internal/api/scim/core/core.go | 8 - internal/api/scim/core/endpoints.go | 6 - internal/api/scim/core/meta.go | 14 - internal/api/scim/core/meta_test.go | 39 - internal/api/scim/core/schemas.go | 13 - .../api/scim/core/service_provider_config.go | 70 -- .../scim/core/service_provider_config_test.go | 66 -- internal/api/scim/protocol/error.go | 24 - internal/api/scim/protocol/error_test.go | 47 - internal/api/scim/protocol/list_response.go | 25 - .../api/scim/protocol/list_response_test.go | 53 -- internal/api/scim/protocol/protocol.go | 18 - internal/api/scim/protocol/protocol_test.go | 23 - internal/api/scim/server.go | 48 -- internal/api/scim/server_test.go | 91 -- .../scim/testdata/empty_list_response.json | 9 - .../api/scim/testdata/not_implemented.json | 7 - internal/api/scim_admin.go | 216 +++++ internal/api/scim_admin_test.go | 580 +++++++++++++ internal/api/scim_filter.go | 52 ++ internal/api/scim_groups.go | 258 ++++++ internal/api/scim_groups_test.go | 712 ++++++++++++++++ internal/api/scim_isolation_test.go | 199 +++++ internal/api/scim_link_test.go | 763 +++++++++++++++++ internal/api/scim_okta_spec_test.go | 228 +++++ internal/api/scim_provider_delete_test.go | 222 +++++ internal/api/scim_test.go | 420 ++++++++- internal/api/scim_users.go | 491 +++++++++++ internal/api/scim_users_test.go | 803 ++++++++++++++++++ internal/api/ssoadmin.go | 3 + .../scim}/filter_forbidden.json | 2 +- .../testdata => testdata/scim}/not_found.json | 0 .../api/testdata/scim/okta_group_patch.json | 157 ++++ .../api/testdata/scim/okta_group_push.json | 375 ++++++++ .../testdata/scim/okta_user_lifecycle.json | 207 +++++ .../scim}/service_provider_config.json | 10 +- internal/conf/configuration.go | 15 +- internal/conf/confload/confload_test.go | 15 +- internal/models/audit_log_entry.go | 31 + internal/models/connection.go | 5 + internal/models/errors.go | 90 +- internal/models/scim.go | 212 +++++ internal/models/scim_group.go | 227 +++++ internal/models/scim_group_test.go | 324 +++++++ internal/models/scim_settings.go | 59 ++ internal/models/scim_settings_test.go | 103 +++ internal/models/scim_token.go | 186 ++++ internal/models/scim_token_test.go | 299 +++++++ internal/models/scim_user.go | 261 ++++++ internal/tokens/service.go | 10 + .../20260929000000_add_scim_groups.up.sql | 40 + .../20260930000000_add_scim_settings.up.sql | 8 + 69 files changed, 8261 insertions(+), 619 deletions(-) create mode 100644 internal/api/scim.go delete mode 100644 internal/api/scim/core/core.go delete mode 100644 internal/api/scim/core/endpoints.go delete mode 100644 internal/api/scim/core/meta.go delete mode 100644 internal/api/scim/core/meta_test.go delete mode 100644 internal/api/scim/core/schemas.go delete mode 100644 internal/api/scim/core/service_provider_config.go delete mode 100644 internal/api/scim/core/service_provider_config_test.go delete mode 100644 internal/api/scim/protocol/error.go delete mode 100644 internal/api/scim/protocol/error_test.go delete mode 100644 internal/api/scim/protocol/list_response.go delete mode 100644 internal/api/scim/protocol/list_response_test.go delete mode 100644 internal/api/scim/protocol/protocol.go delete mode 100644 internal/api/scim/protocol/protocol_test.go delete mode 100644 internal/api/scim/server.go delete mode 100644 internal/api/scim/server_test.go delete mode 100644 internal/api/scim/testdata/empty_list_response.json delete mode 100644 internal/api/scim/testdata/not_implemented.json create mode 100644 internal/api/scim_admin.go create mode 100644 internal/api/scim_admin_test.go create mode 100644 internal/api/scim_filter.go create mode 100644 internal/api/scim_groups.go create mode 100644 internal/api/scim_groups_test.go create mode 100644 internal/api/scim_isolation_test.go create mode 100644 internal/api/scim_link_test.go create mode 100644 internal/api/scim_okta_spec_test.go create mode 100644 internal/api/scim_provider_delete_test.go create mode 100644 internal/api/scim_users.go create mode 100644 internal/api/scim_users_test.go rename internal/api/{scim/testdata => testdata/scim}/filter_forbidden.json (60%) rename internal/api/{scim/testdata => testdata/scim}/not_found.json (100%) create mode 100644 internal/api/testdata/scim/okta_group_patch.json create mode 100644 internal/api/testdata/scim/okta_group_push.json create mode 100644 internal/api/testdata/scim/okta_user_lifecycle.json rename internal/api/{scim/testdata => testdata/scim}/service_provider_config.json (86%) create mode 100644 internal/models/scim.go create mode 100644 internal/models/scim_group.go create mode 100644 internal/models/scim_group_test.go create mode 100644 internal/models/scim_settings.go create mode 100644 internal/models/scim_settings_test.go create mode 100644 internal/models/scim_token.go create mode 100644 internal/models/scim_token_test.go create mode 100644 internal/models/scim_user.go create mode 100644 migrations/20260929000000_add_scim_groups.up.sql create mode 100644 migrations/20260930000000_add_scim_settings.up.sql diff --git a/README.md b/README.md index 7b103730b0..21cfd47f02 100644 --- a/README.md +++ b/README.md @@ -212,6 +212,14 @@ Header on which to rate limit the `/token` endpoint. This header is expected to Rate limit the number of emails sent per hour on the following endpoints: `/signup`, `/invite`, `/magiclink`, `/recover`, `/otp`, & `/user`. +`GOTRUE_SSO_SCIM_ENABLED` - `bool` + +Mounts the SCIM 2.0 routes at `/scim/v2` and the SCIM admin routes at `/admin/sso/providers/{id}/scim`. Defaults to `false`. + +`GOTRUE_RATE_LIMIT_SCIM` - `number` + +Requests per 5 minutes to `/scim/v2`, with a burst of 30. Requests with a valid SCIM token are limited per SSO provider. Requests without a valid token, and requests to `/ServiceProviderConfig` or to an unknown route, are limited per IP. Defaults to 3000. + `GOTRUE_PASSWORD_MIN_LENGTH` - `int` Minimum password length, defaults to 6. diff --git a/go.mod b/go.mod index 50ac8fe71b..2357651739 100644 --- a/go.mod +++ b/go.mod @@ -25,7 +25,6 @@ require ( github.com/consensys/gnark-crypto v0.18.1 // indirect github.com/crate-crypto/go-eth-kzg v1.4.0 // indirect github.com/crewjam/httperr v0.2.0 // indirect - github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect github.com/decred/dcrd/dcrec/secp256k1/v4 v4.3.0 // indirect github.com/dprotaso/go-yit v0.0.0-20220510233725-9ba8df137936 // indirect github.com/ethereum/c-kzg-4844/v2 v2.1.5 // indirect @@ -85,7 +84,6 @@ require ( github.com/onsi/gomega v1.27.6 // indirect github.com/patrickmn/go-cache v2.1.0+incompatible // indirect github.com/philhofer/fwd v1.2.0 // indirect - github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 // indirect github.com/prometheus/client_model v0.6.1 // indirect github.com/prometheus/common v0.48.0 // indirect github.com/prometheus/procfs v0.12.0 // indirect @@ -99,7 +97,7 @@ require ( github.com/speakeasy-api/jsonpath v0.6.3 // indirect github.com/speakeasy-api/openapi v1.24.0 // indirect github.com/spf13/pflag v1.0.6 // indirect - github.com/stretchr/objx v0.5.2 // indirect + github.com/stretchr/objx v0.5.3 // indirect github.com/supranational/blst v0.3.16-0.20250831170142-f48500c1fdbe // indirect github.com/tinylib/msgp v1.6.4 // indirect github.com/vmware-labs/yaml-jsonpath v0.3.2 // indirect @@ -108,7 +106,7 @@ require ( github.com/xeipuuv/gojsonreference v0.0.0-20180127040603-bd5ef7bd5415 // indirect go.opentelemetry.io/auto/sdk v1.2.1 // indirect go.opentelemetry.io/proto/otlp v1.10.0 // indirect - go.yaml.in/yaml/v3 v3.0.4 // indirect + go.yaml.in/yaml/v3 v3.0.5 // indirect golang.org/x/mod v0.40.0 // indirect golang.org/x/net v0.58.0 // indirect golang.org/x/tools v0.49.0 // indirect @@ -164,7 +162,8 @@ require ( github.com/sirupsen/logrus v1.9.3 github.com/spf13/cobra v1.8.1 github.com/standard-webhooks/standard-webhooks/libraries v0.0.0-20240303152453-e0e82adf1721 - github.com/stretchr/testify v1.11.1 + github.com/stretchr/testify v1.12.1 + github.com/supabase-community/scim-go v0.7.5 github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 github.com/xeipuuv/gojsonschema v1.2.0 go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.64.0 diff --git a/go.sum b/go.sum index ef418ed404..ac7525fd17 100644 --- a/go.sum +++ b/go.sum @@ -419,8 +419,6 @@ github.com/pkg/errors v0.8.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINE github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4= github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U= -github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/pquerna/otp v1.4.0 h1:wZvl1TIVxKRThZIBiwOOHOGP/1+nZyWBil9Y2XNEDzg= github.com/pquerna/otp v1.4.0/go.mod h1:dkJfzwRKNiegxyNb54X/3fLwhCynbMspSyWKnvi1AEg= github.com/prometheus/client_golang v1.19.0 h1:ygXvpU1AoN1MhdzckN+PyD9QJOSD4x7kmXYlnfbA6JU= @@ -487,8 +485,8 @@ github.com/stretchr/objx v0.1.1/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+ github.com/stretchr/objx v0.2.0/go.mod h1:qt09Ya8vawLte6SNmTgCsAVtYtaKzEcn8ATUoHMkEqE= github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw= github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= -github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY= -github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA= +github.com/stretchr/objx v0.5.3 h1:jmXUvGomnU1o3W/V5h2VEradbpJDwGrzugQQvL0POH4= +github.com/stretchr/objx v0.5.3/go.mod h1:rDQraq+vQZU7Fde9LOZLr8Tax6zZvy4kuNKF+QYS+U0= github.com/stretchr/testify v1.2.2/go.mod h1:a8OnRcib4nhh0OaRAV+Yts87kKdq0PP7pXfy6kDkUVs= github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4= @@ -498,8 +496,10 @@ github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/ github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU= github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= -github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= -github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= +github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= +github.com/supabase-community/scim-go v0.7.5 h1:bhT3BdYazeGMW2AaiHT06xiWO5weCIr2XV3MJ64v+Ko= +github.com/supabase-community/scim-go v0.7.5/go.mod h1:oEMij9JuKtAl0wl0jeyIHDKqHvPpUmTXegKeCBKyxXw= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 h1:VDuRtwen5Z7QQ5ctuHUse4wAv/JozkKZkdic5vUV4Lg= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869/go.mod h1:eHX5nlSMSnyPjUrbYzeqrA8snCe2SKyfizKjU3dkfOw= github.com/supranational/blst v0.3.16-0.20250831170142-f48500c1fdbe h1:nbdqkIGOGfUAD54q1s2YBcBz/WcsxCO9HUQ4aGV5hUw= @@ -572,8 +572,8 @@ go.uber.org/tools v0.0.0-20190618225709-2cfd321de3ee/go.mod h1:vJERXedbb3MVM5f9E go.uber.org/zap v1.9.1/go.mod h1:vwi/ZaCAaUcBkycHslxD9B2zi4UTXhF60s6SWpuDF0Q= go.uber.org/zap v1.10.0/go.mod h1:vwi/ZaCAaUcBkycHslxD9B2zi4UTXhF60s6SWpuDF0Q= go.uber.org/zap v1.13.0/go.mod h1:zwrFLgMcdUuIBviXEYEH1YKNaOBnKXsx2IPda5bBwHM= -go.yaml.in/yaml/v3 v3.0.4 h1:tfq32ie2Jv2UxXFdLJdh3jXuOzWiL1fo0bu/FbuKpbc= -go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= +go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw= +go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg= golang.org/x/crypto v0.0.0-20170930174604-9419663f5a44/go.mod h1:6SG95UA2DQfeDnfUPMdvaQW0Q7yPrPDi9nlGo2tz2b4= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20190411191339-88737f569e3a/go.mod h1:WFFai1msRO1wXaEeE5yQxYXgSfI8pQAWXbQop6sCtWE= diff --git a/internal/api/admin.go b/internal/api/admin.go index 0827e8ed45..c0172457f1 100644 --- a/internal/api/admin.go +++ b/internal/api/admin.go @@ -627,6 +627,10 @@ func (a *API) adminUserDelete(w http.ResponseWriter, r *http.Request) error { return apierrors.NewInternalServerError("Error soft deleting user").WithInternalError(terr) } + if terr := a.deleteSCIMUsers(tx, r, adminUser, user.ID); terr != nil { + return apierrors.NewInternalServerError("Error deleting user's SCIM users").WithInternalError(terr) + } + if terr := user.SoftDeleteUserIdentities(tx); terr != nil { return apierrors.NewInternalServerError("Error soft deleting user identities").WithInternalError(terr) } @@ -644,6 +648,12 @@ func (a *API) adminUserDelete(w http.ResponseWriter, r *http.Request) error { return apierrors.NewInternalServerError("Error deleting user's sessions").WithInternalError(terr) } } else { + if terr := models.LockUserForSCIM(tx, user.ID); terr != nil { + return apierrors.NewInternalServerError("Error locking user").WithInternalError(terr) + } + if terr := a.deleteSCIMUsers(tx, r, adminUser, user.ID); terr != nil { + return apierrors.NewInternalServerError("Error deleting user's SCIM users").WithInternalError(terr) + } if terr := tx.Destroy(user); terr != nil { return apierrors.NewInternalServerError("Database error deleting user").WithInternalError(terr) } diff --git a/internal/api/admin_test.go b/internal/api/admin_test.go index e29abc8bc5..ad19d8f595 100644 --- a/internal/api/admin_test.go +++ b/internal/api/admin_test.go @@ -864,6 +864,90 @@ func (ts *AdminTestSuite) TestAdminUserDelete() { } } +func (ts *AdminTestSuite) createLinkedSCIMUser(email string) (*models.SCIMUser, *models.User) { + provider := &models.SSOProvider{} + require.NoError(ts.T(), ts.API.db.Create(provider)) + + u, err := models.NewUser("", email, "", ts.Config.JWT.Aud, nil) + require.NoError(ts.T(), err) + u.IsSSOUser = true + require.NoError(ts.T(), ts.API.db.Create(u)) + + scimUser, err := models.CreateSCIMUser(ts.API.db, provider.ID, []byte(`{"userName":"`+email+`"}`)) + require.NoError(ts.T(), err) + require.NoError(ts.T(), models.LinkSCIMUser(ts.API.db, scimUser, u.ID)) + + return scimUser, u +} + +func (ts *AdminTestSuite) findSCIMUserByID(id uuid.UUID) *models.SCIMUser { + var row models.SCIMUser + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&row)) + return &row +} + +func (ts *AdminTestSuite) TestAdminUserDeleteSoftDeletesSCIMUser() { + cases := []struct { + desc string + body map[string]interface{} + wantEmail string + }{ + { + desc: "hard delete", + body: map[string]interface{}{"should_soft_delete": false}, + wantEmail: "scim-hard-delete@example.com", + }, + { + desc: "soft delete", + body: map[string]interface{}{"should_soft_delete": true}, + wantEmail: "scim-soft-delete@example.com", + }, + } + + for _, c := range cases { + ts.Run(c.desc, func() { + scimUser, u := ts.createLinkedSCIMUser(c.wantEmail) + group, err := models.CreateSCIMGroup(ts.API.db, scimUser.SSOProviderID, []byte(`{"displayName":"Engineering"}`)) + require.NoError(ts.T(), err) + _, _, err = models.ReplaceSCIMGroupMembers(ts.API.db, group, []uuid.UUID{scimUser.ID}) + require.NoError(ts.T(), err) + + var buffer bytes.Buffer + require.NoError(ts.T(), json.NewEncoder(&buffer).Encode(c.body)) + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodDelete, fmt.Sprintf("/admin/users/%s", u.ID), &buffer) + req.Header.Set("Authorization", fmt.Sprintf("Bearer %s", ts.token)) + + ts.API.handler.ServeHTTP(w, req) + require.Equal(ts.T(), http.StatusOK, w.Code) + + row := ts.findSCIMUserByID(scimUser.ID) + require.NotNil(ts.T(), row.DeletedAt) + + members, err := ts.API.db.Q().Where("scim_user_id = ?", scimUser.ID).Count(&models.SCIMGroupMember{}) + require.NoError(ts.T(), err) + require.Zero(ts.T(), members) + updated, err := models.FindSCIMGroup(ts.API.db, scimUser.SSOProviderID, group.ID) + require.NoError(ts.T(), err) + require.True(ts.T(), updated.UpdatedAt.After(group.UpdatedAt)) + + entry := models.AuditLogEntry{} + require.NoError(ts.T(), ts.API.db.Q().Where("payload->>'action' = ? AND payload->'traits'->>'scim_user_id' = ?", models.SCIMGroupMemberRemovedAction, scimUser.ID.String()).First(&entry)) + require.Equal(ts.T(), group.ID.String(), entry.Payload["traits"].(map[string]any)["scim_group_id"]) + require.Equal(ts.T(), "supabase_admin", entry.Payload["actor_username"]) + + deleted := []models.AuditLogEntry{} + require.NoError(ts.T(), ts.API.db.Q().Where("payload->>'action' = ? AND payload->'traits'->>'scim_user_id' = ?", models.SCIMUserDeletedAction, scimUser.ID.String()).All(&deleted)) + require.Len(ts.T(), deleted, 1) + traits := deleted[0].Payload["traits"].(map[string]any) + require.Equal(ts.T(), "supabase_admin", deleted[0].Payload["actor_username"]) + require.Equal(ts.T(), scimUser.SSOProviderID.String(), traits["sso_provider_id"]) + require.Equal(ts.T(), u.ID.String(), traits["user_id"]) + require.Equal(ts.T(), "success", traits["outcome"]) + }) + } +} + func (ts *AdminTestSuite) TestAdminUserSoftDeletion() { // create user u, err := models.NewUser("123456789", "test@example.com", "secret", ts.Config.JWT.Aud, map[string]interface{}{"name": "test"}) diff --git a/internal/api/api.go b/internal/api/api.go index 96f6f6d015..f396b0152e 100644 --- a/internal/api/api.go +++ b/internal/api/api.go @@ -8,12 +8,12 @@ import ( "github.com/rs/cors" "github.com/sebest/xff" "github.com/sirupsen/logrus" + "github.com/supabase-community/scim-go/pkg/server" "github.com/supabase/auth/internal/api/apierrors" "github.com/supabase/auth/internal/api/apilimiter" "github.com/supabase/auth/internal/api/apitask" "github.com/supabase/auth/internal/api/oauthserver" "github.com/supabase/auth/internal/api/provider" - "github.com/supabase/auth/internal/api/scim" "github.com/supabase/auth/internal/conf" "github.com/supabase/auth/internal/hooks/hookshttp" "github.com/supabase/auth/internal/hooks/hookspgfunc" @@ -47,7 +47,7 @@ type API struct { hooksMgr *v0hooks.Manager hibpClient *hibp.PwnedClient oauthServer *oauthserver.Server - scim *scim.Server + scim *server.Server tokenService *tokens.Service mailer mailer.Mailer oidcCache *provider.OIDCProviderCache @@ -138,7 +138,12 @@ func NewAPIWithVersion(globalConfig *conf.GlobalConfiguration, db *storage.Conne api.oauthServer = oauthserver.NewServer(globalConfig, db, api.tokenService) } - api.scim = scim.NewServer(globalConfig) + api.scim = newSCIMServer(globalConfig, + api.limitSCIMInvalidToken(newSCIMTokenValidator(db), api.limiterOpts.SCIMIP), + api.limitSCIMProvider(api.limiterOpts.SCIM), + &scimUsers{api: api}, + &scimGroups{api: api}, + ) if api.config.Password.HIBP.Enabled { httpClient := &http.Client{ @@ -404,6 +409,20 @@ func NewAPIWithVersion(globalConfig *conf.GlobalConfiguration, db *storage.Conne r.Get("/", api.adminSSOProvidersGet) r.Put("/", api.adminSSOProvidersUpdate) r.Delete("/", api.adminSSOProvidersDelete) + + r.Route("/scim", func(r *router) { + r.Use(api.requireScimServerEnabled) + + r.Get("/", api.adminSCIMGet) + r.Post("/", api.adminSCIMEnable) + r.Delete("/", api.adminSCIMDisable) + + r.Route("/tokens", func(r *router) { + r.Get("/", api.adminSCIMTokensList) + r.Post("/", api.adminSCIMTokensCreate) + r.Delete("/{prefix}", api.adminSCIMTokensRevoke) + }) + }) }) }) }) @@ -461,13 +480,29 @@ func NewAPIWithVersion(globalConfig *conf.GlobalConfiguration, db *storage.Conne r.With(api.requireAuthentication).Post("/authorizations/{authorization_id}/consent", api.oauthServer.OAuthServerConsent) }) - r.Route(scim.BasePath, func(r *router) { + r.Route(scimBasePath, func(r *router) { r.Use(api.requireScimServerEnabled) - r.NotFound(api.scim.NotFound) - - r.Get("/ServiceProviderConfig", api.scim.ServiceProviderConfig) - r.Get("/ResourceTypes", api.scim.ResourceTypes) - r.Get("/Schemas", api.scim.Schemas) + r.Use(api.withSCIMRequest) + r.UseBypass(api.limitSCIMHandler(api.limiterOpts.SCIMIP, r.chi)) + r.NotFound(scimNotFound) + + r.Method(http.MethodGet, "/ServiceProviderConfig", api.scim) + r.Method(http.MethodGet, "/ResourceTypes", api.scim) + r.Method(http.MethodGet, "/ResourceTypes/{id}", api.scim) + r.Method(http.MethodGet, "/Schemas", api.scim) + r.Method(http.MethodGet, "/Schemas/{id}", api.scim) + r.Method(http.MethodGet, "/Users", api.scim) + r.Method(http.MethodPost, "/Users", api.scim) + r.Method(http.MethodGet, "/Users/{id}", api.scim) + r.Method(http.MethodPut, "/Users/{id}", api.scim) + r.Method(http.MethodPatch, "/Users/{id}", api.scim) + r.Method(http.MethodDelete, "/Users/{id}", api.scim) + r.Method(http.MethodGet, "/Groups", api.scim) + r.Method(http.MethodPost, "/Groups", api.scim) + r.Method(http.MethodGet, "/Groups/{id}", api.scim) + r.Method(http.MethodPut, "/Groups/{id}", api.scim) + r.Method(http.MethodPatch, "/Groups/{id}", api.scim) + r.Method(http.MethodDelete, "/Groups/{id}", api.scim) }) }) diff --git a/internal/api/apierrors/errorcode.go b/internal/api/apierrors/errorcode.go index 7c41718d01..6683be8652 100644 --- a/internal/api/apierrors/errorcode.go +++ b/internal/api/apierrors/errorcode.go @@ -62,6 +62,7 @@ const ( ErrorCodeUserAlreadyExists ErrorCode = "user_already_exists" ErrorCodeSSOProviderNotFound ErrorCode = "sso_provider_not_found" ErrorCodeSSOProviderDisabled ErrorCode = "sso_provider_disabled" + ErrorCodeSCIMTokenNotFound ErrorCode = "scim_token_not_found" // #nosec G101 -- not a credential ErrorCodeSAMLMetadataFetchFailed ErrorCode = "saml_metadata_fetch_failed" ErrorCodeSAMLIdPAlreadyExists ErrorCode = "saml_idp_already_exists" ErrorCodeSSODomainAlreadyExists ErrorCode = "sso_domain_already_exists" diff --git a/internal/api/apilimiter/apilimiter.go b/internal/api/apilimiter/apilimiter.go index 464530c903..c91c763929 100644 --- a/internal/api/apilimiter/apilimiter.go +++ b/internal/api/apilimiter/apilimiter.go @@ -64,6 +64,12 @@ const ( envRateLimitSso = "GOTRUE_RATE_LIMIT_SSO" fieldSSO = "SSO" + // GOTRUE_RATE_LIMIT_SCIM + // -> RateLimitScim + envRateLimitScim = "GOTRUE_RATE_LIMIT_SCIM" + fieldSCIM = "SCIM" + fieldSCIMIP = "SCIMIP" + // GOTRUE_RATE_LIMIT_TOKEN_REFRESH // -> RateLimitTokenRefresh envRateLimitTokenRefresh = "GOTRUE_RATE_LIMIT_TOKEN_REFRESH" // #nosec G101 @@ -99,6 +105,8 @@ var tollboothFieldsToEnv = map[string]string{ fieldPasskeyAuthentication: envRateLimitPasskey, fieldSAMLAssertion: envSAMLRateLimitAssertion, fieldSSO: envRateLimitSso, + fieldSCIM: envRateLimitScim, + fieldSCIMIP: envRateLimitScim, fieldToken: envRateLimitTokenRefresh, fieldVerify: envRateLimitVerify, fieldWeb3: envRateLimitWeb3, @@ -168,6 +176,14 @@ type Limiter struct { // -> RateLimitSso SSO *limiter.Limiter + // GOTRUE_RATE_LIMIT_SCIM + // -> RateLimitScim + SCIM *limiter.Limiter + + // GOTRUE_RATE_LIMIT_SCIM + // -> RateLimitScim + SCIMIP *limiter.Limiter + // GOTRUE_RATE_LIMIT_TOKEN_REFRESH // -> RateLimitTokenRefresh Token *limiter.Limiter @@ -238,6 +254,8 @@ func New(gc *conf.GlobalConfiguration) *Limiter { o.Signups = newLimiterPer5mOver1h(gc.RateLimitOtp) o.OAuthClientRegister = newLimiterPer5mOver1h(gc.RateLimitOAuthDynamicClientRegister) o.PasskeyAuthentication = newLimiterPer5mOver1h(gc.RateLimitPasskey) + o.SCIM = newLimiterPer5mOver1h(gc.RateLimitScim) + o.SCIMIP = newLimiterPer5mOver1h(gc.RateLimitScim) return o } @@ -258,6 +276,8 @@ func (o *Limiter) Copy() *Limiter { Recover: o.Recover, Resend: o.Resend, SAMLAssertion: o.SAMLAssertion, + SCIM: o.SCIM, + SCIMIP: o.SCIMIP, Signups: o.Signups, SSO: o.SSO, Token: o.Token, @@ -332,6 +352,12 @@ func (o *Limiter) Update( logEnvUpdates(le, envRateLimitSso, a, b) } + if a, b := prevCfg.RateLimitScim, nextCfg.RateLimitScim; a != b { + v.SCIM = newLimiterPer5mOver1h(b) + v.SCIMIP = newLimiterPer5mOver1h(b) + logEnvUpdates(le, envRateLimitScim, a, b) + } + if a, b := prevCfg.RateLimitTokenRefresh, nextCfg.RateLimitTokenRefresh; a != b { v.Token = newLimiterPer5mOver1h(b) logEnvUpdates(le, envRateLimitTokenRefresh, a, b) diff --git a/internal/api/apilimiter/apilimiter_test.go b/internal/api/apilimiter/apilimiter_test.go index 8a3b826a31..d6fa85beb7 100644 --- a/internal/api/apilimiter/apilimiter_test.go +++ b/internal/api/apilimiter/apilimiter_test.go @@ -181,6 +181,10 @@ func tollboothByField(o *Limiter, field string) *limiter.Limiter { return o.SAMLAssertion case fieldSSO: return o.SSO + case fieldSCIM: + return o.SCIM + case fieldSCIMIP: + return o.SCIMIP case fieldToken: return o.Token case fieldVerify: @@ -220,6 +224,10 @@ func tollboothCfgByField(gc *conf.GlobalConfiguration, field string) *float64 { return &gc.SAML.RateLimitAssertion case fieldSSO: return &gc.RateLimitSso + case fieldSCIM: + return &gc.RateLimitScim + case fieldSCIMIP: + return &gc.RateLimitScim case fieldToken: return &gc.RateLimitTokenRefresh case fieldVerify: diff --git a/internal/api/context.go b/internal/api/context.go index de87ba009b..ee1a704796 100644 --- a/internal/api/context.go +++ b/internal/api/context.go @@ -2,6 +2,7 @@ package api import ( "context" + "net/http" "net/url" "github.com/gofrs/uuid" @@ -15,20 +16,24 @@ var ( externalProviderTypeKey = ctxkey.New[string]("external_provider_type") externalProviderEmailOptionalKey = ctxkey.New[bool]("external_provider_allow_no_email") - tokenKey = ctxkey.New[*jwt.Token]("jwt") - inviteTokenKey = ctxkey.New[string]("invite_token") - signatureKey = ctxkey.New[string]("signature") - targetUserKey = ctxkey.New[*models.User]("target_user") - factorKey = ctxkey.New[*models.Factor]("factor") - sessionKey = ctxkey.New[*models.Session]("session") - externalReferrerKey = ctxkey.New[string]("external_referrer") - adminUserKey = ctxkey.New[*models.User]("admin_user") - oauthTokenKey = ctxkey.New[string]("oauth_token") // for OAuth1.0, also known as request token - oauthVerifierKey = ctxkey.New[string]("oauth_verifier") - ssoProviderKey = ctxkey.New[*models.SSOProvider]("sso_provider") - externalHostKey = ctxkey.New[*url.URL]("external_host") - oauthClientStateKey = ctxkey.New[uuid.UUID]("oauth_client_state_id") - flowStateContextKey = ctxkey.New[*models.FlowState]("flow_state") + tokenKey = ctxkey.New[*jwt.Token]("jwt") + inviteTokenKey = ctxkey.New[string]("invite_token") + signatureKey = ctxkey.New[string]("signature") + targetUserKey = ctxkey.New[*models.User]("target_user") + factorKey = ctxkey.New[*models.Factor]("factor") + sessionKey = ctxkey.New[*models.Session]("session") + externalReferrerKey = ctxkey.New[string]("external_referrer") + adminUserKey = ctxkey.New[*models.User]("admin_user") + oauthTokenKey = ctxkey.New[string]("oauth_token") // for OAuth1.0, also known as request token + oauthVerifierKey = ctxkey.New[string]("oauth_verifier") + ssoProviderKey = ctxkey.New[*models.SSOProvider]("sso_provider") + externalHostKey = ctxkey.New[*url.URL]("external_host") + oauthClientStateKey = ctxkey.New[uuid.UUID]("oauth_client_state_id") + flowStateContextKey = ctxkey.New[*models.FlowState]("flow_state") + scimRequestKey = ctxkey.New[*http.Request]("scim_request") + scimGetQueryKey = ctxkey.New[url.Values]("scim_get_query") + scimSSOProviderIDKey = ctxkey.New[uuid.UUID]("scim_sso_provider_id") + scimTokenPrefixKey = ctxkey.New[string]("scim_token_prefix") ) // withToken adds the JWT token to the context. diff --git a/internal/api/external.go b/internal/api/external.go index f32420ee75..63207d5187 100644 --- a/internal/api/external.go +++ b/internal/api/external.go @@ -311,11 +311,27 @@ func (a *API) createAccountFromExternalIdentity(tx *storage.Connection, r *http. identityData = structs.Map(userData.Metadata) } + id, isSSO := strings.CutPrefix(providerType, "sso:") + scimOn := isSSO && config.SSO.SCIM.Enabled + + if scimOn && userData.Metadata.Email != "" { + if terr := models.LockAccountLinking(tx, providerType, userData.Metadata.Email); terr != nil { + return 0, nil, terr + } + } + decision, terr := models.DetermineAccountLinking(tx, config, userData.Emails, aud, providerType, userData.Metadata.Subject) if terr != nil { return 0, nil, terr } + var ssoProviderID uuid.UUID + if scimOn { + if ssoProviderID, terr = uuid.FromString(id); terr != nil { + return 0, nil, apierrors.NewInternalServerError("Invalid SSO provider id in provider type").WithInternalError(terr) + } + } + switch decision.Decision { case models.LinkAccount: user = decision.User @@ -407,6 +423,16 @@ func (a *API) createAccountFromExternalIdentity(tx *storage.Connection, r *http. return 0, nil, apierrors.NewForbiddenError(apierrors.ErrorCodeUserBanned, "User is banned") } + if scimOn { + deprovisioned, terr := models.IsSCIMDeprovisioned(tx, ssoProviderID, user.ID) + if terr != nil { + return 0, nil, terr + } + if deprovisioned { + return 0, nil, apierrors.NewForbiddenError(apierrors.ErrorCodeUserBanned, "User is banned") + } + } + hasEmails := providerType != Web3Provider && (!emailOptional || decision.CandidateEmail.Email != "") if hasEmails && !user.IsConfirmed() { diff --git a/internal/api/external_test.go b/internal/api/external_test.go index 914f14785f..f714398927 100644 --- a/internal/api/external_test.go +++ b/internal/api/external_test.go @@ -16,6 +16,7 @@ import ( "github.com/supabase/auth/internal/api/provider" "github.com/supabase/auth/internal/conf" "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/storage" ) type ExternalTestSuite struct { @@ -92,6 +93,67 @@ func (ts *ExternalTestSuite) TestAutomaticLinkIdentityWritesAuditLog() { require.Len(ts.T(), logs, 1, "signing in with an existing identity must not emit another audit log") } +func (ts *ExternalTestSuite) TestSSOConcurrentCreateSameEmailLinksToOneUser() { + ts.API.config.SSO.SCIM.Enabled = true + defer func() { ts.API.config.SSO.SCIM.Enabled = false }() + ssoProvider := &models.SSOProvider{} + require.NoError(ts.T(), ts.API.db.Create(ssoProvider)) + providerType := "sso:" + ssoProvider.ID.String() + + userData := func(sub string) *provider.UserProvidedData { + return &provider.UserProvidedData{ + Metadata: &provider.Claims{ + Subject: sub, + Email: "sso-race@example.com", + EmailVerified: true, + }, + Emails: []provider.Email{{ + Email: "sso-race@example.com", + Primary: true, + Verified: true, + }}, + } + } + r := httptest.NewRequest(http.MethodPost, "/sso/saml/acs", nil) + + subs := []string{"sub-a", "sub-b"} + decisions := make([]models.AccountLinkingDecision, len(subs)) + errs := make([]error, len(subs)) + + var wg sync.WaitGroup + start := make(chan struct{}) + for i, sub := range subs { + wg.Add(1) + go func(i int, sub string) { + defer wg.Done() + <-start + errs[i] = ts.API.db.Transaction(func(tx *storage.Connection) error { + decision, _, terr := ts.API.createAccountFromExternalIdentity(tx, r, userData(sub), providerType, false) + decisions[i] = decision + return terr + }) + }(i, sub) + } + close(start) + wg.Wait() + + for _, err := range errs { + require.NoError(ts.T(), err) + } + + created := 0 + for _, decision := range decisions { + if decision == models.CreateAccount { + created++ + } + } + require.Equal(ts.T(), 1, created, "the lock must serialize concurrent SSO creates for the same email so only one account is created") + + count, err := ts.API.db.Q().Where("email = ?", "sso-race@example.com").Count(&models.User{}) + require.NoError(ts.T(), err) + require.EqualValues(ts.T(), 1, count) +} + func (ts *ExternalTestSuite) createUser(providerId string, email string, name string, avatar string, confirmationToken string) (*models.User, error) { // Cleanup existing user, if they already exist if u, _ := models.FindUserByEmailAndAudience(ts.API.db, email, ts.Config.JWT.Aud); u != nil { diff --git a/internal/api/helpers.go b/internal/api/helpers.go index bd6b0f13be..b1b9f1933d 100644 --- a/internal/api/helpers.go +++ b/internal/api/helpers.go @@ -48,6 +48,7 @@ func (a *API) requestAud(ctx context.Context, r *http.Request) string { type RequestParams interface { AdminUserParams | AdminCustomOAuthProviderParams | + AdminSCIMTokenCreateParams | CreateSSOProviderParams | EnrollFactorParams | GenerateLinkParams | diff --git a/internal/api/identity.go b/internal/api/identity.go index 40136a8fbb..8ddb8a8771 100644 --- a/internal/api/identity.go +++ b/internal/api/identity.go @@ -3,6 +3,7 @@ package api import ( "context" "net/http" + "strings" "github.com/fatih/structs" "github.com/go-chi/chi/v5" @@ -51,10 +52,23 @@ func (a *API) DeleteIdentity(w http.ResponseWriter, r *http.Request) error { if identityToBeDeleted == nil { return apierrors.NewUnprocessableEntityError(apierrors.ErrorCodeIdentityNotFound, "Identity doesn't exist") } - provider := identityToBeDeleted.Provider recipientEmail := user.GetEmail() err = db.Transaction(func(tx *storage.Connection) error { + if id, ok := strings.CutPrefix(identityToBeDeleted.Provider, "sso:"); ok && a.config.SSO.SCIM.Enabled { + if providerID, perr := uuid.FromString(id); perr == nil { + if terr := models.LockUserForSCIM(tx, user.ID); terr != nil { + return apierrors.NewInternalServerError("Database error locking user").WithInternalError(terr) + } + managed, terr := models.IsSCIMManaged(tx, providerID, user.ID) + if terr != nil { + return apierrors.NewInternalServerError("Database error finding SCIM user").WithInternalError(terr) + } + if managed { + return apierrors.NewUnprocessableEntityError(apierrors.ErrorCodeUserSSOManaged, "Identity is managed by SCIM provisioning") + } + } + } if terr := models.NewAuditLogEntry(config.AuditLog, r, tx, user, models.IdentityUnlinkAction, utilities.GetIPAddress(r), map[string]any{ "identity_id": identityToBeDeleted.ID, "provider": identityToBeDeleted.Provider, diff --git a/internal/api/middleware.go b/internal/api/middleware.go index dd61843344..03b5d78958 100644 --- a/internal/api/middleware.go +++ b/internal/api/middleware.go @@ -427,7 +427,7 @@ func (a *API) requirePasskeyEnabled(w http.ResponseWriter, req *http.Request) (c func (a *API) requireScimServerEnabled(w http.ResponseWriter, req *http.Request) (context.Context, error) { ctx := req.Context() - if !a.config.Experimental.ScimEnabled { + if !a.config.SSO.SCIM.Enabled { return nil, apierrors.NewNotFoundError(apierrors.ErrorCodeFeatureDisabled, "SCIM server is disabled") } return ctx, nil diff --git a/internal/api/router.go b/internal/api/router.go index 25fbae3d11..97c047f33c 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -37,6 +37,10 @@ func (r *router) Delete(pattern string, fn apiHandler) { r.chi.Delete(pattern, handler(fn)) } +func (r *router) Method(method, pattern string, h http.Handler) { + r.chi.Method(method, pattern, h) +} + func (r *router) With(fn middlewareHandler) *router { c := r.chi.With(middleware(fn)) return &router{c} diff --git a/internal/api/scim.go b/internal/api/scim.go new file mode 100644 index 0000000000..7d65dd1537 --- /dev/null +++ b/internal/api/scim.go @@ -0,0 +1,377 @@ +package api + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strconv" + "strings" + "time" + + "github.com/didip/tollbooth/v5" + "github.com/didip/tollbooth/v5/limiter" + "github.com/go-chi/chi/v5" + "github.com/gofrs/uuid" + "github.com/supabase-community/scim-go/pkg/core" + "github.com/supabase-community/scim-go/pkg/protocol" + "github.com/supabase-community/scim-go/pkg/scimerrors" + "github.com/supabase-community/scim-go/pkg/server" + "github.com/supabase/auth/internal/api/apierrors" + "github.com/supabase/auth/internal/conf" + "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/observability" + "github.com/supabase/auth/internal/storage" + "github.com/supabase/auth/internal/utilities" +) + +const ( + scimBasePath = "/scim/v2" + scimResourceTypeUser = "User" + scimResourceTypeGroup = "Group" +) + +var errMissingSSOProvider = errors.New("scim: request has no SSO provider") + +func newSCIMServer(config *conf.GlobalConfiguration, validate server.TokenValidator, limit func(http.Handler) http.Handler, users server.Repository[*core.User], groups server.Repository[*core.Group]) *server.Server { + requireToken := server.RequireBearerToken(validate) + authenticate := func(next http.Handler) http.Handler { + next = scimRejectRemoveWithValue(scimRememberGetQuery(next)) + if limit != nil { + next = limit(next) + } + return requireToken(next) + } + return server.New(scimBasePath, + core.NewServiceProviderConfig().Filtering(protocol.DefaultLimits.MaxCount).Patching().Sorting().Versioning(), + server.WithBaseURL(scimBaseURL(config)), + server.ErrorHandler(scimLogError), + server.WithResource(server.NewResource[*core.User](scimResourceTypeUser, "/Users", core.SchemaUser, scimUserSchemas.Base().Attributes...). + WithExtension(core.SchemaEnterpriseUser, scimUserSchemas.Extensions()[0].Attributes...). + WithRepository(users)), + server.WithResource(server.NewResource[*core.Group](scimResourceTypeGroup, "/Groups", core.SchemaGroup, scimGroupSchemas.Base().Attributes...). + WithRepository(groups)), + server.WithAuthentication(core.NewOAuthBearerToken().AsPrimary(), authenticate), + ) +} + +var ( + scimUserSchemas = core.Schemas{ + core.NewSchema(core.SchemaUser).With(core.UserAttributes()...), + core.NewSchema(core.SchemaEnterpriseUser).With(core.EnterpriseUserAttributes()...), + } + scimGroupSchemas = newSCIMGroupSchemas() +) + +func newSCIMGroupSchemas() core.Schemas { + attributes := core.GroupAttributes() + for _, attribute := range attributes { + for _, sub := range attribute.SubAttributes { + switch sub.Name { + case "type": + sub.Suggesting(scimResourceTypeUser) + case "$ref": + sub.Referencing(scimResourceTypeUser) + } + } + } + return core.Schemas{core.NewSchema(core.SchemaGroup).With(attributes...)} +} + +func scimBaseURL(config *conf.GlobalConfiguration) string { + return strings.TrimRight(config.API.ExternalURL, "/") + scimBasePath +} + +func scimNotFound(w http.ResponseWriter, r *http.Request) error { + return protocol.SendError(w, scimerrors.ErrNotFound("Endpoint or resource does not exist")) +} + +func scimTooManyRequests(w http.ResponseWriter, r *http.Request) error { + return protocol.SendError(w, errSCIMTooManyRequests()) +} + +func scimLogError(r *http.Request, err error) { + observability.GetLogEntry(r).Entry.WithError(err).Error("scim: request failed") +} + +// RFC 7644 Section 3.5.2.2 defines no "value" for "remove"; refuse it rather than remove every value at "path". +func scimRejectRemoveWithValue(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPatch { + next.ServeHTTP(w, r) + return + } + var read bytes.Buffer + req, err := protocol.DecodePatchRequest(io.TeeReader(r.Body, &read)) + if err != nil { + _ = protocol.SendError(w, err) + return + } + for _, op := range req.Operations { + if strings.EqualFold(string(op.Op), "remove") && len(op.Value) > 0 && string(bytes.TrimSpace(op.Value)) != "null" { + _ = protocol.SendError(w, scimerrors.ErrInvalidSyntax(`"remove" does not take a "value"`)) + return + } + } + r.Body = io.NopCloser(io.MultiReader(&read, r.Body)) + next.ServeHTTP(w, r) + }) +} + +func scimRememberGetQuery(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + r = r.WithContext(scimGetQueryKey.WithValue(r.Context(), r.URL.Query())) + } + next.ServeHTTP(w, r) + }) +} + +func newSCIMTokenValidator(db *storage.Connection) server.TokenValidator { + return func(ctx context.Context, candidate string) (context.Context, error) { + token, err := models.AuthenticateSCIMToken(db.WithContext(ctx), candidate) + if models.IsNotFoundError(err) { + return ctx, server.ErrInvalidToken + } + if err != nil { + return ctx, err + } + ctx = scimTokenPrefixKey.WithValue(ctx, token.Prefix) + return scimSSOProviderIDKey.WithValue(ctx, token.SSOProviderID), nil + } +} + +func (a *API) withSCIMRequest(w http.ResponseWriter, req *http.Request) (context.Context, error) { + return scimRequestKey.WithValue(req.Context(), req), nil +} + +func (a *API) limitSCIMHandler(lmt *limiter.Limiter, routes chi.Routes) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !(scimHasBearerToken(r) && scimValidatesToken(routes, r)) && a.performRateLimiting(lmt, r) != nil { + handler(scimTooManyRequests)(w, r) + return + } + next.ServeHTTP(w, r) + }) + } +} + +func scimHasBearerToken(r *http.Request) bool { + scheme, token, _ := strings.Cut(r.Header.Get("Authorization"), " ") + return strings.EqualFold(scheme, "Bearer") && token != "" +} + +func scimValidatesToken(routes chi.Routes, r *http.Request) bool { + path := r.URL.Path + if rctx := chi.RouteContext(r.Context()); rctx != nil && rctx.RoutePath != "" { + path = rctx.RoutePath + } + return path != "/ServiceProviderConfig" && routes.Match(chi.NewRouteContext(), r.Method, path) +} + +func (a *API) limitSCIMInvalidToken(validate server.TokenValidator, lmt *limiter.Limiter) server.TokenValidator { + return func(ctx context.Context, candidate string) (context.Context, error) { + next, err := validate(ctx, candidate) + if !errors.Is(err, server.ErrInvalidToken) { + return next, err + } + if r := scimRequestKey.Value(ctx); r != nil && a.performRateLimiting(lmt, r) != nil { + return ctx, errSCIMTooManyRequests() + } + return next, err + } +} + +func (a *API) limitSCIMProvider(lmt *limiter.Limiter) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + providerID, ok := scimSSOProviderIDKey.Lookup(r.Context()) + if !ok { + next.ServeHTTP(w, r) + return + } + if err := tollbooth.LimitByKeys(lmt, []string{providerID.String()}); err != nil { + handler(scimTooManyRequests)(w, r) + return + } + next.ServeHTTP(w, r) + }) + } +} + +func scimSearch(query *protocol.SearchRequest, schemas core.Schemas, name string) (models.SCIMQuery, error) { + search := models.SCIMQuery{Offset: query.Offset(), Limit: query.Count} + if query.Filter != "" { + criteria, err := protocol.Filter(schemas, query.Filter, scimEqFilter{name: name}) + if err != nil { + return search, err + } + search.Filter = criteria + } + if query.SortBy == "" { + return search, nil + } + parent, attribute, err := query.SortAttribute(schemas) + if err != nil { + return search, err + } + key := parent.Name + if attribute != parent { + key += "." + attribute.Name + } + by, ok := map[string]models.SCIMSortKey{ + "id": models.SCIMSortByID, + strings.ToLower(name): models.SCIMSortByName, + "meta.created": models.SCIMSortByCreatedAt, + "meta.lastmodified": models.SCIMSortByUpdatedAt, + }[strings.ToLower(key)] + if !ok { + return search, scimerrors.ErrInvalidValue(fmt.Sprintf(`"sortBy" must be one of "id", %q, "meta.created" or "meta.lastModified"`, name)) + } + search.Order = models.SCIMOrder{By: by, Descending: query.Descending()} + return search, nil +} + +func scimReturns(projection protocol.Projection, name string) bool { + document, err := json.Marshal(projection.Of(map[string]any{name: []any{map[string]any{"value": name}}})) + if err != nil { + return true + } + var projected map[string]any + if err := json.Unmarshal(document, &projected); err != nil { + return true + } + _, ok := projected[name] + return ok +} + +func scimEncode(resource core.Resource, drop ...string) ([]byte, error) { + fields, err := core.NewObject(resource) + if err != nil { + return nil, err + } + for _, key := range drop { + fields.Remove(key) + } + return json.Marshal(fields) +} + +func scimTarget(ctx context.Context, id, version string) (providerID, resourceID uuid.UUID, updatedAt *time.Time, err error) { + if providerID, err = scimProviderID(ctx); err != nil { + return uuid.Nil, uuid.Nil, nil, err + } + if resourceID, err = uuid.FromString(id); err != nil { + return uuid.Nil, uuid.Nil, nil, errSCIMNotFound() + } + if updatedAt, err = scimParseVersion(version); err != nil { + return uuid.Nil, uuid.Nil, nil, err + } + return providerID, resourceID, updatedAt, nil +} + +func scimProviderID(ctx context.Context) (uuid.UUID, error) { + providerID, ok := scimSSOProviderIDKey.Lookup(ctx) + if !ok || providerID == uuid.Nil { + return uuid.Nil, errMissingSSOProvider + } + return providerID, nil +} + +func scimRequest(ctx context.Context) (*http.Request, error) { + r := scimRequestKey.Value(ctx) + if r == nil { + return nil, apierrors.NewInternalServerError("SCIM request missing from context") + } + return r.WithContext(ctx), nil +} + +func scimMeta(resourceType core.ResourceTypeName, location string, created, updated time.Time) core.Meta { + return core.Meta{ + ResourceType: resourceType, + Created: created.UTC(), + LastModified: updated.UTC(), + Location: location, + Version: scimVersion(updated), + } +} + +func scimVersion(updatedAt time.Time) string { + return `W/"` + strconv.FormatInt(updatedAt.UnixMicro(), 10) + `"` +} + +func scimParseVersion(version string) (*time.Time, error) { + if version == "" { + return nil, nil + } + micros, err := strconv.ParseInt(strings.TrimSuffix(strings.TrimPrefix(version, `W/"`), `"`), 10, 64) + if err != nil { + return nil, errSCIMStale() + } + updatedAt := time.UnixMicro(micros) + return &updatedAt, nil +} + +func scimTranslate(err error) error { + switch { + case err == nil: + return nil + case models.IsNotFoundError(err): + return errSCIMNotFound() + case errors.Is(err, models.SCIMUserStaleError{}), errors.Is(err, models.SCIMGroupStaleError{}): + return errSCIMStale() + case errors.Is(err, models.SCIMGroupConflictError{}): + return scimerrors.ErrUniqueness(`"externalId" must be unique`) + case errors.As(err, &models.SCIMGroupMemberNotFoundError{}): + return errSCIMMemberNotFound() + case errors.Is(err, models.SCIMUserConflictError{}): + return scimerrors.ErrUniqueness(`"userName" and "externalId" must be unique`) + case errors.Is(err, models.SCIMUserLinkedError{}): + return scimerrors.ErrUniqueness("user is already provisioned by this provider") + } + return err +} + +func errSCIMNotFound() error { + return scimerrors.ErrNotFound("Resource not found") +} + +func errSCIMStale() error { + return scimerrors.ErrPreconditionFailed("resource has changed on the server") +} + +func errSCIMMemberNotFound() error { + return scimerrors.ErrInvalidValue(`"members.value" must reference a User in this provider`) +} + +func errSCIMEmailRequired() error { + return scimerrors.ErrInvalidValue(`"emails" is required`) +} + +func errSCIMTooManyRequests() error { + return scimerrors.NewError(http.StatusTooManyRequests, "", "Request rate limit reached") +} + +func scimActor(r *http.Request) *models.User { + return &models.User{Email: storage.NullString("scim:" + scimTokenPrefixKey.Value(r.Context()))} +} + +func (a *API) auditSCIM(tx *storage.Connection, r *http.Request, actor *models.User, action models.AuditAction, providerID uuid.UUID, traits map[string]any) error { + traits["sso_provider_id"] = providerID + traits["outcome"] = "success" + return models.NewAuditLogEntry(a.config.AuditLog, r, tx, actor, action, utilities.GetIPAddress(r), traits) +} + +func (a *API) auditSCIMMember(tx *storage.Connection, r *http.Request, actor *models.User, action models.AuditAction, providerID, groupID, scimUserID uuid.UUID, userID *uuid.UUID) error { + traits := map[string]any{ + "scim_group_id": groupID, + "scim_user_id": scimUserID, + } + if userID != nil { + traits["user_id"] = *userID + } + return a.auditSCIM(tx, r, actor, action, providerID, traits) +} diff --git a/internal/api/scim/core/core.go b/internal/api/scim/core/core.go deleted file mode 100644 index d625dab4e1..0000000000 --- a/internal/api/scim/core/core.go +++ /dev/null @@ -1,8 +0,0 @@ -// Package core implements the SCIM 2.0 core schema defined in RFC 7643. -package core - -// SchemaURI identifies a SCIM schema -type SchemaURI string - -// ResourceTypeName names a resource type -type ResourceTypeName string diff --git a/internal/api/scim/core/endpoints.go b/internal/api/scim/core/endpoints.go deleted file mode 100644 index b1f9003dfb..0000000000 --- a/internal/api/scim/core/endpoints.go +++ /dev/null @@ -1,6 +0,0 @@ -package core - -// The resource endpoints of RFC 7644, Section 3.2, relative to the base URL -const ( - EndpointServiceProviderConfig = "/ServiceProviderConfig" -) diff --git a/internal/api/scim/core/meta.go b/internal/api/scim/core/meta.go deleted file mode 100644 index a47e4a4b30..0000000000 --- a/internal/api/scim/core/meta.go +++ /dev/null @@ -1,14 +0,0 @@ -package core - -// Meta is the resource metadata common attribute defined in RFC 7643, Section 3.1. -type Meta struct { - ResourceType ResourceTypeName `json:"resourceType"` - Location string `json:"location,omitempty"` -} - -func NewMeta(baseURL string, resourceType ResourceTypeName, endpoint string) Meta { - return Meta{ - ResourceType: resourceType, - Location: baseURL + endpoint, - } -} diff --git a/internal/api/scim/core/meta_test.go b/internal/api/scim/core/meta_test.go deleted file mode 100644 index 4b7383bd70..0000000000 --- a/internal/api/scim/core/meta_test.go +++ /dev/null @@ -1,39 +0,0 @@ -package core - -import ( - "encoding/json" - "testing" - - "github.com/stretchr/testify/require" -) - -func TestNewMeta(t *testing.T) { - t.Run("locates the resource at its endpoint", func(t *testing.T) { - meta := NewMeta("http://localhost:9999/scim/v2", ResourceTypeServiceProviderConfig, EndpointServiceProviderConfig) - - require.Equal(t, ResourceTypeServiceProviderConfig, meta.ResourceType) - require.Equal(t, "http://localhost:9999/scim/v2/ServiceProviderConfig", meta.Location) - }) -} - -func TestMeta(t *testing.T) { - t.Run("serializes to JSON correctly", func(t *testing.T) { - body, err := json.Marshal(Meta{ - ResourceType: ResourceTypeServiceProviderConfig, - Location: "http://localhost:9999/scim/v2/ServiceProviderConfig", - }) - - require.NoError(t, err) - require.JSONEq(t, `{ - "resourceType": "ServiceProviderConfig", - "location": "http://localhost:9999/scim/v2/ServiceProviderConfig" - }`, string(body)) - }) - - t.Run("omits the location when it is empty", func(t *testing.T) { - body, err := json.Marshal(Meta{ResourceType: ResourceTypeServiceProviderConfig}) - - require.NoError(t, err) - require.JSONEq(t, `{"resourceType": "ServiceProviderConfig"}`, string(body)) - }) -} diff --git a/internal/api/scim/core/schemas.go b/internal/api/scim/core/schemas.go deleted file mode 100644 index 128b2ea719..0000000000 --- a/internal/api/scim/core/schemas.go +++ /dev/null @@ -1,13 +0,0 @@ -package core - -// The schema URIs of RFC 7643 -const ( - schemaRoot = "urn:ietf:params:scim:schemas" - schemaCore = schemaRoot + ":core:2.0" - - SchemaServiceProviderConfig SchemaURI = schemaCore + ":ServiceProviderConfig" -) - -const ( - ResourceTypeServiceProviderConfig ResourceTypeName = "ServiceProviderConfig" -) diff --git a/internal/api/scim/core/service_provider_config.go b/internal/api/scim/core/service_provider_config.go deleted file mode 100644 index 26c64da947..0000000000 --- a/internal/api/scim/core/service_provider_config.go +++ /dev/null @@ -1,70 +0,0 @@ -package core - -type SupportedFeature struct { - Supported bool `json:"supported"` -} - -type BulkFeature struct { - Supported bool `json:"supported"` - MaxOperations int `json:"maxOperations"` - MaxPayloadSize int `json:"maxPayloadSize"` -} - -type FilterFeature struct { - Supported bool `json:"supported"` - MaxResults int `json:"maxResults"` -} - -type AuthenticationSchemeType string - -const ( - AuthenticationSchemeOAuthBearerToken AuthenticationSchemeType = "oauthbearertoken" -) - -// AuthenticationScheme is the authentication scheme of RFC 7643, Section 5. -type AuthenticationScheme struct { - Type AuthenticationSchemeType `json:"type"` - Name string `json:"name"` - Description string `json:"description"` - SpecURI string `json:"specUri,omitempty"` - Primary bool `json:"primary"` -} - -func NewOAuthBearerToken() *AuthenticationScheme { - return &AuthenticationScheme{ - Type: AuthenticationSchemeOAuthBearerToken, - Name: "OAuth Bearer Token", - Description: "Authentication scheme using the OAuth Bearer Token Standard", - SpecURI: "http://www.rfc-editor.org/info/rfc6750", - } -} - -func (scheme *AuthenticationScheme) AsPrimary() *AuthenticationScheme { - scheme.Primary = true - return scheme -} - -// ServiceProviderConfig is the schema defined in RFC 7643, Section 5. -type ServiceProviderConfig struct { - Schemas []SchemaURI `json:"schemas"` - Patch SupportedFeature `json:"patch"` - Bulk BulkFeature `json:"bulk"` - Filter FilterFeature `json:"filter"` - ChangePassword SupportedFeature `json:"changePassword"` - Sort SupportedFeature `json:"sort"` - ETag SupportedFeature `json:"etag"` - AuthenticationSchemes []*AuthenticationScheme `json:"authenticationSchemes"` - Meta Meta `json:"meta"` -} - -func NewServiceProviderConfig(baseURL string, schemes ...*AuthenticationScheme) *ServiceProviderConfig { - if schemes == nil { - schemes = []*AuthenticationScheme{} - } - - return &ServiceProviderConfig{ - Schemas: []SchemaURI{SchemaServiceProviderConfig}, - AuthenticationSchemes: schemes, - Meta: NewMeta(baseURL, ResourceTypeServiceProviderConfig, EndpointServiceProviderConfig), - } -} diff --git a/internal/api/scim/core/service_provider_config_test.go b/internal/api/scim/core/service_provider_config_test.go deleted file mode 100644 index 03ff2dca95..0000000000 --- a/internal/api/scim/core/service_provider_config_test.go +++ /dev/null @@ -1,66 +0,0 @@ -package core - -import ( - "encoding/json" - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -func TestNewServiceProviderConfig(t *testing.T) { - t.Run("advertises the schemes the caller declares", func(t *testing.T) { - scheme := NewOAuthBearerToken().AsPrimary() - - config := NewServiceProviderConfig("", scheme) - - require.Equal(t, []SchemaURI{SchemaServiceProviderConfig}, config.Schemas) - require.Equal(t, []*AuthenticationScheme{scheme}, config.AuthenticationSchemes) - }) - - t.Run("identifies itself with resource metadata", func(t *testing.T) { - baseURL := "http://localhost:9999/scim/v2" - - config := NewServiceProviderConfig(baseURL) - - require.Equal(t, ResourceTypeServiceProviderConfig, config.Meta.ResourceType) - require.Equal(t, baseURL+EndpointServiceProviderConfig, config.Meta.Location) - }) - - t.Run("supports none of the optional protocol features", func(t *testing.T) { - config := NewServiceProviderConfig("") - - assert.False(t, config.Patch.Supported) - assert.False(t, config.Bulk.Supported) - assert.False(t, config.Filter.Supported) - assert.False(t, config.ChangePassword.Supported) - assert.False(t, config.Sort.Supported) - assert.False(t, config.ETag.Supported) - }) - - t.Run("serializes authenticationSchemes as an array", func(t *testing.T) { - body, err := json.Marshal(NewServiceProviderConfig("")) - - require.NoError(t, err) - require.Contains(t, string(body), `"authenticationSchemes":[]`) - }) -} - -func TestAuthenticationScheme(t *testing.T) { - t.Run("NewOAuthBearerToken", func(t *testing.T) { - scheme := NewOAuthBearerToken() - - assert.Equal(t, AuthenticationSchemeOAuthBearerToken, scheme.Type) - assert.Equal(t, "OAuth Bearer Token", scheme.Name) - assert.Equal(t, "Authentication scheme using the OAuth Bearer Token Standard", scheme.Description) - assert.Equal(t, "http://www.rfc-editor.org/info/rfc6750", scheme.SpecURI) - assert.False(t, scheme.Primary) - }) - - t.Run("AsPrimary marks the scheme primary", func(t *testing.T) { - scheme := NewOAuthBearerToken() - - require.Same(t, scheme, scheme.AsPrimary()) - assert.True(t, scheme.Primary) - }) -} diff --git a/internal/api/scim/protocol/error.go b/internal/api/scim/protocol/error.go deleted file mode 100644 index fb183692f1..0000000000 --- a/internal/api/scim/protocol/error.go +++ /dev/null @@ -1,24 +0,0 @@ -package protocol - -import ( - "strconv" -) - -const SchemaError = "urn:ietf:params:scim:api:messages:2.0:Error" - -// Error is the error message form defined in RFC 7644, Section 3.12. -type Error struct { - Schemas []string `json:"schemas"` - ScimType string `json:"scimType,omitempty"` - Detail string `json:"detail,omitempty"` - Status string `json:"status"` -} - -func NewError(status int, scimType string, detail string) *Error { - return &Error{ - Schemas: []string{SchemaError}, - ScimType: scimType, - Detail: detail, - Status: strconv.Itoa(status), - } -} diff --git a/internal/api/scim/protocol/error_test.go b/internal/api/scim/protocol/error_test.go deleted file mode 100644 index aede262a84..0000000000 --- a/internal/api/scim/protocol/error_test.go +++ /dev/null @@ -1,47 +0,0 @@ -package protocol - -import ( - "encoding/json" - "net/http" - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -func TestNewError(t *testing.T) { - t.Run("serializes to JSON correctly", func(t *testing.T) { - body, err := json.Marshal(NewError(http.StatusNotFound, "", "Endpoint or resource does not exist")) - - require.NoError(t, err) - assert.JSONEq(t, `{ - "schemas": [ - "urn:ietf:params:scim:api:messages:2.0:Error" - ], - "status": "404", - "detail": "Endpoint or resource does not exist" - }`, string(body)) - }) - - t.Run("includes the scimType when one is given", func(t *testing.T) { - body, err := json.Marshal(NewError(http.StatusBadRequest, "invalidValue", "A required value was missing")) - - require.NoError(t, err) - assert.JSONEq(t, `{ - "schemas": ["urn:ietf:params:scim:api:messages:2.0:Error"], - "scimType": "invalidValue", - "detail": "A required value was missing", - "status": "400" - }`, string(body)) - }) - - t.Run("omits the optional attributes when they are empty", func(t *testing.T) { - body, err := json.Marshal(NewError(http.StatusBadRequest, "", "")) - - require.NoError(t, err) - assert.JSONEq(t, `{ - "schemas": ["urn:ietf:params:scim:api:messages:2.0:Error"], - "status": "400" - }`, string(body)) - }) -} diff --git a/internal/api/scim/protocol/list_response.go b/internal/api/scim/protocol/list_response.go deleted file mode 100644 index 972229f71f..0000000000 --- a/internal/api/scim/protocol/list_response.go +++ /dev/null @@ -1,25 +0,0 @@ -package protocol - -const SchemaListResponse = "urn:ietf:params:scim:api:messages:2.0:ListResponse" - -type ListResponse[T any] struct { - Schemas []string `json:"schemas"` - TotalResults int `json:"totalResults"` - StartIndex int `json:"startIndex"` - ItemsPerPage int `json:"itemsPerPage"` - Resources []T `json:"Resources"` -} - -func NewListResponse[T any](resources []T) *ListResponse[T] { - if resources == nil { - resources = []T{} - } - n := len(resources) - return &ListResponse[T]{ - Schemas: []string{SchemaListResponse}, - TotalResults: n, - StartIndex: 1, - ItemsPerPage: n, - Resources: resources, - } -} diff --git a/internal/api/scim/protocol/list_response_test.go b/internal/api/scim/protocol/list_response_test.go deleted file mode 100644 index 6c4de87bdd..0000000000 --- a/internal/api/scim/protocol/list_response_test.go +++ /dev/null @@ -1,53 +0,0 @@ -package protocol - -import ( - "encoding/json" - "testing" - - "github.com/stretchr/testify/require" -) - -const emptyListResponse = `{ - "schemas": ["urn:ietf:params:scim:api:messages:2.0:ListResponse"], - "totalResults": 0, - "startIndex": 1, - "itemsPerPage": 0, - "Resources": [] -}` - -func TestNewListResponse(t *testing.T) { - for _, tc := range []struct { - name string - resources []string - expected string - }{ - { - name: "nil resources marshal to an empty array", - resources: nil, - expected: emptyListResponse, - }, - { - name: "empty resources marshal to an empty array", - resources: []string{}, - expected: emptyListResponse, - }, - { - name: "populated resources are counted", - resources: []string{"a", "b"}, - expected: `{ - "schemas": ["urn:ietf:params:scim:api:messages:2.0:ListResponse"], - "totalResults": 2, - "startIndex": 1, - "itemsPerPage": 2, - "Resources": ["a", "b"] - }`, - }, - } { - t.Run(tc.name, func(t *testing.T) { - body, err := json.Marshal(NewListResponse(tc.resources)) - - require.NoError(t, err) - require.JSONEq(t, tc.expected, string(body)) - }) - } -} diff --git a/internal/api/scim/protocol/protocol.go b/internal/api/scim/protocol/protocol.go deleted file mode 100644 index 7e3f4fbffb..0000000000 --- a/internal/api/scim/protocol/protocol.go +++ /dev/null @@ -1,18 +0,0 @@ -// Package protocol implements the SCIM 2.0 protocol defined in RFC 7644. -package protocol - -import ( - "net/http" - - "github.com/supabase/auth/internal/api/shared" -) - -const MediaType = "application/scim+json" - -func Send(w http.ResponseWriter, status int, obj any) error { - return shared.JSON(w).ContentType(MediaType).Status(status).Send(obj) -} - -func SendError(w http.ResponseWriter, status int, scimType string, detail string) error { - return Send(w, status, NewError(status, scimType, detail)) -} diff --git a/internal/api/scim/protocol/protocol_test.go b/internal/api/scim/protocol/protocol_test.go deleted file mode 100644 index a23ec040ef..0000000000 --- a/internal/api/scim/protocol/protocol_test.go +++ /dev/null @@ -1,23 +0,0 @@ -package protocol - -import ( - "net/http" - "net/http/httptest" - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -func TestSend(t *testing.T) { - t.Run("writes a JSON response with a SCIM media type", func(t *testing.T) { - w := httptest.NewRecorder() - - err := Send(w, http.StatusTeapot, map[string]string{"key": "value"}) - require.NoError(t, err) - - assert.Equal(t, http.StatusTeapot, w.Code) - assert.Equal(t, "application/scim+json", w.Header().Get("Content-Type")) - assert.JSONEq(t, `{"key":"value"}`, w.Body.String()) - }) -} diff --git a/internal/api/scim/server.go b/internal/api/scim/server.go deleted file mode 100644 index 4de38e4ce3..0000000000 --- a/internal/api/scim/server.go +++ /dev/null @@ -1,48 +0,0 @@ -package scim - -import ( - "net/http" - "strings" - - "github.com/supabase/auth/internal/api/scim/core" - "github.com/supabase/auth/internal/api/scim/protocol" - "github.com/supabase/auth/internal/conf" -) - -const BasePath = "/scim/v2" - -type Server struct { - serviceProviderConfig *core.ServiceProviderConfig -} - -func NewServer(config *conf.GlobalConfiguration) *Server { - return &Server{ - serviceProviderConfig: core.NewServiceProviderConfig( - strings.TrimRight(config.API.ExternalURL, "/")+BasePath, - core.NewOAuthBearerToken().AsPrimary(), - ), - } -} - -func (srv *Server) ServiceProviderConfig(w http.ResponseWriter, r *http.Request) error { - return protocol.Send(w, http.StatusOK, srv.serviceProviderConfig) -} - -func (srv *Server) ResourceTypes(w http.ResponseWriter, r *http.Request) error { - return list(w, r, []any{}) -} - -func (srv *Server) Schemas(w http.ResponseWriter, r *http.Request) error { - return list(w, r, []any{}) -} - -func (srv *Server) NotFound(w http.ResponseWriter, r *http.Request) error { - return protocol.SendError(w, http.StatusNotFound, "", "Endpoint or resource does not exist") -} - -func list[T any](w http.ResponseWriter, r *http.Request, resources []T) error { - if r.URL.Query().Has("filter") { - return protocol.SendError(w, http.StatusForbidden, "", "Filtering is not supported on this endpoint") - } - return protocol.Send(w, http.StatusOK, protocol.NewListResponse(resources)) -} diff --git a/internal/api/scim/server_test.go b/internal/api/scim/server_test.go deleted file mode 100644 index 773638bcdd..0000000000 --- a/internal/api/scim/server_test.go +++ /dev/null @@ -1,91 +0,0 @@ -package scim - -import ( - "embed" - "net/http" - "net/http/httptest" - "net/url" - "testing" - - "github.com/stretchr/testify/require" - "github.com/supabase/auth/internal/api/scim/protocol" - "github.com/supabase/auth/internal/conf" -) - -//go:embed testdata/* -var fixtures embed.FS - -func testFixture(t *testing.T, file string) string { - data, err := fixtures.ReadFile("testdata/" + file) - require.NoError(t, err) - return string(data) -} - -func newServerFor(externalURL string) *Server { - return NewServer(&conf.GlobalConfiguration{ - API: conf.APIConfiguration{ExternalURL: externalURL}, - }) -} - -func TestServer(t *testing.T) { - srv := newServerFor("http://localhost:9999") - require.NotNil(t, srv) - - t.Run("NewServer trims a trailing slash from the external URL", func(t *testing.T) { - location := newServerFor("https://auth.example.com/").serviceProviderConfig.Meta.Location - - require.Equal(t, "https://auth.example.com"+BasePath+"/ServiceProviderConfig", location) - }) - - t.Run("ServiceProviderConfig", func(t *testing.T) { - r := httptest.NewRequest(http.MethodGet, BasePath+"/ServiceProviderConfig", nil) - w := httptest.NewRecorder() - - require.NoError(t, srv.ServiceProviderConfig(w, r)) - - require.Equal(t, http.StatusOK, w.Code) - require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) - require.JSONEq(t, testFixture(t, "service_provider_config.json"), w.Body.String()) - }) - - for _, tc := range []struct { - path string - handler func(http.ResponseWriter, *http.Request) error - }{ - {"ResourceTypes", srv.ResourceTypes}, - {"Schemas", srv.Schemas}, - } { - t.Run(tc.path, func(t *testing.T) { - r := httptest.NewRequest(http.MethodGet, BasePath+"/"+tc.path, nil) - w := httptest.NewRecorder() - - require.NoError(t, tc.handler(w, r)) - - require.Equal(t, http.StatusOK, w.Code) - require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) - require.JSONEq(t, testFixture(t, "empty_list_response.json"), w.Body.String()) - }) - - t.Run(tc.path+" rejects filter query parameter", func(t *testing.T) { - filter := url.Values{"filter": {`name eq "User"`}}.Encode() - r := httptest.NewRequest(http.MethodGet, BasePath+"/"+tc.path+"?"+filter, nil) - w := httptest.NewRecorder() - - require.NoError(t, tc.handler(w, r)) - - require.Equal(t, http.StatusForbidden, w.Code) - require.JSONEq(t, testFixture(t, "filter_forbidden.json"), w.Body.String()) - }) - } - - t.Run("NotFound", func(t *testing.T) { - r := httptest.NewRequest(http.MethodGet, BasePath+"/Unknown", nil) - w := httptest.NewRecorder() - - require.NoError(t, srv.NotFound(w, r)) - - require.Equal(t, http.StatusNotFound, w.Code) - require.Equal(t, "application/scim+json", w.Header().Get("Content-Type")) - require.JSONEq(t, testFixture(t, "not_found.json"), w.Body.String()) - }) -} diff --git a/internal/api/scim/testdata/empty_list_response.json b/internal/api/scim/testdata/empty_list_response.json deleted file mode 100644 index d13e376c64..0000000000 --- a/internal/api/scim/testdata/empty_list_response.json +++ /dev/null @@ -1,9 +0,0 @@ -{ - "schemas": [ - "urn:ietf:params:scim:api:messages:2.0:ListResponse" - ], - "totalResults": 0, - "startIndex": 1, - "itemsPerPage": 0, - "Resources": [] -} diff --git a/internal/api/scim/testdata/not_implemented.json b/internal/api/scim/testdata/not_implemented.json deleted file mode 100644 index 416f929734..0000000000 --- a/internal/api/scim/testdata/not_implemented.json +++ /dev/null @@ -1,7 +0,0 @@ -{ - "schemas": [ - "urn:ietf:params:scim:api:messages:2.0:Error" - ], - "status": "501", - "detail": "The request endpoint is not implemented" -} diff --git a/internal/api/scim_admin.go b/internal/api/scim_admin.go new file mode 100644 index 0000000000..887c3adb1e --- /dev/null +++ b/internal/api/scim_admin.go @@ -0,0 +1,216 @@ +package api + +import ( + "errors" + "net/http" + "time" + + "github.com/go-chi/chi/v5" + "github.com/gofrs/uuid" + "github.com/supabase/auth/internal/api/apierrors" + "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/storage" + "github.com/supabase/auth/internal/utilities" +) + +const scimProviderDeletedBan = 100 * 365 * 24 * time.Hour + +type AdminSCIMTokenCreateParams struct { + ExpiresAt *time.Time `json:"expires_at"` +} + +type AdminSCIMTokenCreateResponse struct { + BaseURL string `json:"base_url"` + Token string `json:"token"` + *models.SCIMToken +} + +type AdminSCIMTokenListResponse struct { + Tokens []models.SCIMToken `json:"tokens"` +} + +type AdminSCIMStatusResponse struct { + Enabled bool `json:"enabled"` + BaseURL string `json:"base_url"` + Tokens []models.SCIMToken `json:"tokens"` +} + +func (a *API) adminSCIMGet(w http.ResponseWriter, r *http.Request) error { + ctx := r.Context() + return a.sendSCIMStatus(w, a.db.WithContext(ctx), getSSOProvider(ctx)) +} + +func (a *API) adminSCIMEnable(w http.ResponseWriter, r *http.Request) error { + return a.changeSCIMEnabled(w, r, models.EnableSCIM, models.SCIMEnabledAction, "enabling") +} + +func (a *API) adminSCIMDisable(w http.ResponseWriter, r *http.Request) error { + return a.changeSCIMEnabled(w, r, models.DisableSCIM, models.SCIMDisabledAction, "disabling") +} + +func (a *API) changeSCIMEnabled(w http.ResponseWriter, r *http.Request, change func(*storage.Connection, uuid.UUID) (bool, error), action models.AuditAction, verb string) error { + ctx := r.Context() + db := a.db.WithContext(ctx) + provider := getSSOProvider(ctx) + + if err := db.Transaction(func(tx *storage.Connection) error { + changed, err := change(tx, provider.ID) + if err != nil || !changed { + return err + } + return a.auditSCIM(tx, r, getAdminUser(ctx), action, provider.ID, map[string]any{}) + }); err != nil { + return apierrors.NewInternalServerError("Error %s SCIM", verb).WithInternalError(err) + } + + return a.sendSCIMStatus(w, db, provider) +} + +func (a *API) sendSCIMStatus(w http.ResponseWriter, db *storage.Connection, provider *models.SSOProvider) error { + tokens, err := models.FindActiveSCIMTokensBySSOProvider(db, provider.ID) + if err != nil { + return apierrors.NewInternalServerError("Error finding SCIM tokens").WithInternalError(err) + } + enabled, err := a.isSCIMEnabled(db, provider) + if err != nil { + return apierrors.NewInternalServerError("Error finding SCIM settings").WithInternalError(err) + } + + return sendJSON(w, http.StatusOK, &AdminSCIMStatusResponse{ + Enabled: enabled, + BaseURL: scimBaseURL(a.config), + Tokens: tokens, + }) +} + +func (a *API) isSCIMEnabled(db *storage.Connection, provider *models.SSOProvider) (bool, error) { + if !a.config.SSO.SCIM.Enabled || !provider.IsEnabled() { + return false, nil + } + return models.IsSCIMEnabled(db, provider.ID) +} + +func (a *API) deprovisionSCIM(tx *storage.Connection, r *http.Request, provider *models.SSOProvider) error { + disabled, err := models.DisableSCIM(tx, provider.ID) + if err != nil { + return err + } + prefixes, err := a.revokeSCIMTokens(tx, r, provider) + if err != nil { + return err + } + if disabled && a.config.SSO.SCIM.Enabled { + if err := a.auditSCIM(tx, r, getAdminUser(r.Context()), models.SCIMDisabledAction, provider.ID, map[string]any{"token_prefixes": prefixes}); err != nil { + return err + } + } + banned, err := models.BanDeprovisionedSCIMUsers(tx, provider.ID, a.Now().Add(scimProviderDeletedBan)) + if err != nil || banned == 0 { + return err + } + return a.auditSCIM(tx, r, getAdminUser(r.Context()), models.SCIMUsersBannedAction, provider.ID, map[string]any{"banned_user_count": banned}) +} + +func (a *API) revokeSCIMTokens(tx *storage.Connection, r *http.Request, provider *models.SSOProvider) ([]string, error) { + if err := models.LockSCIMTokens(tx, provider.ID); err != nil { + return nil, err + } + tokens, err := models.RevokeSCIMTokensBySSOProvider(tx, provider.ID) + if err != nil { + return nil, err + } + actor := getAdminUser(r.Context()) + prefixes := make([]string, len(tokens)) + for i := range tokens { + prefixes[i] = tokens[i].Prefix + if err := a.auditSCIM(tx, r, actor, models.SCIMTokenRevokedAction, provider.ID, map[string]any{"token_prefix": tokens[i].Prefix}); err != nil { + return nil, err + } + } + return prefixes, nil +} + +func (a *API) adminSCIMTokensCreate(w http.ResponseWriter, r *http.Request) error { + ctx := r.Context() + db := a.db.WithContext(ctx) + provider := getSSOProvider(ctx) + + params := &AdminSCIMTokenCreateParams{} + if body, err := utilities.GetBodyBytes(r); err != nil || len(body) > 0 { + if err := retrieveRequestParams(r, params); err != nil { + return err + } + } + if params.ExpiresAt != nil && !params.ExpiresAt.After(a.Now()) { + return apierrors.NewBadRequestError(apierrors.ErrorCodeValidationFailed, "expires_at must be in the future") + } + + var ( + token *models.SCIMToken + plaintext string + ) + if err := db.Transaction(func(tx *storage.Connection) error { + if err := models.LockSCIMTokens(tx, provider.ID); err != nil { + return err + } + var err error + if token, plaintext, err = models.CreateSCIMToken(tx, provider, params.ExpiresAt); err != nil { + return err + } + return a.auditSCIM(tx, r, getAdminUser(ctx), models.SCIMTokenCreatedAction, provider.ID, map[string]any{"token_prefix": token.Prefix}) + }); err != nil { + if errors.Is(err, models.SCIMTokenExpiryError{}) { + return apierrors.NewBadRequestError(apierrors.ErrorCodeValidationFailed, "expires_at must be in the future") + } + return apierrors.NewInternalServerError("Error creating SCIM token").WithInternalError(err) + } + + return sendJSON(w, http.StatusCreated, &AdminSCIMTokenCreateResponse{ + BaseURL: scimBaseURL(a.config), + Token: plaintext, + SCIMToken: token, + }) +} + +func (a *API) adminSCIMTokensList(w http.ResponseWriter, r *http.Request) error { + ctx := r.Context() + provider := getSSOProvider(ctx) + + tokens, err := models.FindSCIMTokensBySSOProvider(a.db.WithContext(ctx), provider.ID) + if err != nil { + return apierrors.NewInternalServerError("Error listing SCIM tokens").WithInternalError(err) + } + + return sendJSON(w, http.StatusOK, &AdminSCIMTokenListResponse{Tokens: tokens}) +} + +func (a *API) adminSCIMTokensRevoke(w http.ResponseWriter, r *http.Request) error { + ctx := r.Context() + db := a.db.WithContext(ctx) + provider := getSSOProvider(ctx) + + var token *models.SCIMToken + if err := db.Transaction(func(tx *storage.Connection) error { + if err := models.LockSCIMTokens(tx, provider.ID); err != nil { + return err + } + var err error + if token, err = models.FindSCIMTokenByPrefix(tx, provider.ID, chi.URLParam(r, "prefix")); err != nil { + return err + } + if token.IsRevoked() { + return nil + } + if err = token.Revoke(tx); err != nil { + return err + } + return a.auditSCIM(tx, r, getAdminUser(ctx), models.SCIMTokenRevokedAction, provider.ID, map[string]any{"token_prefix": token.Prefix}) + }); err != nil { + if models.IsNotFoundError(err) { + return apierrors.NewNotFoundError(apierrors.ErrorCodeSCIMTokenNotFound, "SCIM token not found") + } + return apierrors.NewInternalServerError("Error revoking SCIM token").WithInternalError(err) + } + + return sendJSON(w, http.StatusOK, token) +} diff --git a/internal/api/scim_admin_test.go b/internal/api/scim_admin_test.go new file mode 100644 index 0000000000..d0368d35fe --- /dev/null +++ b/internal/api/scim_admin_test.go @@ -0,0 +1,580 @@ +package api + +import ( + "bytes" + "context" + "encoding/json" + "maps" + "net/http" + "net/http/httptest" + "slices" + "strings" + "sync" + "testing" + "time" + + "github.com/gofrs/uuid" + jwt "github.com/golang-jwt/jwt/v5" + "github.com/stretchr/testify/require" + "github.com/stretchr/testify/suite" + "github.com/supabase-community/scim-go/pkg/server" + "github.com/supabase/auth/internal/conf" + "github.com/supabase/auth/internal/models" +) + +type SCIMTokensTestSuite struct { + suite.Suite + API *API + Config *conf.GlobalConfiguration + AdminJWT string + Provider *models.SSOProvider +} + +func TestSCIMTokens(t *testing.T) { + api, config := setupSCIMAPI(t, nil) + defer api.db.Close() + + suite.Run(t, &SCIMTokensTestSuite{API: api, Config: config}) +} + +func (ts *SCIMTokensTestSuite) SetupTest() { + require.NoError(ts.T(), models.TruncateAll(ts.API.db)) + ts.API.config.SSO.SCIM.Enabled = true + + token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, &AccessTokenClaims{Role: "supabase_admin"}).SignedString([]byte(ts.Config.JWT.Secret)) + require.NoError(ts.T(), err) + ts.AdminJWT = token + + ts.Provider = ts.createProvider() +} + +func (ts *SCIMTokensTestSuite) createProvider() *models.SSOProvider { + return createSCIMEnabledProvider(ts.T(), ts.API.db) +} + +func (ts *SCIMTokensTestSuite) tokensPath(provider *models.SSOProvider) string { + return "/admin/sso/providers/" + provider.ID.String() + "/scim/tokens" +} + +func (ts *SCIMTokensTestSuite) request(method, path string, body any) *httptest.ResponseRecorder { + var buf bytes.Buffer + if body != nil { + require.NoError(ts.T(), json.NewEncoder(&buf).Encode(body)) + } + r := httptest.NewRequest(method, path, &buf) + r.Header.Set("Authorization", "Bearer "+ts.AdminJWT) + r.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + return w +} + +func (ts *SCIMTokensTestSuite) create(provider *models.SSOProvider, body any) AdminSCIMTokenCreateResponse { + w := ts.request(http.MethodPost, ts.tokensPath(provider), body) + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) + + var response AdminSCIMTokenCreateResponse + require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &response)) + return response +} + +func (ts *SCIMTokensTestSuite) scimRequest(token string) *httptest.ResponseRecorder { + r := httptest.NewRequest(http.MethodGet, "/scim/v2/Users", nil) + r.Header.Set("Authorization", "Bearer "+token) + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + return w +} + +func (ts *SCIMTokensTestSuite) TestCreate() { + w := ts.request(http.MethodPost, ts.tokensPath(ts.Provider), map[string]any{}) + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) + + var body map[string]any + require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &body)) + require.ElementsMatch(ts.T(), []string{"base_url", "token", "prefix", "created_at", "expires_at", "revoked_at", "last_used_at"}, slices.Collect(maps.Keys(body))) + require.Equal(ts.T(), "http://localhost:9999/scim/v2", body["base_url"]) + require.Regexp(ts.T(), `^scim_[0-9a-f]{40}$`, body["token"]) + require.Equal(ts.T(), body["token"].(string)[:12], body["prefix"]) + require.Nil(ts.T(), body["expires_at"]) + require.Nil(ts.T(), body["revoked_at"]) + + require.Equal(ts.T(), http.StatusOK, ts.scimRequest(body["token"].(string)).Code) +} + +func (ts *SCIMTokensTestSuite) TestMultipleActiveTokens() { + first := ts.create(ts.Provider, map[string]any{}) + second := ts.create(ts.Provider, map[string]any{}) + + require.Equal(ts.T(), http.StatusOK, ts.scimRequest(first.Token).Code) + require.Equal(ts.T(), http.StatusOK, ts.scimRequest(second.Token).Code) +} + +func (ts *SCIMTokensTestSuite) TestCreateWithoutBody() { + r := httptest.NewRequest(http.MethodPost, ts.tokensPath(ts.Provider), nil) + r.Header.Set("Authorization", "Bearer "+ts.AdminJWT) + w := httptest.NewRecorder() + + ts.API.handler.ServeHTTP(w, r) + + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) +} + +func (ts *SCIMTokensTestSuite) TestTokenValidatorResolvesSSOProvider() { + created := ts.create(ts.Provider, map[string]any{}) + validate := newSCIMTokenValidator(ts.API.db) + + ctx, err := validate(context.Background(), created.Token) + require.NoError(ts.T(), err) + providerID, ok := scimSSOProviderIDKey.Lookup(ctx) + require.True(ts.T(), ok) + require.Equal(ts.T(), ts.Provider.ID, providerID) + + ctx, err = validate(context.Background(), "scim_invalid") + require.ErrorIs(ts.T(), err, server.ErrInvalidToken) + _, ok = scimSSOProviderIDKey.Lookup(ctx) + require.False(ts.T(), ok) + + cancelled, cancel := context.WithCancel(context.Background()) + cancel() + _, err = validate(cancelled, created.Token) + require.Error(ts.T(), err) + require.NotErrorIs(ts.T(), err, server.ErrInvalidToken) +} + +func (ts *SCIMTokensTestSuite) TestCreateWithExpiry() { + expiresAt := time.Now().Add(time.Hour).UTC().Truncate(time.Second) + + response := ts.create(ts.Provider, map[string]any{"expires_at": expiresAt}) + + require.NotNil(ts.T(), response.ExpiresAt) + require.True(ts.T(), expiresAt.Equal(*response.ExpiresAt)) +} + +func (ts *SCIMTokensTestSuite) TestCreateRejectsPastExpiry() { + w := ts.request(http.MethodPost, ts.tokensPath(ts.Provider), map[string]any{"expires_at": time.Now().Add(-time.Minute)}) + + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + require.Contains(ts.T(), w.Body.String(), "validation_failed") +} + +func (ts *SCIMTokensTestSuite) TestCreateRejectsExpiryBeforeDatabaseClock() { + expiresAt := time.Now().Add(-time.Minute) + ts.API.overrideTime = func() time.Time { return expiresAt.Add(-time.Hour) } + defer func() { ts.API.overrideTime = nil }() + + w := ts.request(http.MethodPost, ts.tokensPath(ts.Provider), map[string]any{"expires_at": expiresAt}) + + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + require.Contains(ts.T(), w.Body.String(), "validation_failed") +} + +func (ts *SCIMTokensTestSuite) TestCreateRejectsOversizedBody() { + w := ts.request(http.MethodPost, ts.tokensPath(ts.Provider), strings.Repeat("a", 1<<20)) + + require.Equal(ts.T(), http.StatusRequestEntityTooLarge, w.Code, w.Body.String()) + require.Contains(ts.T(), w.Body.String(), "request_entity_too_large") +} + +func (ts *SCIMTokensTestSuite) TestCreateForUnknownProvider() { + w := ts.request(http.MethodPost, "/admin/sso/providers/"+uuid.Must(uuid.NewV4()).String()+"/scim/tokens", map[string]any{}) + + require.Equal(ts.T(), http.StatusNotFound, w.Code) + require.Contains(ts.T(), w.Body.String(), "sso_provider_not_found") +} + +func (ts *SCIMTokensTestSuite) TestList() { + first := ts.create(ts.Provider, map[string]any{}) + second := ts.create(ts.Provider, map[string]any{}) + ts.create(ts.createProvider(), map[string]any{}) + require.Equal(ts.T(), http.StatusOK, ts.request(http.MethodDelete, ts.tokensPath(ts.Provider)+"/"+second.Prefix, nil).Code) + + w := ts.request(http.MethodGet, ts.tokensPath(ts.Provider), nil) + require.Equal(ts.T(), http.StatusOK, w.Code) + require.NotContains(ts.T(), w.Body.String(), first.Token) + require.NotContains(ts.T(), w.Body.String(), second.Token) + + var body struct { + Tokens []map[string]any `json:"tokens"` + } + require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &body)) + require.Len(ts.T(), body.Tokens, 2) + require.ElementsMatch(ts.T(), []string{"prefix", "created_at", "expires_at", "revoked_at", "last_used_at"}, slices.Collect(maps.Keys(body.Tokens[0]))) + require.ElementsMatch(ts.T(), []any{first.Prefix, second.Prefix}, []any{body.Tokens[0]["prefix"], body.Tokens[1]["prefix"]}) +} + +func (ts *SCIMTokensTestSuite) TestListEmpty() { + w := ts.request(http.MethodGet, ts.tokensPath(ts.Provider), nil) + + require.Equal(ts.T(), http.StatusOK, w.Code) + require.JSONEq(ts.T(), `{"tokens":[]}`, w.Body.String()) +} + +func (ts *SCIMTokensTestSuite) TestRevoke() { + created := ts.create(ts.Provider, map[string]any{}) + path := ts.tokensPath(ts.Provider) + "/" + created.Prefix + + w := ts.request(http.MethodDelete, path, nil) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + var revoked models.SCIMToken + require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &revoked)) + require.Equal(ts.T(), created.Prefix, revoked.Prefix) + require.NotNil(ts.T(), revoked.RevokedAt) + + require.Equal(ts.T(), http.StatusUnauthorized, ts.scimRequest(created.Token).Code) + + w = ts.request(http.MethodDelete, path, nil) + require.Equal(ts.T(), http.StatusOK, w.Code) + var again models.SCIMToken + require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &again)) + require.True(ts.T(), revoked.RevokedAt.Equal(*again.RevokedAt)) +} + +func (ts *SCIMTokensTestSuite) TestRevokeUnknownPrefix() { + created := ts.create(ts.createProvider(), map[string]any{}) + + for _, prefix := range []string{"scim_0000000", created.Prefix} { + w := ts.request(http.MethodDelete, ts.tokensPath(ts.Provider)+"/"+prefix, nil) + + require.Equal(ts.T(), http.StatusNotFound, w.Code) + require.Contains(ts.T(), w.Body.String(), "scim_token_not_found") + } +} + +func (ts *SCIMTokensTestSuite) TestRequiresAdmin() { + created := ts.create(ts.Provider, nil) + for _, route := range []struct{ method, path string }{ + {http.MethodGet, "/admin/sso/providers/" + ts.Provider.ID.String() + "/scim"}, + {http.MethodPost, "/admin/sso/providers/" + ts.Provider.ID.String() + "/scim"}, + {http.MethodDelete, "/admin/sso/providers/" + ts.Provider.ID.String() + "/scim"}, + {http.MethodGet, ts.tokensPath(ts.Provider)}, + {http.MethodPost, ts.tokensPath(ts.Provider)}, + {http.MethodDelete, ts.tokensPath(ts.Provider) + "/" + created.Prefix}, + } { + r := httptest.NewRequest(route.method, route.path, nil) + w := httptest.NewRecorder() + + ts.API.handler.ServeHTTP(w, r) + + require.Equal(ts.T(), http.StatusUnauthorized, w.Code, route.method+" "+route.path) + } + require.Equal(ts.T(), http.StatusOK, ts.scimRequest(created.Token).Code) +} + +func (ts *SCIMTokensTestSuite) TestSCIMRejectsAdminCredentials() { + for _, role := range []string{"service_role", "supabase_admin"} { + token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, &AccessTokenClaims{Role: role}).SignedString([]byte(ts.Config.JWT.Secret)) + require.NoError(ts.T(), err) + + r := httptest.NewRequest(http.MethodGet, "/scim/v2/Users", nil) + r.Header.Set("Authorization", "Bearer "+token) + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + + require.Equal(ts.T(), http.StatusUnauthorized, w.Code, role) + } +} + +func (ts *SCIMTokensTestSuite) TestDisabled() { + ts.API.config.SSO.SCIM.Enabled = false + + for _, tc := range []struct{ method, path string }{ + {http.MethodGet, ts.tokensPath(ts.Provider)}, + {http.MethodPost, ts.tokensPath(ts.Provider)}, + {http.MethodGet, ts.scimPath(ts.Provider)}, + {http.MethodPost, ts.scimPath(ts.Provider)}, + {http.MethodDelete, ts.scimPath(ts.Provider)}, + } { + w := ts.request(tc.method, tc.path, map[string]any{}) + + require.Equal(ts.T(), http.StatusNotFound, w.Code, tc.method+" "+tc.path) + require.Contains(ts.T(), w.Body.String(), "feature_disabled") + } +} + +func (ts *SCIMTokensTestSuite) scimPath(provider *models.SSOProvider) string { + return "/admin/sso/providers/" + provider.ID.String() + "/scim" +} + +func (ts *SCIMTokensTestSuite) status(method string, provider *models.SSOProvider) AdminSCIMStatusResponse { + w := ts.request(method, ts.scimPath(provider), nil) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + + var response AdminSCIMStatusResponse + require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &response)) + return response +} + +func (ts *SCIMTokensTestSuite) TestStatus() { + status := ts.status(http.MethodGet, createSSOProvider(ts.T(), ts.API.db)) + require.False(ts.T(), status.Enabled) + require.Equal(ts.T(), scimBaseURL(ts.API.config), status.BaseURL) + require.Empty(ts.T(), status.Tokens) + + status = ts.status(http.MethodGet, ts.Provider) + require.True(ts.T(), status.Enabled) + require.Equal(ts.T(), scimBaseURL(ts.API.config), status.BaseURL) + require.Empty(ts.T(), status.Tokens) + + active := ts.create(ts.Provider, map[string]any{}) + revoked := ts.create(ts.Provider, map[string]any{}) + w := ts.request(http.MethodDelete, ts.tokensPath(ts.Provider)+"/"+revoked.Prefix, nil) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + ts.create(ts.createProvider(), map[string]any{}) + + w = ts.request(http.MethodGet, ts.scimPath(ts.Provider), nil) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.NotContains(ts.T(), w.Body.String(), active.Token) + require.NotContains(ts.T(), w.Body.String(), "token_hash") + + status = ts.status(http.MethodGet, ts.Provider) + require.True(ts.T(), status.Enabled) + require.Len(ts.T(), status.Tokens, 1) + require.Equal(ts.T(), active.Prefix, status.Tokens[0].Prefix) +} + +func (ts *SCIMTokensTestSuite) TestEnableWithZeroTokens() { + provider := createSSOProvider(ts.T(), ts.API.db) + + status := ts.status(http.MethodPost, provider) + require.True(ts.T(), status.Enabled) + require.Equal(ts.T(), scimBaseURL(ts.API.config), status.BaseURL) + require.Empty(ts.T(), status.Tokens) + require.True(ts.T(), ts.status(http.MethodGet, provider).Enabled) +} + +func (ts *SCIMTokensTestSuite) TestEnableAndDisableLeaveTokensUnchanged() { + ts.create(ts.Provider, map[string]any{}) + expiring := ts.create(ts.Provider, map[string]any{"expires_at": time.Now().Add(time.Hour)}) + ts.revoke(ts.create(ts.Provider, map[string]any{}).Prefix) + before, err := models.FindSCIMTokensBySSOProvider(ts.API.db, ts.Provider.ID) + require.NoError(ts.T(), err) + require.NotNil(ts.T(), expiring.ExpiresAt) + + for _, method := range []string{http.MethodDelete, http.MethodDelete, http.MethodPost, http.MethodPost} { + ts.status(method, ts.Provider) + after, err := models.FindSCIMTokensBySSOProvider(ts.API.db, ts.Provider.ID) + require.NoError(ts.T(), err) + require.Equal(ts.T(), before, after, method) + } +} + +func (ts *SCIMTokensTestSuite) TestDisableStopsAuthenticatedRequests() { + token := ts.create(ts.Provider, map[string]any{}) + + status := ts.status(http.MethodDelete, ts.Provider) + require.False(ts.T(), status.Enabled) + require.Len(ts.T(), status.Tokens, 1) + + for _, path := range []string{"/scim/v2/Users", "/scim/v2/Groups", "/scim/v2/Schemas", "/scim/v2/ResourceTypes", "/scim/v2/ServiceProviderConfig"} { + r := httptest.NewRequest(http.MethodGet, path, nil) + r.Header.Set("Authorization", "Bearer "+token.Token) + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + + expected := http.StatusUnauthorized + if path == "/scim/v2/ServiceProviderConfig" { + expected = http.StatusOK + } + require.Equal(ts.T(), expected, w.Code, path) + } +} + +func (ts *SCIMTokensTestSuite) TestReenableRestoresExistingTokens() { + token := ts.create(ts.Provider, map[string]any{}) + require.Equal(ts.T(), http.StatusOK, ts.scimRequest(token.Token).Code) + + ts.status(http.MethodDelete, ts.Provider) + require.Equal(ts.T(), http.StatusUnauthorized, ts.scimRequest(token.Token).Code) + + status := ts.status(http.MethodPost, ts.Provider) + require.True(ts.T(), status.Enabled) + require.Equal(ts.T(), http.StatusOK, ts.scimRequest(token.Token).Code) +} + +func (ts *SCIMTokensTestSuite) TestMintAndRevokeWhileDisabled() { + ts.status(http.MethodDelete, ts.Provider) + + token := ts.create(ts.Provider, map[string]any{}) + require.Equal(ts.T(), http.StatusUnauthorized, ts.scimRequest(token.Token).Code) + revoked := ts.create(ts.Provider, map[string]any{}) + ts.revoke(revoked.Prefix) + + ts.status(http.MethodPost, ts.Provider) + require.Equal(ts.T(), http.StatusOK, ts.scimRequest(token.Token).Code) + require.Equal(ts.T(), http.StatusUnauthorized, ts.scimRequest(revoked.Token).Code) + + require.Equal(ts.T(), []scimTokenEvent{ + {string(models.SCIMDisabledAction), ""}, + {string(models.SCIMTokenCreatedAction), token.Prefix}, + {string(models.SCIMTokenCreatedAction), revoked.Prefix}, + {string(models.SCIMTokenRevokedAction), revoked.Prefix}, + {string(models.SCIMEnabledAction), ""}, + }, ts.tokenEvents()) +} + +func (ts *SCIMTokensTestSuite) TestStatusIndependentOfTokens() { + require.True(ts.T(), ts.status(http.MethodGet, ts.Provider).Enabled) + + token := ts.create(ts.Provider, map[string]any{}) + ts.revoke(token.Prefix) + require.True(ts.T(), ts.status(http.MethodGet, ts.Provider).Enabled) + + ts.create(ts.Provider, map[string]any{}) + ts.status(http.MethodDelete, ts.Provider) + status := ts.status(http.MethodGet, ts.Provider) + require.False(ts.T(), status.Enabled) + require.Len(ts.T(), status.Tokens, 1) +} + +func (ts *SCIMTokensTestSuite) TestConcurrentEnableAndDisable() { + provider := createSSOProvider(ts.T(), ts.API.db) + + for _, method := range []string{http.MethodPost, http.MethodDelete} { + var wg sync.WaitGroup + codes := make(chan int, 10) + for range 10 { + wg.Go(func() { + codes <- ts.request(method, ts.scimPath(provider), nil).Code + }) + } + wg.Wait() + close(codes) + for code := range codes { + require.Equal(ts.T(), http.StatusOK, code, method) + } + } + + require.Equal(ts.T(), []string{string(models.SCIMEnabledAction), string(models.SCIMDisabledAction)}, ts.scimActions(provider)) +} + +func (ts *SCIMTokensTestSuite) scimActions(provider *models.SSOProvider) []string { + entries := []models.AuditLogEntry{} + require.NoError(ts.T(), ts.API.db.Q().Where("payload->>'log_type' = ? AND payload->'traits'->>'sso_provider_id' = ?", "scim", provider.ID.String()).Order("created_at asc").All(&entries)) + actions := []string{} + for _, entry := range entries { + actions = append(actions, entry.Payload["action"].(string)) + } + return actions +} + +func (ts *SCIMTokensTestSuite) TestStatusForUnknownProvider() { + for _, method := range []string{http.MethodGet, http.MethodPost, http.MethodDelete} { + w := ts.request(method, "/admin/sso/providers/"+uuid.Must(uuid.NewV4()).String()+"/scim", nil) + require.Equal(ts.T(), http.StatusNotFound, w.Code, method) + require.Contains(ts.T(), w.Body.String(), "sso_provider_not_found", method) + } + require.Empty(ts.T(), ts.tokenEvents()) +} + +func (ts *SCIMTokensTestSuite) setProviderDisabled(disabled bool) { + require.NoError(ts.T(), ts.API.db.RawQuery("UPDATE "+ts.Provider.TableName()+" SET disabled = ? WHERE id = ?", disabled, ts.Provider.ID).Exec()) +} + +func (ts *SCIMTokensTestSuite) TestStatusForDisabledProvider() { + first := ts.create(ts.Provider, map[string]any{}) + ts.setProviderDisabled(true) + + status := ts.status(http.MethodGet, ts.Provider) + require.False(ts.T(), status.Enabled) + require.Len(ts.T(), status.Tokens, 1) + require.Equal(ts.T(), http.StatusUnauthorized, ts.scimRequest(first.Token).Code) + + second := ts.create(ts.Provider, map[string]any{}) + third := ts.create(ts.Provider, map[string]any{}) + status = ts.status(http.MethodGet, ts.Provider) + require.False(ts.T(), status.Enabled) + require.Len(ts.T(), status.Tokens, 3) + + ts.revoke(first.Prefix) + ts.setProviderDisabled(false) + status = ts.status(http.MethodGet, ts.Provider) + require.True(ts.T(), status.Enabled) + require.Len(ts.T(), status.Tokens, 2) + require.Equal(ts.T(), http.StatusOK, ts.scimRequest(second.Token).Code) + + ts.setProviderDisabled(true) + ts.revoke(third.Prefix) + + require.Equal(ts.T(), []scimTokenEvent{ + {string(models.SCIMTokenCreatedAction), first.Prefix}, + {string(models.SCIMTokenCreatedAction), second.Prefix}, + {string(models.SCIMTokenCreatedAction), third.Prefix}, + {string(models.SCIMTokenRevokedAction), first.Prefix}, + {string(models.SCIMTokenRevokedAction), third.Prefix}, + }, ts.tokenEvents()) +} + +type scimTokenEvent struct{ action, prefix string } + +func (ts *SCIMTokensTestSuite) tokenEvents() []scimTokenEvent { + entries := []models.AuditLogEntry{} + require.NoError(ts.T(), ts.API.db.Q().Where("payload->>'log_type' = ?", "scim").Order("created_at asc").All(&entries)) + + events := []scimTokenEvent{} + for _, entry := range entries { + require.Equal(ts.T(), "supabase_admin", entry.Payload["actor_username"]) + traits := entry.Payload["traits"].(map[string]any) + require.Equal(ts.T(), ts.Provider.ID.String(), traits["sso_provider_id"]) + require.Equal(ts.T(), "success", traits["outcome"]) + prefix, _ := traits["token_prefix"].(string) + if prefixes, ok := traits["token_prefixes"].([]any); ok { + require.Len(ts.T(), prefixes, 1) + prefix = prefixes[0].(string) + } + events = append(events, scimTokenEvent{entry.Payload["action"].(string), prefix}) + } + return events +} + +func (ts *SCIMTokensTestSuite) revoke(prefix string) { + w := ts.request(http.MethodDelete, ts.tokensPath(ts.Provider)+"/"+prefix, nil) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) +} + +func (ts *SCIMTokensTestSuite) TestAuditLog() { + first := ts.create(ts.Provider, map[string]any{}) + second := ts.create(ts.Provider, map[string]any{}) + ts.revoke(first.Prefix) + ts.revoke(first.Prefix) + ts.revoke(second.Prefix) + ts.status(http.MethodDelete, ts.Provider) + ts.status(http.MethodDelete, ts.Provider) + ts.status(http.MethodPost, ts.Provider) + ts.status(http.MethodPost, ts.Provider) + + w := ts.request(http.MethodPost, ts.tokensPath(ts.Provider), map[string]any{"expires_at": "2000-01-01T00:00:00Z"}) + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + ts.API.config.SSO.SCIM.Enabled = false + w = ts.request(http.MethodDelete, ts.scimPath(ts.Provider), nil) + require.Equal(ts.T(), http.StatusNotFound, w.Code, w.Body.String()) + ts.API.config.SSO.SCIM.Enabled = true + + require.Equal(ts.T(), []scimTokenEvent{ + {string(models.SCIMTokenCreatedAction), first.Prefix}, + {string(models.SCIMTokenCreatedAction), second.Prefix}, + {string(models.SCIMTokenRevokedAction), first.Prefix}, + {string(models.SCIMTokenRevokedAction), second.Prefix}, + {string(models.SCIMDisabledAction), ""}, + {string(models.SCIMEnabledAction), ""}, + }, ts.tokenEvents()) +} + +func (ts *SCIMTokensTestSuite) TestDisableNeverEnabledWritesNoEvent() { + provider := createSSOProvider(ts.T(), ts.API.db) + + status := ts.status(http.MethodDelete, provider) + require.False(ts.T(), status.Enabled) + require.Empty(ts.T(), ts.scimActions(provider)) +} + +func (ts *SCIMTokensTestSuite) TestEnableSSODisabledProvider() { + provider := createSSOProvider(ts.T(), ts.API.db) + require.NoError(ts.T(), ts.API.db.RawQuery("UPDATE "+provider.TableName()+" SET disabled = true WHERE id = ?", provider.ID).Exec()) + + require.False(ts.T(), ts.status(http.MethodPost, provider).Enabled) + require.Equal(ts.T(), []string{string(models.SCIMEnabledAction)}, ts.scimActions(provider)) + + require.NoError(ts.T(), ts.API.db.RawQuery("UPDATE "+provider.TableName()+" SET disabled = false WHERE id = ?", provider.ID).Exec()) + require.True(ts.T(), ts.status(http.MethodGet, provider).Enabled) +} diff --git a/internal/api/scim_filter.go b/internal/api/scim_filter.go new file mode 100644 index 0000000000..3129bdcfc4 --- /dev/null +++ b/internal/api/scim_filter.go @@ -0,0 +1,52 @@ +package api + +import ( + "fmt" + + "github.com/supabase-community/scim-go/pkg/filter" + "github.com/supabase-community/scim-go/pkg/protocol" + "github.com/supabase-community/scim-go/pkg/scimerrors" + "github.com/supabase/auth/internal/models" +) + +type scimEqFilter struct { + name string +} + +func (f scimEqFilter) Compare(attribute *protocol.Attribute, op filter.Operator, value any) (models.SCIMFilter, error) { + text, ok := value.(string) + if op != filter.OpEquals || attribute.Parent != nil || !ok { + return f.unsupported() + } + switch attribute.Definition.Name { + case f.name: + return models.SCIMFilter{Name: &text}, nil + case "externalId": + return models.SCIMFilter{ExternalID: &text}, nil + } + return f.unsupported() +} + +func (f scimEqFilter) Present(*protocol.Attribute) (models.SCIMFilter, error) { + return f.unsupported() +} + +func (f scimEqFilter) And(models.SCIMFilter, models.SCIMFilter) (models.SCIMFilter, error) { + return f.unsupported() +} + +func (f scimEqFilter) Or(models.SCIMFilter, models.SCIMFilter) (models.SCIMFilter, error) { + return f.unsupported() +} + +func (f scimEqFilter) Not(models.SCIMFilter) (models.SCIMFilter, error) { + return f.unsupported() +} + +func (f scimEqFilter) ValuePath(*protocol.Attribute, func() (models.SCIMFilter, error)) (models.SCIMFilter, error) { + return f.unsupported() +} + +func (f scimEqFilter) unsupported() (models.SCIMFilter, error) { + return models.SCIMFilter{}, scimerrors.ErrInvalidFilter(fmt.Sprintf(`only "%s eq" and "externalId eq" filters are supported`, f.name)) +} diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go new file mode 100644 index 0000000000..0bc224d1e1 --- /dev/null +++ b/internal/api/scim_groups.go @@ -0,0 +1,258 @@ +package api + +import ( + "context" + "encoding/json" + "net/http" + "strings" + + "github.com/gofrs/uuid" + "github.com/supabase-community/scim-go/pkg/core" + "github.com/supabase-community/scim-go/pkg/protocol" + "github.com/supabase-community/scim-go/pkg/scimerrors" + "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/storage" +) + +type scimGroups struct { + api *API +} + +func (s *scimGroups) List(ctx context.Context, query *protocol.SearchRequest) ([]*core.Group, int, error) { + providerID, err := scimProviderID(ctx) + if err != nil { + return nil, 0, err + } + search, err := scimSearch(query, scimGroupSchemas, "displayName") + if err != nil { + return nil, 0, err + } + db := s.api.db.WithContext(ctx) + rows, total, err := models.FindSCIMGroups(db, providerID, search) + if err != nil { + return nil, 0, err + } + projection, err := query.Projection(scimGroupSchemas) + if err != nil { + projection = protocol.Projection{} + } + groups, err := s.render(db, providerID, rows, projection) + if err != nil { + return nil, 0, err + } + return groups, total, nil +} + +func (s *scimGroups) Get(ctx context.Context, id string) (*core.Group, error) { + providerID, resourceID, _, err := scimTarget(ctx, id, "") + if err != nil { + return nil, err + } + db := s.api.db.WithContext(ctx) + row, err := models.FindSCIMGroup(db, providerID, resourceID) + if err != nil { + return nil, scimTranslate(err) + } + projection, err := protocol.ParseProjection(scimGetQueryKey.Value(ctx), scimGroupSchemas) + if err != nil { + projection = protocol.Projection{} + } + return s.renderOne(db, providerID, row, projection) +} + +func (s *scimGroups) Create(ctx context.Context, group *core.Group) (*core.Group, error) { + providerID, err := scimProviderID(ctx) + if err != nil { + return nil, err + } + return s.save(ctx, providerID, models.SCIMGroupCreatedAction, group, func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error) { + row, err := models.CreateSCIMGroup(tx, providerID, resource) + return row, true, err + }) +} + +func (s *scimGroups) Replace(ctx context.Context, group *core.Group) (*core.Group, error) { + providerID, id, updatedAt, err := scimTarget(ctx, group.ID, group.Meta.Version) + if err != nil { + return nil, err + } + return s.save(ctx, providerID, models.SCIMGroupUpdatedAction, group, func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error) { + unchanged, err := models.FindUnchangedSCIMGroup(tx, providerID, id, resource, updatedAt) + if err != nil || unchanged != nil { + return unchanged, false, err + } + row, err := models.ReplaceSCIMGroup(tx, providerID, id, resource, updatedAt) + return row, true, err + }) +} + +func (s *scimGroups) Delete(ctx context.Context, id, version string) error { + providerID, resourceID, updatedAt, err := scimTarget(ctx, id, version) + if err != nil { + return err + } + r, err := scimRequest(ctx) + if err != nil { + return err + } + return scimTranslate(s.api.db.WithContext(ctx).Transaction(func(tx *storage.Connection) error { + row, err := models.FindSCIMGroupForUpdate(tx, providerID, resourceID) + if err != nil { + return err + } + _, removed, err := models.ReplaceSCIMGroupMembers(tx, row, nil) + if err != nil { + return err + } + if row, err = models.DeleteSCIMGroup(tx, providerID, resourceID, updatedAt); err != nil { + return err + } + if err := s.auditMembers(tx, r, row, nil, removed); err != nil { + return err + } + return s.audit(tx, r, models.SCIMGroupDeletedAction, row) + })) +} + +func (s *scimGroups) save(ctx context.Context, providerID uuid.UUID, action models.AuditAction, group *core.Group, write func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error)) (*core.Group, error) { + members, err := scimMemberIDs(group.Members) + if err != nil { + return nil, err + } + resource, err := scimEncode(group, "id", "meta", "members") + if err != nil { + return nil, err + } + r, err := scimRequest(ctx) + if err != nil { + return nil, err + } + db := s.api.db.WithContext(ctx) + var row *models.SCIMGroup + err = db.Transaction(func(tx *storage.Connection) error { + var ( + changed bool + terr error + ) + if row, changed, terr = write(tx, resource); terr != nil { + return terr + } + added, removed, terr := models.ReplaceSCIMGroupMembers(tx, row, members) + if terr != nil { + return terr + } + if !changed { + if len(added) == 0 && len(removed) == 0 { + return nil + } + if row, terr = models.TouchSCIMGroup(tx, row); terr != nil { + return terr + } + } else if terr := s.audit(tx, r, action, row); terr != nil { + return terr + } + return s.auditMembers(tx, r, row, added, removed) + }) + if err != nil { + return nil, scimTranslate(err) + } + return s.renderOne(db, providerID, row, protocol.Projection{}) +} + +func (s *scimGroups) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMGroup, projection protocol.Projection) ([]*core.Group, error) { + memberships := []models.SCIMGroupMembership{} + if scimReturns(projection, "members") { + ids := make([]uuid.UUID, len(rows)) + for i, row := range rows { + ids[i] = row.ID + } + var err error + if memberships, err = models.FindSCIMGroupMembers(tx, providerID, ids); err != nil { + return nil, err + } + } + base := scimBaseURL(s.api.config) + members := map[uuid.UUID][]core.Member{} + for _, m := range memberships { + members[m.GroupID] = append(members[m.GroupID], core.Member{ + Value: m.SCIMUserID.String(), + Ref: base + "/Users/" + m.SCIMUserID.String(), + Type: scimResourceTypeUser, + }) + } + + groups := make([]*core.Group, 0, len(rows)) + for _, row := range rows { + group := &core.Group{} + if err := json.Unmarshal(row.Resource, group); err != nil { + return nil, err + } + group.ID = row.ID.String() + group.Schemas = []core.SchemaURI{core.SchemaGroup} + group.Meta = scimMeta(scimResourceTypeGroup, base+"/Groups/"+group.ID, row.CreatedAt, row.UpdatedAt) + group.Members = members[row.ID] + groups = append(groups, group) + } + return groups, nil +} + +func (s *scimGroups) renderOne(tx *storage.Connection, providerID uuid.UUID, row *models.SCIMGroup, projection protocol.Projection) (*core.Group, error) { + groups, err := s.render(tx, providerID, []models.SCIMGroup{*row}, projection) + if err != nil { + return nil, err + } + return groups[0], nil +} + +func (s *scimGroups) audit(tx *storage.Connection, r *http.Request, action models.AuditAction, row *models.SCIMGroup) error { + var resource struct { + DisplayName string `json:"displayName"` + } + if err := json.Unmarshal(row.Resource, &resource); err != nil { + return err + } + return s.api.auditSCIM(tx, r, scimActor(r), action, row.SSOProviderID, map[string]any{ + "scim_group_id": row.ID, + "display_name": resource.DisplayName, + }) +} + +func (s *scimGroups) auditMembers(tx *storage.Connection, r *http.Request, row *models.SCIMGroup, added, removed []uuid.UUID) error { + links, err := models.FindSCIMUserLinks(tx, append(append([]uuid.UUID{}, added...), removed...)) + if err != nil { + return err + } + for _, change := range []struct { + action models.AuditAction + ids []uuid.UUID + }{ + {models.SCIMGroupMemberAddedAction, added}, + {models.SCIMGroupMemberRemovedAction, removed}, + } { + for _, id := range change.ids { + var userID *uuid.UUID + if linked, ok := links[id]; ok { + userID = &linked + } + if err := s.api.auditSCIMMember(tx, r, scimActor(r), change.action, row.SSOProviderID, row.ID, id, userID); err != nil { + return err + } + } + } + return nil +} + +func scimMemberIDs(members []core.Member) ([]uuid.UUID, error) { + ids := make([]uuid.UUID, 0, len(members)) + for _, member := range members { + if member.Type != "" && !strings.EqualFold(string(member.Type), scimResourceTypeUser) { + return nil, scimerrors.ErrInvalidValue(`nested groups are not supported; "members.type" must be "User"`) + } + id, err := uuid.FromString(member.Value) + if err != nil { + return nil, errSCIMMemberNotFound() + } + ids = append(ids, id) + } + return ids, nil +} diff --git a/internal/api/scim_groups_test.go b/internal/api/scim_groups_test.go new file mode 100644 index 0000000000..8eaadc6dab --- /dev/null +++ b/internal/api/scim_groups_test.go @@ -0,0 +1,712 @@ +package api + +import ( + "fmt" + "net/http" + "net/url" + "strings" + + "github.com/gofrs/uuid" + logrustest "github.com/sirupsen/logrus/hooks/test" + "github.com/stretchr/testify/require" + "github.com/supabase-community/scim-go/pkg/protocol" + "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/storage" +) + +func groupWith(displayName, externalID string, memberIDs ...string) string { + members := make([]string, len(memberIDs)) + for i, id := range memberIDs { + members[i] = `{"value":"` + id + `"}` + } + return `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:Group"],"displayName":"` + displayName + `","externalId":"` + externalID + `","members":[` + strings.Join(members, ",") + `]}` +} + +func (ts *SCIMUsersTestSuite) createGroup(token, body string) string { + w, created := ts.do(token, http.MethodPost, "/Groups", body) + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) + return created["id"].(string) +} + +func (ts *SCIMUsersTestSuite) listGroups(token, filter string) map[string]any { + w, body := ts.do(token, http.MethodGet, "/Groups?"+url.Values{"filter": {filter}}.Encode(), "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + return body +} + +func memberValues(group map[string]any) []string { + values := []string{} + members, _ := group["members"].([]any) + for _, member := range members { + values = append(values, member.(map[string]any)["value"].(string)) + } + return values +} + +func (ts *SCIMUsersTestSuite) TestGroupsLifecycle() { + alice := ts.create(ts.TokenA, userWith("Alice@Example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + + w, created := ts.do(ts.TokenA, http.MethodPost, "/Groups", groupWith("Engineering", "Finance", alice)) + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) + id := created["id"].(string) + location := "http://localhost:9999/scim/v2/Groups/" + id + require.Equal(ts.T(), location, w.Header().Get("Location")) + require.NotEmpty(ts.T(), w.Header().Get("ETag")) + require.Equal(ts.T(), "Engineering", created["displayName"]) + require.Equal(ts.T(), "Finance", created["externalId"]) + meta := created["meta"].(map[string]any) + require.Equal(ts.T(), "Group", meta["resourceType"]) + require.Equal(ts.T(), location, meta["location"]) + member := created["members"].([]any)[0].(map[string]any) + require.Equal(ts.T(), alice, member["value"]) + require.Equal(ts.T(), "User", member["type"]) + require.NotContains(ts.T(), member, "display") + require.Equal(ts.T(), "http://localhost:9999/scim/v2/Users/"+alice, member["$ref"]) + + var stored models.SCIMGroup + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&stored)) + require.Equal(ts.T(), ts.A.ID, stored.SSOProviderID) + require.NotContains(ts.T(), string(stored.Resource), "members") + require.NotContains(ts.T(), string(stored.Resource), `"id"`) + + for _, filter := range []string{`displayName eq "Engineering"`, `displayName eq "engineering"`, `externalId eq "Finance"`} { + found := ts.listGroups(ts.TokenA, filter) + require.EqualValues(ts.T(), 1, found["totalResults"], filter) + require.Equal(ts.T(), id, found["Resources"].([]any)[0].(map[string]any)["id"], filter) + } + + w, replaced := ts.do(ts.TokenA, http.MethodPut, "/Groups/"+id, groupWith("Engineering", "Finance", bob)) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), []string{bob}, memberValues(replaced)) + require.Equal(ts.T(), meta["created"], replaced["meta"].(map[string]any)["created"]) + + w, patched := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ + {"op":"add","path":"members","value":[{"value":"`+alice+`"}]}, + {"op":"replace","path":"displayName","value":"Platform"} + ]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.ElementsMatch(ts.T(), []string{alice, bob}, memberValues(patched)) + require.Equal(ts.T(), "Platform", patched["displayName"]) + + w, patched = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ + {"op":"remove","path":"members[value eq \"`+bob+`\"]"} + ]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), []string{alice}, memberValues(patched)) + + w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Equal(ts.T(), "Platform", got["displayName"]) + require.Equal(ts.T(), []string{alice}, memberValues(got)) + + w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Groups/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + w, _ = ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + require.Equal(ts.T(), http.StatusNotFound, w.Code) + + w, _ = ts.do(ts.TokenA, http.MethodGet, "/Users/"+alice, "") + require.Equal(ts.T(), http.StatusOK, w.Code) +} + +func (ts *SCIMUsersTestSuite) TestGroupsWithoutMembers() { + w, created := ts.do(ts.TokenA, http.MethodPost, "/Groups", `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:Group"],"displayName":"Empty"}`) + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) + require.NotContains(ts.T(), created, "members") + + ts.createGroup(ts.TokenA, groupWith("Empty", "e-2")) + require.EqualValues(ts.T(), 2, ts.listGroups(ts.TokenA, `displayName eq "Empty"`)["totalResults"]) +} + +func (ts *SCIMUsersTestSuite) TestGroupsRejectInvalidMembers() { + outsider := ts.create(ts.TokenB, userWith("mallory@example.com", "m-1")) + deleted := ts.create(ts.TokenA, userWith("gone@example.com", "g-1")) + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+deleted, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + existing := ts.createGroup(ts.TokenA, groupWith("Existing", "")) + hook := logrustest.NewGlobal() + defer hook.Reset() + + for name, body := range map[string]string{ + "other provider": groupWith("Engineering", "", outsider), + "deleted user": groupWith("Engineering", "", deleted), + "unknown id": groupWith("Engineering", "", "00000000-0000-0000-0000-000000000000"), + "not a uuid": groupWith("Engineering", "", "alice"), + "nested group": `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:Group"],"displayName":"Engineering","members":[{"value":"00000000-0000-0000-0000-000000000000","type":"Group"}]}`, + } { + w, body := ts.do(ts.TokenA, http.MethodPost, "/Groups", body) + require.Equal(ts.T(), http.StatusBadRequest, w.Code, name+" "+w.Body.String()) + require.Equal(ts.T(), "invalidValue", body["scimType"], name) + } + w, _ = ts.do(ts.TokenA, http.MethodPut, "/Groups/"+existing, groupWith("Existing", "", outsider)) + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + require.EqualValues(ts.T(), 1, ts.listGroups(ts.TokenA, "")["totalResults"]) + for _, entry := range hook.AllEntries() { + require.NotEqual(ts.T(), "audit_event", entry.Message, entry.Data) + } +} + +func (ts *SCIMUsersTestSuite) TestGroupsMemberTypeIsCaseInsensitive() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + + for _, kind := range []string{"user", "USER", "User"} { + body := fmt.Sprintf(`{"schemas":["urn:ietf:params:scim:schemas:core:2.0:Group"],"displayName":"Engineering %s","members":[{"value":%q,"type":%q}]}`, kind, alice, kind) + w, created := ts.do(ts.TokenA, http.MethodPost, "/Groups", body) + require.Equal(ts.T(), http.StatusCreated, w.Code, kind+" "+w.Body.String()) + require.Equal(ts.T(), []string{alice}, memberValues(created), kind) + } +} + +func (ts *SCIMUsersTestSuite) TestPatchReplaceMembers() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + carol := ts.create(ts.TokenA, userWith("carol@example.com", "c-1")) + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, bob)) + before := len(ts.scimAuditEntries()) + + w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ + {"op":"replace","path":"members","value":[{"value":"`+bob+`"},{"value":"`+carol+`"}]} + ]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.ElementsMatch(ts.T(), []string{bob, carol}, memberValues(got)) + + events := map[string][]string{} + for _, entry := range ts.scimAuditEntries()[before:] { + scimUserID, _ := entry.Payload["traits"].(map[string]any)["scim_user_id"].(string) + events[entry.Payload["action"].(string)] = append(events[entry.Payload["action"].(string)], scimUserID) + } + require.Equal(ts.T(), map[string][]string{ + string(models.SCIMGroupMemberAddedAction): {carol}, + string(models.SCIMGroupMemberRemovedAction): {alice}, + }, events) +} + +func (ts *SCIMUsersTestSuite) TestGroupMemberEventsCarryUserID() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + userIDs := map[string]string{} + for _, id := range []string{alice, bob} { + row, err := models.FindSCIMUser(ts.API.db, ts.A.ID, uuid.FromStringOrNil(id)) + require.NoError(ts.T(), err) + require.NotNil(ts.T(), row.UserID) + userIDs[id] = row.UserID.String() + } + before := len(ts.scimAuditEntries()) + + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, bob)) + w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"remove","path":"members[value eq \"`+alice+`\"]"}]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Users/"+bob, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + + type event struct { + action, scimUserID, userID string + } + events := []event{} + for _, entry := range ts.scimAuditEntries()[before:] { + traits := entry.Payload["traits"].(map[string]any) + if _, ok := traits["scim_group_id"]; !ok { + continue + } + scimUserID, _ := traits["scim_user_id"].(string) + if scimUserID == "" { + continue + } + userID, _ := traits["user_id"].(string) + events = append(events, event{entry.Payload["action"].(string), scimUserID, userID}) + } + require.ElementsMatch(ts.T(), []event{ + {string(models.SCIMGroupMemberAddedAction), alice, userIDs[alice]}, + {string(models.SCIMGroupMemberAddedAction), bob, userIDs[bob]}, + {string(models.SCIMGroupMemberRemovedAction), alice, userIDs[alice]}, + {string(models.SCIMGroupMemberRemovedAction), bob, userIDs[bob]}, + }, events) +} + +func (ts *SCIMUsersTestSuite) TestExcludedMembersKeepsWrites() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + carol := ts.create(ts.TokenA, userWith("carol@example.com", "c-1")) + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, bob)) + + w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id+"?excludedAttributes=members", "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.NotContains(ts.T(), got, "members") + require.Equal(ts.T(), "Engineering", got["displayName"]) + + w, got = ts.do(ts.TokenA, http.MethodGet, "/Groups?excludedAttributes=members", "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.NotContains(ts.T(), got["Resources"].([]any)[0], "members") + + w, got = ts.do(ts.TokenA, http.MethodGet, "/Users/"+alice+"?excludedAttributes=groups", "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.NotContains(ts.T(), got, "groups") + + w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id+"?excludedAttributes=members", `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"add","path":"members","value":[{"value":"`+carol+`"}]}]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + w, _ = ts.do(ts.TokenA, http.MethodPut, "/Groups/"+id+"?excludedAttributes=members", groupWith("Platform", "", alice, bob, carol)) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + + w, got = ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.ElementsMatch(ts.T(), []string{alice, bob, carol}, memberValues(got)) + require.Equal(ts.T(), "Platform", got["displayName"]) +} + +func (ts *SCIMUsersTestSuite) TestPatchRemoveAbsentMember() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice)) + before := len(ts.auditActions(models.SCIMGroupMemberRemovedAction)) + + w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"remove","path":"members[value eq \"`+bob+`\"]"}]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), []string{alice}, memberValues(got)) + require.Len(ts.T(), ts.auditActions(models.SCIMGroupMemberRemovedAction), before) +} + +func (ts *SCIMUsersTestSuite) TestPatchRejectsRemoveWithValue() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, bob)) + + for path, body := range map[string]string{ + "/Groups/" + id: `{"op":"Remove","path":"members","value":[{"$ref":null,"value":"` + bob + `"}]}`, + "/Users/" + alice: `{"op":"remove","path":"emails","value":[{"value":"alice@example.com"}]}`, + } { + w, got := ts.do(ts.TokenA, http.MethodPatch, path, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[`+body+`]}`) + require.Equal(ts.T(), http.StatusBadRequest, w.Code, path+" "+w.Body.String()) + require.Equal(ts.T(), "invalidSyntax", got["scimType"], path) + } + + w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.ElementsMatch(ts.T(), []string{alice, bob}, memberValues(got)) + w, got = ts.do(ts.TokenA, http.MethodGet, "/Users/"+alice, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.NotEmpty(ts.T(), got["emails"]) + + w, got = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ + {"op":"remove","path":"members[value eq \"`+bob+`\"]","value":null} + ]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), []string{alice}, memberValues(got)) +} + +func (ts *SCIMUsersTestSuite) TestGroupsExternalIDUniqueWithinProvider() { + ts.createGroup(ts.TokenA, groupWith("A", "g-1")) + + w, body := ts.do(ts.TokenA, http.MethodPost, "/Groups", groupWith("B", "g-1")) + require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) + require.Equal(ts.T(), "uniqueness", body["scimType"]) + + ts.createGroup(ts.TokenB, groupWith("A", "g-1")) +} + +func (ts *SCIMUsersTestSuite) TestGroupsETagAndIfMatch() { + w, created := ts.do(ts.TokenA, http.MethodPost, "/Groups", groupWith("Engineering", "g-1")) + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) + id := created["id"].(string) + stale := w.Header().Get("ETag") + + patch := `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","path":"displayName","value":"Platform"}]}` + w, _ = ts.doAs(protocol.MediaType, ts.TokenA, http.MethodPatch, "/Groups/"+id, patch, "If-Match", stale) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + current := w.Header().Get("ETag") + require.NotEqual(ts.T(), stale, current) + + for _, tc := range []struct{ method, body string }{ + {http.MethodPut, groupWith("Engineering", "g-1")}, + {http.MethodPatch, patch}, + {http.MethodDelete, ""}, + } { + w, _ := ts.doAs(protocol.MediaType, ts.TokenA, tc.method, "/Groups/"+id, tc.body, "If-Match", stale) + require.Equal(ts.T(), http.StatusPreconditionFailed, w.Code, tc.method+" "+w.Body.String()) + } + + w, _ = ts.doAs(protocol.MediaType, ts.TokenA, http.MethodDelete, "/Groups/"+id, "", "If-Match", current) + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) +} + +func (ts *SCIMUsersTestSuite) TestIdenticalPutChecksIfMatch() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + for path, body := range map[string]string{ + "/Users/" + alice: userWith("alice@example.com", "a-2"), + "/Groups/" + ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)): groupWith("Engineering", "g-2", alice), + } { + w, _ := ts.do(ts.TokenA, http.MethodGet, path, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + stale := w.Header().Get("ETag") + w, _ = ts.do(ts.TokenA, http.MethodPut, path, body) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + current := w.Header().Get("ETag") + require.NotEqual(ts.T(), stale, current, path) + events := len(ts.scimAuditEntries()) + + w, _ = ts.doAs(protocol.MediaType, ts.TokenA, http.MethodPut, path, body, "If-Match", stale) + require.Equal(ts.T(), http.StatusPreconditionFailed, w.Code, path+" "+w.Body.String()) + + w, _ = ts.doAs(protocol.MediaType, ts.TokenA, http.MethodPut, path, body, "If-Match", current) + require.Equal(ts.T(), http.StatusOK, w.Code, path+" "+w.Body.String()) + require.Equal(ts.T(), current, w.Header().Get("ETag"), path) + require.Len(ts.T(), ts.scimAuditEntries(), events, path) + } +} + +func (ts *SCIMUsersTestSuite) TestGroupsRemoveDeletedMembers() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + eng := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice, bob)) + ops := ts.createGroup(ts.TokenA, groupWith("Ops", "g-2", alice)) + before := len(ts.scimAuditEntries()) + + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+alice, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + + w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+eng, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Equal(ts.T(), []string{bob}, memberValues(got)) + w, got = ts.do(ts.TokenA, http.MethodGet, "/Groups/"+ops, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Empty(ts.T(), memberValues(got)) + + count, err := ts.API.db.Q().Where("scim_user_id = ?", alice).Count(&models.SCIMGroupMember{}) + require.NoError(ts.T(), err) + require.Zero(ts.T(), count) + + removed := []string{} + for _, entry := range ts.scimAuditEntries()[before:] { + if entry.Payload["action"] != string(models.SCIMGroupMemberRemovedAction) { + continue + } + traits := entry.Payload["traits"].(map[string]any) + require.Equal(ts.T(), alice, traits["scim_user_id"]) + removed = append(removed, traits["scim_group_id"].(string)) + } + require.ElementsMatch(ts.T(), []string{eng, ops}, removed) +} + +func (ts *SCIMUsersTestSuite) TestGroupsVersionChangesWhenMemberDeleted() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) + w, _ := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + stale := w.Header().Get("ETag") + + w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Users/"+alice, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + + w, _ = ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + require.NotEqual(ts.T(), stale, w.Header().Get("ETag")) + w, _ = ts.doAs(protocol.MediaType, ts.TokenA, http.MethodPut, "/Groups/"+id, groupWith("Engineering", "g-1"), "If-Match", stale) + require.Equal(ts.T(), http.StatusPreconditionFailed, w.Code, w.Body.String()) +} + +func (ts *SCIMUsersTestSuite) TestUserDeleteWaitsForGroupWrite() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice, bob)) + group, err := models.FindSCIMGroup(ts.API.db, ts.A.ID, uuid.FromStringOrNil(id)) + require.NoError(ts.T(), err) + + code, err := ts.whileLocked( + func(tx *storage.Connection) error { + _, err := models.FindSCIMGroupForUpdate(tx, ts.A.ID, group.ID) + return err + }, + func(tx *storage.Connection) error { + _, _, err := models.ReplaceSCIMGroupMembers(tx, group, []uuid.UUID{uuid.FromStringOrNil(bob)}) + return err + }, + http.MethodDelete, "/Users/"+alice, "", + ) + require.NoError(ts.T(), err) + require.Equal(ts.T(), http.StatusNoContent, code) + + w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Equal(ts.T(), []string{bob}, memberValues(got)) +} + +func (ts *SCIMUsersTestSuite) TestGroupsKeepDeactivatedMembers() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) + before := len(ts.scimAuditEntries()) + + w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+alice, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ + {"op":"replace","path":"active","value":false} + ]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + + w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ + {"op":"add","path":"members","value":[{"value":"`+bob+`"}]} + ]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.ElementsMatch(ts.T(), []string{alice, bob}, memberValues(got)) + + w, user := ts.do(ts.TokenA, http.MethodGet, "/Users/"+alice, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), false, user["active"]) + require.Len(ts.T(), user["groups"], 1) + + for _, entry := range ts.scimAuditEntries()[before:] { + require.NotEqual(ts.T(), string(models.SCIMGroupMemberRemovedAction), entry.Payload["action"]) + } +} + +func (ts *SCIMUsersTestSuite) TestGroupsSortAndPaginate() { + ts.createGroup(ts.TokenA, groupWith("beta", "g-2")) + ts.createGroup(ts.TokenA, groupWith("Alpha", "g-1")) + ts.createGroup(ts.TokenA, groupWith("gamma", "g-3")) + + w, page := ts.do(ts.TokenA, http.MethodGet, "/Groups?sortBy=displayName&startIndex=2&count=1", "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.EqualValues(ts.T(), 3, page["totalResults"]) + require.Equal(ts.T(), "beta", page["Resources"].([]any)[0].(map[string]any)["displayName"]) + + w, page = ts.do(ts.TokenA, http.MethodGet, "/Groups?sortBy=displayName&sortOrder=descending", "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), "gamma", page["Resources"].([]any)[0].(map[string]any)["displayName"]) + + w, body := ts.do(ts.TokenA, http.MethodGet, "/Groups?sortBy=members.value", "") + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + require.Equal(ts.T(), "invalidValue", body["scimType"]) +} + +func (ts *SCIMUsersTestSuite) TestGroupsUnsupportedFilters() { + for _, filter := range []string{ + `displayName co "eng"`, + `members.value eq "00000000-0000-0000-0000-000000000000"`, + `members[value eq "00000000-0000-0000-0000-000000000000"]`, + `displayName eq "a" or displayName eq "b"`, + `displayName pr`, + } { + w, body := ts.do(ts.TokenA, http.MethodGet, "/Groups?"+url.Values{"filter": {filter}}.Encode(), "") + require.Equal(ts.T(), http.StatusBadRequest, w.Code, filter) + require.Equal(ts.T(), "invalidFilter", body["scimType"], filter) + } +} + +func (ts *SCIMUsersTestSuite) TestGroupsUnknownID() { + for _, id := range []string{"not-a-uuid", "00000000-0000-0000-0000-000000000000"} { + w, _ := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + require.Equal(ts.T(), http.StatusNotFound, w.Code, id) + } +} + +func (ts *SCIMUsersTestSuite) TestUsersGroupsAttribute() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + ops := ts.createGroup(ts.TokenA, groupWith("Ops", "g-2", alice)) + eng := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) + + groupsOf := func(user map[string]any) []map[string]any { + found := []map[string]any{} + groups, _ := user["groups"].([]any) + for _, group := range groups { + found = append(found, group.(map[string]any)) + } + return found + } + storedResource := func(id string) string { + var stored models.SCIMUser + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&stored)) + return string(stored.Resource) + } + + w, got := ts.do(ts.TokenA, http.MethodGet, "/Users/"+alice, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), []map[string]any{ + {"value": eng, "$ref": "http://localhost:9999/scim/v2/Groups/" + eng, "display": "Engineering", "type": "direct"}, + {"value": ops, "$ref": "http://localhost:9999/scim/v2/Groups/" + ops, "display": "Ops", "type": "direct"}, + }, groupsOf(got)) + + w, got = ts.do(ts.TokenA, http.MethodGet, "/Users/"+bob, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.NotContains(ts.T(), got, "groups") + + listed := ts.list(ts.TokenA, `userName eq "alice@example.com"`) + require.Len(ts.T(), groupsOf(listed["Resources"].([]any)[0].(map[string]any)), 2) + + claimed := `[{"value":"` + ops + `","display":"Forged"}]` + w, created := ts.do(ts.TokenA, http.MethodPost, "/Users", `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"carol@example.com","emails":[{"primary":true,"value":"carol@example.com"}],"groups":`+claimed+`}`) + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) + require.NotContains(ts.T(), created, "groups") + require.NotContains(ts.T(), storedResource(created["id"].(string)), "groups") + + w, replaced := ts.do(ts.TokenA, http.MethodPut, "/Users/"+bob, `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"bob@example.com","groups":`+claimed+`}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.NotContains(ts.T(), replaced, "groups") + require.NotContains(ts.T(), storedResource(bob), "groups") + + w, patched := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+alice, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ + {"op":"replace","path":"displayName","value":"Alice"} + ]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Len(ts.T(), groupsOf(patched), 2) + require.NotContains(ts.T(), storedResource(alice), "groups") + + w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Groups/"+ops, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + w, got = ts.do(ts.TokenA, http.MethodGet, "/Users/"+alice, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Len(ts.T(), groupsOf(got), 1) + require.Equal(ts.T(), eng, groupsOf(got)[0]["value"]) +} + +func (ts *SCIMUsersTestSuite) TestGroupsAuditLog() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + before := len(ts.scimAuditEntries()) + + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) + + w, _ := ts.do(ts.TokenA, http.MethodPost, "/Groups", groupWith("Engineering", "g-1")) + require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) + w, _ = ts.do(ts.TokenA, http.MethodPost, "/Groups", groupWith("Invalid", "g-2", uuid.Must(uuid.NewV4()).String())) + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + + w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ + {"op":"add","path":"members","value":[{"value":"`+bob+`"}]}, + {"op":"remove","path":"members[value eq \"`+alice+`\"]"}, + {"op":"replace","path":"displayName","value":"Platform"} + ]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + + w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ + {"op":"replace","path":"displayName","value":"Rejected"}, + {"op":"add","path":"members","value":[{"value":"`+uuid.Must(uuid.NewV4()).String()+`"}]} + ]}`) + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + + w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Groups/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + + tokens, err := models.FindSCIMTokensBySSOProvider(ts.API.db, ts.A.ID) + require.NoError(ts.T(), err) + + type event struct { + action, displayName, scimUserID string + } + events := []event{} + for _, entry := range ts.scimAuditEntries()[before:] { + require.Equal(ts.T(), uuid.Nil.String(), entry.Payload["actor_id"]) + require.Equal(ts.T(), "scim:"+tokens[0].Prefix, entry.Payload["actor_username"]) + traits := entry.Payload["traits"].(map[string]any) + require.Equal(ts.T(), ts.A.ID.String(), traits["sso_provider_id"]) + require.Equal(ts.T(), id, traits["scim_group_id"]) + require.Equal(ts.T(), "success", traits["outcome"]) + displayName, _ := traits["display_name"].(string) + scimUserID, _ := traits["scim_user_id"].(string) + events = append(events, event{entry.Payload["action"].(string), displayName, scimUserID}) + } + require.ElementsMatch(ts.T(), []event{ + {string(models.SCIMGroupCreatedAction), "Engineering", ""}, + {string(models.SCIMGroupMemberAddedAction), "", alice}, + {string(models.SCIMGroupUpdatedAction), "Platform", ""}, + {string(models.SCIMGroupMemberAddedAction), "", bob}, + {string(models.SCIMGroupMemberRemovedAction), "", alice}, + {string(models.SCIMGroupMemberRemovedAction), "", bob}, + {string(models.SCIMGroupDeletedAction), "Platform", ""}, + }, events) +} + +func (ts *SCIMUsersTestSuite) TestGroupsPushReplay() { + bjensen := ts.create(ts.TokenA, userWith("bjensen@example.com", "bjensen")) + jsmith := ts.create(ts.TokenA, userWith("jsmith@example.com", "701984")) + + type state struct { + displayName string + members []string + bjensenActive bool + } + expected := map[string]state{ + "push group": {"Tour Guides", []string{}, true}, + "add bjensen (sent twice)": {"Tour Guides", []string{bjensen}, true}, + "retry: bjensen and jsmith": {"Tour Guides", []string{bjensen, jsmith}, true}, + "remove bjensen": {"Tour Guides", []string{jsmith}, true}, + "remove jsmith": {"Tour Guides", []string{}, true}, + "rename": {"Group A", []string{}, true}, + "re-add bjensen (sent twice)": {"Group A", []string{bjensen}, true}, + "deactivate bjensen": {"Group A", []string{bjensen}, false}, + "reactivate bjensen and reassign app": {"Group A", []string{bjensen}, true}, + } + before := len(ts.scimAuditEntries()) + last := map[string]struct { + body string + version any + }{} + onRequest := func(step string, request replayRequest, got map[string]any, _ string) { + version := got["meta"].(map[string]any)["version"] + if prev, ok := last[request.Path]; ok && request.Method == http.MethodPut && prev.body == string(request.Body) { + require.Equal(ts.T(), prev.version, version, step) + } + last[request.Path] = struct { + body string + version any + }{string(request.Body), version} + } + played := ts.replay("okta_group_push.json", rfcGroup, []string{rfcBjensen, bjensen, rfcJsmith, jsmith}, onRequest, func(step, group string) { + want, ok := expected[step] + require.True(ts.T(), ok, step) + ts.requireGroup(step, group, want.displayName, want.members) + w, user := ts.do(ts.TokenA, http.MethodGet, "/Users/"+bjensen, "") + require.Equal(ts.T(), http.StatusOK, w.Code, step) + require.Equal(ts.T(), want.bjensenActive, user["active"], step) + }) + require.Equal(ts.T(), len(expected), played) + + type event struct { + action, subject string + } + events := []event{} + for _, entry := range ts.scimAuditEntries()[before:] { + traits := entry.Payload["traits"].(map[string]any) + subject, _ := traits["scim_user_id"].(string) + if name, ok := traits["display_name"].(string); ok { + subject = name + } + events = append(events, event{entry.Payload["action"].(string), subject}) + } + require.Equal(ts.T(), []event{ + {string(models.SCIMGroupCreatedAction), "Tour Guides"}, + {string(models.SCIMGroupMemberAddedAction), bjensen}, + {string(models.SCIMGroupMemberAddedAction), jsmith}, + {string(models.SCIMGroupMemberRemovedAction), bjensen}, + {string(models.SCIMGroupMemberRemovedAction), jsmith}, + {string(models.SCIMGroupUpdatedAction), "Group A"}, + {string(models.SCIMGroupMemberAddedAction), bjensen}, + {string(models.SCIMUserDeactivatedAction), bjensen}, + {string(models.SCIMUserReactivatedAction), bjensen}, + }, events) +} + +func (ts *SCIMUsersTestSuite) TestGroupsPatchReplay() { + bjensen := ts.create(ts.TokenA, userWith("bjensen@example.com", "bjensen")) + jsmith := ts.create(ts.TokenA, userWith("jsmith@example.com", "701984")) + babs := ts.create(ts.TokenA, userWith("babs@jensen.org", "babs")) + + type state struct { + displayName string + members []string + } + expected := map[string]state{ + "push group": {"Tour Guides", []string{bjensen, jsmith}}, + "remove jsmith": {"Tour Guides", []string{bjensen}}, + "add babs": {"Tour Guides", []string{bjensen, babs}}, + "rename": {"Group B", []string{bjensen, babs}}, + } + played := ts.replay("okta_group_patch.json", rfcGroup, []string{rfcBjensen, bjensen, rfcJsmith, jsmith, rfcBabs, babs}, nil, func(step, group string) { + want, ok := expected[step] + require.True(ts.T(), ok, step) + ts.requireGroup(step, group, want.displayName, want.members) + }) + require.Equal(ts.T(), len(expected), played) +} + +func (ts *SCIMUsersTestSuite) requireGroup(step, group, displayName string, members []string) { + w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+group, "") + require.Equal(ts.T(), http.StatusOK, w.Code, step) + require.Equal(ts.T(), displayName, got["displayName"], step) + require.ElementsMatch(ts.T(), members, memberValues(got), step) +} diff --git a/internal/api/scim_isolation_test.go b/internal/api/scim_isolation_test.go new file mode 100644 index 0000000000..2ab39830b0 --- /dev/null +++ b/internal/api/scim_isolation_test.go @@ -0,0 +1,199 @@ +package api + +import ( + "net/http" + "strings" + "time" + + "github.com/stretchr/testify/require" + "github.com/supabase-community/scim-go/pkg/core" + "github.com/supabase/auth/internal/models" +) + +func (ts *SCIMUsersTestSuite) TestTenantIsolation() { + idA := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + idB := ts.create(ts.TokenB, userWith("bob@example.com", "b-1")) + + listA := ts.list(ts.TokenA, "") + require.EqualValues(ts.T(), 1, listA["totalResults"]) + require.Equal(ts.T(), idA, listA["Resources"].([]any)[0].(map[string]any)["id"]) + require.EqualValues(ts.T(), 0, ts.list(ts.TokenA, `userName eq "bob@example.com"`)["totalResults"]) + require.EqualValues(ts.T(), 0, ts.list(ts.TokenA, `externalId eq "b-1"`)["totalResults"]) + + for _, tc := range []struct{ method, body string }{ + {http.MethodGet, ""}, + {http.MethodPut, userWith("bob@example.com", "b-1")}, + {http.MethodPatch, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","value":{"active":false}}]}`}, + {http.MethodDelete, ""}, + } { + w, _ := ts.do(ts.TokenA, tc.method, "/Users/"+idB, tc.body) + require.Equal(ts.T(), http.StatusNotFound, w.Code, tc.method) + } + + w, got := ts.do(ts.TokenB, http.MethodGet, "/Users/"+idB, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Equal(ts.T(), "bob@example.com", got["userName"]) + require.Equal(ts.T(), true, got["active"]) +} + +func (ts *SCIMUsersTestSuite) TestTenantIsolationWithSameEmail() { + body := userWith("shared@example.com", "shared-1") + ids := map[string]string{ts.TokenA: ts.create(ts.TokenA, body), ts.TokenB: ts.create(ts.TokenB, body)} + require.NotEqual(ts.T(), ids[ts.TokenA], ids[ts.TokenB]) + require.NotEqual(ts.T(), ts.linkedUser(ids[ts.TokenA]).ID, ts.linkedUser(ids[ts.TokenB]).ID) + + for token, other := range map[string]string{ts.TokenA: ts.TokenB, ts.TokenB: ts.TokenA} { + for _, filter := range []string{"", `userName eq "shared@example.com"`, `externalId eq "shared-1"`} { + found := ts.list(token, filter) + require.EqualValues(ts.T(), 1, found["totalResults"], filter) + require.Equal(ts.T(), ids[token], found["Resources"].([]any)[0].(map[string]any)["id"], filter) + } + w, _ := ts.do(token, http.MethodGet, "/Users/"+ids[other], "") + require.Equal(ts.T(), http.StatusNotFound, w.Code) + } + + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+ids[ts.TokenA], "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + w, got := ts.do(ts.TokenB, http.MethodGet, "/Users/"+ids[ts.TokenB], "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Equal(ts.T(), true, got["active"]) +} + +func (ts *SCIMUsersTestSuite) TestTombstonedUsersInvisibleToBothProviders() { + id := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + group := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", id)) + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + + for _, token := range []string{ts.TokenA, ts.TokenB} { + w, _ := ts.do(token, http.MethodGet, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNotFound, w.Code) + for _, filter := range []string{"", `userName eq "alice@example.com"`, `externalId eq "a-1"`} { + require.EqualValues(ts.T(), 0, ts.list(token, filter)["totalResults"], filter) + } + } + + w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+group, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Empty(ts.T(), memberValues(got)) +} + +func (ts *SCIMUsersTestSuite) TestRevokedAndExpiredTokensRefusedEverywhere() { + user := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + group := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", user)) + + tokens, err := models.FindSCIMTokensBySSOProvider(ts.API.db, ts.A.ID) + require.NoError(ts.T(), err) + require.Len(ts.T(), tokens, 1) + require.NoError(ts.T(), tokens[0].Revoke(ts.API.db)) + + expiresAt := time.Now().Add(time.Hour) + expired, expiredToken, err := models.CreateSCIMToken(ts.API.db, ts.A, &expiresAt) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.RawQuery( + "UPDATE "+expired.TableName()+" SET created_at = now() - interval '2 hours', expires_at = now() - interval '1 hour' WHERE id = ?", expired.ID, + ).Exec()) + + routes := []struct{ method, path, body string }{ + {http.MethodGet, "/ResourceTypes", ""}, + {http.MethodGet, "/ResourceTypes/User", ""}, + {http.MethodGet, "/ResourceTypes/Group", ""}, + {http.MethodGet, "/Schemas", ""}, + {http.MethodGet, "/Schemas/" + string(core.SchemaUser), ""}, + {http.MethodGet, "/Schemas/" + string(core.SchemaGroup), ""}, + {http.MethodGet, "/Users", ""}, + {http.MethodPost, "/Users", userWith("bob@example.com", "b-1")}, + {http.MethodGet, "/Users/" + user, ""}, + {http.MethodPut, "/Users/" + user, userWith("alice@example.com", "a-2")}, + {http.MethodPatch, "/Users/" + user, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","value":{"active":false}}]}`}, + {http.MethodDelete, "/Users/" + user, ""}, + {http.MethodGet, "/Groups", ""}, + {http.MethodPost, "/Groups", groupWith("Platform", "g-2")}, + {http.MethodGet, "/Groups/" + group, ""}, + {http.MethodPut, "/Groups/" + group, groupWith("Owned", "g-1")}, + {http.MethodPatch, "/Groups/" + group, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","path":"displayName","value":"Owned"}]}`}, + {http.MethodDelete, "/Groups/" + group, ""}, + } + for name, token := range map[string]string{"revoked": ts.TokenA, "expired": expiredToken} { + w, _ := ts.do(token, http.MethodGet, "/ServiceProviderConfig", "") + require.Equal(ts.T(), http.StatusOK, w.Code, name+" GET /ServiceProviderConfig") + for _, route := range routes { + w, _ := ts.do(token, route.method, route.path, route.body) + require.Equal(ts.T(), http.StatusUnauthorized, w.Code, name+" "+route.method+" "+route.path) + require.True(ts.T(), strings.HasPrefix(w.Header().Get("WWW-Authenticate"), "Bearer"), name+" "+route.method+" "+route.path) + } + } + + var row models.SCIMUser + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", user).First(&row)) + require.True(ts.T(), row.Active) + require.Nil(ts.T(), row.DeletedAt) + require.Contains(ts.T(), string(row.Resource), `"a-1"`) + users, err := ts.API.db.Q().Where("sso_provider_id = ?", ts.A.ID).Count(&models.SCIMUser{}) + require.NoError(ts.T(), err) + require.Equal(ts.T(), 1, users) + + var stored models.SCIMGroup + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", group).First(&stored)) + require.Contains(ts.T(), string(stored.Resource), "Engineering") + groups, err := ts.API.db.Q().Where("sso_provider_id = ?", ts.A.ID).Count(&models.SCIMGroup{}) + require.NoError(ts.T(), err) + require.Equal(ts.T(), 1, groups) + members, err := ts.API.db.Q().Where("group_id = ?", group).Count(&models.SCIMGroupMember{}) + require.NoError(ts.T(), err) + require.Equal(ts.T(), 1, members) + + w, _ := ts.do(ts.TokenB, http.MethodGet, "/Users", "") + require.Equal(ts.T(), http.StatusOK, w.Code) +} + +func (ts *SCIMUsersTestSuite) TestGroupsTenantIsolation() { + aliceA := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + groupA := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", aliceA)) + + require.EqualValues(ts.T(), 0, ts.listGroups(ts.TokenB, "")["totalResults"]) + require.EqualValues(ts.T(), 0, ts.listGroups(ts.TokenB, `displayName eq "Engineering"`)["totalResults"]) + + for _, tc := range []struct{ method, body string }{ + {http.MethodGet, ""}, + {http.MethodPut, groupWith("Engineering", "g-1")}, + {http.MethodPatch, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","path":"displayName","value":"Owned"}]}`}, + {http.MethodDelete, ""}, + } { + w, _ := ts.do(ts.TokenB, tc.method, "/Groups/"+groupA, tc.body) + require.Equal(ts.T(), http.StatusNotFound, w.Code, tc.method) + } + + w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+groupA, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Equal(ts.T(), "Engineering", got["displayName"]) + require.Equal(ts.T(), []string{aliceA}, memberValues(got)) + + bobB := ts.create(ts.TokenB, userWith("bob@example.com", "b-1")) + groupB := ts.createGroup(ts.TokenB, groupWith("Engineering", "g-1", bobB)) + for _, filter := range []string{"", `displayName eq "Engineering"`, `externalId eq "g-1"`} { + found := ts.listGroups(ts.TokenB, filter) + require.EqualValues(ts.T(), 1, found["totalResults"], filter) + require.Equal(ts.T(), groupB, found["Resources"].([]any)[0].(map[string]any)["id"], filter) + } + + for _, tc := range []struct{ method, body string }{ + {http.MethodPut, groupWith("Engineering", "g-1", aliceA, bobB)}, + {http.MethodPatch, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"add","path":"members","value":[{"value":"` + bobB + `"}]}]}`}, + } { + w, body := ts.do(ts.TokenA, tc.method, "/Groups/"+groupA, tc.body) + require.Equal(ts.T(), http.StatusBadRequest, w.Code, tc.method+" "+w.Body.String()) + require.Equal(ts.T(), "invalidValue", body["scimType"], tc.method) + } + + w, _ = ts.do(ts.TokenB, http.MethodDelete, "/Users/"+bobB, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + + w, got = ts.do(ts.TokenA, http.MethodGet, "/Groups/"+groupA, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Equal(ts.T(), []string{aliceA}, memberValues(got)) + w, user := ts.do(ts.TokenA, http.MethodGet, "/Users/"+aliceA, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Len(ts.T(), user["groups"], 1) + require.Equal(ts.T(), groupA, user["groups"].([]any)[0].(map[string]any)["value"]) +} diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go new file mode 100644 index 0000000000..bed0cd3789 --- /dev/null +++ b/internal/api/scim_link_test.go @@ -0,0 +1,763 @@ +package api + +import ( + "errors" + "net/http" + "net/http/httptest" + "strconv" + "strings" + "sync" + "time" + + "github.com/sirupsen/logrus" + logrustest "github.com/sirupsen/logrus/hooks/test" + "github.com/stretchr/testify/require" + "github.com/supabase-community/scim-go/pkg/protocol" + "github.com/supabase/auth/internal/api/apierrors" + "github.com/supabase/auth/internal/api/provider" + "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/storage" +) + +func (ts *SCIMUsersTestSuite) ssoUser(provider *models.SSOProvider, sub, email string) *models.User { + user, err := models.NewUser("", email, "", ts.API.config.JWT.Aud, nil) + require.NoError(ts.T(), err) + user.IsSSOUser = true + require.NoError(ts.T(), ts.API.db.Create(user)) + identity, err := models.NewIdentity(user, "sso:"+provider.ID.String(), map[string]any{"sub": sub, "email": email}) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Create(identity)) + return user +} + +func (ts *SCIMUsersTestSuite) linkedUser(id string) *models.User { + var row models.SCIMUser + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&row)) + require.NotNil(ts.T(), row.UserID) + user, err := models.FindUserByID(ts.API.db, *row.UserID) + require.NoError(ts.T(), err) + return user +} + +func (ts *SCIMUsersTestSuite) identities(user *models.User) []*models.Identity { + identities, err := models.FindIdentitiesByUserID(ts.API.db, user.ID) + require.NoError(ts.T(), err) + return identities +} + +func (ts *SCIMUsersTestSuite) TestCreateProvisionsSSOUser() { + user := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + + require.True(ts.T(), user.IsSSOUser) + require.Equal(ts.T(), "alice@example.com", user.GetEmail()) + require.Equal(ts.T(), ts.API.config.JWT.Aud, user.Aud) + require.NotNil(ts.T(), user.EmailConfirmedAt) + require.False(ts.T(), user.IsBanned()) + require.Equal(ts.T(), []any{"sso:" + ts.A.ID.String()}, user.AppMetaData["providers"]) + + identities := ts.identities(user) + require.Len(ts.T(), identities, 1) + require.Equal(ts.T(), "sso:"+ts.A.ID.String(), identities[0].Provider) + require.Equal(ts.T(), "Alice@Example.com", identities[0].ProviderID) +} + +func (ts *SCIMUsersTestSuite) TestCreateDoesNotLinkOutsideProvider() { + password, err := models.NewUser("", "alice@example.com", "", ts.API.config.JWT.Aud, nil) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Create(password)) + other := ts.ssoUser(ts.B, "Alice@Example.com", "alice@example.com") + + user := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + + require.NotEqual(ts.T(), password.ID, user.ID) + require.NotEqual(ts.T(), other.ID, user.ID) + require.Len(ts.T(), ts.identities(password), 0) + require.Len(ts.T(), ts.identities(other), 1) +} + +func (ts *SCIMUsersTestSuite) TestCreateReusesSAMLIdentity() { + existing := ts.ssoUser(ts.A, "Alice@Example.com", "alice@example.com") + + user := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + + require.Equal(ts.T(), existing.ID, user.ID) + require.Len(ts.T(), ts.identities(user), 1) +} + +func (ts *SCIMUsersTestSuite) TestCreateLinksByEmailWithinProvider() { + existing := ts.ssoUser(ts.A, "saml-name-id", "alice@example.com") + + user := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + + require.Equal(ts.T(), existing.ID, user.ID) + require.Len(ts.T(), ts.identities(user), 2) +} + +func (ts *SCIMUsersTestSuite) TestCreateInactiveLogsOutWithoutBanning() { + existing := ts.ssoUser(ts.A, "Alice@Example.com", "alice@example.com") + ts.session(existing) + + user := ts.linkedUser(ts.create(ts.TokenA, strings.Replace(oktaUser, `"active": true`, `"active": false`, 1))) + + require.False(ts.T(), user.IsBanned()) + require.Zero(ts.T(), ts.sessions(user)) +} + +func (ts *SCIMUsersTestSuite) TestCreateRejectsSharedUser() { + ts.create(ts.TokenA, oktaUser) + + w, body := ts.do(ts.TokenA, http.MethodPost, "/Users", strings.Replace(oktaUser, `"userName": "Alice@Example.com"`, `"userName": "alice.smith"`, 1)) + require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) + require.Equal(ts.T(), "uniqueness", body["scimType"]) + require.EqualValues(ts.T(), 0, ts.list(ts.TokenA, `userName eq "alice.smith"`)["totalResults"]) +} + +func (ts *SCIMUsersTestSuite) TestCreateRequiresEmail() { + w, body := ts.do(ts.TokenA, http.MethodPost, "/Users", `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"alice"}`) + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + require.Equal(ts.T(), "invalidValue", body["scimType"]) +} + +func (ts *SCIMUsersTestSuite) TestCreateKeepsAdminBan() { + existing := ts.ssoUser(ts.A, "Alice@Example.com", "alice@example.com") + require.NoError(ts.T(), existing.Ban(ts.API.db, time.Hour)) + + require.True(ts.T(), ts.linkedUser(ts.create(ts.TokenA, oktaUser)).IsBanned()) +} + +func (ts *SCIMUsersTestSuite) TestCreateLeavesNoUserOnConflict() { + ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + + w, _ := ts.do(ts.TokenA, http.MethodPost, "/Users", userWith("bob@example.com", "a-1")) + require.Equal(ts.T(), http.StatusConflict, w.Code) + count, err := ts.API.db.Q().Where("email = ?", "bob@example.com").Count(&models.User{}) + require.NoError(ts.T(), err) + require.Zero(ts.T(), count) +} + +func (ts *SCIMUsersTestSuite) session(user *models.User) { + session, err := models.NewSession(user.ID, nil) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Create(session)) +} + +func (ts *SCIMUsersTestSuite) refreshToken(user *models.User) string { + token, err := models.GrantAuthenticatedUser(ts.API.db, user, models.GrantParams{}) + require.NoError(ts.T(), err) + return token.Token +} + +func (ts *SCIMUsersTestSuite) refresh(token string) int { + r := httptest.NewRequest(http.MethodPost, "/token?grant_type=refresh_token", strings.NewReader(`{"refresh_token":"`+token+`"}`)) + r.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + return w.Code +} + +func (ts *SCIMUsersTestSuite) sessions(user *models.User) int { + count, err := ts.API.db.Q().Where("user_id = ?", user.ID).Count(&models.Session{}) + require.NoError(ts.T(), err) + return count +} + +func (ts *SCIMUsersTestSuite) samlLogin(ssoProvider *models.SSOProvider, sub, email string) (*models.User, error) { + userData := &provider.UserProvidedData{ + Metadata: &provider.Claims{ + Subject: sub, + Email: email, + EmailVerified: true, + }, + Emails: []provider.Email{{ + Email: email, + Primary: true, + Verified: true, + }}, + } + r := httptest.NewRequest(http.MethodPost, "/sso/saml/acs", nil) + + var user *models.User + err := ts.API.db.Transaction(func(tx *storage.Connection) error { + var terr error + _, user, terr = ts.API.createAccountFromExternalIdentity(tx, r, userData, "sso:"+ssoProvider.ID.String(), false) + return terr + }) + return user, err +} + +func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowedForActiveSCIMUser() { + id := ts.create(ts.TokenA, oktaUser) + linked := ts.linkedUser(id) + + user, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") + + require.NoError(ts.T(), err) + require.Equal(ts.T(), linked.ID, user.ID) +} + +func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowedForDeprovisionedUserWhileSCIMFlagOff() { + id := ts.create(ts.TokenA, oktaUser) + linked := ts.linkedUser(id) + ts.setActive(id, false) + ts.API.config.SSO.SCIM.Enabled = false + defer func() { ts.API.config.SSO.SCIM.Enabled = true }() + + user, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") + require.NoError(ts.T(), err) + require.Equal(ts.T(), linked.ID, user.ID) +} + +func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedWhilePATCHedInactive() { + id := ts.create(ts.TokenA, oktaUser) + ts.setActive(id, false) + + _, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") + require.Error(ts.T(), err) + + ts.setActive(id, true) + linked := ts.linkedUser(id) + + user, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") + require.NoError(ts.T(), err) + require.Equal(ts.T(), linked.ID, user.ID) +} + +func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedAfterDelete() { + id := ts.create(ts.TokenA, oktaUser) + + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + + _, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") + require.Error(ts.T(), err) + + relinked := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + user, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") + require.NoError(ts.T(), err) + require.Equal(ts.T(), relinked.ID, user.ID) +} + +func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedWhenCreatedInactive() { + id := ts.create(ts.TokenA, strings.Replace(oktaUser, `"active": true`, `"active": false`, 1)) + ts.linkedUser(id) + + _, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") + require.Error(ts.T(), err) +} + +func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowedWhenActiveOmitted() { + id := ts.create(ts.TokenA, `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"Alice@Example.com","emails":[{"primary":true,"value":"alice@example.com"}]}`) + linked := ts.linkedUser(id) + + user, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") + require.NoError(ts.T(), err) + require.Equal(ts.T(), linked.ID, user.ID) +} + +func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowedWithoutSCIMRow() { + existing := ts.ssoUser(ts.A, "jit-user", "jit@example.com") + + user, err := ts.samlLogin(ts.A, "jit-user", "jit@example.com") + + require.NoError(ts.T(), err) + require.Equal(ts.T(), existing.ID, user.ID) +} + +func (ts *SCIMUsersTestSuite) TestSAMLLoginLinksDivergedNameIDToSCIMUser() { + linked := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + + for range 2 { + user, err := ts.samlLogin(ts.A, "saml-name-id", "alice@example.com") + require.NoError(ts.T(), err) + require.Equal(ts.T(), linked.ID, user.ID) + } + + count, err := ts.API.db.Q().Where("email = ?", "alice@example.com").Count(&models.User{}) + require.NoError(ts.T(), err) + require.Equal(ts.T(), 1, count) + + providerIDs := []string{} + for _, identity := range ts.identities(linked) { + require.Equal(ts.T(), "sso:"+ts.A.ID.String(), identity.Provider) + providerIDs = append(providerIDs, identity.ProviderID) + } + require.ElementsMatch(ts.T(), []string{"Alice@Example.com", "saml-name-id"}, providerIDs) +} + +func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedForDivergedNameIDWhileInactive() { + id := ts.create(ts.TokenA, oktaUser) + ts.setActive(id, false) + linked := ts.linkedUser(id) + + _, err := ts.samlLogin(ts.A, "saml-name-id", "alice@example.com") + require.Error(ts.T(), err) + require.Len(ts.T(), ts.identities(linked), 1) +} + +func (ts *SCIMUsersTestSuite) users(email string) int { + count, err := ts.API.db.Q().Where("email = ?", email).Count(&models.User{}) + require.NoError(ts.T(), err) + return count +} + +func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowsJITWithoutSCIMToken() { + provider := createSSOProvider(ts.T(), ts.API.db) + + user, err := ts.samlLogin(provider, "jit-user", "jit@example.com") + + require.NoError(ts.T(), err) + require.Equal(ts.T(), "jit@example.com", user.GetEmail()) +} + +func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowsJITWhileSCIMFlagOff() { + ts.API.config.SSO.SCIM.Enabled = false + defer func() { ts.API.config.SSO.SCIM.Enabled = true }() + + user, err := ts.samlLogin(ts.A, "jit-user", "jit@example.com") + + require.NoError(ts.T(), err) + require.Equal(ts.T(), "jit@example.com", user.GetEmail()) +} + +func (ts *SCIMUsersTestSuite) TestSAMLLoginNotBlockedByOtherProvider() { + id := ts.create(ts.TokenB, oktaUser) + w, _ := ts.do(ts.TokenB, http.MethodPatch, "/Users/"+id, `{ + "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + "Operations": [{"op": "replace", "value": {"active": false}}] + }`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + + existing := ts.ssoUser(ts.A, "Alice@Example.com", "alice@example.com") + user, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") + + require.NoError(ts.T(), err) + require.Equal(ts.T(), existing.ID, user.ID) +} + +func (ts *SCIMUsersTestSuite) issueSession(conn *storage.Connection, user *models.User) error { + r := httptest.NewRequest(http.MethodPost, "/token", nil) + _, err := ts.API.tokenService.IssueRefreshToken(r, http.Header{}, conn, user, models.OAuth, models.GrantParams{}) + return err +} + +func (ts *SCIMUsersTestSuite) requireBanned(err error) { + var httpErr *apierrors.HTTPError + require.True(ts.T(), errors.As(err, &httpErr), err) + require.Equal(ts.T(), http.StatusForbidden, httpErr.HTTPStatus) + require.Equal(ts.T(), apierrors.ErrorCodeUserBanned, httpErr.ErrorCode) +} + +func (ts *SCIMUsersTestSuite) TestSessionRefusedWhileDeprovisioned() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + require.NoError(ts.T(), ts.issueSession(ts.API.db, user)) + + ts.setActive(id, false) + ts.requireBanned(ts.issueSession(ts.API.db, user)) + require.Zero(ts.T(), ts.sessions(user)) + + ts.setActive(id, true) + require.NoError(ts.T(), ts.issueSession(ts.API.db, user)) + + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + ts.requireBanned(ts.issueSession(ts.API.db, user)) + + relinked := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + require.Equal(ts.T(), user.ID, relinked.ID) + require.NoError(ts.T(), ts.issueSession(ts.API.db, relinked)) +} + +func (ts *SCIMUsersTestSuite) TestSessionAllowedWhileSCIMFlagOff() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + ts.setActive(id, false) + ts.API.config.SSO.SCIM.Enabled = false + defer func() { ts.API.config.SSO.SCIM.Enabled = true }() + + require.NoError(ts.T(), ts.issueSession(ts.API.db, user)) +} + +func (ts *SCIMUsersTestSuite) TestSessionRefusedForLinkedOAuthIdentityWhileDeprovisioned() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + identity, err := models.NewIdentity(user, "google", map[string]any{"sub": "google-sub", "email": "alice@example.com"}) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Create(identity)) + ts.setActive(id, false) + + userData := &provider.UserProvidedData{ + Metadata: &provider.Claims{Subject: "google-sub", Email: "alice@example.com", EmailVerified: true}, + Emails: []provider.Email{{Email: "alice@example.com", Primary: true, Verified: true}}, + } + err = ts.API.db.Transaction(func(tx *storage.Connection) error { + _, found, terr := ts.API.createAccountFromExternalIdentity(tx, httptest.NewRequest(http.MethodGet, "/callback", nil), userData, "google", false) + if terr != nil { + return terr + } + require.Equal(ts.T(), user.ID, found.ID) + return ts.issueSession(tx, found) + }) + ts.requireBanned(err) + require.Zero(ts.T(), ts.sessions(user)) +} + +func (ts *SCIMUsersTestSuite) TestSessionWaitsForConcurrentDeactivation() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + locked, release := make(chan struct{}), make(chan struct{}) + deactivated := make(chan error, 1) + go func() { + deactivated <- ts.API.db.Transaction(func(tx *storage.Connection) error { + if err := models.LockUserForSCIM(tx, user.ID); err != nil { + return err + } + close(locked) + <-release + if err := tx.RawQuery("UPDATE "+(&models.SCIMUser{}).TableName()+" SET deleted_at = now() WHERE id = ?", id).Exec(); err != nil { + return err + } + return models.Logout(tx, user.ID) + }) + }() + <-locked + + issued := make(chan error, 1) + go func() { issued <- ts.issueSession(ts.API.db, user) }() + select { + case err := <-issued: + ts.T().Fatalf("session issued while deactivation held the user lock: %v", err) + case <-time.After(200 * time.Millisecond): + } + close(release) + + require.NoError(ts.T(), <-deactivated) + ts.requireBanned(<-issued) + require.Zero(ts.T(), ts.sessions(user)) +} + +func (ts *SCIMUsersTestSuite) TestWritesWaitForAdminUserDelete() { + for _, method := range []string{http.MethodDelete, http.MethodPut} { + body := userWith(strings.ToLower(method)+"@example.com", method) + id := ts.create(ts.TokenA, body) + user := ts.linkedUser(id) + code, err := ts.whileLocked( + func(tx *storage.Connection) error { return models.LockUserForSCIM(tx, user.ID) }, + func(tx *storage.Connection) error { + _, err := models.SoftDeleteSCIMUsersByUserID(tx, user.ID) + return err + }, + method, "/Users/"+id, body, + ) + require.NoError(ts.T(), err, method) + require.Equal(ts.T(), http.StatusNotFound, code, method) + } +} + +func (ts *SCIMUsersTestSuite) TestSessionAllowedForSSOUserWithoutSCIMRow() { + user := ts.ssoUser(ts.A, "saml-sub", "carol@example.com") + require.NoError(ts.T(), ts.issueSession(ts.API.db, user)) +} + +func (ts *SCIMUsersTestSuite) setActive(id string, active bool) { + ts.setActiveAs(ts.TokenA, id, active) +} + +func (ts *SCIMUsersTestSuite) setActiveAs(token, id string, active bool) { + w, _ := ts.do(token, http.MethodPatch, "/Users/"+id, `{ + "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + "Operations": [{"op": "replace", "value": {"active": `+strconv.FormatBool(active)+`}}] + }`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) +} + +func (ts *SCIMUsersTestSuite) TestReplaceDeactivatesAndReactivates() { + id := ts.create(ts.TokenA, oktaUser) + require.Equal(ts.T(), http.StatusOK, ts.refresh(ts.refreshToken(ts.linkedUser(id)))) + refreshToken := ts.refreshToken(ts.linkedUser(id)) + + ts.setActive(id, false) + user := ts.linkedUser(id) + require.False(ts.T(), user.IsBanned()) + require.Zero(ts.T(), ts.sessions(user)) + require.Equal(ts.T(), http.StatusBadRequest, ts.refresh(refreshToken)) + + listed := ts.list(ts.TokenA, "") + require.EqualValues(ts.T(), 1, listed["totalResults"]) + require.Equal(ts.T(), id, listed["Resources"].([]any)[0].(map[string]any)["id"]) + require.Equal(ts.T(), false, listed["Resources"].([]any)[0].(map[string]any)["active"]) + + ts.setActive(id, true) + require.False(ts.T(), ts.linkedUser(id).IsBanned()) + require.Zero(ts.T(), ts.sessions(user)) + require.Equal(ts.T(), http.StatusBadRequest, ts.refresh(refreshToken)) +} + +func (ts *SCIMUsersTestSuite) TestPutInactiveRevokesSessions() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + refreshToken := ts.refreshToken(user) + + w, got := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, strings.Replace(oktaUser, `"active": true`, `"active": false`, 1)) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), false, got["active"]) + require.Zero(ts.T(), ts.sessions(user)) + require.Equal(ts.T(), http.StatusBadRequest, ts.refresh(refreshToken)) + require.Len(ts.T(), ts.auditActions(models.SCIMUserDeactivatedAction), 1) +} + +func (ts *SCIMUsersTestSuite) TestReplaceKeepsAdminBanWhenActiveDoesNotChange() { + id := ts.create(ts.TokenA, oktaUser) + require.NoError(ts.T(), ts.linkedUser(id).Ban(ts.API.db, time.Hour)) + + ts.setActive(id, true) + require.True(ts.T(), ts.linkedUser(id).IsBanned()) +} + +func (ts *SCIMUsersTestSuite) TestReplaceLinksUnlinkedRow() { + row, err := models.CreateSCIMUser(ts.API.db, ts.A.ID, []byte(`{"userName":"Alice@Example.com"}`)) + require.NoError(ts.T(), err) + + w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+row.ID.String(), strings.Replace(oktaUser, `"active": true`, `"active": false`, 1)) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + + user := ts.linkedUser(row.ID.String()) + require.Equal(ts.T(), "alice@example.com", user.GetEmail()) + require.False(ts.T(), user.IsBanned()) +} + +func (ts *SCIMUsersTestSuite) TestDeleteLogsOutWithoutBanning() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + refreshToken := ts.refreshToken(user) + + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + + user, err := models.FindUserByID(ts.API.db, user.ID) + require.NoError(ts.T(), err) + require.False(ts.T(), user.IsBanned()) + require.Zero(ts.T(), ts.sessions(user)) + require.Equal(ts.T(), http.StatusBadRequest, ts.refresh(refreshToken)) +} + +func (ts *SCIMUsersTestSuite) TestCreateDoesNotBanAfterDelete() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + + relinked := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + + require.Equal(ts.T(), user.ID, relinked.ID) + require.False(ts.T(), relinked.IsBanned()) +} + +func (ts *SCIMUsersTestSuite) TestCreateConcurrentSameEmailLinksToOneUser() { + body := func(userName, externalID string) string { + return `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"` + userName + `","externalId":"` + externalID + `","emails":[{"primary":true,"value":"race@example.com"}]}` + } + bodies := []string{body("race-a", "race-a"), body("race-b", "race-b")} + + var wg sync.WaitGroup + start := make(chan struct{}) + codes := make([]int, len(bodies)) + for i := range bodies { + wg.Add(1) + go func(i int) { + defer wg.Done() + <-start + r := httptest.NewRequest(http.MethodPost, "/scim/v2/Users", strings.NewReader(bodies[i])) + r.Header.Set("Authorization", "Bearer "+ts.TokenA) + r.Header.Set("Content-Type", protocol.MediaType) + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + codes[i] = w.Code + }(i) + } + close(start) + wg.Wait() + + // The lock serializes the two creates: whichever commits first creates the + // user, the other observes that account under the same provider and is + // rejected as already linked -- never silently creating a second user. + created := 0 + for _, code := range codes { + if code == http.StatusCreated { + created++ + } else { + require.Equal(ts.T(), http.StatusConflict, code) + } + } + require.Equal(ts.T(), 1, created) + + count, err := ts.API.db.Q().Where("email = ?", "race@example.com").Count(&models.User{}) + require.NoError(ts.T(), err) + require.EqualValues(ts.T(), 1, count) +} + +func (ts *SCIMUsersTestSuite) rename(id, userName string) (int, string) { + w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, strings.Replace(oktaUser, `"userName": "Alice@Example.com"`, `"userName": "`+userName+`"`, 1)) + return w.Code, w.Body.String() +} + +func (ts *SCIMUsersTestSuite) TestReplaceRenamesSSOIdentity() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + + code, body := ts.rename(id, "alice2@example.com") + require.Equal(ts.T(), http.StatusOK, code, body) + + identity, err := models.FindIdentityByIdAndProvider(ts.API.db, "alice2@example.com", "sso:"+ts.A.ID.String()) + require.NoError(ts.T(), err) + require.Equal(ts.T(), user.ID, identity.UserID) + require.Equal(ts.T(), "alice2@example.com", identity.IdentityData["sub"]) + require.Len(ts.T(), ts.identities(user), 1) +} + +func (ts *SCIMUsersTestSuite) providerIDs(user *models.User) []string { + ids := []string{} + for _, identity := range ts.identities(user) { + ids = append(ids, identity.ProviderID) + } + return ids +} + +func (ts *SCIMUsersTestSuite) TestReplaceRenamesCaseOnly() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + _, err := ts.samlLogin(ts.A, "alice@example.com", "alice@example.com") + require.NoError(ts.T(), err) + require.ElementsMatch(ts.T(), []string{"Alice@Example.com", "alice@example.com"}, ts.providerIDs(user)) + + for range 2 { + code, body := ts.rename(id, "alice@example.com") + require.Equal(ts.T(), http.StatusOK, code, body) + } + require.Equal(ts.T(), []string{"alice@example.com"}, ts.providerIDs(user)) + + signedIn, err := ts.samlLogin(ts.A, "alice@example.com", "alice@example.com") + require.NoError(ts.T(), err) + require.Equal(ts.T(), user.ID, signedIn.ID) +} + +func (ts *SCIMUsersTestSuite) TestReplaceRenameRemovesOldNameIDIdentity() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + _, err := ts.samlLogin(ts.A, "alice@example.com", "alice@example.com") + require.NoError(ts.T(), err) + _, err = ts.samlLogin(ts.A, "saml-name-id", "alice@example.com") + require.NoError(ts.T(), err) + + code, body := ts.rename(id, "bob@example.com") + require.Equal(ts.T(), http.StatusOK, code, body) + require.ElementsMatch(ts.T(), []string{"bob@example.com", "saml-name-id"}, ts.providerIDs(user)) + + signedIn, err := ts.samlLogin(ts.A, "alice@example.com", "new-hire@example.com") + require.NoError(ts.T(), err) + require.NotEqual(ts.T(), user.ID, signedIn.ID) +} + +func (ts *SCIMUsersTestSuite) TestReplaceRejectsRenameToTakenIdentity() { + id := ts.create(ts.TokenA, oktaUser) + ts.ssoUser(ts.A, "bob@example.com", "bob@example.com") + + code, body := ts.rename(id, "bob@example.com") + require.Equal(ts.T(), http.StatusConflict, code, body) + + var row models.SCIMUser + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&row)) + require.Equal(ts.T(), "alice@example.com", row.UserName) + _, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", "sso:"+ts.A.ID.String()) + require.NoError(ts.T(), err) +} + +func (ts *SCIMUsersTestSuite) unlink(user *models.User, identity *models.Identity) *httptest.ResponseRecorder { + session, err := models.NewSession(user.ID, nil) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Create(session)) + token, _, err := ts.API.generateAccessToken(httptest.NewRequest(http.MethodPost, "/token", nil), ts.API.db, user, &session.ID, models.PasswordGrant) + require.NoError(ts.T(), err) + r := httptest.NewRequest(http.MethodDelete, "/user/identities/"+identity.ID.String(), nil) + r.Header.Set("Authorization", "Bearer "+token) + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + return w +} + +func (ts *SCIMUsersTestSuite) TestUnlinkRefusedForSCIMManagedIdentity() { + ts.API.config.Security.ManualLinkingEnabled = true + defer func() { ts.API.config.Security.ManualLinkingEnabled = false }() + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + google, err := models.NewIdentity(user, "google", map[string]any{"sub": "google-1", "email": "alice@example.com"}) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Create(google)) + sso, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", "sso:"+ts.A.ID.String()) + require.NoError(ts.T(), err) + + w := ts.unlink(user, sso) + require.Equal(ts.T(), http.StatusUnprocessableEntity, w.Code, w.Body.String()) + require.Contains(ts.T(), w.Body.String(), string(apierrors.ErrorCodeUserSSOManaged)) + require.Len(ts.T(), ts.identities(user), 2) + + code, body := ts.rename(id, "alice2@example.com") + require.Equal(ts.T(), http.StatusOK, code, body) + sso, err = models.FindIdentityByIdAndProvider(ts.API.db, "alice2@example.com", "sso:"+ts.A.ID.String()) + require.NoError(ts.T(), err) + + w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + w = ts.unlink(user, sso) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Len(ts.T(), ts.identities(user), 1) +} + +func (ts *SCIMUsersTestSuite) TestUnlinkAllowedWhileSCIMFlagOff() { + ts.API.config.Security.ManualLinkingEnabled = true + defer func() { ts.API.config.Security.ManualLinkingEnabled = false }() + user := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + google, err := models.NewIdentity(user, "google", map[string]any{"sub": "google-1", "email": "alice@example.com"}) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Create(google)) + sso, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", "sso:"+ts.A.ID.String()) + require.NoError(ts.T(), err) + ts.API.config.SSO.SCIM.Enabled = false + defer func() { ts.API.config.SSO.SCIM.Enabled = true }() + + w := ts.unlink(user, sso) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Len(ts.T(), ts.identities(user), 1) +} + +func (ts *SCIMUsersTestSuite) TestRenameSkippedWhenIdentityMissing() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + sso, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", "sso:"+ts.A.ID.String()) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Destroy(sso)) + before := len(ts.scimAuditEntries()) + hook := logrustest.NewGlobal() + defer hook.Reset() + + code, body := ts.rename(id, "alice2@example.com") + require.Equal(ts.T(), http.StatusOK, code, body) + + var row models.SCIMUser + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&row)) + require.Equal(ts.T(), "alice2@example.com", row.UserName) + require.Empty(ts.T(), ts.identities(user)) + entries := ts.scimAuditEntries()[before:] + require.Len(ts.T(), entries, 1) + require.Equal(ts.T(), string(models.SCIMUserUpdatedAction), entries[0].Payload["action"]) + warned := false + for _, entry := range hook.AllEntries() { + warned = warned || entry.Level == logrus.WarnLevel && strings.Contains(entry.Message, "SCIM identity") + } + require.True(ts.T(), warned) + + signedIn, err := ts.samlLogin(ts.A, "alice2@example.com", "alice@example.com") + require.NoError(ts.T(), err) + require.NotEqual(ts.T(), user.ID, signedIn.ID) + require.Equal(ts.T(), 2, ts.users("alice@example.com")) +} diff --git a/internal/api/scim_okta_spec_test.go b/internal/api/scim_okta_spec_test.go new file mode 100644 index 0000000000..c36f30870c --- /dev/null +++ b/internal/api/scim_okta_spec_test.go @@ -0,0 +1,228 @@ +package api + +import ( + "encoding/json" + "io/fs" + "net/http" + "net/url" + "os" + "strings" + + "github.com/stretchr/testify/require" + "github.com/supabase-community/scim-go/pkg/core" + "github.com/supabase-community/scim-go/pkg/protocol" + "github.com/supabase/auth/internal/models" +) + +func (ts *SCIMUsersTestSuite) okta(method, path, body string, headers ...string) (int, map[string]any) { + contentType := "application/scim+json; charset=utf-8" + if method == http.MethodPost { + contentType = "application/json" + } + headers = append([]string{"Accept", "application/scim+json", "Accept-Charset", "utf-8", "User-Agent", "OKTA SCIM Integration"}, headers...) + w, got := ts.doAs(contentType, ts.TokenA, method, path, body, headers...) + return w.Code, got +} + +const ( + rfcBjensen = "2819c223-7f76-453a-919d-413861904646" + rfcJsmith = "c75ad752-64ae-4823-840d-ffa80929976c" + rfcBabs = "6c5bb468-14b2-4183-baf2-06d523e03bd3" + rfcGroup = "e9e30dba-f08f-4109-8486-d5c6a331660a" +) + +type replayRequest struct { + Method string `json:"method"` + Path string `json:"path"` + Body json.RawMessage `json:"body"` +} + +func (ts *SCIMUsersTestSuite) replay(file, created string, ids []string, onRequest func(step string, request replayRequest, got map[string]any, id string), afterStep func(step, id string)) int { + raw, err := fs.ReadFile(os.DirFS("testdata/scim"), file) + require.NoError(ts.T(), err) + var steps []struct { + Step string `json:"step"` + Requests []replayRequest `json:"requests"` + } + require.NoError(ts.T(), json.Unmarshal([]byte(strings.NewReplacer(ids...).Replace(string(raw))), &steps)) + + id := created + for _, step := range steps { + for _, request := range step.Requests { + request.Path = strings.ReplaceAll(strings.TrimPrefix(request.Path, "/scim/v2"), created, id) + request.Body = json.RawMessage(strings.ReplaceAll(string(request.Body), created, id)) + status, got := ts.okta(request.Method, request.Path, string(request.Body)) + require.Less(ts.T(), status, 300, "%s: %s %s: %v", step.Step, request.Method, request.Path, got) + if request.Method == http.MethodPost { + id = got["id"].(string) + } + if onRequest != nil { + onRequest(step.Step, request, got, id) + } + } + afterStep(step.Step, id) + } + return len(steps) +} + +func oktaFilter(userName string) string { + return "/Users?" + url.Values{"filter": {`userName eq "` + userName + `"`}}.Encode() +} + +func (ts *SCIMUsersTestSuite) TestOktaSpec() { + const ( + userName = "okta.spec.user@example.com" + givenName = "Okta" + familyName = "Spec" + ) + create := func(email string) string { + return `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"` + userName + `","name":{"givenName":"` + givenName + `","familyName":"` + familyName + `"},"emails":[{"primary":true,"value":"` + email + `","type":"work"}],"displayName":"` + givenName + " " + familyName + `","active":true}` + } + requireError := func(got map[string]any, status string) { + require.NotEmpty(ts.T(), got["detail"]) + require.Equal(ts.T(), status, got["status"]) + require.Contains(ts.T(), got["schemas"], string(protocol.SchemaError)) + } + ts.create(ts.TokenA, oktaUser) + + status, got := ts.okta(http.MethodGet, "/Users?count=1&startIndex=1", "") + require.Equal(ts.T(), http.StatusOK, status) + require.Contains(ts.T(), got["schemas"], string(protocol.SchemaListResponse)) + require.IsType(ts.T(), float64(0), got["itemsPerPage"]) + require.IsType(ts.T(), float64(0), got["startIndex"]) + require.IsType(ts.T(), float64(0), got["totalResults"]) + require.NotEmpty(ts.T(), got["Resources"]) + first := got["Resources"].([]any)[0].(map[string]any) + require.NotEmpty(ts.T(), first["id"]) + require.NotEmpty(ts.T(), first["name"].(map[string]any)["familyName"]) + require.NotEmpty(ts.T(), first["name"].(map[string]any)["givenName"]) + require.NotEmpty(ts.T(), first["userName"]) + require.NotNil(ts.T(), first["active"]) + require.NotEmpty(ts.T(), first["emails"].([]any)[0].(map[string]any)["value"]) + id := first["id"].(string) + + status, got = ts.okta(http.MethodGet, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusOK, status) + require.Equal(ts.T(), id, got["id"]) + require.NotEmpty(ts.T(), got["name"].(map[string]any)["familyName"]) + require.NotEmpty(ts.T(), got["name"].(map[string]any)["givenName"]) + require.NotEmpty(ts.T(), got["userName"]) + require.NotNil(ts.T(), got["active"]) + require.NotEmpty(ts.T(), got["emails"].([]any)[0].(map[string]any)["value"]) + + for _, missing := range []string{"invalid.user@example.com", userName} { + status, got = ts.okta(http.MethodGet, oktaFilter(missing), "") + require.Equal(ts.T(), http.StatusOK, status, missing) + require.Contains(ts.T(), got["schemas"], string(protocol.SchemaListResponse), missing) + require.EqualValues(ts.T(), 0, got["totalResults"], missing) + } + + status, got = ts.okta(http.MethodGet, "/Users/010101", "") + require.Equal(ts.T(), http.StatusNotFound, status) + requireError(got, "404") + + status, got = ts.okta(http.MethodPost, "/Users", create(userName)) + require.Equal(ts.T(), http.StatusCreated, status) + require.Equal(ts.T(), true, got["active"]) + require.NotEmpty(ts.T(), got["id"]) + require.Equal(ts.T(), familyName, got["name"].(map[string]any)["familyName"]) + require.Equal(ts.T(), givenName, got["name"].(map[string]any)["givenName"]) + require.Contains(ts.T(), got["schemas"], string(core.SchemaUser)) + require.Equal(ts.T(), userName, got["userName"]) + created := got["id"].(string) + + status, got = ts.okta(http.MethodGet, "/Users/"+created, "") + require.Equal(ts.T(), http.StatusOK, status) + require.Equal(ts.T(), userName, got["userName"]) + require.Equal(ts.T(), familyName, got["name"].(map[string]any)["familyName"]) + require.Equal(ts.T(), givenName, got["name"].(map[string]any)["givenName"]) + + status, _ = ts.okta(http.MethodPost, "/Users", create(userName)) + require.Equal(ts.T(), http.StatusConflict, status) + + status, got = ts.okta(http.MethodGet, oktaFilter(strings.ToUpper(userName)), "") + require.Equal(ts.T(), http.StatusOK, status) + require.EqualValues(ts.T(), 1, got["totalResults"]) + require.Equal(ts.T(), created, got["Resources"].([]any)[0].(map[string]any)["id"]) + + status, got = ts.okta(http.MethodGet, "/Groups", "") + require.Equal(ts.T(), http.StatusOK, status) + require.EqualValues(ts.T(), 0, got["totalResults"]) + + status, got = ts.okta(http.MethodGet, oktaFilter(strings.ToUpper(userName)), "", "Authorization", "non-token") + require.Equal(ts.T(), http.StatusUnauthorized, status) + requireError(got, "401") + + status, got = ts.okta(http.MethodGet, "/Users/00919288221112222", "") + require.Equal(ts.T(), http.StatusNotFound, status) + requireError(got, "404") +} + +func (ts *SCIMUsersTestSuite) TestOktaUserLifecycleReplay() { + const password = "okta-generated-password" + + type state struct { + familyName string + active bool + found int + } + expected := map[string]state{ + "assign new user": {"Smith", true, 0}, + "edit last name": {"Jensen", true, -1}, + "unassign": {"Jensen", false, -1}, + "reassign (PUT sent twice)": {"Jensen", true, 1}, + } + before := len(ts.scimAuditEntries()) + + id, version := "", "" + played := ts.replay("okta_user_lifecycle.json", rfcBjensen, nil, func(step string, request replayRequest, got map[string]any, created string) { + if strings.Contains(request.Path, "filter=") { + want := expected[step] + require.EqualValues(ts.T(), want.found, got["totalResults"], step) + if want.found > 0 { + require.Equal(ts.T(), created, got["Resources"].([]any)[0].(map[string]any)["id"], step) + } + } + if request.Method == http.MethodPost { + require.Contains(ts.T(), string(request.Body), password) + require.NotContains(ts.T(), got, "password") + } + }, func(step, created string) { + want, ok := expected[step] + require.True(ts.T(), ok, step) + id = created + status, got := ts.okta(http.MethodGet, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusOK, status, step) + require.Equal(ts.T(), want.familyName, got["name"].(map[string]any)["familyName"], step) + require.Equal(ts.T(), want.active, got["active"], step) + current := got["meta"].(map[string]any)["version"].(string) + require.NotEqual(ts.T(), version, current, step) + version = current + + count, err := ts.API.db.Q().Where("sso_provider_id = ?", ts.A.ID).Count(&models.SCIMUser{}) + require.NoError(ts.T(), err) + require.Equal(ts.T(), 1, count, step) + }) + require.Equal(ts.T(), len(expected), played) + + var stored models.SCIMUser + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&stored)) + require.NotContains(ts.T(), string(stored.Resource), "password") + user, err := models.FindUserByID(ts.API.db, *stored.UserID) + require.NoError(ts.T(), err) + require.False(ts.T(), user.HasPassword()) + + actions := []string{} + for _, entry := range ts.scimAuditEntries()[before:] { + payload, err := json.Marshal(entry.Payload) + require.NoError(ts.T(), err) + require.NotContains(ts.T(), string(payload), password) + actions = append(actions, entry.Payload["action"].(string)) + } + require.Equal(ts.T(), []string{ + string(models.SCIMUserCreatedAction), + string(models.SCIMUserUpdatedAction), + string(models.SCIMUserDeactivatedAction), + string(models.SCIMUserReactivatedAction), + }, actions) +} diff --git a/internal/api/scim_provider_delete_test.go b/internal/api/scim_provider_delete_test.go new file mode 100644 index 0000000000..b1b37489e2 --- /dev/null +++ b/internal/api/scim_provider_delete_test.go @@ -0,0 +1,222 @@ +package api + +import ( + "net/http" + "net/http/httptest" + "time" + + "github.com/gofrs/uuid" + jwt "github.com/golang-jwt/jwt/v5" + "github.com/stretchr/testify/require" + "github.com/supabase/auth/internal/api/provider" + "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/storage" +) + +func scimUser(name string) string { + return userWith(name+"@example.com", name) +} + +func (ts *SCIMUsersTestSuite) deleteProvider(p *models.SSOProvider) { + token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, &AccessTokenClaims{Role: "supabase_admin"}).SignedString([]byte(ts.API.config.JWT.Secret)) + require.NoError(ts.T(), err) + r := httptest.NewRequest(http.MethodDelete, "/admin/sso/providers/"+p.ID.String(), nil) + r.Header.Set("Authorization", "Bearer "+token) + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) +} + +func (ts *SCIMUsersTestSuite) reloadUser(id uuid.UUID) *models.User { + user, err := models.FindUserByID(ts.API.db, id) + require.NoError(ts.T(), err) + return user +} + +func (ts *SCIMUsersTestSuite) countRows(model any, where string, args ...any) int { + count, err := ts.API.db.Q().Where(where, args...).Count(model) + require.NoError(ts.T(), err) + return count +} + +func (ts *SCIMUsersTestSuite) auditActions(action models.AuditAction) []models.AuditLogEntry { + entries := []models.AuditLogEntry{} + require.NoError(ts.T(), ts.API.db.Q().Where("payload->>'action' = ?", string(action)).All(&entries)) + return entries +} + +func (ts *SCIMUsersTestSuite) TestProviderDeleteBansDeprovisionedUsers() { + active := ts.linkedUser(ts.create(ts.TokenA, scimUser("active"))) + + deactivatedID := ts.create(ts.TokenA, scimUser("deactivated")) + deactivated := ts.linkedUser(deactivatedID) + ts.setActive(deactivatedID, false) + + deletedID := ts.create(ts.TokenA, scimUser("deleted")) + deleted := ts.linkedUser(deletedID) + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+deletedID, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + + recreatedID := ts.create(ts.TokenA, scimUser("recreated")) + recreated := ts.linkedUser(recreatedID) + w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Users/"+recreatedID, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + require.Equal(ts.T(), recreated.ID, ts.linkedUser(ts.create(ts.TokenA, scimUser("recreated"))).ID) + + ts.createGroup(ts.TokenA, groupWith("A", "", deactivatedID)) + + otherID := ts.create(ts.TokenB, scimUser("other")) + other := ts.linkedUser(otherID) + ts.setActiveAs(ts.TokenB, otherID, false) + + ts.deleteProvider(ts.A) + + require.False(ts.T(), ts.reloadUser(active.ID).IsBanned()) + require.True(ts.T(), ts.reloadUser(deleted.ID).IsBanned()) + require.False(ts.T(), ts.reloadUser(recreated.ID).IsBanned()) + require.False(ts.T(), ts.reloadUser(other.ID).IsBanned()) + require.True(ts.T(), ts.reloadUser(deactivated.ID).BannedUntil.After(time.Now().Add(99*365*24*time.Hour))) + + require.Zero(ts.T(), ts.countRows(&models.SCIMUser{}, "sso_provider_id = ?", ts.A.ID)) + require.Zero(ts.T(), ts.countRows(&models.SCIMGroup{}, "sso_provider_id = ?", ts.A.ID)) + require.Zero(ts.T(), ts.countRows(&models.SCIMToken{}, "sso_provider_id = ?", ts.A.ID)) + require.Zero(ts.T(), ts.countRows(&models.SCIMGroupMember{}, "1 = 1")) + require.Equal(ts.T(), 1, ts.countRows(&models.SCIMUser{}, "sso_provider_id = ?", ts.B.ID)) + w, _ = ts.do(ts.TokenB, http.MethodGet, "/Users", "") + require.Equal(ts.T(), http.StatusOK, w.Code) +} + +func (ts *SCIMUsersTestSuite) TestProviderDeleteKeepsLongerBan() { + id := ts.create(ts.TokenA, scimUser("banned")) + user := ts.linkedUser(id) + ts.setActive(id, false) + require.NoError(ts.T(), user.Ban(ts.API.db, 200*365*24*time.Hour)) + until := *ts.reloadUser(user.ID).BannedUntil + + ts.deleteProvider(ts.A) + + require.True(ts.T(), until.Equal(*ts.reloadUser(user.ID).BannedUntil)) + require.Empty(ts.T(), ts.auditActions(models.SCIMUsersBannedAction)) +} + +func (ts *SCIMUsersTestSuite) TestProviderDeleteClosesOAuthBypass() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + identity, err := models.NewIdentity(user, "google", map[string]any{"sub": "google-sub", "email": "alice@example.com"}) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Create(identity)) + ts.setActive(id, false) + + ts.deleteProvider(ts.A) + + userData := &provider.UserProvidedData{ + Metadata: &provider.Claims{Subject: "google-sub", Email: "alice@example.com", EmailVerified: true}, + Emails: []provider.Email{{Email: "alice@example.com", Primary: true, Verified: true}}, + } + err = ts.API.db.Transaction(func(tx *storage.Connection) error { + _, found, terr := ts.API.createAccountFromExternalIdentity(tx, httptest.NewRequest(http.MethodGet, "/callback", nil), userData, "google", false) + if terr != nil { + return terr + } + require.Equal(ts.T(), user.ID, found.ID) + return ts.issueSession(tx, found) + }) + ts.requireBanned(err) + require.Zero(ts.T(), ts.sessions(user)) +} + +func (ts *SCIMUsersTestSuite) TestProviderDeleteAudit() { + ts.setActive(ts.create(ts.TokenA, scimUser("audited")), false) + tokens, err := models.FindActiveSCIMTokensBySSOProvider(ts.API.db, ts.A.ID) + require.NoError(ts.T(), err) + require.Len(ts.T(), tokens, 1) + + ts.deleteProvider(ts.A) + + disabled := ts.auditActions(models.SCIMDisabledAction) + require.Len(ts.T(), disabled, 1) + traits := disabled[0].Payload["traits"].(map[string]any) + require.Equal(ts.T(), []any{tokens[0].Prefix}, traits["token_prefixes"]) + require.Equal(ts.T(), ts.A.ID.String(), traits["sso_provider_id"]) + + revoked := ts.auditActions(models.SCIMTokenRevokedAction) + require.Len(ts.T(), revoked, 1) + traits = revoked[0].Payload["traits"].(map[string]any) + require.Equal(ts.T(), tokens[0].Prefix, traits["token_prefix"]) + require.Equal(ts.T(), ts.A.ID.String(), traits["sso_provider_id"]) + + banned := ts.auditActions(models.SCIMUsersBannedAction) + require.Len(ts.T(), banned, 1) + traits = banned[0].Payload["traits"].(map[string]any) + require.EqualValues(ts.T(), 1, traits["banned_user_count"]) + require.Equal(ts.T(), ts.A.ID.String(), traits["sso_provider_id"]) +} + +func (ts *SCIMUsersTestSuite) TestProviderDeleteWritesNoGroupEvents() { + alice := ts.create(ts.TokenA, scimUser("alice")) + ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) + before := ts.countRows(&models.AuditLogEntry{}, "payload->>'action' LIKE 'scim_group_%'") + + ts.deleteProvider(ts.A) + + require.Equal(ts.T(), before, ts.countRows(&models.AuditLogEntry{}, "payload->>'action' LIKE 'scim_group_%'")) + require.Zero(ts.T(), ts.countRows(&models.SCIMGroup{}, "sso_provider_id = ?", ts.A.ID)) +} + +func (ts *SCIMUsersTestSuite) TestProviderDeleteAuditWithExpiredTokens() { + ts.setActive(ts.create(ts.TokenA, scimUser("expired")), false) + require.NoError(ts.T(), ts.API.db.RawQuery( + "UPDATE "+(&models.SCIMToken{}).TableName()+" SET created_at = now() - interval '2 hours', expires_at = now() - interval '1 hour' WHERE sso_provider_id = ?", ts.A.ID, + ).Exec()) + + ts.deleteProvider(ts.A) + + disabled := ts.auditActions(models.SCIMDisabledAction) + require.Len(ts.T(), disabled, 1) + require.Equal(ts.T(), []any{}, disabled[0].Payload["traits"].(map[string]any)["token_prefixes"]) + require.Empty(ts.T(), ts.auditActions(models.SCIMTokenRevokedAction)) + require.Len(ts.T(), ts.auditActions(models.SCIMUsersBannedAction), 1) +} + +func (ts *SCIMUsersTestSuite) TestProviderDeleteWithoutSCIMEnabled() { + provider := createSSOProvider(ts.T(), ts.API.db) + token, _, err := models.CreateSCIMToken(ts.API.db, provider, nil) + require.NoError(ts.T(), err) + + ts.deleteProvider(provider) + + require.Empty(ts.T(), ts.auditActions(models.SCIMDisabledAction)) + revoked := ts.auditActions(models.SCIMTokenRevokedAction) + require.Len(ts.T(), revoked, 1) + require.Equal(ts.T(), token.Prefix, revoked[0].Payload["traits"].(map[string]any)["token_prefix"]) +} + +func (ts *SCIMUsersTestSuite) TestProviderDeleteAfterSCIMDisabled() { + id := ts.create(ts.TokenA, scimUser("disabled")) + user := ts.linkedUser(id) + ts.setActive(id, false) + _, err := models.DisableSCIM(ts.API.db, ts.A.ID) + require.NoError(ts.T(), err) + + ts.deleteProvider(ts.A) + + require.Empty(ts.T(), ts.auditActions(models.SCIMDisabledAction)) + require.Len(ts.T(), ts.auditActions(models.SCIMTokenRevokedAction), 1) + require.True(ts.T(), ts.reloadUser(user.ID).IsBanned()) +} + +func (ts *SCIMUsersTestSuite) TestProviderDeleteStillBansWhileSCIMFlagOff() { + id := ts.create(ts.TokenA, scimUser("flagoff")) + user := ts.linkedUser(id) + ts.setActive(id, false) + before := len(ts.scimAuditEntries()) + ts.API.config.SSO.SCIM.Enabled = false + defer func() { ts.API.config.SSO.SCIM.Enabled = true }() + + ts.deleteProvider(ts.A) + + require.True(ts.T(), ts.reloadUser(user.ID).IsBanned()) + require.Greater(ts.T(), len(ts.scimAuditEntries()), before) + require.Empty(ts.T(), ts.auditActions(models.SCIMDisabledAction)) + require.Zero(ts.T(), ts.countRows(&models.SCIMUser{}, "sso_provider_id = ?", ts.A.ID)) +} diff --git a/internal/api/scim_test.go b/internal/api/scim_test.go index a6a966823d..753467aa9f 100644 --- a/internal/api/scim_test.go +++ b/internal/api/scim_test.go @@ -1,15 +1,27 @@ package api import ( + "context" + "encoding/json" + "io/fs" "net/http" "net/http/httptest" "net/url" + "os" + "strings" "testing" + "time" + "github.com/pkg/errors" + "github.com/sirupsen/logrus" + logrustest "github.com/sirupsen/logrus/hooks/test" "github.com/stretchr/testify/require" - scimCore "github.com/supabase/auth/internal/api/scim/core" - scimProtocol "github.com/supabase/auth/internal/api/scim/protocol" + scimCore "github.com/supabase-community/scim-go/pkg/core" + scimProtocol "github.com/supabase-community/scim-go/pkg/protocol" + "github.com/supabase-community/scim-go/pkg/server" "github.com/supabase/auth/internal/conf" + "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/observability" "github.com/supabase/auth/internal/storage" ) @@ -17,6 +29,7 @@ const ( scimServiceProviderConfigPath = "/scim/v2/ServiceProviderConfig" scimResourceTypesPath = "/scim/v2/ResourceTypes" scimSchemasPath = "/scim/v2/Schemas" + scimUsersPath = "/scim/v2/Users" ) var scimPaths = []string{ @@ -30,7 +43,7 @@ func TestSCIM(t *testing.T) { api, _, err := setupAPIForTest() require.NoError(t, err) - require.False(t, api.config.Experimental.ScimEnabled) + require.False(t, api.config.SSO.SCIM.Enabled) for _, path := range scimPaths { r := httptest.NewRequest(http.MethodGet, path, nil) @@ -57,17 +70,22 @@ func TestSCIM(t *testing.T) { t.Run("Can be enabled", func(t *testing.T) { api, _, err := setupAPIForTestWithCallback(func(config *conf.GlobalConfiguration, conn *storage.Connection) { if config != nil { - config.Experimental.ScimEnabled = true + config.SSO.SCIM.Enabled = true } }) require.NoError(t, err) - require.True(t, api.config.Experimental.ScimEnabled) + require.True(t, api.config.SSO.SCIM.Enabled) + + provider := createSCIMEnabledProvider(t, api.db) + _, token, err := models.CreateSCIMToken(api.db, provider, nil) + require.NoError(t, err) t.Run(scimServiceProviderConfigPath, func(t *testing.T) { r := httptest.NewRequest(http.MethodGet, scimServiceProviderConfigPath, nil) w := httptest.NewRecorder() + r.Header.Set("Authorization", "Bearer "+token) api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusOK, w.Code) @@ -80,6 +98,7 @@ func TestSCIM(t *testing.T) { r := httptest.NewRequest(http.MethodGet, path, nil) w := httptest.NewRecorder() + r.Header.Set("Authorization", "Bearer "+token) api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusOK, w.Code) @@ -92,6 +111,7 @@ func TestSCIM(t *testing.T) { r := httptest.NewRequest(http.MethodGet, path+"?"+filter, nil) w := httptest.NewRecorder() + r.Header.Set("Authorization", "Bearer "+token) api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusForbidden, w.Code) @@ -100,10 +120,101 @@ func TestSCIM(t *testing.T) { }) } + t.Run("Every route is served by the SCIM server", func(t *testing.T) { + for _, tc := range []struct{ method, path string }{ + {http.MethodGet, scimServiceProviderConfigPath}, + {http.MethodGet, scimResourceTypesPath}, + {http.MethodGet, scimResourceTypesPath + "/User"}, + {http.MethodGet, scimSchemasPath}, + {http.MethodGet, scimSchemasPath + "/" + string(scimCore.SchemaUser)}, + {http.MethodGet, scimUsersPath}, + {http.MethodPost, scimUsersPath}, + {http.MethodGet, scimUsersPath + "/missing"}, + {http.MethodPut, scimUsersPath + "/missing"}, + {http.MethodPatch, scimUsersPath + "/missing"}, + {http.MethodDelete, scimUsersPath + "/missing"}, + } { + t.Run(tc.method+" "+tc.path, func(t *testing.T) { + r := httptest.NewRequest(tc.method, tc.path, strings.NewReader(`{}`)) + w := httptest.NewRecorder() + + r.Header.Set("Authorization", "Bearer "+token) + api.handler.ServeHTTP(w, r) + + require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type"), w.Body.String()) + }) + } + }) + + t.Run("Requires an active SCIM token", func(t *testing.T) { + revoked, revokedToken, err := models.CreateSCIMToken(api.db, provider, nil) + require.NoError(t, err) + require.NoError(t, revoked.Revoke(api.db)) + + expiresAt := time.Now().Add(time.Hour) + expired, expiredToken, err := models.CreateSCIMToken(api.db, provider, &expiresAt) + require.NoError(t, err) + require.NoError(t, api.db.RawQuery( + "UPDATE "+expired.TableName()+" SET created_at = now() - interval '2 hours', expires_at = now() - interval '1 hour' WHERE id = ?", expired.ID, + ).Exec()) + + for _, tc := range []struct{ name, authorization string }{ + {"missing", ""}, + {"basic", "Basic " + token}, + {"malformed", "Bearer notatoken"}, + {"unknown", "Bearer scim_0000000000000000000000000000000000000000"}, + {"revoked", "Bearer " + revokedToken}, + {"expired", "Bearer " + expiredToken}, + } { + for _, path := range []string{scimResourceTypesPath, scimSchemasPath, scimUsersPath} { + t.Run(tc.name+" "+path, func(t *testing.T) { + r := httptest.NewRequest(http.MethodGet, path, nil) + if tc.authorization != "" { + r.Header.Set("Authorization", tc.authorization) + } + w := httptest.NewRecorder() + + api.handler.ServeHTTP(w, r) + + require.Equal(t, http.StatusUnauthorized, w.Code) + require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + require.True(t, strings.HasPrefix(w.Header().Get("WWW-Authenticate"), "Bearer")) + }) + } + + t.Run(tc.name+" "+scimServiceProviderConfigPath, func(t *testing.T) { + r := httptest.NewRequest(http.MethodGet, scimServiceProviderConfigPath, nil) + if tc.authorization != "" { + r.Header.Set("Authorization", tc.authorization) + } + w := httptest.NewRecorder() + + api.handler.ServeHTTP(w, r) + + require.Equal(t, http.StatusOK, w.Code) + require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + }) + } + }) + + t.Run("Records when a token is used", func(t *testing.T) { + r := httptest.NewRequest(http.MethodGet, scimUsersPath, nil) + r.Header.Set("Authorization", "Bearer "+token) + w := httptest.NewRecorder() + + api.handler.ServeHTTP(w, r) + require.Equal(t, http.StatusOK, w.Code) + + found, err := models.FindSCIMTokenByPrefix(api.db, provider.ID, token[:12]) + require.NoError(t, err) + require.NotNil(t, found.LastUsedAt) + }) + t.Run("Returns a SCIM 404 for an unknown endpoint", func(t *testing.T) { r := httptest.NewRequest(http.MethodGet, "/scim/v2/Unknown", nil) w := httptest.NewRecorder() + r.Header.Set("Authorization", "Bearer "+token) api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusNotFound, w.Code) @@ -118,6 +229,7 @@ func TestSCIM(t *testing.T) { r := httptest.NewRequest(method, path, nil) w := httptest.NewRecorder() + r.Header.Set("Authorization", "Bearer "+token) api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusMethodNotAllowed, w.Code) @@ -125,6 +237,304 @@ func TestSCIM(t *testing.T) { }) } } + + for _, tc := range []struct { + method, path string + allow []string + }{ + {http.MethodPut, scimUsersPath, []string{http.MethodGet, http.MethodPost}}, + {http.MethodPost, scimUsersPath + "/missing", []string{http.MethodGet, http.MethodPut, http.MethodPatch, http.MethodDelete}}, + } { + t.Run(tc.method+" "+tc.path, func(t *testing.T) { + r := httptest.NewRequest(tc.method, tc.path, nil) + w := httptest.NewRecorder() + + r.Header.Set("Authorization", "Bearer "+token) + api.handler.ServeHTTP(w, r) + + require.Equal(t, http.StatusMethodNotAllowed, w.Code) + require.ElementsMatch(t, tc.allow, w.Header().Values("Allow")) + }) + } }) }) } + +const scimValidToken = "scim_valid" + +func scimFixture(t *testing.T, file string) string { + data, err := fs.ReadFile(os.DirFS("testdata/scim"), file) + require.NoError(t, err) + return string(data) +} + +func newSCIMServerFor(externalURL string) *server.Server { + validate := func(ctx context.Context, candidate string) (context.Context, error) { + if candidate != scimValidToken { + return ctx, server.ErrInvalidToken + } + return ctx, nil + } + return newSCIMServer(&conf.GlobalConfiguration{API: conf.APIConfiguration{ExternalURL: externalURL}}, validate, nil, nil, nil) +} + +func scimServe(t *testing.T, srv *server.Server, method, path, body string, headers ...string) *httptest.ResponseRecorder { + r := httptest.NewRequest(method, path, strings.NewReader(body)) + r.Header.Set("Content-Type", scimProtocol.MediaType) + r.Header.Set("Authorization", "Bearer "+scimValidToken) + for i := 0; i+1 < len(headers); i += 2 { + r.Header.Set(headers[i], headers[i+1]) + } + w := httptest.NewRecorder() + srv.ServeHTTP(w, r) + return w +} + +func scimDecode(t *testing.T, w *httptest.ResponseRecorder) map[string]any { + var body map[string]any + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body)) + return body +} + +func TestSCIMServer(t *testing.T) { + srv := newSCIMServerFor("http://localhost:9999") + require.NotNil(t, srv) + + t.Run("NewServer trims a trailing slash from the external URL", func(t *testing.T) { + w := scimServe(t, newSCIMServerFor("https://auth.example.com/"), http.MethodGet, scimBasePath+"/ServiceProviderConfig", "") + + meta := scimDecode(t, w)["meta"].(map[string]any) + require.Equal(t, "https://auth.example.com"+scimBasePath+"/ServiceProviderConfig", meta["location"]) + }) + + t.Run("ServiceProviderConfig", func(t *testing.T) { + w := scimServe(t, srv, http.MethodGet, scimBasePath+"/ServiceProviderConfig", "") + + require.Equal(t, http.StatusOK, w.Code) + require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + require.JSONEq(t, scimFixture(t, "service_provider_config.json"), w.Body.String()) + }) + + t.Run("ResourceTypes", func(t *testing.T) { + w := scimServe(t, srv, http.MethodGet, scimBasePath+"/ResourceTypes", "") + + require.Equal(t, http.StatusOK, w.Code) + require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + body := scimDecode(t, w) + require.EqualValues(t, 2, body["totalResults"]) + resources := map[string]map[string]any{} + for _, resource := range body["Resources"].([]any) { + resources[resource.(map[string]any)["id"].(string)] = resource.(map[string]any) + } + user := resources["User"] + require.Equal(t, "/Users", user["endpoint"]) + require.Equal(t, string(scimCore.SchemaUser), user["schema"]) + extension := user["schemaExtensions"].([]any)[0].(map[string]any) + require.Equal(t, string(scimCore.SchemaEnterpriseUser), extension["schema"]) + group := resources["Group"] + require.Equal(t, "/Groups", group["endpoint"]) + require.Equal(t, string(scimCore.SchemaGroup), group["schema"]) + require.Empty(t, group["schemaExtensions"]) + }) + + for _, id := range []string{"User", "Group"} { + t.Run("ResourceTypes/"+id, func(t *testing.T) { + w := scimServe(t, srv, http.MethodGet, scimBasePath+"/ResourceTypes/"+id, "") + + require.Equal(t, http.StatusOK, w.Code) + require.Equal(t, id, scimDecode(t, w)["id"]) + }) + } + + t.Run("Schemas", func(t *testing.T) { + w := scimServe(t, srv, http.MethodGet, scimBasePath+"/Schemas", "") + + require.Equal(t, http.StatusOK, w.Code) + require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + body := scimDecode(t, w) + require.EqualValues(t, 3, body["totalResults"]) + ids := []string{} + for _, resource := range body["Resources"].([]any) { + ids = append(ids, resource.(map[string]any)["id"].(string)) + } + require.ElementsMatch(t, []string{string(scimCore.SchemaUser), string(scimCore.SchemaEnterpriseUser), string(scimCore.SchemaGroup)}, ids) + }) + + t.Run("Schemas/{id}", func(t *testing.T) { + for _, id := range []scimCore.SchemaURI{scimCore.SchemaUser, scimCore.SchemaEnterpriseUser, scimCore.SchemaGroup} { + w := scimServe(t, srv, http.MethodGet, scimBasePath+"/Schemas/"+string(id), "") + + require.Equal(t, http.StatusOK, w.Code) + body := scimDecode(t, w) + require.Equal(t, string(id), body["id"]) + location := "http://localhost:9999" + scimBasePath + "/Schemas/" + string(id) + require.Equal(t, location, body["meta"].(map[string]any)["location"]) + require.Equal(t, location, w.Header().Get("Content-Location")) + } + }) + + t.Run("Group schema members reference only Users", func(t *testing.T) { + w := scimServe(t, srv, http.MethodGet, scimBasePath+"/Schemas/"+string(scimCore.SchemaGroup), "") + + require.Equal(t, http.StatusOK, w.Code) + sub := map[string]map[string]any{} + for _, attribute := range scimDecode(t, w)["attributes"].([]any) { + if attribute.(map[string]any)["name"] != "members" { + continue + } + for _, s := range attribute.(map[string]any)["subAttributes"].([]any) { + sub[s.(map[string]any)["name"].(string)] = s.(map[string]any) + } + } + require.Equal(t, []any{"User"}, sub["type"]["canonicalValues"]) + require.Equal(t, []any{"User"}, sub["$ref"]["referenceTypes"]) + }) + + t.Run("Schemas/{id} location uses the external URL prefix", func(t *testing.T) { + w := scimServe(t, newSCIMServerFor("https://project.supabase.co/auth/v1"), http.MethodGet, scimBasePath+"/Schemas/"+string(scimCore.SchemaUser), "") + + require.Equal(t, http.StatusOK, w.Code) + location := "https://project.supabase.co/auth/v1" + scimBasePath + "/Schemas/" + string(scimCore.SchemaUser) + require.Equal(t, location, scimDecode(t, w)["meta"].(map[string]any)["location"]) + require.Equal(t, location, w.Header().Get("Content-Location")) + }) + + t.Run("Schemas/User advertises the full RFC 7643 User attributes", func(t *testing.T) { + w := scimServe(t, srv, http.MethodGet, scimBasePath+"/Schemas/"+string(scimCore.SchemaUser), "") + + require.Equal(t, http.StatusOK, w.Code) + names := []string{} + for _, attribute := range scimDecode(t, w)["attributes"].([]any) { + names = append(names, attribute.(map[string]any)["name"].(string)) + } + for _, name := range []string{"userName", "name", "displayName", "title", "active", "emails", "phoneNumbers", "groups", "roles"} { + require.Contains(t, names, name) + } + }) + + for _, path := range []string{"/ResourceTypes", "/Schemas"} { + t.Run(path+" rejects filter query parameter", func(t *testing.T) { + query := url.Values{"filter": {`name eq "User"`}}.Encode() + w := scimServe(t, srv, http.MethodGet, scimBasePath+path+"?"+query, "") + + require.Equal(t, http.StatusForbidden, w.Code) + require.JSONEq(t, scimFixture(t, "filter_forbidden.json"), w.Body.String()) + }) + } + + t.Run("logError logs through the request log entry", func(t *testing.T) { + logger, hook := logrustest.NewNullLogger() + r := httptest.NewRequest(http.MethodGet, scimBasePath+"/Users", nil) + entry := observability.NewLogEntry(logger.WithField("request_id", "req-1")) + r = r.WithContext(observability.SetLogEntryWithContext(r.Context(), entry)) + + scimLogError(r, errors.New("broken pipe")) + + require.Len(t, hook.Entries, 1) + require.Equal(t, logrus.ErrorLevel, hook.LastEntry().Level) + require.Equal(t, "req-1", hook.LastEntry().Data["request_id"]) + require.EqualError(t, hook.LastEntry().Data[logrus.ErrorKey].(error), "broken pipe") + }) + + t.Run("requires a bearer token", func(t *testing.T) { + for _, tc := range []struct { + name, authorization string + status int + challenge string + }{ + {"missing header", "", http.StatusUnauthorized, `Bearer realm="scim"`}, + {"wrong scheme", "Basic " + scimValidToken, http.StatusUnauthorized, `Bearer realm="scim"`}, + {"empty token", "Bearer ", http.StatusBadRequest, `Bearer realm="scim", error="invalid_request", error_description="missing bearer token"`}, + {"invalid token", "Bearer scim_invalid", http.StatusUnauthorized, `Bearer realm="scim", error="invalid_token", error_description="The access token is invalid"`}, + } { + t.Run(tc.name, func(t *testing.T) { + w := scimServe(t, srv, http.MethodGet, scimBasePath+"/Users", "", "Authorization", tc.authorization) + + require.Equal(t, tc.status, w.Code) + require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + require.Equal(t, tc.challenge, w.Header().Get("WWW-Authenticate")) + }) + } + }) + + t.Run("NotFound", func(t *testing.T) { + r := httptest.NewRequest(http.MethodGet, scimBasePath+"/Unknown", nil) + w := httptest.NewRecorder() + + require.NoError(t, scimNotFound(w, r)) + + require.Equal(t, http.StatusNotFound, w.Code) + require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + require.JSONEq(t, scimFixture(t, "not_found.json"), w.Body.String()) + }) +} + +func TestSCIMReturns(t *testing.T) { + schemas := scimGroupSchemas + for query, want := range map[string]bool{ + "": true, + "excludedAttributes=members": false, + "excludedAttributes=MEMBERS": false, + "excludedAttributes=urn:ietf:params:scim:schemas:core:2.0:Group:members": false, + "excludedAttributes=displayName": true, + "excludedAttributes=members.display": true, + "attributes=displayName": false, + "attributes=members": true, + "attributes=members.value": true, + "attributes=urn:ietf:params:scim:schemas:core:2.0:Group:Members": true, + } { + values, err := url.ParseQuery(query) + require.NoError(t, err) + projection, err := scimProtocol.ParseProjection(values, schemas) + require.NoError(t, err, query) + require.Equal(t, want, scimReturns(projection, "members"), query) + } + require.True(t, scimReturns(scimProtocol.Projection{}, "members")) +} + +func TestSCIMRememberGetQuery(t *testing.T) { + for method, want := range map[string]url.Values{ + http.MethodGet: {"excludedAttributes": {"members"}}, + http.MethodPatch: nil, + http.MethodPut: nil, + } { + var got url.Values + scimRememberGetQuery(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { + got = scimGetQueryKey.Value(r.Context()) + })).ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(method, "/Groups/x?excludedAttributes=members", nil)) + require.Equal(t, want, got, method) + } +} + +func TestSCIMUserFields(t *testing.T) { + t.Run("has no email without emails", func(t *testing.T) { + require.Empty(t, scimPrimaryEmail(nil)) + }) + + t.Run("prefers the primary email", func(t *testing.T) { + require.Equal(t, "home@example.com", scimPrimaryEmail([]scimCore.Email{ + {Value: "work@example.com"}, + {Value: "home@example.com", Primary: new(true)}, + })) + }) + + t.Run("falls back to the first email", func(t *testing.T) { + require.Equal(t, "work@example.com", scimPrimaryEmail([]scimCore.Email{{Value: "work@example.com"}, {Value: "home@example.com"}})) + }) + + t.Run("drops id, meta and password from the resource", func(t *testing.T) { + user := &scimCore.User{UserName: "alice", Password: "secret"} + user.ID = "abc" + user.Meta = scimCore.Meta{Version: `W/"1"`} + + encoded, err := scimUserResource(user) + require.NoError(t, err) + + resource := map[string]any{} + require.NoError(t, json.Unmarshal(encoded, &resource)) + require.NotContains(t, resource, "id") + require.NotContains(t, resource, "meta") + require.NotContains(t, resource, "password") + require.Equal(t, "alice", resource["userName"]) + }) +} diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go new file mode 100644 index 0000000000..8cea8ff71e --- /dev/null +++ b/internal/api/scim_users.go @@ -0,0 +1,491 @@ +package api + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "time" + + "github.com/gofrs/uuid" + "github.com/sirupsen/logrus" + "github.com/supabase-community/scim-go/pkg/core" + "github.com/supabase-community/scim-go/pkg/protocol" + "github.com/supabase-community/scim-go/pkg/scimerrors" + "github.com/supabase/auth/internal/api/apierrors" + "github.com/supabase/auth/internal/api/provider" + "github.com/supabase/auth/internal/hooks/v0hooks" + "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/storage" +) + +type scimUsers struct { + api *API +} + +func (s *scimUsers) List(ctx context.Context, query *protocol.SearchRequest) ([]*core.User, int, error) { + providerID, err := scimProviderID(ctx) + if err != nil { + return nil, 0, err + } + search, err := scimSearch(query, scimUserSchemas, "userName") + if err != nil { + return nil, 0, err + } + db := s.api.db.WithContext(ctx) + rows, total, err := models.FindSCIMUsers(db, providerID, search) + if err != nil { + return nil, 0, err + } + projection, err := query.Projection(scimUserSchemas) + if err != nil { + projection = protocol.Projection{} + } + users, err := s.render(db, providerID, rows, projection) + if err != nil { + return nil, 0, err + } + return users, total, nil +} + +func (s *scimUsers) Get(ctx context.Context, id string) (*core.User, error) { + providerID, resourceID, _, err := scimTarget(ctx, id, "") + if err != nil { + return nil, err + } + db := s.api.db.WithContext(ctx) + row, err := models.FindSCIMUser(db, providerID, resourceID) + if err != nil { + return nil, scimTranslate(err) + } + projection, err := protocol.ParseProjection(scimGetQueryKey.Value(ctx), scimUserSchemas) + if err != nil { + projection = protocol.Projection{} + } + return s.renderOne(db, providerID, row, projection) +} + +func (s *scimUsers) Create(ctx context.Context, user *core.User) (*core.User, error) { + providerID, err := scimProviderID(ctx) + if err != nil { + return nil, err + } + resource, err := scimUserResource(user) + if err != nil { + return nil, err + } + if scimPrimaryEmail(user.Emails) == "" { + return nil, errSCIMEmailRequired() + } + r, err := scimRequest(ctx) + if err != nil { + return nil, err + } + db := s.api.db.WithContext(ctx) + if err := s.beforeCreate(r, db, providerID, user); err != nil { + return nil, scimTranslate(err) + } + + var row *models.SCIMUser + var created *models.User + err = db.Transaction(func(tx *storage.Connection) error { + if terr := models.LockAccountLinking(tx, "sso:"+providerID.String(), scimPrimaryEmail(user.Emails)); terr != nil { + return terr + } + var terr error + if row, terr = models.CreateSCIMUser(tx, providerID, resource); terr != nil { + return terr + } + if created, terr = s.linkNew(tx, row, user); terr != nil { + return terr + } + return s.audit(tx, r, models.SCIMUserCreatedAction, row) + }) + if err != nil { + return nil, scimTranslate(err) + } + s.afterCreate(r, db, created) + return s.renderOne(db, providerID, row, protocol.Projection{}) +} + +func (s *scimUsers) Replace(ctx context.Context, user *core.User) (*core.User, error) { + providerID, id, updatedAt, err := scimTarget(ctx, user.ID, user.Meta.Version) + if err != nil { + return nil, err + } + resource, err := scimUserResource(user) + if err != nil { + return nil, err + } + email := scimPrimaryEmail(user.Emails) + r, err := scimRequest(ctx) + if err != nil { + return nil, err + } + db := s.api.db.WithContext(ctx) + existing, err := models.FindSCIMUser(db, providerID, id) + if err != nil { + return nil, scimTranslate(err) + } + if existing.UserID == nil { + if email == "" { + return nil, errSCIMEmailRequired() + } + if err := s.beforeCreate(r, db, providerID, user); err != nil { + return nil, scimTranslate(err) + } + } + + var row *models.SCIMUser + var created *models.User + err = db.Transaction(func(tx *storage.Connection) error { + if terr := models.LockAccountLinking(tx, "sso:"+providerID.String(), email); terr != nil { + return terr + } + if existing.UserID != nil { + if terr := models.LockUserForSCIM(tx, *existing.UserID); terr != nil { + return terr + } + } + old, terr := models.FindSCIMUserForUpdate(tx, providerID, id) + if terr != nil { + return terr + } + if old.UserID != nil { + if terr := models.LockUserForSCIM(tx, *old.UserID); terr != nil { + return terr + } + if row, terr = models.FindUnchangedSCIMUser(tx, providerID, id, resource, updatedAt); terr != nil || row != nil { + return terr + } + } + if row, terr = models.ReplaceSCIMUser(tx, providerID, id, resource, updatedAt); terr != nil { + return terr + } + if created, terr = s.sync(tx, providerID, old, row, user); terr != nil { + return terr + } + return s.audit(tx, r, scimUserAuditAction(old, row), row) + }) + if err != nil { + return nil, scimTranslate(err) + } + s.afterCreate(r, db, created) + return s.renderOne(db, providerID, row, protocol.Projection{}) +} + +func (s *scimUsers) Delete(ctx context.Context, id, version string) error { + providerID, resourceID, updatedAt, err := scimTarget(ctx, id, version) + if err != nil { + return err + } + r, err := scimRequest(ctx) + if err != nil { + return err + } + db := s.api.db.WithContext(ctx) + existing, err := models.FindSCIMUser(db, providerID, resourceID) + if err != nil { + return scimTranslate(err) + } + return scimTranslate(db.Transaction(func(tx *storage.Connection) error { + if existing.UserID != nil { + if err := models.LockUserForSCIM(tx, *existing.UserID); err != nil { + return err + } + } + row, err := models.DeleteSCIMUser(tx, providerID, resourceID, updatedAt) + if err != nil { + return err + } + if row.UserID != nil { + if err := scimDeactivate(tx, *row.UserID); err != nil { + return err + } + } + if err := s.api.removeSCIMUserFromGroups(tx, r, scimActor(r), row); err != nil { + return err + } + return s.audit(tx, r, models.SCIMUserDeletedAction, row) + })) +} + +func (s *scimUsers) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMUser, projection protocol.Projection) ([]*core.User, error) { + memberships := []models.SCIMGroupMembership{} + if scimReturns(projection, "groups") { + ids := make([]uuid.UUID, len(rows)) + for i, row := range rows { + ids[i] = row.ID + } + var err error + if memberships, err = models.FindSCIMGroupsForUsers(tx, providerID, ids); err != nil { + return nil, err + } + } + base := scimBaseURL(s.api.config) + groups := map[uuid.UUID][]core.GroupMembership{} + for _, m := range memberships { + groups[m.SCIMUserID] = append(groups[m.SCIMUserID], core.GroupMembership{ + Value: m.GroupID.String(), + Ref: base + "/Groups/" + m.GroupID.String(), + Display: m.Display, + Type: "direct", + }) + } + + users := make([]*core.User, 0, len(rows)) + for _, row := range rows { + user := &core.User{} + if err := json.Unmarshal(row.Resource, user); err != nil { + return nil, err + } + user.Active = &row.Active + user.ID = row.ID.String() + user.Schemas = []core.SchemaURI{core.SchemaUser} + if user.EnterpriseUser != nil { + user.Schemas = append(user.Schemas, core.SchemaEnterpriseUser) + } + user.Meta = scimMeta(scimResourceTypeUser, base+"/Users/"+user.ID, row.CreatedAt, row.UpdatedAt) + user.Groups = groups[row.ID] + users = append(users, user) + } + return users, nil +} + +func (s *scimUsers) renderOne(tx *storage.Connection, providerID uuid.UUID, row *models.SCIMUser, projection protocol.Projection) (*core.User, error) { + users, err := s.render(tx, providerID, []models.SCIMUser{*row}, projection) + if err != nil { + return nil, err + } + return users[0], nil +} + +func (s *scimUsers) sync(tx *storage.Connection, providerID uuid.UUID, old, row *models.SCIMUser, user *core.User) (*models.User, error) { + if old.UserID == nil { + if scimPrimaryEmail(user.Emails) == "" { + return nil, errSCIMEmailRequired() + } + return s.linkNew(tx, row, user) + } + + linked, err := models.FindUserByID(tx, *old.UserID) + if err != nil { + return nil, err + } + if from := scimUserName(old.Resource); from != user.UserName { + data := map[string]any{"sub": user.UserName} + if email := scimPrimaryEmail(user.Emails); email != "" { + data["email"] = email + } + err := models.RenameSCIMIdentity(tx, linked.ID, "sso:"+providerID.String(), from, user.UserName, data) + if errors.Is(err, models.SCIMIdentityNotFoundError{}) { + logrus.WithField("user_id", linked.ID).WithField("sso_provider_id", providerID).Warn("scim: SCIM identity not found, rename skipped") + } else if err != nil { + return nil, err + } + } + if old.Active && !row.Active { + return nil, scimDeactivate(tx, linked.ID) + } + return nil, nil +} + +func (s *scimUsers) linkNew(tx *storage.Connection, row *models.SCIMUser, user *core.User) (*models.User, error) { + linked, isNew, err := s.link(tx, row, user) + if err != nil { + return nil, err + } + var created *models.User + if isNew { + created = linked + } + if !row.Active { + return created, scimDeactivate(tx, linked.ID) + } + return created, nil +} + +func (s *scimUsers) link(tx *storage.Connection, row *models.SCIMUser, user *core.User) (*models.User, bool, error) { + providerType := "sso:" + row.SSOProviderID.String() + decision, err := s.decide(tx, providerType, user) + if err != nil { + return nil, false, err + } + + linked := decision.User + switch decision.Decision { + case models.AccountExists: + case models.LinkAccount: + if _, err = s.api.createNewIdentity(tx, linked, providerType, scimIdentityData(user)); err != nil { + return nil, false, err + } + if err = linked.UpdateAppMetaDataProviders(tx); err != nil { + return nil, false, err + } + case models.CreateAccount: + if linked, err = s.newUser(providerType, decision, user); err != nil { + return nil, false, err + } + if linked, err = s.api.signupNewUser(tx, linked); err != nil { + return nil, false, err + } + if _, err = s.api.createNewIdentity(tx, linked, providerType, scimIdentityData(user)); err != nil { + return nil, false, err + } + return linked, true, models.LinkSCIMUser(tx, row, linked.ID) + case models.MultipleAccounts: + return nil, false, scimerrors.ErrUniqueness("multiple users share this email in the SSO provider") + default: + return nil, false, apierrors.NewInternalServerError("Unknown automatic linking decision: %v", decision.Decision) + } + + return linked, false, models.LinkSCIMUser(tx, row, linked.ID) +} + +func (s *scimUsers) beforeCreate(r *http.Request, db *storage.Connection, providerID uuid.UUID, user *core.User) error { + if !s.api.hooksMgr.Enabled(v0hooks.BeforeUserCreated) { + return nil + } + providerType := "sso:" + providerID.String() + decision, err := s.decide(db, providerType, user) + if err != nil || decision.Decision != models.CreateAccount { + return err + } + candidate, err := s.newUser(providerType, decision, user) + if err != nil { + return err + } + return scimHookError(s.api.triggerBeforeUserCreated(r, db, candidate)) +} + +func (s *scimUsers) afterCreate(r *http.Request, db *storage.Connection, user *models.User) { + if user == nil { + return + } + if err := s.api.triggerAfterUserCreated(r, db, user); err != nil { + logrus.WithError(err).WithField("user_id", user.ID).Error("scim: after user created hook failed") + } +} + +func (s *scimUsers) decide(conn *storage.Connection, providerType string, user *core.User) (models.AccountLinkingResult, error) { + emails := []provider.Email{{Email: scimPrimaryEmail(user.Emails), Verified: true, Primary: true}} + return models.DetermineAccountLinking(conn, s.api.config, emails, s.api.config.JWT.Aud, providerType, user.UserName) +} + +func (s *scimUsers) newUser(providerType string, decision models.AccountLinkingResult, user *core.User) (*models.User, error) { + params := &SignupParams{ + Provider: providerType, + Email: decision.CandidateEmail.Email, + Aud: s.api.config.JWT.Aud, + Data: scimIdentityData(user), + } + candidate, err := params.ToUserModel(true) + if err != nil { + return nil, err + } + now := time.Now() + candidate.EmailConfirmedAt = &now + return candidate, nil +} + +func (s *scimUsers) audit(tx *storage.Connection, r *http.Request, action models.AuditAction, row *models.SCIMUser) error { + return s.api.auditSCIM(tx, r, scimActor(r), action, row.SSOProviderID, scimUserTraits(row)) +} + +func (a *API) deleteSCIMUsers(tx *storage.Connection, r *http.Request, actor *models.User, userID uuid.UUID) error { + rows, err := models.SoftDeleteSCIMUsersByUserID(tx, userID) + if err != nil { + return err + } + for i := range rows { + if err := a.removeSCIMUserFromGroups(tx, r, actor, &rows[i]); err != nil { + return err + } + if err := a.auditSCIM(tx, r, actor, models.SCIMUserDeletedAction, rows[i].SSOProviderID, scimUserTraits(&rows[i])); err != nil { + return err + } + } + return nil +} + +func (a *API) removeSCIMUserFromGroups(tx *storage.Connection, r *http.Request, actor *models.User, row *models.SCIMUser) error { + groupIDs, err := models.RemoveSCIMUserFromGroups(tx, row.ID) + if err != nil { + return err + } + for _, groupID := range groupIDs { + if err := a.auditSCIMMember(tx, r, actor, models.SCIMGroupMemberRemovedAction, row.SSOProviderID, groupID, row.ID, row.UserID); err != nil { + return err + } + } + return nil +} + +func scimUserResource(user *core.User) ([]byte, error) { + return scimEncode(user, "id", "meta", "password", "groups") +} + +func scimPrimaryEmail(emails []core.Email) string { + for _, email := range emails { + if email.Primary != nil && *email.Primary { + return email.Value + } + } + if len(emails) > 0 { + return emails[0].Value + } + return "" +} + +func scimIdentityData(user *core.User) map[string]any { + return map[string]any{ + "sub": user.UserName, + "email": scimPrimaryEmail(user.Emails), + "email_verified": true, + } +} + +func scimUserName(resource []byte) string { + var r struct { + UserName string `json:"userName"` + } + _ = json.Unmarshal(resource, &r) + return r.UserName +} + +func scimUserTraits(row *models.SCIMUser) map[string]any { + traits := map[string]any{ + "scim_user_id": row.ID, + "user_name": row.UserName, + "active": row.Active, + } + if row.UserID != nil { + traits["user_id"] = *row.UserID + } + return traits +} + +func scimUserAuditAction(before, after *models.SCIMUser) models.AuditAction { + switch { + case before.Active && !after.Active: + return models.SCIMUserDeactivatedAction + case !before.Active && after.Active: + return models.SCIMUserReactivatedAction + } + return models.SCIMUserUpdatedAction +} + +func scimDeactivate(tx *storage.Connection, userID uuid.UUID) error { + if err := models.LockUserForSCIM(tx, userID); err != nil { + return err + } + return models.Logout(tx, userID) +} + +func scimHookError(err error) error { + var httpErr *apierrors.HTTPError + if errors.As(err, &httpErr) && httpErr.HTTPStatus < http.StatusInternalServerError { + return scimerrors.NewError(httpErr.HTTPStatus, "", httpErr.Message) + } + return err +} diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go new file mode 100644 index 0000000000..67b586279c --- /dev/null +++ b/internal/api/scim_users_test.go @@ -0,0 +1,803 @@ +package api + +import ( + "context" + "encoding/json" + "fmt" + "maps" + "net/http" + "net/http/httptest" + "net/url" + "slices" + "strconv" + "strings" + "sync" + "testing" + "time" + + "github.com/gofrs/uuid" + "github.com/stretchr/testify/require" + "github.com/stretchr/testify/suite" + "github.com/supabase-community/scim-go/pkg/core" + "github.com/supabase-community/scim-go/pkg/protocol" + "github.com/supabase-community/scim-go/pkg/scimerrors" + "github.com/supabase-community/scim-go/pkg/server" + "github.com/supabase/auth/internal/conf" + "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/storage" +) + +const oktaUser = `{ + "schemas": ["urn:ietf:params:scim:schemas:core:2.0:User"], + "userName": "Alice@Example.com", + "name": {"givenName": "Alice", "familyName": "Smith"}, + "emails": [{"primary": true, "value": "alice@example.com", "type": "work"}], + "displayName": "Alice Smith", + "locale": "en-US", + "externalId": "00u1abcd", + "groups": [], + "password": "hunter2hunter2", + "active": true +}` + +type SCIMUsersTestSuite struct { + suite.Suite + API *API + TokenA string + TokenB string + A *models.SSOProvider + B *models.SSOProvider +} + +func TestSCIMUsers(t *testing.T) { + api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { + config.RateLimitScim = 1_000_000 + }) + defer api.db.Close() + + suite.Run(t, &SCIMUsersTestSuite{API: api}) +} + +func (ts *SCIMUsersTestSuite) SetupTest() { + require.NoError(ts.T(), models.TruncateAll(ts.API.db)) + ts.A, ts.TokenA = ts.provider() + ts.B, ts.TokenB = ts.provider() +} + +func (ts *SCIMUsersTestSuite) provider() (*models.SSOProvider, string) { + return createSSOProviderWithSCIMToken(ts.T(), ts.API.db) +} + +func setupSCIMAPI(t *testing.T, tweak func(*conf.GlobalConfiguration)) (*API, *conf.GlobalConfiguration) { + api, config, err := setupAPIForTestWithCallback(func(config *conf.GlobalConfiguration, conn *storage.Connection) { + if config != nil { + config.SSO.SCIM.Enabled = true + if tweak != nil { + tweak(config) + } + } + }) + require.NoError(t, err) + return api, config +} + +func createSSOProvider(t require.TestingT, db *storage.Connection) *models.SSOProvider { + provider := &models.SSOProvider{} + require.NoError(t, db.Create(provider)) + return provider +} + +func createSCIMEnabledProvider(t require.TestingT, db *storage.Connection) *models.SSOProvider { + provider := createSSOProvider(t, db) + _, err := models.EnableSCIM(db, provider.ID) + require.NoError(t, err) + return provider +} + +func createSSOProviderWithSCIMToken(t require.TestingT, db *storage.Connection) (*models.SSOProvider, string) { + provider := createSCIMEnabledProvider(t, db) + _, token, err := models.CreateSCIMToken(db, provider, nil) + require.NoError(t, err) + return provider, token +} + +func (ts *SCIMUsersTestSuite) do(token, method, path, body string) (*httptest.ResponseRecorder, map[string]any) { + return ts.doAs(protocol.MediaType, token, method, path, body) +} + +func (ts *SCIMUsersTestSuite) doAs(contentType, token, method, path, body string, headers ...string) (*httptest.ResponseRecorder, map[string]any) { + r := httptest.NewRequest(method, "/scim/v2"+path, strings.NewReader(body)) + r.Header.Set("Authorization", "Bearer "+token) + r.Header.Set("Content-Type", contentType) + for i := 0; i+1 < len(headers); i += 2 { + r.Header.Set(headers[i], headers[i+1]) + } + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + + var decoded map[string]any + if w.Body.Len() > 0 { + require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &decoded), w.Body.String()) + } + return w, decoded +} + +func (ts *SCIMUsersTestSuite) create(token, body string) string { + w, created := ts.do(token, http.MethodPost, "/Users", body) + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) + return created["id"].(string) +} + +func (ts *SCIMUsersTestSuite) list(token, filter string) map[string]any { + w, body := ts.do(token, http.MethodGet, "/Users?"+url.Values{"filter": {filter}, "startIndex": {"1"}, "count": {"100"}}.Encode(), "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + return body +} + +func userWith(userName, externalID string) string { + return `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"` + userName + `","externalId":"` + externalID + `","emails":[{"primary":true,"value":"` + userName + `"}]}` +} + +func (ts *SCIMUsersTestSuite) repository() (context.Context, server.Repository[*core.User]) { + ctx, err := newSCIMTokenValidator(ts.API.db)(context.Background(), ts.TokenA) + require.NoError(ts.T(), err) + ctx = scimRequestKey.WithValue(ctx, httptest.NewRequest(http.MethodPost, "/scim/v2/Users", nil)) + return ctx, &scimUsers{api: ts.API} +} + +func emails(value string) []core.Email { + return []core.Email{{Value: value, Primary: new(true)}} +} + +func (ts *SCIMUsersTestSuite) TestOktaLifecycle() { + require.EqualValues(ts.T(), 0, ts.list(ts.TokenA, `userName eq "alice@example.com"`)["totalResults"]) + + w, created := ts.do(ts.TokenA, http.MethodPost, "/Users", oktaUser) + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) + id := created["id"].(string) + location := "http://localhost:9999/scim/v2/Users/" + id + require.Equal(ts.T(), location, w.Header().Get("Location")) + require.Equal(ts.T(), location, created["meta"].(map[string]any)["location"]) + require.Equal(ts.T(), "Alice@Example.com", created["userName"]) + require.Equal(ts.T(), "00u1abcd", created["externalId"]) + require.Equal(ts.T(), "Alice Smith", created["displayName"]) + require.NotContains(ts.T(), w.Body.String(), "hunter2") + + var stored models.SCIMUser + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&stored)) + require.Equal(ts.T(), ts.A.ID, stored.SSOProviderID) + require.Equal(ts.T(), "alice@example.com", stored.UserName) + require.NotContains(ts.T(), string(stored.Resource), "hunter2") + require.NotContains(ts.T(), string(stored.Resource), `"id"`) + + for _, filter := range []string{`userName eq "alice@example.com"`, `userName eq "ALICE@EXAMPLE.COM"`, `externalId eq "00u1abcd"`} { + found := ts.list(ts.TokenA, filter) + require.EqualValues(ts.T(), 1, found["totalResults"], filter) + require.Equal(ts.T(), id, found["Resources"].([]any)[0].(map[string]any)["id"], filter) + } + require.EqualValues(ts.T(), 0, ts.list(ts.TokenA, `externalId eq "00U1ABCD"`)["totalResults"]) + + w, got := ts.do(ts.TokenA, http.MethodGet, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Equal(ts.T(), "Alice", got["name"].(map[string]any)["givenName"]) + + w, replaced := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, strings.Replace(oktaUser, `"givenName": "Alice"`, `"givenName": "Alicia"`, 1)) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), "Alicia", replaced["name"].(map[string]any)["givenName"]) + require.Equal(ts.T(), created["meta"].(map[string]any)["created"], replaced["meta"].(map[string]any)["created"]) + + w, patched := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id, `{ + "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + "Operations": [{"op": "replace", "value": {"active": false}}] + }`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), false, patched["active"]) + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&stored)) + require.False(ts.T(), stored.Active) + + w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + + for _, tc := range []struct{ method, body string }{ + {http.MethodGet, ""}, + {http.MethodPut, oktaUser}, + {http.MethodPatch, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","value":{"active":true}}]}`}, + {http.MethodDelete, ""}, + } { + w, _ = ts.do(ts.TokenA, tc.method, "/Users/"+id, tc.body) + require.Equal(ts.T(), http.StatusNotFound, w.Code, tc.method) + } + for _, filter := range []string{"", `userName eq "alice@example.com"`, `externalId eq "00u1abcd"`} { + require.EqualValues(ts.T(), 0, ts.list(ts.TokenA, filter)["totalResults"], filter) + } + + reprovisioned := ts.create(ts.TokenA, oktaUser) + require.NotEqual(ts.T(), id, reprovisioned) + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&stored)) + require.NotNil(ts.T(), stored.DeletedAt) +} + +func (ts *SCIMUsersTestSuite) TestOktaContentTypesAndReactivate() { + for i, contentType := range []string{"application/scim+json; charset=utf-8", "application/json", "application/json; charset=utf-8"} { + name := string(rune('a'+i)) + "@example.com" + w, created := ts.doAs(contentType, ts.TokenA, http.MethodPost, "/Users", userWith(name, name)) + require.Equal(ts.T(), http.StatusCreated, w.Code, contentType+" "+w.Body.String()) + id := created["id"].(string) + + for _, active := range []bool{false, true} { + w, patched := ts.doAs(contentType, ts.TokenA, http.MethodPatch, "/Users/"+id, `{ + "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + "Operations": [{"op": "replace", "value": {"active": `+strconv.FormatBool(active)+`}}] + }`) + require.Equal(ts.T(), http.StatusOK, w.Code, contentType+" "+w.Body.String()) + require.Equal(ts.T(), active, patched["active"], contentType) + } + } +} + +func (ts *SCIMUsersTestSuite) TestUniquenessWithinProvider() { + ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + + for _, body := range []string{userWith("ALICE@example.com", "a-2"), userWith("bob@example.com", "a-1")} { + w, response := ts.do(ts.TokenA, http.MethodPost, "/Users", body) + require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) + require.Equal(ts.T(), "uniqueness", response["scimType"]) + } + + ts.create(ts.TokenB, userWith("alice@example.com", "a-1")) +} + +func (ts *SCIMUsersTestSuite) TestUniqueIndexIsTheBackstop() { + ctx, users := ts.repository() + + _, err := users.Create(ctx, &core.User{UserName: "alice@example.com", Emails: emails("alice@example.com")}) + require.NoError(ts.T(), err) + + _, err = users.Create(ctx, &core.User{UserName: "Alice@Example.com", Emails: emails("alice@example.com")}) + var scimErr *scimerrors.Error + require.ErrorAs(ts.T(), err, &scimErr) + require.Equal(ts.T(), http.StatusConflict, scimErr.StatusCode()) +} + +func (ts *SCIMUsersTestSuite) TestReplaceRejectsStaleVersion() { + ctx, users := ts.repository() + + created, err := users.Create(ctx, &core.User{UserName: "alice@example.com", Emails: emails("alice@example.com")}) + require.NoError(ts.T(), err) + read, err := users.Get(ctx, created.ID) + require.NoError(ts.T(), err) + require.Equal(ts.T(), created.Meta.Version, read.Meta.Version) + + winner := &core.User{UserName: "alice@example.com", Title: "winner"} + winner.ID = read.ID + winner.Meta = core.Meta{Version: read.Meta.Version} + replaced, err := users.Replace(ctx, winner) + require.NoError(ts.T(), err) + require.NotEqual(ts.T(), read.Meta.Version, replaced.Meta.Version) + + for _, version := range []string{read.Meta.Version, `W/"garbage"`} { + loser := &core.User{UserName: "alice@example.com", Title: "loser"} + loser.ID = read.ID + loser.Meta = core.Meta{Version: version} + _, err = users.Replace(ctx, loser) + var scimErr *scimerrors.Error + require.ErrorAs(ts.T(), err, &scimErr, version) + require.Equal(ts.T(), http.StatusPreconditionFailed, scimErr.StatusCode(), version) + } + + missing := &core.User{UserName: "bob@example.com"} + missing.ID = uuid.Must(uuid.NewV4()).String() + missing.Meta = core.Meta{Version: read.Meta.Version} + _, err = users.Replace(ctx, missing) + var scimErr *scimerrors.Error + require.ErrorAs(ts.T(), err, &scimErr) + require.Equal(ts.T(), http.StatusNotFound, scimErr.StatusCode()) + + var stored models.SCIMUser + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", read.ID).First(&stored)) + require.Contains(ts.T(), string(stored.Resource), "winner") +} + +func (ts *SCIMUsersTestSuite) TestDeleteRejectsStaleVersion() { + ctx, users := ts.repository() + + created, err := users.Create(ctx, &core.User{UserName: "alice@example.com", Emails: emails("alice@example.com")}) + require.NoError(ts.T(), err) + + updated := &core.User{UserName: "alice@example.com", Title: "renamed"} + updated.ID = created.ID + updated.Meta = core.Meta{Version: created.Meta.Version} + replaced, err := users.Replace(ctx, updated) + require.NoError(ts.T(), err) + + for _, version := range []string{created.Meta.Version, `W/"garbage"`} { + err := users.Delete(ctx, created.ID, version) + var scimErr *scimerrors.Error + require.ErrorAs(ts.T(), err, &scimErr, version) + require.Equal(ts.T(), http.StatusPreconditionFailed, scimErr.StatusCode(), version) + } + + var stored models.SCIMUser + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", created.ID).First(&stored)) + require.Nil(ts.T(), stored.DeletedAt) + + missing := uuid.Must(uuid.NewV4()).String() + var scimErr *scimerrors.Error + err = users.Delete(ctx, missing, "") + require.ErrorAs(ts.T(), err, &scimErr) + require.Equal(ts.T(), http.StatusNotFound, scimErr.StatusCode()) + + require.NoError(ts.T(), users.Delete(ctx, created.ID, replaced.Meta.Version)) + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", created.ID).First(&stored)) + require.NotNil(ts.T(), stored.DeletedAt) +} + +func (ts *SCIMUsersTestSuite) whileLocked(lock, finish func(tx *storage.Connection) error, method, path, body string) (int, error) { + locked, release := make(chan struct{}), make(chan struct{}) + held := make(chan error, 1) + go func() { + held <- ts.API.db.Transaction(func(tx *storage.Connection) error { + err := lock(tx) + close(locked) + if err != nil { + return err + } + <-release + return finish(tx) + }) + }() + <-locked + + code := make(chan int, 1) + go func() { + w, _ := ts.do(ts.TokenA, method, path, body) + code <- w.Code + }() + select { + case c := <-code: + ts.T().Fatalf("%s %s finished while the lock was held: %d", method, path, c) + case <-time.After(200 * time.Millisecond): + } + close(release) + return <-code, <-held +} + +func (ts *SCIMUsersTestSuite) scimAuditEntries() []models.AuditLogEntry { + entries := []models.AuditLogEntry{} + require.NoError(ts.T(), ts.API.db.Q().Where("payload->>'log_type' = ?", "scim").Order("created_at asc").All(&entries)) + return entries +} + +func (ts *SCIMUsersTestSuite) TestAuditLog() { + id := ts.create(ts.TokenA, oktaUser) + + w, _ := ts.do(ts.TokenA, http.MethodPost, "/Users", oktaUser) + require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) + + for _, operation := range []string{ + `{"op":"replace","path":"displayName","value":"Alice S."}`, + `{"op":"replace","path":"active","value":false}`, + `{"op":"replace","path":"active","value":true}`, + } { + w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[`+operation+`]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + } + + w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + + tokens, err := models.FindSCIMTokensBySSOProvider(ts.API.db, ts.A.ID) + require.NoError(ts.T(), err) + var row models.SCIMUser + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&row)) + + actions := []string{} + for _, entry := range ts.scimAuditEntries() { + actions = append(actions, entry.Payload["action"].(string)) + require.Equal(ts.T(), uuid.Nil.String(), entry.Payload["actor_id"]) + require.Equal(ts.T(), "scim:"+tokens[0].Prefix, entry.Payload["actor_username"]) + traits := entry.Payload["traits"].(map[string]any) + require.Equal(ts.T(), ts.A.ID.String(), traits["sso_provider_id"]) + require.Equal(ts.T(), id, traits["scim_user_id"]) + require.Equal(ts.T(), row.UserID.String(), traits["user_id"]) + require.Equal(ts.T(), "success", traits["outcome"]) + } + require.Equal(ts.T(), []string{ + string(models.SCIMUserCreatedAction), + string(models.SCIMUserUpdatedAction), + string(models.SCIMUserDeactivatedAction), + string(models.SCIMUserReactivatedAction), + string(models.SCIMUserDeletedAction), + }, actions) +} + +func (ts *SCIMUsersTestSuite) TestRolesRoundTrip() { + body := `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"alice@example.com","emails":[{"primary":true,"value":"alice@example.com"}],"roles":[{"value":"admin","primary":true},{"value":"billing"}]}` + id := ts.create(ts.TokenA, body) + + w, read := ts.do(ts.TokenA, http.MethodGet, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + roles := []string{} + for _, role := range read["roles"].([]any) { + roles = append(roles, role.(map[string]any)["value"].(string)) + } + require.Equal(ts.T(), []string{"admin", "billing"}, roles) +} + +func (ts *SCIMUsersTestSuite) TestConcurrentCreateWithinProvider() { + const attempts = 8 + codes := make(chan int, attempts) + start := make(chan struct{}) + var wg sync.WaitGroup + for range attempts { + wg.Go(func() { + <-start + r := httptest.NewRequest(http.MethodPost, "/scim/v2/Users", strings.NewReader(userWith("race@example.com", ""))) + r.Header.Set("Authorization", "Bearer "+ts.TokenA) + r.Header.Set("Content-Type", protocol.MediaType) + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + codes <- w.Code + }) + } + close(start) + wg.Wait() + close(codes) + + counts := map[int]int{} + for code := range codes { + counts[code]++ + } + require.Equal(ts.T(), map[int]int{http.StatusCreated: 1, http.StatusConflict: attempts - 1}, counts) + require.EqualValues(ts.T(), 1, ts.list(ts.TokenA, `userName eq "race@example.com"`)["totalResults"]) +} + +func (ts *SCIMUsersTestSuite) TestSort() { + ids := map[string]string{} + for _, name := range []string{"carol@example.com", "Alice@example.com", "bob@example.com"} { + ids[name] = ts.create(ts.TokenA, userWith(name, name)) + } + ts.create(ts.TokenB, userWith("aaron@example.com", "b")) + + sorted := func(params url.Values) []string { + w, body := ts.do(ts.TokenA, http.MethodGet, "/Users?"+params.Encode(), "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + names := []string{} + for _, resource := range body["Resources"].([]any) { + names = append(names, resource.(map[string]any)["userName"].(string)) + } + return names + } + + require.Equal(ts.T(), []string{"Alice@example.com", "bob@example.com", "carol@example.com"}, sorted(url.Values{"sortBy": {"userName"}})) + require.Equal(ts.T(), []string{"carol@example.com", "bob@example.com", "Alice@example.com"}, sorted(url.Values{"sortBy": {"userName"}, "sortOrder": {"descending"}})) + require.Equal(ts.T(), []string{"carol@example.com", "Alice@example.com", "bob@example.com"}, sorted(url.Values{"sortBy": {"meta.created"}})) + require.Equal(ts.T(), []string{"carol@example.com", "Alice@example.com", "bob@example.com"}, sorted(url.Values{"sortOrder": {"descending"}})) + require.Equal(ts.T(), []string{"Alice@example.com"}, sorted(url.Values{"sortBy": {"userName"}, "count": {"1"}})) + require.Equal(ts.T(), []string{"bob@example.com"}, sorted(url.Values{"sortBy": {"userName"}, "startIndex": {"2"}, "count": {"1"}})) + + _, body := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+ids["carol@example.com"], `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","path":"title","value":"Lead"}]}`) + require.Equal(ts.T(), "Lead", body["title"]) + require.Equal(ts.T(), "carol@example.com", sorted(url.Values{"sortBy": {"meta.lastModified"}, "sortOrder": {"descending"}})[0]) + + byID := sorted(url.Values{"sortBy": {"id"}}) + require.Len(ts.T(), byID, 3) + expected := []string{ids["carol@example.com"], ids["Alice@example.com"], ids["bob@example.com"]} + slices.Sort(expected) + names := map[string]string{} + for name, id := range ids { + names[id] = name + } + require.Equal(ts.T(), []string{names[expected[0]], names[expected[1]], names[expected[2]]}, byID) + + for _, sortBy := range []string{"displayName", "emails.value", "password"} { + w, body := ts.do(ts.TokenA, http.MethodGet, "/Users?"+url.Values{"sortBy": {sortBy}}.Encode(), "") + require.Equal(ts.T(), http.StatusBadRequest, w.Code, sortBy) + require.Equal(ts.T(), string(scimerrors.InvalidValue), body["scimType"], sortBy) + } +} + +func (ts *SCIMUsersTestSuite) TestAttributeProjection() { + id := ts.create(ts.TokenA, oktaUser) + + for _, path := range []string{"/Users/" + id, "/Users"} { + w, body := ts.do(ts.TokenA, http.MethodGet, path+"?attributes=userName", "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + user := body + if path == "/Users" { + user = body["Resources"].([]any)[0].(map[string]any) + } + require.Equal(ts.T(), "Alice@Example.com", user["userName"]) + require.Equal(ts.T(), id, user["id"]) + require.NotNil(ts.T(), user["schemas"]) + require.NotContains(ts.T(), user, "meta") + require.NotContains(ts.T(), user, "displayName") + require.NotContains(ts.T(), user, "emails") + + w, body = ts.do(ts.TokenA, http.MethodGet, path+"?excludedAttributes=displayName,emails", "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + user = body + if path == "/Users" { + user = body["Resources"].([]any)[0].(map[string]any) + } + require.Equal(ts.T(), "Alice@Example.com", user["userName"]) + require.NotContains(ts.T(), user, "displayName") + require.NotContains(ts.T(), user, "emails") + require.Contains(ts.T(), user, "name") + + w, body = ts.do(ts.TokenA, http.MethodGet, path+"?attributes=userName&excludedAttributes=emails", "") + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + require.Equal(ts.T(), string(scimerrors.InvalidValue), body["scimType"]) + } +} + +func (ts *SCIMUsersTestSuite) TestWriteResponseProjection() { + w, created := ts.do(ts.TokenA, http.MethodPost, "/Users?attributes=userName", userWith("alice@example.com", "a-1")) + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) + id := created["id"].(string) + require.ElementsMatch(ts.T(), []string{"id", "schemas", "userName"}, slices.Collect(maps.Keys(created))) + + w, got := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id+"?excludedAttributes=emails", userWith("alice@example.com", "a-2")) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.NotContains(ts.T(), got, "emails") + require.Equal(ts.T(), "a-2", got["externalId"]) + + w, got = ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id+"?excludedAttributes=emails", `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","path":"externalId","value":"a-3"}]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.NotContains(ts.T(), got, "emails") + require.Equal(ts.T(), "a-3", got["externalId"]) + + group := ts.createGroup(ts.TokenA, groupWith("Engineering", "")) + w, got = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+group+"?attributes=displayName", `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"add","path":"members","value":[{"value":"`+id+`"}]}]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.ElementsMatch(ts.T(), []string{"id", "schemas", "displayName"}, slices.Collect(maps.Keys(got))) + + w, got = ts.do(ts.TokenA, http.MethodGet, "/Groups/"+group, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), []string{id}, memberValues(got)) +} + +func (ts *SCIMUsersTestSuite) TestETagAndIfMatch() { + w, created := ts.do(ts.TokenA, http.MethodPost, "/Users", oktaUser) + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) + id := created["id"].(string) + stale := w.Header().Get("ETag") + require.NotEmpty(ts.T(), stale) + require.Equal(ts.T(), created["meta"].(map[string]any)["version"], stale) + + w, _ = ts.do(ts.TokenA, http.MethodGet, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), stale, w.Header().Get("ETag")) + + patch := `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","path":"displayName","value":"Alice S."}]}` + w, patched := ts.doAs(protocol.MediaType, ts.TokenA, http.MethodPatch, "/Users/"+id, patch, "If-Match", stale) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), "Alice S.", patched["displayName"]) + current := w.Header().Get("ETag") + require.NotEqual(ts.T(), stale, current) + + for _, tc := range []struct{ method, body string }{ + {http.MethodPut, oktaUser}, + {http.MethodPatch, patch}, + {http.MethodDelete, ""}, + } { + w, _ := ts.doAs(protocol.MediaType, ts.TokenA, tc.method, "/Users/"+id, tc.body, "If-Match", stale) + require.Equal(ts.T(), http.StatusPreconditionFailed, w.Code, tc.method+" "+w.Body.String()) + } + + w, replaced := ts.doAs(protocol.MediaType, ts.TokenA, http.MethodPut, "/Users/"+id, oktaUser, "If-Match", current) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), "Alice Smith", replaced["displayName"]) + + w, _ = ts.doAs(protocol.MediaType, ts.TokenA, http.MethodDelete, "/Users/"+id, "", "If-Match", w.Header().Get("ETag")) + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) +} + +func (ts *SCIMUsersTestSuite) TestPatchAttributesOutsideTheMinimalSchema() { + id := ts.create(ts.TokenA, oktaUser) + + patch := `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ + {"op":"replace","path":"title","value":"Engineer"}, + {"op":"add","path":"phoneNumbers","value":[{"value":"555-0100","type":"work"}]}, + {"op":"replace","path":"urn:ietf:params:scim:schemas:extension:enterprise:2.0:User:department","value":"Auth"} + ]}` + w, patched := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id, patch) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), "Engineer", patched["title"]) + + w, read := ts.do(ts.TokenA, http.MethodGet, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), "Engineer", read["title"]) + require.Equal(ts.T(), "555-0100", read["phoneNumbers"].([]any)[0].(map[string]any)["value"]) + require.Equal(ts.T(), "Auth", read[string(core.SchemaEnterpriseUser)].(map[string]any)["department"]) + require.ElementsMatch(ts.T(), []any{string(core.SchemaUser), string(core.SchemaEnterpriseUser)}, read["schemas"]) +} + +func (ts *SCIMUsersTestSuite) TestActiveDefaultsToTrue() { + w, created := ts.do(ts.TokenA, http.MethodPost, "/Users", userWith("alice@example.com", "a-1")) + require.Equal(ts.T(), http.StatusCreated, w.Code) + require.Equal(ts.T(), true, created["active"]) + + w, replaced := ts.do(ts.TokenA, http.MethodPut, "/Users/"+created["id"].(string), userWith("alice@example.com", "a-1")) + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Equal(ts.T(), true, replaced["active"]) +} + +func (ts *SCIMUsersTestSuite) TestPagination() { + for _, name := range []string{"a", "b", "c"} { + ts.create(ts.TokenA, userWith(name+"@example.com", name)) + } + + w, page := ts.do(ts.TokenA, http.MethodGet, "/Users?startIndex=2&count=1", "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.EqualValues(ts.T(), 3, page["totalResults"]) + require.EqualValues(ts.T(), 2, page["startIndex"]) + require.Len(ts.T(), page["Resources"], 1) + require.Equal(ts.T(), "b@example.com", page["Resources"].([]any)[0].(map[string]any)["userName"]) + + w, page = ts.do(ts.TokenA, http.MethodGet, "/Users?count=0", "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.EqualValues(ts.T(), 3, page["totalResults"]) + require.Empty(ts.T(), page["Resources"]) + + w, page = ts.do(ts.TokenA, http.MethodGet, "/Users?startIndex=10&count=5", "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.EqualValues(ts.T(), 3, page["totalResults"]) + require.Empty(ts.T(), page["Resources"]) +} + +func (ts *SCIMUsersTestSuite) TestPageSizeCap() { + for i := range 101 { + _, err := models.CreateSCIMUser(ts.API.db, ts.A.ID, []byte(`{"userName":"user`+strconv.Itoa(i)+`@example.com"}`)) + require.NoError(ts.T(), err) + } + + for _, query := range []string{"", "?count=200"} { + w, page := ts.do(ts.TokenA, http.MethodGet, "/Users"+query, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), []any{string(protocol.SchemaListResponse)}, page["schemas"], query) + require.EqualValues(ts.T(), 101, page["totalResults"], query) + require.EqualValues(ts.T(), 1, page["startIndex"], query) + require.EqualValues(ts.T(), 100, page["itemsPerPage"], query) + require.Len(ts.T(), page["Resources"], 100, query) + } +} + +func (ts *SCIMUsersTestSuite) TestSortTieBreaksOnID() { + ids := []string{ + ts.create(ts.TokenA, userWith("alice@example.com", "a-1")), + ts.create(ts.TokenA, userWith("bob@example.com", "b-1")), + ts.create(ts.TokenA, userWith("carol@example.com", "c-1")), + } + require.NoError(ts.T(), ts.API.db.RawQuery( + "UPDATE "+(&models.SCIMUser{}).TableName()+" SET created_at = '2026-01-01T00:00:00Z', updated_at = '2026-01-01T00:00:00Z' WHERE sso_provider_id = ?", ts.A.ID, + ).Exec()) + slices.Sort(ids) + descending := slices.Clone(ids) + slices.Reverse(descending) + + for _, sortBy := range []string{"meta.created", "meta.lastModified"} { + for order, want := range map[string][]string{"ascending": ids, "descending": descending} { + got := []string{} + for startIndex := 1; startIndex <= len(ids); startIndex++ { + params := url.Values{"sortBy": {sortBy}, "sortOrder": {order}, "startIndex": {strconv.Itoa(startIndex)}, "count": {"1"}} + w, body := ts.do(ts.TokenA, http.MethodGet, "/Users?"+params.Encode(), "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + got = append(got, body["Resources"].([]any)[0].(map[string]any)["id"].(string)) + } + require.Equal(ts.T(), want, got, sortBy+" "+order) + } + } +} + +func (ts *SCIMUsersTestSuite) TestUnsupportedFilters() { + for _, filter := range []string{ + `userName co "alice"`, + `userName ne "alice"`, + `name.givenName eq "Alice"`, + `emails[value eq "alice@example.com"]`, + `userName eq "a" or userName eq "b"`, + `userName eq "a" and externalId eq "b"`, + `userName pr`, + `not (userName eq "a")`, + } { + w, body := ts.do(ts.TokenA, http.MethodGet, "/Users?"+url.Values{"filter": {filter}}.Encode(), "") + require.Equal(ts.T(), http.StatusBadRequest, w.Code, filter) + require.Equal(ts.T(), "invalidFilter", body["scimType"], filter) + } +} + +func (ts *SCIMUsersTestSuite) TestUnknownID() { + for _, id := range []string{"not-a-uuid", "00000000-0000-0000-0000-000000000000"} { + w, _ := ts.do(ts.TokenA, http.MethodGet, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNotFound, w.Code, id) + } +} + +func (ts *SCIMUsersTestSuite) TestRequiresSSOProviderOnContext() { + users := &scimUsers{api: ts.API} + + _, _, err := users.List(context.Background(), &protocol.SearchRequest{Count: 10}) + require.Error(ts.T(), err) + _, err = users.Get(context.Background(), "00000000-0000-0000-0000-000000000000") + require.Error(ts.T(), err) +} + +func TestSCIMRateLimit(t *testing.T) { + api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { + config.RateLimitScim = 1 + }) + defer api.db.Close() + require.NoError(t, models.TruncateAll(api.db)) + + token := func() string { + _, token := createSSOProviderWithSCIMToken(t, api.db) + return token + } + tokenA, tokenB := token(), token() + + get := func(token, ip string) *httptest.ResponseRecorder { + r := httptest.NewRequest(http.MethodGet, "/scim/v2/Users", nil) + if token != "" { + r.Header.Set("Authorization", "Bearer "+token) + } + r.Header.Set(api.config.RateLimitHeader, ip) + w := httptest.NewRecorder() + api.handler.ServeHTTP(w, r) + return w + } + + limited := `{"schemas":["urn:ietf:params:scim:api:messages:2.0:Error"],"detail":"Request rate limit reached","status":"429"}` + const ip = "192.0.2.1" + + for range 30 { + w := get(tokenA, ip) + require.Equal(t, http.StatusOK, w.Code, w.Body.String()) + } + w := get(tokenA, ip) + require.Equal(t, http.StatusTooManyRequests, w.Code) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) + require.JSONEq(t, limited, w.Body.String()) + w = get(tokenB, ip) + require.Equal(t, http.StatusOK, w.Code, w.Body.String()) + + for range 30 { + w := get("scim_invalid", ip) + require.Equal(t, http.StatusUnauthorized, w.Code, w.Body.String()) + } + for _, token := range []string{"scim_invalid", ""} { + w := get(token, ip) + require.Equal(t, http.StatusTooManyRequests, w.Code, w.Body.String()) + require.JSONEq(t, limited, w.Body.String()) + } + w = get(tokenB, ip) + require.Equal(t, http.StatusOK, w.Code, w.Body.String()) + w = get("scim_invalid", "198.51.100.1") + require.Equal(t, http.StatusUnauthorized, w.Code, w.Body.String()) + + for i, tc := range []struct { + method, path string + status int + }{ + {http.MethodGet, "/scim/v2/ServiceProviderConfig", http.StatusOK}, + {http.MethodGet, "/scim/v2/Unknown", http.StatusNotFound}, + {http.MethodDelete, "/scim/v2/Users", http.StatusMethodNotAllowed}, + } { + ip := fmt.Sprintf("203.0.113.%d", i+1) + send := func() *httptest.ResponseRecorder { + r := httptest.NewRequest(tc.method, tc.path, nil) + r.Header.Set("Authorization", "Bearer scim_invalid") + r.Header.Set(api.config.RateLimitHeader, ip) + w := httptest.NewRecorder() + api.handler.ServeHTTP(w, r) + return w + } + for range 30 { + w := send() + require.Equal(t, tc.status, w.Code, tc.method+" "+tc.path) + } + w := send() + require.Equal(t, http.StatusTooManyRequests, w.Code, tc.method+" "+tc.path) + require.JSONEq(t, limited, w.Body.String()) + } +} diff --git a/internal/api/ssoadmin.go b/internal/api/ssoadmin.go index 5f0c73741a..1e76019e6d 100644 --- a/internal/api/ssoadmin.go +++ b/internal/api/ssoadmin.go @@ -453,6 +453,9 @@ func (a *API) adminSSOProvidersDelete(w http.ResponseWriter, r *http.Request) er provider := getSSOProvider(ctx) if err := db.Transaction(func(tx *storage.Connection) error { + if err := a.deprovisionSCIM(tx, r, provider); err != nil { + return err + } return tx.Eager().Destroy(provider) }); err != nil { return err diff --git a/internal/api/scim/testdata/filter_forbidden.json b/internal/api/testdata/scim/filter_forbidden.json similarity index 60% rename from internal/api/scim/testdata/filter_forbidden.json rename to internal/api/testdata/scim/filter_forbidden.json index 5f060363c6..c834b9fb8e 100644 --- a/internal/api/scim/testdata/filter_forbidden.json +++ b/internal/api/testdata/scim/filter_forbidden.json @@ -2,6 +2,6 @@ "schemas": [ "urn:ietf:params:scim:api:messages:2.0:Error" ], - "detail": "Filtering is not supported on this endpoint", + "detail": "\"filter\" is not supported on this endpoint", "status": "403" } diff --git a/internal/api/scim/testdata/not_found.json b/internal/api/testdata/scim/not_found.json similarity index 100% rename from internal/api/scim/testdata/not_found.json rename to internal/api/testdata/scim/not_found.json diff --git a/internal/api/testdata/scim/okta_group_patch.json b/internal/api/testdata/scim/okta_group_patch.json new file mode 100644 index 0000000000..c85716c9ec --- /dev/null +++ b/internal/api/testdata/scim/okta_group_patch.json @@ -0,0 +1,157 @@ +[ + { + "step": "push group", + "requests": [ + { + "method": "POST", + "path": "/scim/v2/Groups", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "displayName": "Tour Guides", + "members": [ + { + "value": "c75ad752-64ae-4823-840d-ffa80929976c", + "display": "jsmith@example.com" + }, + { + "value": "2819c223-7f76-453a-919d-413861904646", + "display": "bjensen@example.com" + } + ] + } + }, + { + "method": "PATCH", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:api:messages:2.0:PatchOp" + ], + "Operations": [ + { + "op": "add", + "path": "members", + "value": [ + { + "value": "c75ad752-64ae-4823-840d-ffa80929976c", + "display": "jsmith@example.com" + }, + { + "value": "2819c223-7f76-453a-919d-413861904646", + "display": "bjensen@example.com" + } + ] + } + ] + } + } + ] + }, + { + "step": "remove jsmith", + "requests": [ + { + "method": "PATCH", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:api:messages:2.0:PatchOp" + ], + "Operations": [ + { + "op": "replace", + "value": { + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Tour Guides" + } + } + ] + } + }, + { + "method": "PATCH", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:api:messages:2.0:PatchOp" + ], + "Operations": [ + { + "op": "remove", + "path": "members[value eq \"c75ad752-64ae-4823-840d-ffa80929976c\"]" + } + ] + } + } + ] + }, + { + "step": "add babs", + "requests": [ + { + "method": "PATCH", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:api:messages:2.0:PatchOp" + ], + "Operations": [ + { + "op": "replace", + "value": { + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Tour Guides" + } + } + ] + } + }, + { + "method": "PATCH", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:api:messages:2.0:PatchOp" + ], + "Operations": [ + { + "op": "add", + "path": "members", + "value": [ + { + "value": "6c5bb468-14b2-4183-baf2-06d523e03bd3", + "display": "babs@jensen.org" + } + ] + } + ] + } + } + ] + }, + { + "step": "rename", + "requests": [ + { + "method": "PATCH", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:api:messages:2.0:PatchOp" + ], + "Operations": [ + { + "op": "replace", + "value": { + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Group B" + } + } + ] + } + } + ] + } +] diff --git a/internal/api/testdata/scim/okta_group_push.json b/internal/api/testdata/scim/okta_group_push.json new file mode 100644 index 0000000000..abca70307c --- /dev/null +++ b/internal/api/testdata/scim/okta_group_push.json @@ -0,0 +1,375 @@ +[ + { + "step": "push group", + "requests": [ + { + "method": "POST", + "path": "/scim/v2/Groups", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "displayName": "Tour Guides", + "members": [] + } + }, + { + "method": "PUT", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Tour Guides", + "members": [] + } + } + ] + }, + { + "step": "add bjensen (sent twice)", + "requests": [ + { + "method": "PUT", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Tour Guides", + "members": [ + { + "value": "2819c223-7f76-453a-919d-413861904646", + "display": "bjensen@example.com" + } + ] + } + }, + { + "method": "PUT", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Tour Guides", + "members": [ + { + "value": "2819c223-7f76-453a-919d-413861904646", + "display": "bjensen@example.com" + } + ] + } + } + ] + }, + { + "step": "retry: bjensen and jsmith", + "requests": [ + { + "method": "PUT", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Tour Guides", + "members": [ + { + "value": "2819c223-7f76-453a-919d-413861904646", + "display": null + }, + { + "value": "c75ad752-64ae-4823-840d-ffa80929976c", + "display": "jsmith@example.com" + } + ] + } + }, + { + "method": "PUT", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Tour Guides", + "members": [ + { + "value": "2819c223-7f76-453a-919d-413861904646", + "display": "bjensen@example.com" + }, + { + "value": "c75ad752-64ae-4823-840d-ffa80929976c", + "display": "jsmith@example.com" + } + ] + } + } + ] + }, + { + "step": "remove bjensen", + "requests": [ + { + "method": "PUT", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Tour Guides", + "members": [ + { + "value": "2819c223-7f76-453a-919d-413861904646", + "display": null + }, + { + "value": "c75ad752-64ae-4823-840d-ffa80929976c", + "display": null + } + ] + } + }, + { + "method": "PUT", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Tour Guides", + "members": [ + { + "value": "c75ad752-64ae-4823-840d-ffa80929976c", + "display": null + } + ] + } + } + ] + }, + { + "step": "remove jsmith", + "requests": [ + { + "method": "PUT", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Tour Guides", + "members": [ + { + "value": "c75ad752-64ae-4823-840d-ffa80929976c", + "display": null + } + ] + } + }, + { + "method": "PUT", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Tour Guides", + "members": [] + } + } + ] + }, + { + "step": "rename", + "requests": [ + { + "method": "PUT", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Group A", + "members": [] + } + } + ] + }, + { + "step": "re-add bjensen (sent twice)", + "requests": [ + { + "method": "PUT", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Group A", + "members": [ + { + "value": "2819c223-7f76-453a-919d-413861904646", + "display": "bjensen@example.com" + } + ] + } + }, + { + "method": "PUT", + "path": "/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:Group" + ], + "id": "e9e30dba-f08f-4109-8486-d5c6a331660a", + "displayName": "Group A", + "members": [ + { + "value": "2819c223-7f76-453a-919d-413861904646", + "display": "bjensen@example.com" + } + ] + } + } + ] + }, + { + "step": "deactivate bjensen", + "requests": [ + { + "method": "PUT", + "path": "/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "body": { + "active": false, + "displayName": "Babs Jensen", + "emails": [ + { + "primary": true, + "type": "work", + "value": "bjensen@example.com" + } + ], + "externalId": "bjensen", + "groups": [ + { + "$ref": "http://localhost:9999/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "display": "Group A", + "type": "direct", + "value": "e9e30dba-f08f-4109-8486-d5c6a331660a" + } + ], + "id": "2819c223-7f76-453a-919d-413861904646", + "locale": "en-US", + "meta": { + "created": "2026-09-28T17:52:36.295891Z", + "lastModified": "2026-09-28T17:52:36.295891Z", + "location": "http://localhost:9999/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "resourceType": "User", + "version": "W/\"1790617956295891\"" + }, + "name": { + "familyName": "Jensen", + "givenName": "Barbara" + }, + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:User" + ], + "userName": "bjensen@example.com" + } + } + ] + }, + { + "step": "reactivate bjensen and reassign app", + "requests": [ + { + "method": "PUT", + "path": "/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "body": { + "active": true, + "displayName": "Babs Jensen", + "emails": [ + { + "primary": true, + "type": "work", + "value": "bjensen@example.com" + } + ], + "externalId": "bjensen", + "groups": [], + "id": "2819c223-7f76-453a-919d-413861904646", + "locale": "en-US", + "meta": { + "created": "2026-09-28T17:52:36.295891Z", + "lastModified": "2026-09-28T18:15:25.689427Z", + "location": "http://localhost:9999/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "resourceType": "User", + "version": "W/\"1790619325689427\"" + }, + "name": { + "familyName": "Jensen", + "givenName": "Barbara" + }, + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:User" + ], + "userName": "bjensen@example.com" + } + }, + { + "method": "PUT", + "path": "/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "body": { + "active": true, + "displayName": "Babs Jensen", + "emails": [ + { + "primary": true, + "type": "work", + "value": "bjensen@example.com" + } + ], + "externalId": "bjensen", + "groups": [ + { + "$ref": "http://localhost:9999/scim/v2/Groups/e9e30dba-f08f-4109-8486-d5c6a331660a", + "display": "Group A", + "type": "direct", + "value": "e9e30dba-f08f-4109-8486-d5c6a331660a" + } + ], + "id": "2819c223-7f76-453a-919d-413861904646", + "locale": "en-US", + "meta": { + "created": "2026-09-28T17:52:36.295891Z", + "lastModified": "2026-09-28T18:13:23.777587Z", + "location": "http://localhost:9999/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "resourceType": "User", + "version": "W/\"1790619203777587\"" + }, + "name": { + "familyName": "Jensen", + "givenName": "Barbara" + }, + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:User" + ], + "userName": "bjensen@example.com" + } + } + ] + } +] diff --git a/internal/api/testdata/scim/okta_user_lifecycle.json b/internal/api/testdata/scim/okta_user_lifecycle.json new file mode 100644 index 0000000000..4c3b3e7fae --- /dev/null +++ b/internal/api/testdata/scim/okta_user_lifecycle.json @@ -0,0 +1,207 @@ +[ + { + "step": "assign new user", + "requests": [ + { + "method": "GET", + "path": "/scim/v2/Users?filter=userName%20eq%20%22jsmith%40example.com%22&startIndex=1&count=100" + }, + { + "method": "POST", + "path": "/scim/v2/Users", + "body": { + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:User" + ], + "userName": "jsmith@example.com", + "name": { + "givenName": "John", + "familyName": "Smith" + }, + "emails": [ + { + "primary": true, + "value": "jsmith@example.com", + "type": "work" + } + ], + "displayName": "John Smith", + "locale": "en-US", + "externalId": "701984", + "groups": [], + "password": "okta-generated-password", + "active": true + } + } + ] + }, + { + "step": "edit last name", + "requests": [ + { + "method": "GET", + "path": "/scim/v2/Users/2819c223-7f76-453a-919d-413861904646" + }, + { + "method": "PUT", + "path": "/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "body": { + "active": true, + "displayName": "John Jensen", + "emails": [ + { + "primary": true, + "type": "work", + "value": "jsmith@example.com" + } + ], + "externalId": "701984", + "id": "2819c223-7f76-453a-919d-413861904646", + "locale": "en-US", + "meta": { + "created": "2026-09-28T19:13:58.892385Z", + "lastModified": "2026-09-28T19:13:58.892385Z", + "location": "http://localhost:9999/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "resourceType": "User", + "version": "W/\"1790622838892385\"" + }, + "name": { + "familyName": "Jensen", + "givenName": "John" + }, + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:User" + ], + "userName": "jsmith@example.com", + "groups": [] + } + } + ] + }, + { + "step": "unassign", + "requests": [ + { + "method": "GET", + "path": "/scim/v2/Users/2819c223-7f76-453a-919d-413861904646" + }, + { + "method": "PUT", + "path": "/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "body": { + "active": false, + "displayName": "John Jensen", + "emails": [ + { + "primary": true, + "type": "work", + "value": "jsmith@example.com" + } + ], + "externalId": "701984", + "id": "2819c223-7f76-453a-919d-413861904646", + "locale": "en-US", + "meta": { + "created": "2026-09-28T19:13:58.892385Z", + "lastModified": "2026-09-28T19:14:54.181071Z", + "location": "http://localhost:9999/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "resourceType": "User", + "version": "W/\"1790622894181071\"" + }, + "name": { + "familyName": "Jensen", + "givenName": "John" + }, + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:User" + ], + "userName": "jsmith@example.com" + } + } + ] + }, + { + "step": "reassign (PUT sent twice)", + "requests": [ + { + "method": "GET", + "path": "/scim/v2/Users?filter=userName%20eq%20%22jsmith%40example.com%22&startIndex=1&count=100" + }, + { + "method": "GET", + "path": "/scim/v2/Users/2819c223-7f76-453a-919d-413861904646" + }, + { + "method": "PUT", + "path": "/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "body": { + "active": true, + "displayName": "John Jensen", + "emails": [ + { + "primary": true, + "type": "work", + "value": "jsmith@example.com" + } + ], + "externalId": "701984", + "id": "2819c223-7f76-453a-919d-413861904646", + "locale": "en-US", + "meta": { + "created": "2026-09-28T19:13:58.892385Z", + "lastModified": "2026-09-28T19:15:39.083683Z", + "location": "http://localhost:9999/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "resourceType": "User", + "version": "W/\"1790622939083683\"" + }, + "name": { + "familyName": "Jensen", + "givenName": "John" + }, + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:User" + ], + "userName": "jsmith@example.com" + } + }, + { + "method": "GET", + "path": "/scim/v2/Users/2819c223-7f76-453a-919d-413861904646" + }, + { + "method": "PUT", + "path": "/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "body": { + "active": true, + "displayName": "John Jensen", + "emails": [ + { + "primary": true, + "type": "work", + "value": "jsmith@example.com" + } + ], + "externalId": "701984", + "id": "2819c223-7f76-453a-919d-413861904646", + "locale": "en-US", + "meta": { + "created": "2026-09-28T19:13:58.892385Z", + "lastModified": "2026-09-28T19:16:02.349821Z", + "location": "http://localhost:9999/scim/v2/Users/2819c223-7f76-453a-919d-413861904646", + "resourceType": "User", + "version": "W/\"1790622962349821\"" + }, + "name": { + "familyName": "Jensen", + "givenName": "John" + }, + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:User" + ], + "userName": "jsmith@example.com", + "groups": [] + } + } + ] + } +] diff --git a/internal/api/scim/testdata/service_provider_config.json b/internal/api/testdata/scim/service_provider_config.json similarity index 86% rename from internal/api/scim/testdata/service_provider_config.json rename to internal/api/testdata/scim/service_provider_config.json index 22b2337714..7c1ecb08d0 100644 --- a/internal/api/scim/testdata/service_provider_config.json +++ b/internal/api/testdata/scim/service_provider_config.json @@ -3,7 +3,7 @@ "urn:ietf:params:scim:schemas:core:2.0:ServiceProviderConfig" ], "patch": { - "supported": false + "supported": true }, "bulk": { "supported": false, @@ -11,17 +11,17 @@ "maxPayloadSize": 0 }, "filter": { - "supported": false, - "maxResults": 0 + "supported": true, + "maxResults": 100 }, "changePassword": { "supported": false }, "sort": { - "supported": false + "supported": true }, "etag": { - "supported": false + "supported": true }, "authenticationSchemes": [ { diff --git a/internal/conf/configuration.go b/internal/conf/configuration.go index 1e5fb3af79..6ffb3ae5a9 100644 --- a/internal/conf/configuration.go +++ b/internal/conf/configuration.go @@ -409,11 +409,6 @@ type ExperimentalConfiguration struct { // Env: GOTRUE_EXPERIMENTAL_CURSOR_PAGINATION_ENABLED=true CursorPaginationEnabled bool `split_words:"true" default:"false"` - // ScimEnabled gates the /scim/v2 router. Ships dark: no per-provider - // enablement yet, just a kill switch for internal verification. - // Env: GOTRUE_EXPERIMENTAL_SCIM_ENABLED=true - ScimEnabled bool `split_words:"true" default:"false"` - // CreateEmailIdentityOnPasswordSetEnabled creates the missing email provider // identity for a user when a password is added to an account that didn't have // one (e.g. a user who signed up with an external provider and later sets a password). @@ -487,6 +482,7 @@ type GlobalConfiguration struct { RateLimitVerify float64 `split_words:"true" default:"30"` RateLimitTokenRefresh float64 `split_words:"true" default:"150"` RateLimitSso float64 `split_words:"true" default:"30"` + RateLimitScim float64 `split_words:"true" default:"3000"` RateLimitAnonymousUsers float64 `split_words:"true" default:"30"` RateLimitOtp float64 `split_words:"true" default:"30"` RateLimitWeb3 float64 `split_words:"true" default:"30"` @@ -506,6 +502,7 @@ type GlobalConfiguration struct { Sessions SessionsConfiguration `json:"sessions"` MFA MFAConfiguration `json:"MFA"` SAML SAMLConfiguration `json:"saml"` + SSO SSOConfiguration `json:"sso"` WebAuthn WebAuthnConfiguration `json:"webauthn"` Passkey PasskeyConfiguration `json:"passkey"` CORS CORSConfiguration `json:"cors"` @@ -515,6 +512,14 @@ type GlobalConfiguration struct { Reloading ReloadingConfiguration `json:"reloading"` } +type SSOConfiguration struct { + SCIM SCIMConfiguration `json:"scim"` +} + +type SCIMConfiguration struct { + Enabled bool `json:"enabled" default:"false"` +} + type CORSConfiguration struct { AllowedHeaders []string `json:"allowed_headers" split_words:"true"` } diff --git a/internal/conf/confload/confload_test.go b/internal/conf/confload/confload_test.go index e7193ebee8..443802c5a8 100644 --- a/internal/conf/confload/confload_test.go +++ b/internal/conf/confload/confload_test.go @@ -266,7 +266,7 @@ func TestExperimentalCursorPaginationEnabled(t *testing.T) { } } -func TestExperimentalScimEndpoints(t *testing.T) { +func TestSCIMEnabled(t *testing.T) { baseEnv := func() { os.Clearenv() os.Setenv("GOTRUE_SITE_URL", "http://localhost:8080") @@ -281,7 +281,16 @@ func TestExperimentalScimEndpoints(t *testing.T) { cfg, err := LoadGlobalFromEnv() require.NoError(t, err) require.NotNil(t, cfg) - assert.Equal(t, false, cfg.Experimental.ScimEnabled) + assert.Equal(t, false, cfg.SSO.SCIM.Enabled) + } + + { + baseEnv() + os.Setenv("GOTRUE_SSO_SCIM_ENABLED", "true") + cfg, err := LoadGlobalFromEnv() + require.NoError(t, err) + require.NotNil(t, cfg) + assert.Equal(t, true, cfg.SSO.SCIM.Enabled) } { @@ -290,7 +299,7 @@ func TestExperimentalScimEndpoints(t *testing.T) { cfg, err := LoadGlobalFromEnv() require.NoError(t, err) require.NotNil(t, cfg) - assert.Equal(t, true, cfg.Experimental.ScimEnabled) + assert.Equal(t, false, cfg.SSO.SCIM.Enabled) } } diff --git a/internal/models/audit_log_entry.go b/internal/models/audit_log_entry.go index f8b493a783..76b4291183 100644 --- a/internal/models/audit_log_entry.go +++ b/internal/models/audit_log_entry.go @@ -50,12 +50,28 @@ const ( RecoveryCodesVerifiedAction AuditAction = "recovery_codes_verified" RecoveryCodesRegeneratedAction AuditAction = "recovery_codes_regenerated" RecoveryCodesDeletedAction AuditAction = "recovery_codes_deleted" + SCIMUserCreatedAction AuditAction = "scim_user_created" + SCIMUserUpdatedAction AuditAction = "scim_user_updated" + SCIMUserDeactivatedAction AuditAction = "scim_user_deactivated" + SCIMUserReactivatedAction AuditAction = "scim_user_reactivated" + SCIMUserDeletedAction AuditAction = "scim_user_deleted" + SCIMGroupCreatedAction AuditAction = "scim_group_created" + SCIMGroupUpdatedAction AuditAction = "scim_group_updated" + SCIMGroupDeletedAction AuditAction = "scim_group_deleted" + SCIMGroupMemberAddedAction AuditAction = "scim_group_member_added" + SCIMGroupMemberRemovedAction AuditAction = "scim_group_member_removed" + SCIMEnabledAction AuditAction = "scim_enabled" + SCIMDisabledAction AuditAction = "scim_disabled" + SCIMTokenCreatedAction AuditAction = "scim_token_created" // #nosec G101 + SCIMTokenRevokedAction AuditAction = "scim_token_revoked" // #nosec G101 + SCIMUsersBannedAction AuditAction = "scim_users_banned" account auditLogType = "account" team auditLogType = "team" token auditLogType = "token" user auditLogType = "user" factor auditLogType = "factor" + scim auditLogType = "scim" ) var ActionLogTypeMap = map[AuditAction]auditLogType{ @@ -88,6 +104,21 @@ var ActionLogTypeMap = map[AuditAction]auditLogType{ PasskeyCreatedAction: user, PasskeyUpdatedAction: user, PasskeyDeletedAction: user, + SCIMUserCreatedAction: scim, + SCIMUserUpdatedAction: scim, + SCIMUserDeactivatedAction: scim, + SCIMUserReactivatedAction: scim, + SCIMUserDeletedAction: scim, + SCIMGroupCreatedAction: scim, + SCIMGroupUpdatedAction: scim, + SCIMGroupDeletedAction: scim, + SCIMGroupMemberAddedAction: scim, + SCIMGroupMemberRemovedAction: scim, + SCIMEnabledAction: scim, + SCIMDisabledAction: scim, + SCIMTokenCreatedAction: scim, + SCIMTokenRevokedAction: scim, + SCIMUsersBannedAction: scim, } // AuditLogEntry is the database model for audit log entries. diff --git a/internal/models/connection.go b/internal/models/connection.go index 1a1009db6a..d7529a3960 100644 --- a/internal/models/connection.go +++ b/internal/models/connection.go @@ -64,6 +64,11 @@ func TruncateAll(conn *storage.Connection) error { (&pop.Model{Value: RecoveryCodeEntry{}}).TableName(), (&pop.Model{Value: Challenge{}}).TableName(), (&pop.Model{Value: AMRClaim{}}).TableName(), + (&pop.Model{Value: SCIMGroupMember{}}).TableName(), + (&pop.Model{Value: SCIMGroup{}}).TableName(), + (&pop.Model{Value: SCIMUser{}}).TableName(), + (&pop.Model{Value: SCIMToken{}}).TableName(), + (&pop.Model{Value: SCIMSettings{}}).TableName(), (&pop.Model{Value: SSOProvider{}}).TableName(), (&pop.Model{Value: SSODomain{}}).TableName(), (&pop.Model{Value: SAMLProvider{}}).TableName(), diff --git a/internal/models/errors.go b/internal/models/errors.go index 149541d44f..04176e11a5 100644 --- a/internal/models/errors.go +++ b/internal/models/errors.go @@ -1,6 +1,10 @@ package models -import "errors" +import ( + "errors" + + "github.com/gofrs/uuid" +) // sentinel error for all not found errors. var errNotFound = errors.New("not found") @@ -211,3 +215,87 @@ type RecoveryCodeAlreadyConsumedError struct{} func (e RecoveryCodeAlreadyConsumedError) Error() string { return "Recovery code already consumed" } + +type SCIMTokenNotFoundError struct{} + +func (e SCIMTokenNotFoundError) Error() string { + return "SCIM token not found" +} + +func (e SCIMTokenNotFoundError) Is(target error) bool { + return target == errNotFound +} + +type SCIMTokenExpiryError struct{} + +func (e SCIMTokenExpiryError) Error() string { + return "SCIM token must expire after it is created" +} + +type SCIMUserNotFoundError struct{} + +func (e SCIMUserNotFoundError) Error() string { + return "SCIM user not found" +} + +func (e SCIMUserNotFoundError) Is(target error) bool { + return target == errNotFound +} + +type SCIMUserStaleError struct{} + +func (e SCIMUserStaleError) Error() string { + return "SCIM user has changed since it was read" +} + +type SCIMUserConflictError struct{} + +func (e SCIMUserConflictError) Error() string { + return "SCIM user conflicts with an existing user" +} + +type SCIMUserLinkedError struct{} + +func (e SCIMUserLinkedError) Error() string { + return "user is already linked to a SCIM user in this provider" +} + +type SCIMGroupNotFoundError struct{} + +func (e SCIMGroupNotFoundError) Error() string { + return "SCIM group not found" +} + +func (e SCIMGroupNotFoundError) Is(target error) bool { + return target == errNotFound +} + +type SCIMGroupStaleError struct{} + +func (e SCIMGroupStaleError) Error() string { + return "SCIM group has changed since it was read" +} + +type SCIMGroupConflictError struct{} + +func (e SCIMGroupConflictError) Error() string { + return "SCIM group conflicts with an existing group" +} + +type SCIMGroupMemberNotFoundError struct { + IDs []uuid.UUID +} + +func (e SCIMGroupMemberNotFoundError) Error() string { + return "SCIM group member is not a user in this provider" +} + +type SCIMIdentityNotFoundError struct{} + +func (e SCIMIdentityNotFoundError) Error() string { + return "SCIM identity not found" +} + +func (e SCIMIdentityNotFoundError) Is(target error) bool { + return target == errNotFound +} diff --git a/internal/models/scim.go b/internal/models/scim.go new file mode 100644 index 0000000000..5b55c34359 --- /dev/null +++ b/internal/models/scim.go @@ -0,0 +1,212 @@ +package models + +import ( + "database/sql" + "fmt" + "strings" + "time" + + "github.com/gofrs/uuid" + "github.com/jackc/pgconn" + "github.com/jackc/pgerrcode" + "github.com/pkg/errors" + "github.com/supabase/auth/internal/storage" +) + +type SCIMFilter struct { + Name *string + ExternalID *string +} + +type SCIMSortKey int + +const ( + SCIMSortByCreatedAt SCIMSortKey = iota + SCIMSortByID + SCIMSortByName + SCIMSortByUpdatedAt +) + +type SCIMOrder struct { + By SCIMSortKey + Descending bool +} + +type SCIMQuery struct { + Filter SCIMFilter + Order SCIMOrder + Offset int + Limit int +} + +type scimTable struct { + name string + label string + columns string + nameCol string + live string + notFound error + stale error + conflict error +} + +func findSCIMPage[T any](tx *storage.Connection, table scimTable, providerID uuid.UUID, query SCIMQuery) ([]T, int, error) { + where, args := table.where(providerID, query.Filter) + total, err := tx.Q().Where(where, args...).Count(new(T)) + if err != nil { + return nil, 0, errors.Wrapf(err, "error counting %s", table.name) + } + rows := []T{} + if query.Limit <= 0 || query.Offset >= total { + return rows, total, nil + } + if err := tx.RawQuery( + fmt.Sprintf("SELECT %s FROM %q WHERE %s ORDER BY %s OFFSET ? LIMIT ?", table.columns, table.name, where, table.orderBy(query.Order)), + append(args, query.Offset, query.Limit)..., + ).All(&rows); err != nil { + return nil, 0, errors.Wrapf(err, "error finding %s", table.name) + } + return rows, total, nil +} + +func createSCIMRow[T any](tx *storage.Connection, table scimTable, providerID uuid.UUID, resource []byte) (*T, error) { + row := new(T) + if err := tx.RawQuery( + fmt.Sprintf("INSERT INTO %q (id, sso_provider_id, resource) VALUES (?, ?, ?::jsonb) RETURNING "+table.columns, table.name), + uuid.Must(uuid.NewV4()), providerID, string(resource), + ).First(row); err != nil { + return nil, table.error(err, "creating") + } + return row, nil +} + +func findSCIMRow[T any](tx *storage.Connection, table scimTable, providerID, id uuid.UUID, lock string) (*T, error) { + where := "id = ? AND sso_provider_id = ?" + if table.live != "" { + where += " AND " + table.live + } + row := new(T) + if err := tx.RawQuery( + fmt.Sprintf("SELECT %s FROM %q WHERE %s%s", table.columns, table.name, where, lock), + id, providerID, + ).First(row); err != nil { + return nil, table.error(err, "finding") + } + return row, nil +} + +func findUnchangedSCIMRow[T any](tx *storage.Connection, table scimTable, providerID, id uuid.UUID, resource []byte, updatedAt *time.Time) (*T, error) { + where := "id = ? AND sso_provider_id = ? AND resource = ?::jsonb AND (?::timestamptz IS NULL OR updated_at = ?)" + if table.live != "" { + where += " AND " + table.live + } + row := new(T) + if err := tx.RawQuery( + fmt.Sprintf("SELECT %s FROM %q WHERE %s FOR UPDATE", table.columns, table.name, where), + id, providerID, string(resource), updatedAt, updatedAt, + ).First(row); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, table.error(err, "finding") + } + return row, nil +} + +func scimWriteError[T any](tx *storage.Connection, table scimTable, err error, providerID, id uuid.UUID, updatedAt *time.Time, verb string) error { + if errors.Is(err, sql.ErrNoRows) && updatedAt != nil { + if _, findErr := findSCIMRow[T](tx, table, providerID, id, ""); findErr != nil { + return findErr + } + return table.stale + } + return table.error(err, verb) +} + +func (t scimTable) where(providerID uuid.UUID, filter SCIMFilter) (string, []any) { + clauses := []string{"sso_provider_id = ?"} + args := []any{providerID} + if t.live != "" { + clauses = append(clauses, t.live) + } + if filter.Name != nil { + clauses = append(clauses, t.nameCol+` COLLATE "C" = lower(?)`) + args = append(args, *filter.Name) + } + if filter.ExternalID != nil { + clauses = append(clauses, "external_id = ?") + args = append(args, *filter.ExternalID) + } + return strings.Join(clauses, " AND "), args +} + +func (t scimTable) orderBy(order SCIMOrder) string { + direction := "ASC" + if order.Descending { + direction = "DESC" + } + switch order.By { + case SCIMSortByID: + return "id " + direction + case SCIMSortByName: + return t.nameCol + ` COLLATE "C" ` + direction + ", id " + direction + case SCIMSortByUpdatedAt: + return "updated_at " + direction + ", id " + direction + } + return "created_at " + direction + ", id " + direction +} + +func (t scimTable) error(err error, verb string) error { + switch { + case errors.Is(err, sql.ErrNoRows): + return t.notFound + case isUniqueViolation(err): + return t.conflict + } + return errors.Wrapf(err, "error %s %s", verb, t.label) +} + +func isUniqueViolation(err error) bool { + var pgErr *pgconn.PgError + return errors.As(err, &pgErr) && pgErr.Code == pgerrcode.UniqueViolation +} + +func isCheckViolation(err error, constraint string) bool { + var pgErr *pgconn.PgError + return errors.As(err, &pgErr) && pgErr.Code == pgerrcode.CheckViolation && pgErr.ConstraintName == constraint +} + +func uuidArray(ids []uuid.UUID) string { + values := make([]string, len(ids)) + for i, id := range ids { + values[i] = id.String() + } + return "{" + strings.Join(values, ",") + "}" +} + +func dedupeUUIDs(ids []uuid.UUID) []uuid.UUID { + seen := make(map[uuid.UUID]struct{}, len(ids)) + unique := make([]uuid.UUID, 0, len(ids)) + for _, id := range ids { + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + unique = append(unique, id) + } + return unique +} + +func differenceUUIDs(from, subtract []uuid.UUID) []uuid.UUID { + exclude := make(map[uuid.UUID]struct{}, len(subtract)) + for _, id := range subtract { + exclude[id] = struct{}{} + } + result := []uuid.UUID{} + for _, id := range from { + if _, ok := exclude[id]; !ok { + result = append(result, id) + } + } + return result +} diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go new file mode 100644 index 0000000000..bcd136bf82 --- /dev/null +++ b/internal/models/scim_group.go @@ -0,0 +1,227 @@ +package models + +import ( + "fmt" + "time" + + "github.com/gofrs/uuid" + "github.com/pkg/errors" + "github.com/supabase/auth/internal/storage" +) + +type SCIMGroup struct { + ID uuid.UUID `db:"id"` + SSOProviderID uuid.UUID `db:"sso_provider_id"` + Resource []byte `db:"resource"` + DisplayName string `db:"display_name"` + ExternalID *string `db:"external_id"` + CreatedAt time.Time `db:"created_at"` + UpdatedAt time.Time `db:"updated_at"` +} + +const scimGroupColumns = "id, sso_provider_id, resource, display_name, external_id, created_at, updated_at" + +func (SCIMGroup) TableName() string { + return "scim_groups" +} + +type SCIMGroupMember struct { + GroupID uuid.UUID `db:"group_id"` + SCIMUserID uuid.UUID `db:"scim_user_id"` + CreatedAt time.Time `db:"created_at"` +} + +func (SCIMGroupMember) TableName() string { + return "scim_group_members" +} + +type SCIMGroupMembership struct { + GroupID uuid.UUID `db:"group_id"` + SCIMUserID uuid.UUID `db:"scim_user_id"` + Display string `db:"display"` +} + +var scimGroupsTable = scimTable{ + name: SCIMGroup{}.TableName(), + label: "SCIM group", + columns: scimGroupColumns, + nameCol: "display_name", + notFound: SCIMGroupNotFoundError{}, + stale: SCIMGroupStaleError{}, + conflict: SCIMGroupConflictError{}, +} + +func CreateSCIMGroup(tx *storage.Connection, providerID uuid.UUID, resource []byte) (*SCIMGroup, error) { + return createSCIMRow[SCIMGroup](tx, scimGroupsTable, providerID, resource) +} + +func FindSCIMGroup(tx *storage.Connection, providerID, id uuid.UUID) (*SCIMGroup, error) { + return findSCIMRow[SCIMGroup](tx, scimGroupsTable, providerID, id, "") +} + +func FindSCIMGroupForUpdate(tx *storage.Connection, providerID, id uuid.UUID) (*SCIMGroup, error) { + return findSCIMRow[SCIMGroup](tx, scimGroupsTable, providerID, id, " FOR UPDATE") +} + +func FindSCIMGroups(tx *storage.Connection, providerID uuid.UUID, query SCIMQuery) ([]SCIMGroup, int, error) { + return findSCIMPage[SCIMGroup](tx, scimGroupsTable, providerID, query) +} + +func ReplaceSCIMGroup(tx *storage.Connection, providerID, id uuid.UUID, resource []byte, updatedAt *time.Time) (*SCIMGroup, error) { + group := &SCIMGroup{} + err := tx.RawQuery( + fmt.Sprintf("UPDATE %q SET resource = ?::jsonb, updated_at = now() WHERE id = ? AND sso_provider_id = ? AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimGroupColumns, group.TableName()), + string(resource), id, providerID, updatedAt, updatedAt, + ).First(group) + if err != nil { + return nil, scimWriteError[SCIMGroup](tx, scimGroupsTable, err, providerID, id, updatedAt, "replacing") + } + return group, nil +} + +func FindUnchangedSCIMGroup(tx *storage.Connection, providerID, id uuid.UUID, resource []byte, updatedAt *time.Time) (*SCIMGroup, error) { + return findUnchangedSCIMRow[SCIMGroup](tx, scimGroupsTable, providerID, id, resource, updatedAt) +} + +func TouchSCIMGroup(tx *storage.Connection, group *SCIMGroup) (*SCIMGroup, error) { + touched := &SCIMGroup{} + if err := tx.RawQuery( + fmt.Sprintf("UPDATE %q SET updated_at = now() WHERE id = ? RETURNING "+scimGroupColumns, group.TableName()), + group.ID, + ).First(touched); err != nil { + return nil, errors.Wrap(err, "error updating SCIM group") + } + return touched, nil +} + +func DeleteSCIMGroup(tx *storage.Connection, providerID, id uuid.UUID, updatedAt *time.Time) (*SCIMGroup, error) { + group := &SCIMGroup{} + err := tx.RawQuery( + fmt.Sprintf("DELETE FROM %q WHERE id = ? AND sso_provider_id = ? AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimGroupColumns, group.TableName()), + id, providerID, updatedAt, updatedAt, + ).First(group) + if err != nil { + return nil, scimWriteError[SCIMGroup](tx, scimGroupsTable, err, providerID, id, updatedAt, "deleting") + } + return group, nil +} + +func FindSCIMGroupMembers(tx *storage.Connection, providerID uuid.UUID, groupIDs []uuid.UUID) ([]SCIMGroupMembership, error) { + members := []SCIMGroupMembership{} + if len(groupIDs) == 0 { + return members, nil + } + err := tx.RawQuery( + fmt.Sprintf("SELECT m.group_id, m.scim_user_id, u.resource->>'userName' AS display FROM %q m JOIN %q u ON u.id = m.scim_user_id WHERE m.group_id = ANY(?::uuid[]) AND u.sso_provider_id = ? AND u.deleted_at IS NULL ORDER BY m.group_id, m.created_at, m.scim_user_id", (&SCIMGroupMember{}).TableName(), (&SCIMUser{}).TableName()), + uuidArray(groupIDs), providerID, + ).All(&members) + if err != nil { + return nil, errors.Wrap(err, "error finding SCIM group members") + } + return members, nil +} + +func FindSCIMGroupsForUsers(tx *storage.Connection, providerID uuid.UUID, scimUserIDs []uuid.UUID) ([]SCIMGroupMembership, error) { + groups := []SCIMGroupMembership{} + if len(scimUserIDs) == 0 { + return groups, nil + } + err := tx.RawQuery( + fmt.Sprintf("SELECT m.group_id, m.scim_user_id, g.resource->>'displayName' AS display FROM %q m JOIN %q g ON g.id = m.group_id WHERE m.scim_user_id = ANY(?::uuid[]) AND g.sso_provider_id = ? ORDER BY m.scim_user_id, g.display_name COLLATE \"C\", g.id", (&SCIMGroupMember{}).TableName(), (&SCIMGroup{}).TableName()), + uuidArray(scimUserIDs), providerID, + ).All(&groups) + if err != nil { + return nil, errors.Wrap(err, "error finding SCIM groups for users") + } + return groups, nil +} + +func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserIDs []uuid.UUID) (added, removed []uuid.UUID, err error) { + wanted := dedupeUUIDs(scimUserIDs) + + foundIDs := []uuid.UUID{} + if len(wanted) > 0 { + if err := tx.RawQuery( + fmt.Sprintf("SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL", (&SCIMUser{}).TableName()), + uuidArray(wanted), group.SSOProviderID, + ).All(&foundIDs); err != nil { + return nil, nil, errors.Wrap(err, "error finding SCIM group members") + } + } + if missing := differenceUUIDs(wanted, foundIDs); len(missing) > 0 { + return nil, nil, SCIMGroupMemberNotFoundError{IDs: missing} + } + + current := []SCIMGroupMember{} + if err := tx.RawQuery( + fmt.Sprintf("SELECT group_id, scim_user_id, created_at FROM %q WHERE group_id = ?", (&SCIMGroupMember{}).TableName()), + group.ID, + ).All(¤t); err != nil { + return nil, nil, errors.Wrap(err, "error finding SCIM group members") + } + currentIDs := make([]uuid.UUID, len(current)) + for i := range current { + currentIDs[i] = current[i].SCIMUserID + } + + removed = differenceUUIDs(currentIDs, wanted) + added = differenceUUIDs(wanted, currentIDs) + + if len(removed) > 0 { + if err := tx.RawQuery( + fmt.Sprintf("DELETE FROM %q WHERE group_id = ? AND scim_user_id = ANY(?::uuid[])", (&SCIMGroupMember{}).TableName()), + group.ID, uuidArray(removed), + ).Exec(); err != nil { + return nil, nil, errors.Wrap(err, "error removing SCIM group members") + } + } + if len(added) > 0 { + locked := []uuid.UUID{} + if err := tx.RawQuery( + fmt.Sprintf("SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL ORDER BY id FOR SHARE", (&SCIMUser{}).TableName()), + uuidArray(added), group.SSOProviderID, + ).All(&locked); err != nil { + return nil, nil, errors.Wrap(err, "error locking SCIM group members") + } + if missing := differenceUUIDs(added, locked); len(missing) > 0 { + return nil, nil, SCIMGroupMemberNotFoundError{IDs: missing} + } + if err := tx.RawQuery( + fmt.Sprintf("INSERT INTO %q (group_id, scim_user_id) SELECT ?, unnest(?::uuid[])", (&SCIMGroupMember{}).TableName()), + group.ID, uuidArray(locked), + ).Exec(); err != nil { + return nil, nil, errors.Wrap(err, "error adding SCIM group members") + } + } + return added, removed, nil +} + +func RemoveSCIMUserFromGroups(tx *storage.Connection, scimUserID uuid.UUID) ([]uuid.UUID, error) { + groups, members := (&SCIMGroup{}).TableName(), (&SCIMGroupMember{}).TableName() + if err := tx.RawQuery( + fmt.Sprintf("SELECT id FROM %q WHERE id IN (SELECT group_id FROM %q WHERE scim_user_id = ?) ORDER BY id FOR UPDATE", groups, members), + scimUserID, + ).Exec(); err != nil { + return nil, errors.Wrap(err, "error locking SCIM groups") + } + removed := []SCIMGroupMember{} + if err := tx.RawQuery( + fmt.Sprintf("DELETE FROM %q WHERE scim_user_id = ? RETURNING group_id, scim_user_id, created_at", members), + scimUserID, + ).All(&removed); err != nil { + return nil, errors.Wrap(err, "error removing SCIM user from groups") + } + groupIDs := make([]uuid.UUID, len(removed)) + for i := range removed { + groupIDs[i] = removed[i].GroupID + } + if len(groupIDs) > 0 { + if err := tx.RawQuery( + fmt.Sprintf("UPDATE %q SET updated_at = now() WHERE id = ANY(?::uuid[])", groups), + uuidArray(groupIDs), + ).Exec(); err != nil { + return nil, errors.Wrap(err, "error updating SCIM groups") + } + } + return groupIDs, nil +} diff --git a/internal/models/scim_group_test.go b/internal/models/scim_group_test.go new file mode 100644 index 0000000000..1fbf03dcec --- /dev/null +++ b/internal/models/scim_group_test.go @@ -0,0 +1,324 @@ +package models + +import ( + "fmt" + "testing" + "time" + + "github.com/gofrs/uuid" + "github.com/stretchr/testify/require" + "github.com/stretchr/testify/suite" + "github.com/supabase/auth/internal/conf/confload" + "github.com/supabase/auth/internal/storage" + "github.com/supabase/auth/internal/storage/test" +) + +type SCIMGroupTestSuite struct { + suite.Suite + db *storage.Connection + provider *SSOProvider +} + +func TestSCIMGroup(t *testing.T) { + globalConfig, err := confload.LoadGlobal(modelsTestConfig) + require.NoError(t, err) + conn, err := test.SetupDBConnection(globalConfig) + require.NoError(t, err) + ts := &SCIMGroupTestSuite{db: conn} + defer ts.db.Close() + suite.Run(t, ts) +} + +func (ts *SCIMGroupTestSuite) SetupTest() { + require.NoError(ts.T(), TruncateAll(ts.db)) + ts.provider = ts.createProvider() +} + +func (ts *SCIMGroupTestSuite) createProvider() *SSOProvider { + return createSCIMTestProvider(ts.T(), ts.db) +} + +func (ts *SCIMGroupTestSuite) createGroup(providerID uuid.UUID, displayName string) *SCIMGroup { + group, err := CreateSCIMGroup(ts.db, providerID, []byte(fmt.Sprintf(`{"displayName":%q}`, displayName))) + require.NoError(ts.T(), err) + return group +} + +func (ts *SCIMGroupTestSuite) createUser(providerID uuid.UUID, userName string) *SCIMUser { + user, err := CreateSCIMUser(ts.db, providerID, []byte(fmt.Sprintf(`{"userName":%q}`, userName))) + require.NoError(ts.T(), err) + return user +} + +func (ts *SCIMGroupTestSuite) TestCreate() { + group, err := CreateSCIMGroup(ts.db, ts.provider.ID, []byte(`{"displayName":"Engineering","externalId":"ext-1"}`)) + require.NoError(ts.T(), err) + + require.Equal(ts.T(), ts.provider.ID, group.SSOProviderID) + require.Equal(ts.T(), "engineering", group.DisplayName) + require.Equal(ts.T(), "ext-1", *group.ExternalID) + require.JSONEq(ts.T(), `{"displayName":"Engineering","externalId":"ext-1"}`, string(group.Resource)) + require.False(ts.T(), group.CreatedAt.IsZero()) +} + +func (ts *SCIMGroupTestSuite) TestCreateAllowsDuplicateDisplayName() { + ts.createGroup(ts.provider.ID, "Engineering") + ts.createGroup(ts.provider.ID, "engineering") +} + +func (ts *SCIMGroupTestSuite) TestCreateRejectsDuplicateExternalID() { + _, err := CreateSCIMGroup(ts.db, ts.provider.ID, []byte(`{"displayName":"A","externalId":"ext-1"}`)) + require.NoError(ts.T(), err) + + _, err = CreateSCIMGroup(ts.db, ts.provider.ID, []byte(`{"displayName":"B","externalId":"ext-1"}`)) + require.ErrorIs(ts.T(), err, SCIMGroupConflictError{}) + + other := ts.createProvider() + _, err = CreateSCIMGroup(ts.db, other.ID, []byte(`{"displayName":"C","externalId":"ext-1"}`)) + require.NoError(ts.T(), err) +} + +func (ts *SCIMGroupTestSuite) TestFindIsScopedToProvider() { + group := ts.createGroup(ts.provider.ID, "Engineering") + + found, err := FindSCIMGroup(ts.db, ts.provider.ID, group.ID) + require.NoError(ts.T(), err) + require.Equal(ts.T(), group.ID, found.ID) + + other := ts.createProvider() + _, err = FindSCIMGroup(ts.db, other.ID, group.ID) + require.ErrorIs(ts.T(), err, SCIMGroupNotFoundError{}) + require.True(ts.T(), IsNotFoundError(err)) +} + +func (ts *SCIMGroupTestSuite) TestFindGroupsFiltersAndSorts() { + ts.createGroup(ts.provider.ID, "Beta") + ts.createGroup(ts.provider.ID, "alpha") + ts.createGroup(ts.createProvider().ID, "Alpha") + + displayName := "ALPHA" + groups, total, err := FindSCIMGroups(ts.db, ts.provider.ID, SCIMQuery{Filter: SCIMFilter{Name: &displayName}, Limit: 10}) + require.NoError(ts.T(), err) + require.Equal(ts.T(), 1, total) + require.Equal(ts.T(), "alpha", groups[0].DisplayName) + + groups, total, err = FindSCIMGroups(ts.db, ts.provider.ID, SCIMQuery{Order: SCIMOrder{By: SCIMSortByName, Descending: true}, Limit: 10}) + require.NoError(ts.T(), err) + require.Equal(ts.T(), 2, total) + require.Equal(ts.T(), "beta", groups[0].DisplayName) + require.Equal(ts.T(), "alpha", groups[1].DisplayName) + + groups, total, err = FindSCIMGroups(ts.db, ts.provider.ID, SCIMQuery{}) + require.NoError(ts.T(), err) + require.Equal(ts.T(), 2, total) + require.Empty(ts.T(), groups) +} + +func (ts *SCIMGroupTestSuite) TestReplaceChecksVersion() { + group := ts.createGroup(ts.provider.ID, "Engineering") + + replaced, err := ReplaceSCIMGroup(ts.db, ts.provider.ID, group.ID, []byte(`{"displayName":"Platform"}`), &group.UpdatedAt) + require.NoError(ts.T(), err) + require.Equal(ts.T(), "platform", replaced.DisplayName) + + _, err = ReplaceSCIMGroup(ts.db, ts.provider.ID, group.ID, []byte(`{"displayName":"Stale"}`), &group.UpdatedAt) + require.ErrorIs(ts.T(), err, SCIMGroupStaleError{}) + + _, err = ReplaceSCIMGroup(ts.db, ts.createProvider().ID, group.ID, []byte(`{"displayName":"Other"}`), nil) + require.ErrorIs(ts.T(), err, SCIMGroupNotFoundError{}) +} + +func (ts *SCIMGroupTestSuite) TestDeleteRemovesMembers() { + group := ts.createGroup(ts.provider.ID, "Engineering") + user := ts.createUser(ts.provider.ID, "alice") + _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{user.ID}) + require.NoError(ts.T(), err) + + _, err = DeleteSCIMGroup(ts.db, ts.provider.ID, group.ID, nil) + require.NoError(ts.T(), err) + + _, err = FindSCIMGroup(ts.db, ts.provider.ID, group.ID) + require.ErrorIs(ts.T(), err, SCIMGroupNotFoundError{}) + count, err := ts.db.Q().Where("group_id = ?", group.ID).Count(&SCIMGroupMember{}) + require.NoError(ts.T(), err) + require.Zero(ts.T(), count) + + _, err = DeleteSCIMGroup(ts.db, ts.provider.ID, group.ID, nil) + require.ErrorIs(ts.T(), err, SCIMGroupNotFoundError{}) +} + +func (ts *SCIMGroupTestSuite) TestReplaceMembersDiffs() { + group := ts.createGroup(ts.provider.ID, "Engineering") + alice := ts.createUser(ts.provider.ID, "Alice") + bob := ts.createUser(ts.provider.ID, "bob") + carol := ts.createUser(ts.provider.ID, "carol") + + added, removed, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, bob.ID, alice.ID}) + require.NoError(ts.T(), err) + require.ElementsMatch(ts.T(), []uuid.UUID{alice.ID, bob.ID}, added) + require.Empty(ts.T(), removed) + + added, removed, err = ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{bob.ID, carol.ID}) + require.NoError(ts.T(), err) + require.Equal(ts.T(), []uuid.UUID{carol.ID}, added) + require.Equal(ts.T(), []uuid.UUID{alice.ID}, removed) + + members, err := FindSCIMGroupMembers(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) + require.NoError(ts.T(), err) + require.Len(ts.T(), members, 2) + displays := []string{members[0].Display, members[1].Display} + require.ElementsMatch(ts.T(), []string{"bob", "carol"}, displays) + + added, removed, err = ReplaceSCIMGroupMembers(ts.db, group, nil) + require.NoError(ts.T(), err) + require.Empty(ts.T(), added) + require.ElementsMatch(ts.T(), []uuid.UUID{bob.ID, carol.ID}, removed) +} + +func (ts *SCIMGroupTestSuite) TestReplaceMembersRejectsOtherProviderUsers() { + group := ts.createGroup(ts.provider.ID, "Engineering") + alice := ts.createUser(ts.provider.ID, "alice") + outsider := ts.createUser(ts.createProvider().ID, "mallory") + + _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, outsider.ID}) + require.Equal(ts.T(), SCIMGroupMemberNotFoundError{IDs: []uuid.UUID{outsider.ID}}, err) + + members, err := FindSCIMGroupMembers(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) + require.NoError(ts.T(), err) + require.Empty(ts.T(), members) +} + +func (ts *SCIMGroupTestSuite) TestReplaceMembersRejectsDeletedUsers() { + group := ts.createGroup(ts.provider.ID, "Engineering") + alice := ts.createUser(ts.provider.ID, "alice") + _, err := DeleteSCIMUser(ts.db, ts.provider.ID, alice.ID, nil) + require.NoError(ts.T(), err) + + _, _, err = ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) + require.Equal(ts.T(), SCIMGroupMemberNotFoundError{IDs: []uuid.UUID{alice.ID}}, err) + + _, _, err = ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{uuid.Must(uuid.NewV4())}) + require.ErrorAs(ts.T(), err, &SCIMGroupMemberNotFoundError{}) +} + +func (ts *SCIMGroupTestSuite) TestReplaceMembersWaitsForConcurrentUserDelete() { + group := ts.createGroup(ts.provider.ID, "Engineering") + alice := ts.createUser(ts.provider.ID, "alice") + + deleting := ts.beginTx() + defer func() { _ = deleting.TX.Rollback() }() + _, err := DeleteSCIMUser(deleting, ts.provider.ID, alice.ID, nil) + require.NoError(ts.T(), err) + _, err = RemoveSCIMUserFromGroups(deleting, alice.ID) + require.NoError(ts.T(), err) + + result := make(chan error, 1) + go func() { + result <- ts.db.Transaction(func(tx *storage.Connection) error { + _, _, err := ReplaceSCIMGroupMembers(tx, group, []uuid.UUID{alice.ID}) + return err + }) + }() + require.Eventually(ts.T(), func() bool { return ts.lockWaiters() > 0 }, 5*time.Second, 10*time.Millisecond) + require.NoError(ts.T(), deleting.TX.Commit()) + + require.Equal(ts.T(), SCIMGroupMemberNotFoundError{IDs: []uuid.UUID{alice.ID}}, <-result) + count, err := ts.db.Q().Where("group_id = ?", group.ID).Count(&SCIMGroupMember{}) + require.NoError(ts.T(), err) + require.Zero(ts.T(), count) +} + +func (ts *SCIMGroupTestSuite) TestReplaceMembersDoesNotLockExistingMembers() { + group := ts.createGroup(ts.provider.ID, "Engineering") + alice := ts.createUser(ts.provider.ID, "alice") + bob := ts.createUser(ts.provider.ID, "bob") + _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) + require.NoError(ts.T(), err) + + deleting := ts.beginTx() + defer func() { _ = deleting.TX.Rollback() }() + _, err = DeleteSCIMUser(deleting, ts.provider.ID, alice.ID, nil) + require.NoError(ts.T(), err) + + result := make(chan error, 1) + go func() { + result <- ts.db.Transaction(func(tx *storage.Connection) error { + locked, err := FindSCIMGroupForUpdate(tx, ts.provider.ID, group.ID) + if err != nil { + return err + } + _, _, err = ReplaceSCIMGroupMembers(tx, locked, []uuid.UUID{alice.ID, bob.ID}) + return err + }) + }() + select { + case err := <-result: + require.NoError(ts.T(), err) + case <-time.After(5 * time.Second): + ts.T().Fatal("group replace waited on a member that was not added") + } + + removed, err := RemoveSCIMUserFromGroups(deleting, alice.ID) + require.NoError(ts.T(), err) + require.Equal(ts.T(), []uuid.UUID{group.ID}, removed) + require.NoError(ts.T(), deleting.TX.Commit()) + + members, err := FindSCIMGroupMembers(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) + require.NoError(ts.T(), err) + require.Len(ts.T(), members, 1) + require.Equal(ts.T(), bob.ID, members[0].SCIMUserID) +} + +func (ts *SCIMGroupTestSuite) TestFindMembersHidesDeletedUsers() { + group := ts.createGroup(ts.provider.ID, "Engineering") + alice := ts.createUser(ts.provider.ID, "alice") + _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) + require.NoError(ts.T(), err) + + _, err = DeleteSCIMUser(ts.db, ts.provider.ID, alice.ID, nil) + require.NoError(ts.T(), err) + + members, err := FindSCIMGroupMembers(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) + require.NoError(ts.T(), err) + require.Empty(ts.T(), members) +} + +func (ts *SCIMGroupTestSuite) TestFindGroupsForUsers() { + engineering := ts.createGroup(ts.provider.ID, "Engineering") + admins := ts.createGroup(ts.provider.ID, "Admins") + alice := ts.createUser(ts.provider.ID, "alice") + bob := ts.createUser(ts.provider.ID, "bob") + _, _, err := ReplaceSCIMGroupMembers(ts.db, engineering, []uuid.UUID{alice.ID, bob.ID}) + require.NoError(ts.T(), err) + _, _, err = ReplaceSCIMGroupMembers(ts.db, admins, []uuid.UUID{alice.ID}) + require.NoError(ts.T(), err) + + groups, err := FindSCIMGroupsForUsers(ts.db, ts.provider.ID, []uuid.UUID{alice.ID, bob.ID}) + require.NoError(ts.T(), err) + require.Len(ts.T(), groups, 3) + + byUser := map[uuid.UUID][]string{} + for _, g := range groups { + byUser[g.SCIMUserID] = append(byUser[g.SCIMUserID], g.Display) + } + require.Equal(ts.T(), []string{"Admins", "Engineering"}, byUser[alice.ID]) + require.Equal(ts.T(), []string{"Engineering"}, byUser[bob.ID]) + + groups, err = FindSCIMGroupsForUsers(ts.db, ts.createProvider().ID, []uuid.UUID{alice.ID}) + require.NoError(ts.T(), err) + require.Empty(ts.T(), groups) +} + +func (ts *SCIMGroupTestSuite) beginTx() *storage.Connection { + tx, err := ts.db.NewTransaction() + require.NoError(ts.T(), err) + return &storage.Connection{Connection: tx} +} + +func (ts *SCIMGroupTestSuite) lockWaiters() int { + row := struct { + Count int `db:"count"` + }{} + require.NoError(ts.T(), ts.db.RawQuery("SELECT count(*) AS count FROM pg_stat_activity WHERE datname = current_database() AND wait_event_type = 'Lock'").First(&row)) + return row.Count +} diff --git a/internal/models/scim_settings.go b/internal/models/scim_settings.go new file mode 100644 index 0000000000..c47deb0c8e --- /dev/null +++ b/internal/models/scim_settings.go @@ -0,0 +1,59 @@ +package models + +import ( + "fmt" + "time" + + "github.com/gobuffalo/pop/v6" + "github.com/gofrs/uuid" + "github.com/pkg/errors" + "github.com/supabase/auth/internal/storage" +) + +type SCIMSettings struct { + SSOProviderID uuid.UUID `json:"-" db:"sso_provider_id"` + Enabled bool `json:"enabled" db:"enabled"` + CreatedAt time.Time `json:"created_at" db:"created_at"` + UpdatedAt time.Time `json:"updated_at" db:"updated_at"` +} + +func (SCIMSettings) TableName() string { + return "scim_settings" +} + +func (s *SCIMSettings) AfterFind(*pop.Connection) error { + s.CreatedAt = s.CreatedAt.UTC() + s.UpdatedAt = s.UpdatedAt.UTC() + return nil +} + +func EnableSCIM(tx *storage.Connection, providerID uuid.UUID) (bool, error) { + table := (&SCIMSettings{}).TableName() + rows := []SCIMSettings{} + if err := tx.RawQuery( + fmt.Sprintf("INSERT INTO %[1]q (sso_provider_id, enabled) VALUES (?, true) ON CONFLICT (sso_provider_id) DO UPDATE SET enabled = true, updated_at = now() WHERE %[1]q.enabled = false RETURNING *", table), + providerID, + ).All(&rows); err != nil { + return false, errors.Wrap(err, "error enabling SCIM") + } + return len(rows) > 0, nil +} + +func DisableSCIM(tx *storage.Connection, providerID uuid.UUID) (bool, error) { + rows := []SCIMSettings{} + if err := tx.RawQuery( + fmt.Sprintf("UPDATE %q SET enabled = false, updated_at = now() WHERE sso_provider_id = ? AND enabled RETURNING *", (&SCIMSettings{}).TableName()), + providerID, + ).All(&rows); err != nil { + return false, errors.Wrap(err, "error disabling SCIM") + } + return len(rows) > 0, nil +} + +func IsSCIMEnabled(tx *storage.Connection, providerID uuid.UUID) (bool, error) { + enabled, err := tx.Q().Where("sso_provider_id = ? AND enabled", providerID).Exists(&SCIMSettings{}) + if err != nil { + return false, errors.Wrap(err, "error finding SCIM settings") + } + return enabled, nil +} diff --git a/internal/models/scim_settings_test.go b/internal/models/scim_settings_test.go new file mode 100644 index 0000000000..6c37570286 --- /dev/null +++ b/internal/models/scim_settings_test.go @@ -0,0 +1,103 @@ +package models + +import ( + "sync" + "testing" + + "github.com/gofrs/uuid" + "github.com/stretchr/testify/require" + "github.com/stretchr/testify/suite" + "github.com/supabase/auth/internal/conf/confload" + "github.com/supabase/auth/internal/storage" + "github.com/supabase/auth/internal/storage/test" +) + +type SCIMSettingsTestSuite struct { + suite.Suite + db *storage.Connection + provider *SSOProvider +} + +func TestSCIMSettings(t *testing.T) { + globalConfig, err := confload.LoadGlobal(modelsTestConfig) + require.NoError(t, err) + conn, err := test.SetupDBConnection(globalConfig) + require.NoError(t, err) + ts := &SCIMSettingsTestSuite{db: conn} + defer ts.db.Close() + suite.Run(t, ts) +} + +func (ts *SCIMSettingsTestSuite) SetupTest() { + require.NoError(ts.T(), TruncateAll(ts.db)) + ts.provider = &SSOProvider{} + require.NoError(ts.T(), ts.db.Create(ts.provider)) +} + +func (ts *SCIMSettingsTestSuite) TestDisabledByDefault() { + require.False(ts.T(), ts.enabled()) +} + +func (ts *SCIMSettingsTestSuite) TestTransitions() { + for _, step := range []struct { + name string + apply func(*storage.Connection, uuid.UUID) (bool, error) + changed bool + enabled bool + }{ + {"disable never enabled", DisableSCIM, false, false}, + {"enable", EnableSCIM, true, true}, + {"enable again", EnableSCIM, false, true}, + {"disable", DisableSCIM, true, false}, + {"disable again", DisableSCIM, false, false}, + {"re-enable", EnableSCIM, true, true}, + } { + changed, err := step.apply(ts.db, ts.provider.ID) + require.NoError(ts.T(), err, step.name) + require.Equal(ts.T(), step.changed, changed, step.name) + require.Equal(ts.T(), step.enabled, ts.enabled(), step.name) + } +} + +func (ts *SCIMSettingsTestSuite) TestConcurrentEnableChangesOnce() { + type result struct { + changed bool + err error + } + var wg sync.WaitGroup + results := make(chan result, 10) + for range 10 { + wg.Go(func() { + changed, err := EnableSCIM(ts.db, ts.provider.ID) + results <- result{changed, err} + }) + } + wg.Wait() + close(results) + + changes := 0 + for r := range results { + require.NoError(ts.T(), r.err) + if r.changed { + changes++ + } + } + require.Equal(ts.T(), 1, changes) + require.True(ts.T(), ts.enabled()) +} + +func (ts *SCIMSettingsTestSuite) TestDeletedWithProvider() { + _, err := EnableSCIM(ts.db, ts.provider.ID) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.db.Destroy(ts.provider)) + + count, err := ts.db.Q().Where("sso_provider_id = ?", ts.provider.ID).Count(&SCIMSettings{}) + require.NoError(ts.T(), err) + require.Zero(ts.T(), count) +} + +func (ts *SCIMSettingsTestSuite) enabled() bool { + enabled, err := IsSCIMEnabled(ts.db, ts.provider.ID) + require.NoError(ts.T(), err) + return enabled +} diff --git a/internal/models/scim_token.go b/internal/models/scim_token.go new file mode 100644 index 0000000000..701d8858c2 --- /dev/null +++ b/internal/models/scim_token.go @@ -0,0 +1,186 @@ +package models + +import ( + "crypto/rand" + "crypto/sha256" + "database/sql" + "encoding/hex" + "fmt" + "time" + + "github.com/gobuffalo/pop/v6" + "github.com/gofrs/uuid" + "github.com/pkg/errors" + "github.com/supabase/auth/internal/storage" +) + +const ( + SCIMTokenMarker = "scim_" + + scimTokenBytes = 20 + scimTokenPrefixLength = len(SCIMTokenMarker) + 7 +) + +type SCIMToken struct { + ID uuid.UUID `json:"-" db:"id"` + SSOProviderID uuid.UUID `json:"-" db:"sso_provider_id"` + TokenHash string `json:"-" db:"token_hash"` + Prefix string `json:"prefix" db:"prefix"` + CreatedAt time.Time `json:"created_at" db:"created_at"` + ExpiresAt *time.Time `json:"expires_at" db:"expires_at"` + RevokedAt *time.Time `json:"revoked_at" db:"revoked_at"` + LastUsedAt *time.Time `json:"last_used_at" db:"last_used_at"` +} + +func (SCIMToken) TableName() string { + return "scim_tokens" +} + +func (t *SCIMToken) AfterFind(*pop.Connection) error { + t.CreatedAt = t.CreatedAt.UTC() + for _, at := range []*time.Time{t.ExpiresAt, t.RevokedAt, t.LastUsedAt} { + if at != nil { + *at = at.UTC() + } + } + return nil +} + +func (t *SCIMToken) IsRevoked() bool { + return t.RevokedAt != nil +} + +func HashSCIMToken(token string) string { + sum := sha256.Sum256([]byte(token)) + return hex.EncodeToString(sum[:]) +} + +func CreateSCIMToken(tx *storage.Connection, provider *SSOProvider, expiresAt *time.Time) (*SCIMToken, string, error) { + plaintext, err := generateSCIMToken() + if err != nil { + return nil, "", errors.Wrap(err, "error generating SCIM token") + } + + token := &SCIMToken{ + ID: uuid.Must(uuid.NewV4()), + SSOProviderID: provider.ID, + TokenHash: HashSCIMToken(plaintext), + Prefix: plaintext[:scimTokenPrefixLength], + ExpiresAt: expiresAt, + } + if err := tx.RawQuery( + fmt.Sprintf("INSERT INTO %q (id, sso_provider_id, token_hash, prefix, expires_at) VALUES (?, ?, ?, ?, ?) RETURNING *", token.TableName()), + token.ID, token.SSOProviderID, token.TokenHash, token.Prefix, token.ExpiresAt, + ).First(token); err != nil { + if isCheckViolation(err, "scim_tokens_expires_at_future") { + return nil, "", SCIMTokenExpiryError{} + } + return nil, "", errors.Wrap(err, "error creating SCIM token") + } + return token, plaintext, nil +} + +func FindSCIMTokensBySSOProvider(tx *storage.Connection, providerID uuid.UUID) ([]SCIMToken, error) { + tokens := []SCIMToken{} + if err := tx.Q().Where("sso_provider_id = ?", providerID).Order("created_at asc, id asc").All(&tokens); err != nil { + return nil, errors.Wrap(err, "error finding SCIM tokens") + } + return tokens, nil +} + +const activeSCIMTokenClause = "revoked_at IS NULL AND (expires_at IS NULL OR expires_at > now())" // #nosec G101 + +func FindActiveSCIMTokensBySSOProvider(tx *storage.Connection, providerID uuid.UUID) ([]SCIMToken, error) { + tokens := []SCIMToken{} + if err := tx.Q().Where("sso_provider_id = ? AND "+activeSCIMTokenClause, providerID).Order("created_at asc, id asc").All(&tokens); err != nil { + return nil, errors.Wrap(err, "error finding active SCIM tokens") + } + return tokens, nil +} + +func LockSCIMTokens(tx *storage.Connection, providerID uuid.UUID) error { + if err := tx.RawQuery("SELECT pg_advisory_xact_lock(hashtextextended(?, 0))", "scim_tokens|"+providerID.String()).Exec(); err != nil { + return errors.Wrap(err, "error locking SCIM tokens") + } + return nil +} + +func RevokeSCIMTokensBySSOProvider(tx *storage.Connection, providerID uuid.UUID) ([]SCIMToken, error) { + tokens := []SCIMToken{} + if err := tx.RawQuery( + fmt.Sprintf("WITH revoked AS (UPDATE %q SET revoked_at = now() WHERE sso_provider_id = ? AND "+activeSCIMTokenClause+" RETURNING *) SELECT * FROM revoked ORDER BY created_at ASC, id ASC", (&SCIMToken{}).TableName()), + providerID, + ).All(&tokens); err != nil { + return nil, errors.Wrap(err, "error revoking SCIM tokens") + } + return tokens, nil +} + +func FindSCIMTokenByPrefix(tx *storage.Connection, providerID uuid.UUID, prefix string) (*SCIMToken, error) { + tokens := []SCIMToken{} + if err := tx.Q().Where("sso_provider_id = ? AND prefix = ?", providerID, prefix).Limit(2).All(&tokens); err != nil { + return nil, errors.Wrap(err, "error finding SCIM token") + } + + switch len(tokens) { + case 0: + return nil, SCIMTokenNotFoundError{} + case 1: + return &tokens[0], nil + default: + return nil, errors.Errorf("error finding SCIM token: prefix %q is ambiguous", prefix) + } +} + +func (t *SCIMToken) Revoke(tx *storage.Connection) error { + if t.IsRevoked() { + return nil + } + if err := tx.RawQuery( + fmt.Sprintf("UPDATE %q SET revoked_at = now() WHERE id = ? RETURNING *", t.TableName()), + t.ID, + ).First(t); err != nil { + return errors.Wrap(err, "error revoking SCIM token") + } + return nil +} + +func AuthenticateSCIMToken(tx *storage.Connection, plaintext string) (*SCIMToken, error) { + token := &SCIMToken{} + err := tx.RawQuery( + fmt.Sprintf(`WITH authenticated AS ( + SELECT t.* FROM %[1]q AS t + JOIN %[2]q AS p ON p.id = t.sso_provider_id + JOIN %[3]q AS s ON s.sso_provider_id = t.sso_provider_id AND s.enabled + WHERE (p.disabled IS NULL OR p.disabled = false) + AND t.token_hash = ? + AND t.revoked_at IS NULL + AND (t.expires_at IS NULL OR t.expires_at > now()) +), touched AS ( + UPDATE %[1]q AS t SET last_used_at = now() + FROM authenticated AS a + WHERE t.id = a.id + AND (a.last_used_at IS NULL OR a.last_used_at < now() - interval '1 minute') + RETURNING t.* +) +SELECT * FROM touched +UNION ALL +SELECT * FROM authenticated WHERE NOT EXISTS (SELECT 1 FROM touched)`, token.TableName(), (&SSOProvider{}).TableName(), (&SCIMSettings{}).TableName()), + HashSCIMToken(plaintext), + ).First(token) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, SCIMTokenNotFoundError{} + } + return nil, errors.Wrap(err, "error authenticating SCIM token") + } + return token, nil +} + +func generateSCIMToken() (string, error) { + b := make([]byte, scimTokenBytes) + if _, err := rand.Read(b); err != nil { + return "", err + } + return SCIMTokenMarker + hex.EncodeToString(b), nil +} diff --git a/internal/models/scim_token_test.go b/internal/models/scim_token_test.go new file mode 100644 index 0000000000..5a295f37cd --- /dev/null +++ b/internal/models/scim_token_test.go @@ -0,0 +1,299 @@ +package models + +import ( + "regexp" + "testing" + "time" + + "github.com/gofrs/uuid" + "github.com/stretchr/testify/require" + "github.com/stretchr/testify/suite" + "github.com/supabase/auth/internal/conf/confload" + "github.com/supabase/auth/internal/storage" + "github.com/supabase/auth/internal/storage/test" +) + +type SCIMTokenTestSuite struct { + suite.Suite + db *storage.Connection + provider *SSOProvider +} + +func TestSCIMToken(t *testing.T) { + globalConfig, err := confload.LoadGlobal(modelsTestConfig) + require.NoError(t, err) + conn, err := test.SetupDBConnection(globalConfig) + require.NoError(t, err) + ts := &SCIMTokenTestSuite{db: conn} + defer ts.db.Close() + suite.Run(t, ts) +} + +func (ts *SCIMTokenTestSuite) SetupTest() { + require.NoError(ts.T(), TruncateAll(ts.db)) + ts.provider = ts.createProvider() +} + +func (ts *SCIMTokenTestSuite) createProvider() *SSOProvider { + return createSCIMTestProvider(ts.T(), ts.db) +} + +func createSCIMTestProvider(t require.TestingT, db *storage.Connection) *SSOProvider { + provider := &SSOProvider{} + require.NoError(t, db.Create(provider)) + _, err := EnableSCIM(db, provider.ID) + require.NoError(t, err) + return provider +} + +func (ts *SCIMTokenTestSuite) createToken(expiresAt *time.Time) (*SCIMToken, string) { + token, plaintext, err := CreateSCIMToken(ts.db, ts.provider, expiresAt) + require.NoError(ts.T(), err) + return token, plaintext +} + +func (ts *SCIMTokenTestSuite) TestCreate() { + token, plaintext := ts.createToken(nil) + + require.Regexp(ts.T(), regexp.MustCompile(`^scim_[0-9a-f]{40}$`), plaintext) + require.Equal(ts.T(), plaintext[:12], token.Prefix) + require.Equal(ts.T(), HashSCIMToken(plaintext), token.TokenHash) + require.NotEqual(ts.T(), plaintext, token.TokenHash) + require.Equal(ts.T(), ts.provider.ID, token.SSOProviderID) + require.False(ts.T(), token.CreatedAt.IsZero()) + require.Nil(ts.T(), token.ExpiresAt) + require.Nil(ts.T(), token.RevokedAt) + require.Nil(ts.T(), token.LastUsedAt) +} + +func (ts *SCIMTokenTestSuite) TestCreateWithExpiry() { + expiresAt := time.Now().Add(time.Hour).UTC().Truncate(time.Microsecond) + token, _ := ts.createToken(&expiresAt) + + require.NotNil(ts.T(), token.ExpiresAt) + require.True(ts.T(), expiresAt.Equal(*token.ExpiresAt)) +} + +func (ts *SCIMTokenTestSuite) TestCreateWithPastExpiry() { + expiresAt := time.Now().Add(-time.Minute) + _, _, err := CreateSCIMToken(ts.db, ts.createProvider(), &expiresAt) + + require.ErrorIs(ts.T(), err, SCIMTokenExpiryError{}) +} + +func (ts *SCIMTokenTestSuite) TestTimestampsAreUTC() { + local := time.Local + time.Local = time.FixedZone("UTC-7", -7*60*60) + defer func() { time.Local = local }() + + expiresAt := time.Now().Add(time.Hour) + token, plaintext := ts.createToken(&expiresAt) + authenticated, err := AuthenticateSCIMToken(ts.db, plaintext) + require.NoError(ts.T(), err) + require.NoError(ts.T(), authenticated.Revoke(ts.db)) + found, err := FindSCIMTokensBySSOProvider(ts.db, ts.provider.ID) + require.NoError(ts.T(), err) + require.Len(ts.T(), found, 1) + + for _, t := range []*SCIMToken{token, authenticated, &found[0]} { + require.Equal(ts.T(), time.UTC, t.CreatedAt.Location()) + require.Equal(ts.T(), time.UTC, t.ExpiresAt.Location()) + } + for _, t := range []*SCIMToken{authenticated, &found[0]} { + require.Equal(ts.T(), time.UTC, t.LastUsedAt.Location()) + require.Equal(ts.T(), time.UTC, t.RevokedAt.Location()) + } +} + +func (ts *SCIMTokenTestSuite) TestCreateForMissingProvider() { + _, _, err := CreateSCIMToken(ts.db, &SSOProvider{ID: uuid.Must(uuid.NewV4())}, nil) + + require.Error(ts.T(), err) +} + +func (ts *SCIMTokenTestSuite) TestFindBySSOProvider() { + first, _ := ts.createToken(nil) + second, _ := ts.createToken(nil) + require.NoError(ts.T(), second.Revoke(ts.db)) + + other := ts.createProvider() + _, _, err := CreateSCIMToken(ts.db, other, nil) + require.NoError(ts.T(), err) + + tokens, err := FindSCIMTokensBySSOProvider(ts.db, ts.provider.ID) + require.NoError(ts.T(), err) + require.Len(ts.T(), tokens, 2) + require.ElementsMatch(ts.T(), []uuid.UUID{first.ID, second.ID}, []uuid.UUID{tokens[0].ID, tokens[1].ID}) +} + +func (ts *SCIMTokenTestSuite) TestFindByPrefix() { + token, _ := ts.createToken(nil) + + found, err := FindSCIMTokenByPrefix(ts.db, ts.provider.ID, token.Prefix) + require.NoError(ts.T(), err) + require.Equal(ts.T(), token.ID, found.ID) + + _, err = FindSCIMTokenByPrefix(ts.db, ts.createProvider().ID, token.Prefix) + require.True(ts.T(), IsNotFoundError(err)) + + _, err = FindSCIMTokenByPrefix(ts.db, ts.provider.ID, "scim_0000000") + require.True(ts.T(), IsNotFoundError(err)) +} + +func (ts *SCIMTokenTestSuite) TestFindByPrefixAmbiguous() { + token, _ := ts.createToken(nil) + duplicate := &SCIMToken{ + ID: uuid.Must(uuid.NewV4()), + SSOProviderID: ts.provider.ID, + TokenHash: HashSCIMToken("duplicate"), + Prefix: token.Prefix, + } + require.NoError(ts.T(), ts.db.RawQuery( + "INSERT INTO scim_tokens (id, sso_provider_id, token_hash, prefix) VALUES (?, ?, ?, ?)", + duplicate.ID, duplicate.SSOProviderID, duplicate.TokenHash, duplicate.Prefix, + ).Exec()) + + _, err := FindSCIMTokenByPrefix(ts.db, ts.provider.ID, token.Prefix) + require.Error(ts.T(), err) + require.False(ts.T(), IsNotFoundError(err)) +} + +func (ts *SCIMTokenTestSuite) TestRevoke() { + token, _ := ts.createToken(nil) + + require.NoError(ts.T(), token.Revoke(ts.db)) + require.NotNil(ts.T(), token.RevokedAt) + revokedAt := *token.RevokedAt + + require.NoError(ts.T(), token.Revoke(ts.db)) + require.True(ts.T(), revokedAt.Equal(*token.RevokedAt)) + + reloaded, err := FindSCIMTokenByPrefix(ts.db, ts.provider.ID, token.Prefix) + require.NoError(ts.T(), err) + require.True(ts.T(), revokedAt.Equal(*reloaded.RevokedAt)) +} + +func (ts *SCIMTokenTestSuite) TestAuthenticate() { + token, plaintext := ts.createToken(nil) + + authenticated, err := AuthenticateSCIMToken(ts.db, plaintext) + require.NoError(ts.T(), err) + require.Equal(ts.T(), token.ID, authenticated.ID) + require.Equal(ts.T(), ts.provider.ID, authenticated.SSOProviderID) + require.NotNil(ts.T(), authenticated.LastUsedAt) + first := *authenticated.LastUsedAt + + again, err := AuthenticateSCIMToken(ts.db, plaintext) + require.NoError(ts.T(), err) + require.True(ts.T(), first.Equal(*again.LastUsedAt)) + + require.NoError(ts.T(), ts.db.RawQuery("UPDATE scim_tokens SET last_used_at = now() - interval '2 minutes' WHERE id = ?", token.ID).Exec()) + stale, err := AuthenticateSCIMToken(ts.db, plaintext) + require.NoError(ts.T(), err) + require.True(ts.T(), stale.LastUsedAt.After(first.Add(-time.Minute))) +} + +func (ts *SCIMTokenTestSuite) TestAuthenticateRejects() { + disabled := true + + for _, tc := range []struct { + name string + setup func() string + }{ + {"unknown token", func() string { return "scim_" + "00000000000000000000000000000000000000000" }}, + {"revoked token", func() string { + token, plaintext := ts.createToken(nil) + require.NoError(ts.T(), token.Revoke(ts.db)) + return plaintext + }}, + {"expired token", func() string { + token, plaintext := ts.createToken(nil) + require.NoError(ts.T(), ts.db.RawQuery( + "UPDATE scim_tokens SET created_at = now() - interval '2 hours', expires_at = now() - interval '1 hour' WHERE id = ?", token.ID, + ).Exec()) + return plaintext + }}, + {"disabled provider", func() string { + _, plaintext := ts.createToken(nil) + ts.provider.Disabled = &disabled + require.NoError(ts.T(), ts.db.UpdateOnly(ts.provider, "disabled")) + return plaintext + }}, + {"scim never enabled", func() string { + provider := &SSOProvider{} + require.NoError(ts.T(), ts.db.Create(provider)) + _, plaintext, err := CreateSCIMToken(ts.db, provider, nil) + require.NoError(ts.T(), err) + return plaintext + }}, + {"scim disabled", func() string { + provider := ts.createProvider() + _, plaintext, err := CreateSCIMToken(ts.db, provider, nil) + require.NoError(ts.T(), err) + _, err = DisableSCIM(ts.db, provider.ID) + require.NoError(ts.T(), err) + return plaintext + }}, + } { + ts.Run(tc.name, func() { + _, err := AuthenticateSCIMToken(ts.db, tc.setup()) + require.True(ts.T(), IsNotFoundError(err), "%v", err) + }) + } +} + +func (ts *SCIMTokenTestSuite) expire(token *SCIMToken) { + require.NoError(ts.T(), ts.db.RawQuery( + "UPDATE scim_tokens SET created_at = now() - interval '2 hours', expires_at = now() - interval '1 hour' WHERE id = ?", token.ID, + ).Exec()) +} + +func (ts *SCIMTokenTestSuite) TestFindActiveBySSOProvider() { + active, _ := ts.createToken(nil) + revoked, _ := ts.createToken(nil) + require.NoError(ts.T(), revoked.Revoke(ts.db)) + expired, _ := ts.createToken(nil) + ts.expire(expired) + _, _, err := CreateSCIMToken(ts.db, ts.createProvider(), nil) + require.NoError(ts.T(), err) + + tokens, err := FindActiveSCIMTokensBySSOProvider(ts.db, ts.provider.ID) + require.NoError(ts.T(), err) + require.Len(ts.T(), tokens, 1) + require.Equal(ts.T(), active.ID, tokens[0].ID) + + tokens, err = FindActiveSCIMTokensBySSOProvider(ts.db, uuid.Must(uuid.NewV4())) + require.NoError(ts.T(), err) + require.Empty(ts.T(), tokens) +} + +func (ts *SCIMTokenTestSuite) TestRevokeBySSOProvider() { + first, _ := ts.createToken(nil) + second, _ := ts.createToken(nil) + revoked, _ := ts.createToken(nil) + require.NoError(ts.T(), revoked.Revoke(ts.db)) + expired, _ := ts.createToken(nil) + ts.expire(expired) + other, _, err := CreateSCIMToken(ts.db, ts.createProvider(), nil) + require.NoError(ts.T(), err) + + tokens, err := RevokeSCIMTokensBySSOProvider(ts.db, ts.provider.ID) + require.NoError(ts.T(), err) + require.Len(ts.T(), tokens, 2) + require.Equal(ts.T(), []uuid.UUID{first.ID, second.ID}, []uuid.UUID{tokens[0].ID, tokens[1].ID}) + for _, token := range tokens { + require.NotNil(ts.T(), token.RevokedAt) + } + + active, err := FindActiveSCIMTokensBySSOProvider(ts.db, ts.provider.ID) + require.NoError(ts.T(), err) + require.Empty(ts.T(), active) + + tokens, err = RevokeSCIMTokensBySSOProvider(ts.db, ts.provider.ID) + require.NoError(ts.T(), err) + require.Empty(ts.T(), tokens) + + still, err := FindActiveSCIMTokensBySSOProvider(ts.db, other.SSOProviderID) + require.NoError(ts.T(), err) + require.Len(ts.T(), still, 1) +} diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go new file mode 100644 index 0000000000..8f132f3cc2 --- /dev/null +++ b/internal/models/scim_user.go @@ -0,0 +1,261 @@ +package models + +import ( + "encoding/json" + "fmt" + "strings" + "time" + + "github.com/gofrs/uuid" + "github.com/pkg/errors" + "github.com/supabase/auth/internal/storage" +) + +type SCIMUser struct { + ID uuid.UUID `db:"id"` + SSOProviderID uuid.UUID `db:"sso_provider_id"` + UserID *uuid.UUID `db:"user_id"` + Resource []byte `db:"resource"` + UserName string `db:"user_name"` + ExternalID *string `db:"external_id"` + Active bool `db:"active"` + CreatedAt time.Time `db:"created_at"` + UpdatedAt time.Time `db:"updated_at"` + DeletedAt *time.Time `db:"deleted_at"` +} + +const scimUserColumns = "id, sso_provider_id, user_id, resource, user_name, external_id, active, created_at, updated_at, deleted_at" + +func (SCIMUser) TableName() string { + return "scim_users" +} + +var scimUsersTable = scimTable{ + name: SCIMUser{}.TableName(), + label: "SCIM user", + columns: scimUserColumns, + nameCol: "user_name", + live: "deleted_at IS NULL", + notFound: SCIMUserNotFoundError{}, + stale: SCIMUserStaleError{}, + conflict: SCIMUserConflictError{}, +} + +func CreateSCIMUser(tx *storage.Connection, providerID uuid.UUID, resource []byte) (*SCIMUser, error) { + return createSCIMRow[SCIMUser](tx, scimUsersTable, providerID, resource) +} + +func FindSCIMUser(tx *storage.Connection, providerID, id uuid.UUID) (*SCIMUser, error) { + return findSCIMRow[SCIMUser](tx, scimUsersTable, providerID, id, "") +} + +func FindSCIMUserForUpdate(tx *storage.Connection, providerID, id uuid.UUID) (*SCIMUser, error) { + return findSCIMRow[SCIMUser](tx, scimUsersTable, providerID, id, " FOR UPDATE") +} + +func FindUnchangedSCIMUser(tx *storage.Connection, providerID, id uuid.UUID, resource []byte, updatedAt *time.Time) (*SCIMUser, error) { + return findUnchangedSCIMRow[SCIMUser](tx, scimUsersTable, providerID, id, resource, updatedAt) +} + +func FindSCIMUsers(tx *storage.Connection, providerID uuid.UUID, query SCIMQuery) ([]SCIMUser, int, error) { + return findSCIMPage[SCIMUser](tx, scimUsersTable, providerID, query) +} + +func ReplaceSCIMUser(tx *storage.Connection, providerID, id uuid.UUID, resource []byte, updatedAt *time.Time) (*SCIMUser, error) { + user := &SCIMUser{} + err := tx.RawQuery( + fmt.Sprintf("UPDATE %q SET resource = ?::jsonb, updated_at = now() WHERE id = ? AND sso_provider_id = ? AND deleted_at IS NULL AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimUserColumns, user.TableName()), + string(resource), id, providerID, updatedAt, updatedAt, + ).First(user) + if err != nil { + return nil, scimWriteError[SCIMUser](tx, scimUsersTable, err, providerID, id, updatedAt, "replacing") + } + return user, nil +} + +func DeleteSCIMUser(tx *storage.Connection, providerID, id uuid.UUID, updatedAt *time.Time) (*SCIMUser, error) { + user := &SCIMUser{} + err := tx.RawQuery( + fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE id = ? AND sso_provider_id = ? AND deleted_at IS NULL AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimUserColumns, user.TableName()), + id, providerID, updatedAt, updatedAt, + ).First(user) + if err != nil { + return nil, scimWriteError[SCIMUser](tx, scimUsersTable, err, providerID, id, updatedAt, "deleting") + } + return user, nil +} + +func FindSCIMUserLinks(tx *storage.Connection, ids []uuid.UUID) (map[uuid.UUID]uuid.UUID, error) { + links := map[uuid.UUID]uuid.UUID{} + if len(ids) == 0 { + return links, nil + } + rows := []struct { + ID uuid.UUID `db:"id"` + UserID uuid.UUID `db:"user_id"` + }{} + if err := tx.RawQuery( + fmt.Sprintf("SELECT id, user_id FROM %q WHERE id = ANY(?::uuid[]) AND user_id IS NOT NULL", (&SCIMUser{}).TableName()), + uuidArray(ids), + ).All(&rows); err != nil { + return nil, errors.Wrap(err, "error finding SCIM user links") + } + for _, row := range rows { + links[row.ID] = row.UserID + } + return links, nil +} + +func LockUserForSCIM(tx *storage.Connection, userID uuid.UUID) error { + if err := tx.RawQuery( + fmt.Sprintf("SELECT id FROM %q WHERE id = ? FOR UPDATE", (&User{}).TableName()), + userID, + ).Exec(); err != nil { + return errors.Wrap(err, "error locking user") + } + return nil +} + +func SoftDeleteSCIMUsersByUserID(tx *storage.Connection, userID uuid.UUID) ([]SCIMUser, error) { + rows := []SCIMUser{} + if err := tx.RawQuery( + fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE user_id = ? AND deleted_at IS NULL RETURNING "+scimUserColumns, (&SCIMUser{}).TableName()), + userID, + ).All(&rows); err != nil { + return nil, errors.Wrap(err, "error deleting SCIM users by user id") + } + return rows, nil +} + +func BanDeprovisionedSCIMUsers(tx *storage.Connection, providerID uuid.UUID, until time.Time) (int, error) { + users, scimUsers := (&User{}).TableName(), (&SCIMUser{}).TableName() + count, err := tx.RawQuery( + fmt.Sprintf( + "UPDATE %[1]q u SET banned_until = ?, updated_at = now() "+ + "WHERE u.id IN (SELECT user_id FROM %[2]q WHERE sso_provider_id = ? AND (deleted_at IS NOT NULL OR NOT active)) "+ + "AND NOT EXISTS (SELECT 1 FROM %[2]q l WHERE l.user_id = u.id AND l.deleted_at IS NULL AND l.active) "+ + "AND (u.banned_until IS NULL OR u.banned_until < ?)", + users, scimUsers, + ), + until, providerID, until, + ).ExecWithCount() + if err != nil { + return 0, errors.Wrap(err, "error banning deprovisioned SCIM users") + } + return count, nil +} + +func LinkSCIMUser(tx *storage.Connection, user *SCIMUser, userID uuid.UUID) error { + if err := LockUserForSCIM(tx, userID); err != nil { + return err + } + + linked, err := tx.Q().Where("sso_provider_id = ? AND user_id = ? AND deleted_at IS NULL", user.SSOProviderID, userID).Exists(&SCIMUser{}) + if err != nil { + return errors.Wrap(err, "error finding linked SCIM user") + } + if linked { + return SCIMUserLinkedError{} + } + + if err := tx.RawQuery( + fmt.Sprintf("UPDATE %q SET user_id = ? WHERE id = ?", user.TableName()), + userID, user.ID, + ).Exec(); err != nil { + return errors.Wrap(err, "error linking SCIM user") + } + user.UserID = &userID + return nil +} + +func IsSCIMProvisioned(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { + provisioned, err := tx.Q().Where("sso_provider_id = ? AND user_id = ?", providerID, userID).Exists(&SCIMUser{}) + if err != nil { + return false, errors.Wrap(err, "error finding SCIM user") + } + return provisioned, nil +} + +func IsSCIMManaged(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { + managed, err := tx.Q().Where("sso_provider_id = ? AND user_id = ? AND deleted_at IS NULL", providerID, userID).Exists(&SCIMUser{}) + if err != nil { + return false, errors.Wrap(err, "error finding SCIM user") + } + return managed, nil +} + +func IsSCIMDeprovisioned(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { + provisioned, err := IsSCIMProvisioned(tx, providerID, userID) + if err != nil { + return false, err + } + if !provisioned { + return false, nil + } + live, err := tx.Q().Where("sso_provider_id = ? AND user_id = ? AND deleted_at IS NULL AND active", providerID, userID).Exists(&SCIMUser{}) + if err != nil { + return false, errors.Wrap(err, "error finding live SCIM user") + } + return !live, nil +} + +func IsSCIMUserDeprovisionedForUpdate(tx *storage.Connection, userID uuid.UUID) (bool, error) { + if err := tx.RawQuery( + fmt.Sprintf("SELECT id FROM %q WHERE id = ? FOR NO KEY UPDATE", (&User{}).TableName()), + userID, + ).Exec(); err != nil { + return false, errors.Wrap(err, "error locking user") + } + rows := []struct { + Active bool `db:"active"` + DeletedAt *time.Time `db:"deleted_at"` + }{} + if err := tx.RawQuery( + fmt.Sprintf("SELECT active, deleted_at FROM %q WHERE user_id = ?", (&SCIMUser{}).TableName()), + userID, + ).All(&rows); err != nil { + return false, errors.Wrap(err, "error finding SCIM users") + } + for _, row := range rows { + if row.DeletedAt == nil && row.Active { + return false, nil + } + } + return len(rows) > 0, nil +} + +func RenameSCIMIdentity(tx *storage.Connection, userID uuid.UUID, provider, from, to string, data map[string]any) error { + encoded, err := json.Marshal(data) + if err != nil { + return errors.Wrap(err, "error encoding identity data") + } + table := (&Identity{}).TableName() + if err := tx.RawQuery( + fmt.Sprintf("DELETE FROM %[1]q WHERE user_id = ? AND provider = ? AND provider_id <> ? AND (lower(provider_id) = lower(?) OR provider_id = ?) AND EXISTS (SELECT 1 FROM %[1]q WHERE user_id = ? AND provider = ? AND provider_id = ?)", table), + userID, provider, from, from, to, userID, provider, from, + ).Exec(); err != nil { + return errors.Wrap(err, "error removing stale SCIM identities") + } + count, err := tx.RawQuery( + fmt.Sprintf("UPDATE %q SET provider_id = ?, identity_data = identity_data || ?::jsonb, updated_at = now() WHERE user_id = ? AND provider = ? AND provider_id = ?", table), + to, string(encoded), userID, provider, from, + ).ExecWithCount() + if err != nil { + if isUniqueViolation(err) { + return SCIMUserConflictError{} + } + return errors.Wrap(err, "error renaming SCIM identity") + } + if count == 0 { + return SCIMIdentityNotFoundError{} + } + return nil +} + +func LockAccountLinking(tx *storage.Connection, providerType, email string) error { + key := providerType + "|" + strings.ToLower(email) + if err := tx.RawQuery("SELECT pg_advisory_xact_lock(hashtextextended(?, 0))", key).Exec(); err != nil { + return errors.Wrap(err, "error locking account linking") + } + return nil +} diff --git a/internal/tokens/service.go b/internal/tokens/service.go index 6cf445cf11..1b868526d0 100644 --- a/internal/tokens/service.go +++ b/internal/tokens/service.go @@ -887,6 +887,16 @@ func (s *Service) IssueRefreshToken(r *http.Request, responseHeaders http.Header err := conn.Transaction(func(tx *storage.Connection) error { var terr error + if config.SSO.SCIM.Enabled && user.IsSSOUser { + deprovisioned, terr := models.IsSCIMUserDeprovisionedForUpdate(tx, user.ID) + if terr != nil { + return apierrors.NewInternalServerError("Database error checking SCIM user").WithInternalError(terr) + } + if deprovisioned { + return apierrors.NewForbiddenError(apierrors.ErrorCodeUserBanned, "User is banned") + } + } + if config.Security.RefreshTokenAlgorithmVersion == 2 { session, terr := models.NewSession(user.ID, grantParams.FactorID) if terr != nil { diff --git a/migrations/20260929000000_add_scim_groups.up.sql b/migrations/20260929000000_add_scim_groups.up.sql new file mode 100644 index 0000000000..6297142406 --- /dev/null +++ b/migrations/20260929000000_add_scim_groups.up.sql @@ -0,0 +1,40 @@ +/* auth_migration: 20260929000000 */ +create table if not exists {{ index .Options "Namespace" }}.scim_groups ( + id uuid not null, + sso_provider_id uuid not null references {{ index .Options "Namespace" }}.sso_providers (id) on delete cascade, + resource jsonb not null, + display_name text not null generated always as (lower(resource->>'displayName')) stored, + external_id text generated always as (resource->>'externalId') stored, + created_at timestamptz not null default now(), + updated_at timestamptz not null default now(), + constraint scim_groups_pkey primary key (id) +); + +/* auth_migration: 20260929000000 */ +create index if not exists scim_groups_display_name_idx + on {{ index .Options "Namespace" }}.scim_groups (sso_provider_id, display_name collate "C", id); + +/* auth_migration: 20260929000000 */ +create unique index if not exists scim_groups_external_id_key + on {{ index .Options "Namespace" }}.scim_groups (sso_provider_id, external_id) + where external_id is not null; + +/* auth_migration: 20260929000000 */ +create index if not exists scim_groups_created_at_idx + on {{ index .Options "Namespace" }}.scim_groups (sso_provider_id, created_at, id); + +/* auth_migration: 20260929000000 */ +create index if not exists scim_groups_updated_at_idx + on {{ index .Options "Namespace" }}.scim_groups (sso_provider_id, updated_at, id); + +/* auth_migration: 20260929000000 */ +create table if not exists {{ index .Options "Namespace" }}.scim_group_members ( + group_id uuid not null references {{ index .Options "Namespace" }}.scim_groups (id) on delete cascade, + scim_user_id uuid not null references {{ index .Options "Namespace" }}.scim_users (id) on delete cascade, + created_at timestamptz not null default now(), + constraint scim_group_members_pkey primary key (group_id, scim_user_id) +); + +/* auth_migration: 20260929000000 */ +create index if not exists scim_group_members_scim_user_id_idx + on {{ index .Options "Namespace" }}.scim_group_members (scim_user_id); diff --git a/migrations/20260930000000_add_scim_settings.up.sql b/migrations/20260930000000_add_scim_settings.up.sql new file mode 100644 index 0000000000..1b6843a9ec --- /dev/null +++ b/migrations/20260930000000_add_scim_settings.up.sql @@ -0,0 +1,8 @@ +/* auth_migration: 20260930000000 */ +create table if not exists {{ index .Options "Namespace" }}.scim_settings ( + sso_provider_id uuid not null references {{ index .Options "Namespace" }}.sso_providers (id) on delete cascade, + enabled boolean not null default false, + created_at timestamptz not null default now(), + updated_at timestamptz not null default now(), + constraint scim_settings_pkey primary key (sso_provider_id) +); From 88a1a1bd6862311019a775280f9c2c2b24f9d7ac Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 17:04:23 -0600 Subject: [PATCH 02/88] chore(scim): use scim-go v0.8.0 remove-with-value check --- go.mod | 2 +- go.sum | 4 ++-- internal/api/scim.go | 28 +--------------------------- 3 files changed, 4 insertions(+), 30 deletions(-) diff --git a/go.mod b/go.mod index 2357651739..5579bd69ad 100644 --- a/go.mod +++ b/go.mod @@ -163,7 +163,7 @@ require ( github.com/spf13/cobra v1.8.1 github.com/standard-webhooks/standard-webhooks/libraries v0.0.0-20240303152453-e0e82adf1721 github.com/stretchr/testify v1.12.1 - github.com/supabase-community/scim-go v0.7.5 + github.com/supabase-community/scim-go v0.8.0 github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 github.com/xeipuuv/gojsonschema v1.2.0 go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.64.0 diff --git a/go.sum b/go.sum index ac7525fd17..b0f8ea7ef3 100644 --- a/go.sum +++ b/go.sum @@ -498,8 +498,8 @@ github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= -github.com/supabase-community/scim-go v0.7.5 h1:bhT3BdYazeGMW2AaiHT06xiWO5weCIr2XV3MJ64v+Ko= -github.com/supabase-community/scim-go v0.7.5/go.mod h1:oEMij9JuKtAl0wl0jeyIHDKqHvPpUmTXegKeCBKyxXw= +github.com/supabase-community/scim-go v0.8.0 h1:eeLUw37qFBerMY86JLfMS6mV1yFVxUj9EFleEGpoXEw= +github.com/supabase-community/scim-go v0.8.0/go.mod h1:oEMij9JuKtAl0wl0jeyIHDKqHvPpUmTXegKeCBKyxXw= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 h1:VDuRtwen5Z7QQ5ctuHUse4wAv/JozkKZkdic5vUV4Lg= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869/go.mod h1:eHX5nlSMSnyPjUrbYzeqrA8snCe2SKyfizKjU3dkfOw= github.com/supranational/blst v0.3.16-0.20250831170142-f48500c1fdbe h1:nbdqkIGOGfUAD54q1s2YBcBz/WcsxCO9HUQ4aGV5hUw= diff --git a/internal/api/scim.go b/internal/api/scim.go index 7d65dd1537..5347bb1e5b 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -1,12 +1,10 @@ package api import ( - "bytes" "context" "encoding/json" "errors" "fmt" - "io" "net/http" "strconv" "strings" @@ -39,7 +37,7 @@ var errMissingSSOProvider = errors.New("scim: request has no SSO provider") func newSCIMServer(config *conf.GlobalConfiguration, validate server.TokenValidator, limit func(http.Handler) http.Handler, users server.Repository[*core.User], groups server.Repository[*core.Group]) *server.Server { requireToken := server.RequireBearerToken(validate) authenticate := func(next http.Handler) http.Handler { - next = scimRejectRemoveWithValue(scimRememberGetQuery(next)) + next = scimRememberGetQuery(next) if limit != nil { next = limit(next) } @@ -97,30 +95,6 @@ func scimLogError(r *http.Request, err error) { observability.GetLogEntry(r).Entry.WithError(err).Error("scim: request failed") } -// RFC 7644 Section 3.5.2.2 defines no "value" for "remove"; refuse it rather than remove every value at "path". -func scimRejectRemoveWithValue(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.Method != http.MethodPatch { - next.ServeHTTP(w, r) - return - } - var read bytes.Buffer - req, err := protocol.DecodePatchRequest(io.TeeReader(r.Body, &read)) - if err != nil { - _ = protocol.SendError(w, err) - return - } - for _, op := range req.Operations { - if strings.EqualFold(string(op.Op), "remove") && len(op.Value) > 0 && string(bytes.TrimSpace(op.Value)) != "null" { - _ = protocol.SendError(w, scimerrors.ErrInvalidSyntax(`"remove" does not take a "value"`)) - return - } - } - r.Body = io.NopCloser(io.MultiReader(&read, r.Body)) - next.ServeHTTP(w, r) - }) -} - func scimRememberGetQuery(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Method == http.MethodGet { From 49684a2b7baf4c18aff9cc8e3688669c6ad55766 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 17:05:34 -0600 Subject: [PATCH 03/88] chore(scim): read the request projection from scim-go context --- internal/api/context.go | 1 - internal/api/scim.go | 23 ----------------------- internal/api/scim_groups.go | 14 +++----------- internal/api/scim_test.go | 37 ------------------------------------- internal/api/scim_users.go | 14 +++----------- 5 files changed, 6 insertions(+), 83 deletions(-) diff --git a/internal/api/context.go b/internal/api/context.go index ee1a704796..2179f1b19d 100644 --- a/internal/api/context.go +++ b/internal/api/context.go @@ -31,7 +31,6 @@ var ( oauthClientStateKey = ctxkey.New[uuid.UUID]("oauth_client_state_id") flowStateContextKey = ctxkey.New[*models.FlowState]("flow_state") scimRequestKey = ctxkey.New[*http.Request]("scim_request") - scimGetQueryKey = ctxkey.New[url.Values]("scim_get_query") scimSSOProviderIDKey = ctxkey.New[uuid.UUID]("scim_sso_provider_id") scimTokenPrefixKey = ctxkey.New[string]("scim_token_prefix") ) diff --git a/internal/api/scim.go b/internal/api/scim.go index 5347bb1e5b..6820f2ef94 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -37,7 +37,6 @@ var errMissingSSOProvider = errors.New("scim: request has no SSO provider") func newSCIMServer(config *conf.GlobalConfiguration, validate server.TokenValidator, limit func(http.Handler) http.Handler, users server.Repository[*core.User], groups server.Repository[*core.Group]) *server.Server { requireToken := server.RequireBearerToken(validate) authenticate := func(next http.Handler) http.Handler { - next = scimRememberGetQuery(next) if limit != nil { next = limit(next) } @@ -95,15 +94,6 @@ func scimLogError(r *http.Request, err error) { observability.GetLogEntry(r).Entry.WithError(err).Error("scim: request failed") } -func scimRememberGetQuery(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.Method == http.MethodGet { - r = r.WithContext(scimGetQueryKey.WithValue(r.Context(), r.URL.Query())) - } - next.ServeHTTP(w, r) - }) -} - func newSCIMTokenValidator(db *storage.Connection) server.TokenValidator { return func(ctx context.Context, candidate string) (context.Context, error) { token, err := models.AuthenticateSCIMToken(db.WithContext(ctx), candidate) @@ -210,19 +200,6 @@ func scimSearch(query *protocol.SearchRequest, schemas core.Schemas, name string return search, nil } -func scimReturns(projection protocol.Projection, name string) bool { - document, err := json.Marshal(projection.Of(map[string]any{name: []any{map[string]any{"value": name}}})) - if err != nil { - return true - } - var projected map[string]any - if err := json.Unmarshal(document, &projected); err != nil { - return true - } - _, ok := projected[name] - return ok -} - func scimEncode(resource core.Resource, drop ...string) ([]byte, error) { fields, err := core.NewObject(resource) if err != nil { diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index 0bc224d1e1..0c8b64fc05 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -32,11 +32,7 @@ func (s *scimGroups) List(ctx context.Context, query *protocol.SearchRequest) ([ if err != nil { return nil, 0, err } - projection, err := query.Projection(scimGroupSchemas) - if err != nil { - projection = protocol.Projection{} - } - groups, err := s.render(db, providerID, rows, projection) + groups, err := s.render(db, providerID, rows, protocol.ProjectionFrom(ctx)) if err != nil { return nil, 0, err } @@ -53,11 +49,7 @@ func (s *scimGroups) Get(ctx context.Context, id string) (*core.Group, error) { if err != nil { return nil, scimTranslate(err) } - projection, err := protocol.ParseProjection(scimGetQueryKey.Value(ctx), scimGroupSchemas) - if err != nil { - projection = protocol.Projection{} - } - return s.renderOne(db, providerID, row, projection) + return s.renderOne(db, providerID, row, protocol.ProjectionFrom(ctx)) } func (s *scimGroups) Create(ctx context.Context, group *core.Group) (*core.Group, error) { @@ -161,7 +153,7 @@ func (s *scimGroups) save(ctx context.Context, providerID uuid.UUID, action mode func (s *scimGroups) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMGroup, projection protocol.Projection) ([]*core.Group, error) { memberships := []models.SCIMGroupMembership{} - if scimReturns(projection, "members") { + if projection.Returns("members") { ids := make([]uuid.UUID, len(rows)) for i, row := range rows { ids[i] = row.ID diff --git a/internal/api/scim_test.go b/internal/api/scim_test.go index 753467aa9f..3ddfb2b1e9 100644 --- a/internal/api/scim_test.go +++ b/internal/api/scim_test.go @@ -469,43 +469,6 @@ func TestSCIMServer(t *testing.T) { }) } -func TestSCIMReturns(t *testing.T) { - schemas := scimGroupSchemas - for query, want := range map[string]bool{ - "": true, - "excludedAttributes=members": false, - "excludedAttributes=MEMBERS": false, - "excludedAttributes=urn:ietf:params:scim:schemas:core:2.0:Group:members": false, - "excludedAttributes=displayName": true, - "excludedAttributes=members.display": true, - "attributes=displayName": false, - "attributes=members": true, - "attributes=members.value": true, - "attributes=urn:ietf:params:scim:schemas:core:2.0:Group:Members": true, - } { - values, err := url.ParseQuery(query) - require.NoError(t, err) - projection, err := scimProtocol.ParseProjection(values, schemas) - require.NoError(t, err, query) - require.Equal(t, want, scimReturns(projection, "members"), query) - } - require.True(t, scimReturns(scimProtocol.Projection{}, "members")) -} - -func TestSCIMRememberGetQuery(t *testing.T) { - for method, want := range map[string]url.Values{ - http.MethodGet: {"excludedAttributes": {"members"}}, - http.MethodPatch: nil, - http.MethodPut: nil, - } { - var got url.Values - scimRememberGetQuery(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) { - got = scimGetQueryKey.Value(r.Context()) - })).ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(method, "/Groups/x?excludedAttributes=members", nil)) - require.Equal(t, want, got, method) - } -} - func TestSCIMUserFields(t *testing.T) { t.Run("has no email without emails", func(t *testing.T) { require.Empty(t, scimPrimaryEmail(nil)) diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index 8cea8ff71e..9ffb4af3b6 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -37,11 +37,7 @@ func (s *scimUsers) List(ctx context.Context, query *protocol.SearchRequest) ([] if err != nil { return nil, 0, err } - projection, err := query.Projection(scimUserSchemas) - if err != nil { - projection = protocol.Projection{} - } - users, err := s.render(db, providerID, rows, projection) + users, err := s.render(db, providerID, rows, protocol.ProjectionFrom(ctx)) if err != nil { return nil, 0, err } @@ -58,11 +54,7 @@ func (s *scimUsers) Get(ctx context.Context, id string) (*core.User, error) { if err != nil { return nil, scimTranslate(err) } - projection, err := protocol.ParseProjection(scimGetQueryKey.Value(ctx), scimUserSchemas) - if err != nil { - projection = protocol.Projection{} - } - return s.renderOne(db, providerID, row, projection) + return s.renderOne(db, providerID, row, protocol.ProjectionFrom(ctx)) } func (s *scimUsers) Create(ctx context.Context, user *core.User) (*core.User, error) { @@ -212,7 +204,7 @@ func (s *scimUsers) Delete(ctx context.Context, id, version string) error { func (s *scimUsers) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMUser, projection protocol.Projection) ([]*core.User, error) { memberships := []models.SCIMGroupMembership{} - if scimReturns(projection, "groups") { + if projection.Returns("groups") { ids := make([]uuid.UUID, len(rows)) for i, row := range rows { ids[i] = row.ID From 878794ced76db5acc15319b5451e57f380f87379 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 17:06:31 -0600 Subject: [PATCH 04/88] chore(scim): rely on scim-go canonical values for members.type --- internal/api/scim_groups.go | 5 ----- 1 file changed, 5 deletions(-) diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index 0c8b64fc05..d0e066afc2 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -4,12 +4,10 @@ import ( "context" "encoding/json" "net/http" - "strings" "github.com/gofrs/uuid" "github.com/supabase-community/scim-go/pkg/core" "github.com/supabase-community/scim-go/pkg/protocol" - "github.com/supabase-community/scim-go/pkg/scimerrors" "github.com/supabase/auth/internal/models" "github.com/supabase/auth/internal/storage" ) @@ -237,9 +235,6 @@ func (s *scimGroups) auditMembers(tx *storage.Connection, r *http.Request, row * func scimMemberIDs(members []core.Member) ([]uuid.UUID, error) { ids := make([]uuid.UUID, 0, len(members)) for _, member := range members { - if member.Type != "" && !strings.EqualFold(string(member.Type), scimResourceTypeUser) { - return nil, scimerrors.ErrInvalidValue(`nested groups are not supported; "members.type" must be "User"`) - } id, err := uuid.FromString(member.Value) if err != nil { return nil, errSCIMMemberNotFound() From df4b38e3996e375d01484e9f4389d8c1dcc4cfc2 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 20:17:01 -0600 Subject: [PATCH 05/88] chore(scim): test that SCIM users cannot register passkeys or sign in with an old email --- internal/api/scim_link_test.go | 86 ++++++++++++++++++++++++++++++++++ 1 file changed, 86 insertions(+) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index bed0cd3789..b16dd948cb 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -15,6 +15,7 @@ import ( "github.com/supabase-community/scim-go/pkg/protocol" "github.com/supabase/auth/internal/api/apierrors" "github.com/supabase/auth/internal/api/provider" + "github.com/supabase/auth/internal/conf" "github.com/supabase/auth/internal/models" "github.com/supabase/auth/internal/storage" ) @@ -93,6 +94,91 @@ func (ts *SCIMUsersTestSuite) TestCreateLinksByEmailWithinProvider() { require.Len(ts.T(), ts.identities(user), 2) } +func (ts *SCIMUsersTestSuite) passkeyRegistrationOptions(user *models.User) int { + passkey, webauthn := ts.API.config.Passkey, ts.API.config.WebAuthn + defer func() { ts.API.config.Passkey, ts.API.config.WebAuthn = passkey, webauthn }() + ts.API.config.Passkey.Enabled = true + ts.API.config.WebAuthn = conf.WebAuthnConfiguration{ + RPID: "localhost", + RPDisplayName: "Test App", + RPOrigins: []string{"http://localhost:3000"}, + ChallengeExpiryDuration: 5 * time.Minute, + } + + session, err := models.NewSession(user.ID, nil) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Create(session)) + token, _, err := ts.API.generateAccessToken(httptest.NewRequest(http.MethodPost, "/passkeys", nil), ts.API.db, user, &session.ID, models.PasswordGrant) + require.NoError(ts.T(), err) + + r := httptest.NewRequest(http.MethodPost, "/passkeys/registration/options", nil) + r.Header.Set("Authorization", "Bearer "+token) + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + return w.Code +} + +func (ts *SCIMUsersTestSuite) TestPasswordUserWithSameEmailIsNeverLinked() { + password, err := models.NewUser("", "alice@example.com", "", ts.API.config.JWT.Aud, nil) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Create(password)) + + user := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + + require.NotEqual(ts.T(), password.ID, user.ID) + require.True(ts.T(), user.IsSSOUser) + require.Equal(ts.T(), http.StatusUnprocessableEntity, ts.passkeyRegistrationOptions(user)) + reloaded, err := models.FindUserByID(ts.API.db, password.ID) + require.NoError(ts.T(), err) + require.False(ts.T(), reloaded.IsSSOUser) + require.Len(ts.T(), ts.identities(reloaded), 0) +} + +func (ts *SCIMUsersTestSuite) TestLinkAccountKeepsUserSSO() { + password, err := models.NewUser("", "alice@example.com", "", ts.API.config.JWT.Aud, nil) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Create(password)) + existing := ts.ssoUser(ts.A, "saml-name-id", "alice@example.com") + + user := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + + require.Equal(ts.T(), existing.ID, user.ID) + require.Len(ts.T(), ts.identities(user), 2) + require.True(ts.T(), user.IsSSOUser) + require.Equal(ts.T(), http.StatusUnprocessableEntity, ts.passkeyRegistrationOptions(user)) + require.Len(ts.T(), ts.identities(password), 0) +} + +func (ts *SCIMUsersTestSuite) TestOldEmailCannotSignInAfterEmailChange() { + id := ts.create(ts.TokenA, oktaUser) + w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, strings.Replace(oktaUser, `"value": "alice@example.com"`, `"value": "alice.smith@example.com"`, 1)) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + user := ts.linkedUser(id) + require.Equal(ts.T(), "alice@example.com", user.GetEmail()) + + for _, req := range []struct{ path, body string }{ + {"/recover", `{"email":"alice@example.com"}`}, + {"/otp", `{"email":"alice@example.com","create_user":false}`}, + {"/magiclink", `{"email":"alice@example.com"}`}, + {"/token?grant_type=password", `{"email":"alice@example.com","password":"hunter2hunter2"}`}, + } { + r := httptest.NewRequest(http.MethodPost, req.path, strings.NewReader(req.body)) + r.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + require.NotContains(ts.T(), w.Body.String(), "access_token", req.path) + + reloaded := ts.linkedUser(id) + require.Nil(ts.T(), reloaded.RecoverySentAt, req.path) + require.Empty(ts.T(), reloaded.RecoveryToken, req.path) + require.Empty(ts.T(), reloaded.ConfirmationToken, req.path) + count, err := ts.API.db.Q().Where("user_id = ?", user.ID).Count(&models.OneTimeToken{}) + require.NoError(ts.T(), err) + require.Zero(ts.T(), count, req.path) + require.Zero(ts.T(), ts.sessions(user), req.path) + } +} + func (ts *SCIMUsersTestSuite) TestCreateInactiveLogsOutWithoutBanning() { existing := ts.ssoUser(ts.A, "Alice@Example.com", "alice@example.com") ts.session(existing) From 29e6b50f0522c78e60051bd3894346a16d6c13ab Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 20:32:25 -0600 Subject: [PATCH 06/88] chore(scim): test that removing emails keeps the user's email --- internal/api/scim_link_test.go | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index b16dd948cb..10ef7cfb66 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -179,6 +179,25 @@ func (ts *SCIMUsersTestSuite) TestOldEmailCannotSignInAfterEmailChange() { } } +func (ts *SCIMUsersTestSuite) TestRemovingEmailsKeepsUserEmail() { + id := ts.create(ts.TokenA, strings.Replace(oktaUser, `"userName": "Alice@Example.com"`, `"userName": "alice.smith"`, 1)) + user := ts.linkedUser(id) + + for _, userName := range []string{"alice.smith", "asmith"} { + w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"`+userName+`","active":true}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + + require.Equal(ts.T(), "alice@example.com", ts.linkedUser(id).GetEmail(), userName) + identity, err := models.FindIdentityByIdAndProvider(ts.API.db, userName, "sso:"+ts.A.ID.String()) + require.NoError(ts.T(), err, userName) + require.Equal(ts.T(), "alice@example.com", identity.IdentityData["email"], userName) + + signedIn, err := ts.samlLogin(ts.A, "saml-name-id-"+userName, "alice@example.com") + require.NoError(ts.T(), err, userName) + require.Equal(ts.T(), user.ID, signedIn.ID, userName) + } +} + func (ts *SCIMUsersTestSuite) TestCreateInactiveLogsOutWithoutBanning() { existing := ts.ssoUser(ts.A, "Alice@Example.com", "alice@example.com") ts.session(existing) From 644ef6be0ae9a027dd13aa42b7b3990f4e39f37b Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 20:39:16 -0600 Subject: [PATCH 07/88] fix(scim): validate emails and fall back to an email userName --- internal/api/scim.go | 6 ++++- internal/api/scim_link_test.go | 30 ++++++++++++++++++++++++ internal/api/scim_users.go | 42 ++++++++++++++++++++++++++++------ 3 files changed, 70 insertions(+), 8 deletions(-) diff --git a/internal/api/scim.go b/internal/api/scim.go index 6820f2ef94..7675f74b3f 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -299,7 +299,11 @@ func errSCIMMemberNotFound() error { } func errSCIMEmailRequired() error { - return scimerrors.ErrInvalidValue(`"emails" is required`) + return scimerrors.ErrInvalidValue(`"emails" or an email address "userName" is required`) +} + +func errSCIMEmailInvalid() error { + return scimerrors.ErrInvalidValue(`"emails" value must be an email address`) } func errSCIMTooManyRequests() error { diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index 10ef7cfb66..d36319293a 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -223,6 +223,36 @@ func (ts *SCIMUsersTestSuite) TestCreateRequiresEmail() { require.Equal(ts.T(), "invalidValue", body["scimType"]) } +func (ts *SCIMUsersTestSuite) TestCreateFallsBackToEmailUserName() { + id := ts.create(ts.TokenA, `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"Alice@Example.com"}`) + user := ts.linkedUser(id) + require.Equal(ts.T(), "alice@example.com", user.GetEmail()) + + signedIn, err := ts.samlLogin(ts.A, "saml-name-id", "alice@example.com") + require.NoError(ts.T(), err) + require.Equal(ts.T(), user.ID, signedIn.ID) +} + +func (ts *SCIMUsersTestSuite) TestRejectsInvalidEmailsValue() { + invalid := strings.Replace(oktaUser, `"value": "alice@example.com"`, `"value": "not-an-email"`, 1) + w, body := ts.do(ts.TokenA, http.MethodPost, "/Users", invalid) + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + require.Equal(ts.T(), "invalidValue", body["scimType"]) + + id := ts.create(ts.TokenA, oktaUser) + w, body = ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, invalid) + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + require.Equal(ts.T(), "invalidValue", body["scimType"]) + + w, body = ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id, `{ + "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + "Operations": [{"op": "replace", "path": "emails", "value": [{"value": "not-an-email", "primary": true}]}] + }`) + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + require.Equal(ts.T(), "invalidValue", body["scimType"]) + require.Equal(ts.T(), "alice@example.com", ts.linkedUser(id).GetEmail()) +} + func (ts *SCIMUsersTestSuite) TestCreateKeepsAdminBan() { existing := ts.ssoUser(ts.A, "Alice@Example.com", "alice@example.com") require.NoError(ts.T(), existing.Ban(ts.API.db, time.Hour)) diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index 9ffb4af3b6..5d99f52d7d 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -7,6 +7,7 @@ import ( "net/http" "time" + "github.com/badoux/checkmail" "github.com/gofrs/uuid" "github.com/sirupsen/logrus" "github.com/supabase-community/scim-go/pkg/core" @@ -66,7 +67,10 @@ func (s *scimUsers) Create(ctx context.Context, user *core.User) (*core.User, er if err != nil { return nil, err } - if scimPrimaryEmail(user.Emails) == "" { + if err := scimValidateEmails(user); err != nil { + return nil, err + } + if scimEmail(user) == "" { return nil, errSCIMEmailRequired() } r, err := scimRequest(ctx) @@ -81,7 +85,7 @@ func (s *scimUsers) Create(ctx context.Context, user *core.User) (*core.User, er var row *models.SCIMUser var created *models.User err = db.Transaction(func(tx *storage.Connection) error { - if terr := models.LockAccountLinking(tx, "sso:"+providerID.String(), scimPrimaryEmail(user.Emails)); terr != nil { + if terr := models.LockAccountLinking(tx, "sso:"+providerID.String(), scimEmail(user)); terr != nil { return terr } var terr error @@ -109,7 +113,10 @@ func (s *scimUsers) Replace(ctx context.Context, user *core.User) (*core.User, e if err != nil { return nil, err } - email := scimPrimaryEmail(user.Emails) + if err := scimValidateEmails(user); err != nil { + return nil, err + } + email := scimEmail(user) r, err := scimRequest(ctx) if err != nil { return nil, err @@ -254,7 +261,7 @@ func (s *scimUsers) renderOne(tx *storage.Connection, providerID uuid.UUID, row func (s *scimUsers) sync(tx *storage.Connection, providerID uuid.UUID, old, row *models.SCIMUser, user *core.User) (*models.User, error) { if old.UserID == nil { - if scimPrimaryEmail(user.Emails) == "" { + if scimEmail(user) == "" { return nil, errSCIMEmailRequired() } return s.linkNew(tx, row, user) @@ -266,7 +273,7 @@ func (s *scimUsers) sync(tx *storage.Connection, providerID uuid.UUID, old, row } if from := scimUserName(old.Resource); from != user.UserName { data := map[string]any{"sub": user.UserName} - if email := scimPrimaryEmail(user.Emails); email != "" { + if email := scimEmail(user); email != "" { data["email"] = email } err := models.RenameSCIMIdentity(tx, linked.ID, "sso:"+providerID.String(), from, user.UserName, data) @@ -360,7 +367,7 @@ func (s *scimUsers) afterCreate(r *http.Request, db *storage.Connection, user *m } func (s *scimUsers) decide(conn *storage.Connection, providerType string, user *core.User) (models.AccountLinkingResult, error) { - emails := []provider.Email{{Email: scimPrimaryEmail(user.Emails), Verified: true, Primary: true}} + emails := []provider.Email{{Email: scimEmail(user), Verified: true, Primary: true}} return models.DetermineAccountLinking(conn, s.api.config, emails, s.api.config.JWT.Aud, providerType, user.UserName) } @@ -417,6 +424,27 @@ func scimUserResource(user *core.User) ([]byte, error) { return scimEncode(user, "id", "meta", "password", "groups") } +func scimEmail(user *core.User) string { + if email := scimPrimaryEmail(user.Emails); email != "" { + return email + } + if isEmailAddress(user.UserName) { + return user.UserName + } + return "" +} + +func scimValidateEmails(user *core.User) error { + if email := scimPrimaryEmail(user.Emails); email != "" && !isEmailAddress(email) { + return errSCIMEmailInvalid() + } + return nil +} + +func isEmailAddress(value string) bool { + return len(value) <= 255 && checkmail.ValidateFormat(value) == nil +} + func scimPrimaryEmail(emails []core.Email) string { for _, email := range emails { if email.Primary != nil && *email.Primary { @@ -432,7 +460,7 @@ func scimPrimaryEmail(emails []core.Email) string { func scimIdentityData(user *core.User) map[string]any { return map[string]any{ "sub": user.UserName, - "email": scimPrimaryEmail(user.Emails), + "email": scimEmail(user), "email_verified": true, } } From 6b46026696e9bb5ae19fea22691e894e6c679d96 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 20:47:06 -0600 Subject: [PATCH 08/88] chore(scim): parse the SSO provider id from an identity provider in one place --- internal/api/external.go | 12 ++++-------- internal/api/identity.go | 23 ++++++++++------------- internal/models/identity.go | 13 +++++++++++++ internal/models/identity_test.go | 25 +++++++++++++++++++++++++ 4 files changed, 52 insertions(+), 21 deletions(-) diff --git a/internal/api/external.go b/internal/api/external.go index 63207d5187..0fdb5f1a80 100644 --- a/internal/api/external.go +++ b/internal/api/external.go @@ -311,8 +311,11 @@ func (a *API) createAccountFromExternalIdentity(tx *storage.Connection, r *http. identityData = structs.Map(userData.Metadata) } - id, isSSO := strings.CutPrefix(providerType, "sso:") + ssoProviderID, isSSO, perr := models.SSOProviderID(providerType) scimOn := isSSO && config.SSO.SCIM.Enabled + if scimOn && perr != nil { + return 0, nil, apierrors.NewInternalServerError("Invalid SSO provider id in provider type").WithInternalError(perr) + } if scimOn && userData.Metadata.Email != "" { if terr := models.LockAccountLinking(tx, providerType, userData.Metadata.Email); terr != nil { @@ -325,13 +328,6 @@ func (a *API) createAccountFromExternalIdentity(tx *storage.Connection, r *http. return 0, nil, terr } - var ssoProviderID uuid.UUID - if scimOn { - if ssoProviderID, terr = uuid.FromString(id); terr != nil { - return 0, nil, apierrors.NewInternalServerError("Invalid SSO provider id in provider type").WithInternalError(terr) - } - } - switch decision.Decision { case models.LinkAccount: user = decision.User diff --git a/internal/api/identity.go b/internal/api/identity.go index 8ddb8a8771..68a37ac1b9 100644 --- a/internal/api/identity.go +++ b/internal/api/identity.go @@ -3,7 +3,6 @@ package api import ( "context" "net/http" - "strings" "github.com/fatih/structs" "github.com/go-chi/chi/v5" @@ -55,18 +54,16 @@ func (a *API) DeleteIdentity(w http.ResponseWriter, r *http.Request) error { provider := identityToBeDeleted.Provider recipientEmail := user.GetEmail() err = db.Transaction(func(tx *storage.Connection) error { - if id, ok := strings.CutPrefix(identityToBeDeleted.Provider, "sso:"); ok && a.config.SSO.SCIM.Enabled { - if providerID, perr := uuid.FromString(id); perr == nil { - if terr := models.LockUserForSCIM(tx, user.ID); terr != nil { - return apierrors.NewInternalServerError("Database error locking user").WithInternalError(terr) - } - managed, terr := models.IsSCIMManaged(tx, providerID, user.ID) - if terr != nil { - return apierrors.NewInternalServerError("Database error finding SCIM user").WithInternalError(terr) - } - if managed { - return apierrors.NewUnprocessableEntityError(apierrors.ErrorCodeUserSSOManaged, "Identity is managed by SCIM provisioning") - } + if providerID, ok, perr := identityToBeDeleted.SSOProviderID(); ok && perr == nil && a.config.SSO.SCIM.Enabled { + if terr := models.LockUserForSCIM(tx, user.ID); terr != nil { + return apierrors.NewInternalServerError("Database error locking user").WithInternalError(terr) + } + managed, terr := models.IsSCIMManaged(tx, providerID, user.ID) + if terr != nil { + return apierrors.NewInternalServerError("Database error finding SCIM user").WithInternalError(terr) + } + if managed { + return apierrors.NewUnprocessableEntityError(apierrors.ErrorCodeUserSSOManaged, "Identity is managed by SCIM provisioning") } } if terr := models.NewAuditLogEntry(config.AuditLog, r, tx, user, models.IdentityUnlinkAction, utilities.GetIPAddress(r), map[string]any{ diff --git a/internal/models/identity.go b/internal/models/identity.go index 1f5ee5f853..f7113596b1 100644 --- a/internal/models/identity.go +++ b/internal/models/identity.go @@ -81,6 +81,19 @@ func (i *Identity) IsForSSOProvider() bool { return strings.HasPrefix(i.Provider, "sso:") } +func (i *Identity) SSOProviderID() (uuid.UUID, bool, error) { + return SSOProviderID(i.Provider) +} + +func SSOProviderID(provider string) (uuid.UUID, bool, error) { + id, ok := strings.CutPrefix(provider, "sso:") + if !ok { + return uuid.Nil, false, nil + } + providerID, err := uuid.FromString(id) + return providerID, true, err +} + // FindIdentityById searches for an identity with the matching id and provider given. func FindIdentityByIdAndProvider(tx *storage.Connection, providerId, provider string) (*Identity, error) { identity := &Identity{} diff --git a/internal/models/identity_test.go b/internal/models/identity_test.go index ddf1881a33..b7bc70515c 100644 --- a/internal/models/identity_test.go +++ b/internal/models/identity_test.go @@ -115,3 +115,28 @@ func (ts *IdentityTestSuite) createUserWithIdentity(email string) *User { return user } + +func TestSSOProviderID(t *testing.T) { + id := uuid.Must(uuid.NewV4()) + + providerID, ok, err := SSOProviderID("sso:" + id.String()) + require.True(t, ok) + require.NoError(t, err) + require.Equal(t, id, providerID) + + providerID, ok, err = (&Identity{Provider: "sso:" + id.String()}).SSOProviderID() + require.True(t, ok) + require.NoError(t, err) + require.Equal(t, id, providerID) + + _, ok, err = SSOProviderID("sso:not-a-uuid") + require.True(t, ok) + require.Error(t, err) + + for _, provider := range []string{"email", "google", "", "SSO:" + id.String()} { + providerID, ok, err = SSOProviderID(provider) + require.False(t, ok, provider) + require.NoError(t, err, provider) + require.Equal(t, uuid.Nil, providerID, provider) + } +} From 554d684a6063cceb9c63fcb42516b787373c767a Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 20:49:47 -0600 Subject: [PATCH 09/88] fix(scim): lock every verified email that SSO account linking reads --- internal/api/external.go | 4 ++-- internal/api/scim_link_test.go | 30 ++++++++++++++++++++++++++++++ internal/models/linking.go | 15 +++++++++++---- internal/models/scim_user.go | 17 +++++++++++++++++ 4 files changed, 60 insertions(+), 6 deletions(-) diff --git a/internal/api/external.go b/internal/api/external.go index 0fdb5f1a80..0ce737a853 100644 --- a/internal/api/external.go +++ b/internal/api/external.go @@ -317,8 +317,8 @@ func (a *API) createAccountFromExternalIdentity(tx *storage.Connection, r *http. return 0, nil, apierrors.NewInternalServerError("Invalid SSO provider id in provider type").WithInternalError(perr) } - if scimOn && userData.Metadata.Email != "" { - if terr := models.LockAccountLinking(tx, providerType, userData.Metadata.Email); terr != nil { + if scimOn { + if terr := models.LockAccountLinkingEmails(tx, providerType, models.VerifiedEmails(config, userData.Emails)); terr != nil { return 0, nil, terr } } diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index d36319293a..5cd799d404 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -309,6 +309,10 @@ func (ts *SCIMUsersTestSuite) samlLogin(ssoProvider *models.SSOProvider, sub, em Verified: true, }}, } + return ts.samlLoginWith(ssoProvider, userData) +} + +func (ts *SCIMUsersTestSuite) samlLoginWith(ssoProvider *models.SSOProvider, userData *provider.UserProvidedData) (*models.User, error) { r := httptest.NewRequest(http.MethodPost, "/sso/saml/acs", nil) var user *models.User @@ -320,6 +324,32 @@ func (ts *SCIMUsersTestSuite) samlLogin(ssoProvider *models.SSOProvider, sub, em return user, err } +func (ts *SCIMUsersTestSuite) TestSAMLLoginLocksVerifiedEmailWithoutMetadataEmail() { + userData := &provider.UserProvidedData{ + Metadata: &provider.Claims{Subject: "saml-name-id", EmailVerified: true}, + Emails: []provider.Email{{Email: "Alice@Example.com", Primary: true, Verified: true}}, + } + conn, err := ts.API.db.NewTransaction() + require.NoError(ts.T(), err) + tx := &storage.Connection{Connection: conn} + require.NoError(ts.T(), models.LockAccountLinking(tx, "sso:"+ts.A.ID.String(), "alice@example.com")) + + done := make(chan error, 1) + go func() { + _, err := ts.samlLoginWith(ts.A, userData) + done <- err + }() + + select { + case err := <-done: + require.NoError(ts.T(), tx.TX.Rollback()) + require.FailNow(ts.T(), "SAML login did not wait for the account linking lock", "%v", err) + case <-time.After(200 * time.Millisecond): + } + require.NoError(ts.T(), tx.TX.Rollback()) + require.NoError(ts.T(), <-done) +} + func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowedForActiveSCIMUser() { id := ts.create(ts.TokenA, oktaUser) linked := ts.linkedUser(id) diff --git a/internal/models/linking.go b/internal/models/linking.go index 5f5f2d0cd7..84ccacb50c 100644 --- a/internal/models/linking.go +++ b/internal/models/linking.go @@ -53,6 +53,16 @@ type AccountLinkingResult struct { CandidateEmail provider.Email } +func VerifiedEmails(config *conf.GlobalConfiguration, emails []provider.Email) []string { + var verified []string + for _, email := range emails { + if email.Verified || config.Mailer.Autoconfirm { + verified = append(verified, strings.ToLower(email.Email)) + } + } + return verified +} + // DetermineAccountLinking uses the provided data and database state to compute a decision on whether: // - A new User should be created (CreateAccount) // - A new Identity should be created (LinkAccount) with a UserID pointing to an existing user account @@ -61,12 +71,9 @@ type AccountLinkingResult struct { // // Errors signal failure in processing only, like database access errors. func DetermineAccountLinking(tx *storage.Connection, config *conf.GlobalConfiguration, emails []provider.Email, aud, providerName, sub string) (AccountLinkingResult, error) { - var verifiedEmails []string + verifiedEmails := VerifiedEmails(config, emails) var candidateEmail provider.Email for _, email := range emails { - if email.Verified || config.Mailer.Autoconfirm { - verifiedEmails = append(verifiedEmails, strings.ToLower(email.Email)) - } if email.Primary { candidateEmail = email candidateEmail.Email = strings.ToLower(email.Email) diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index 8f132f3cc2..633ccc8163 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -3,6 +3,7 @@ package models import ( "encoding/json" "fmt" + "slices" "strings" "time" @@ -252,6 +253,22 @@ func RenameSCIMIdentity(tx *storage.Connection, userID uuid.UUID, provider, from return nil } +func LockAccountLinkingEmails(tx *storage.Connection, providerType string, emails []string) error { + keys := make([]string, 0, len(emails)) + for _, email := range emails { + if email != "" { + keys = append(keys, strings.ToLower(email)) + } + } + slices.Sort(keys) + for _, email := range slices.Compact(keys) { + if err := LockAccountLinking(tx, providerType, email); err != nil { + return err + } + } + return nil +} + func LockAccountLinking(tx *storage.Connection, providerType, email string) error { key := providerType + "|" + strings.ToLower(email) if err := tx.RawQuery("SELECT pg_advisory_xact_lock(hashtextextended(?, 0))", key).Exec(); err != nil { From d3c0ea966ca0eead4f9df4f25bad7de5b55cc958 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 20:51:36 -0600 Subject: [PATCH 10/88] chore(scim): keep the authenticated token in one context key --- internal/api/context.go | 33 ++++++++++++++++----------------- internal/api/scim.go | 19 +++++++++++-------- internal/api/scim_admin_test.go | 8 ++++---- 3 files changed, 31 insertions(+), 29 deletions(-) diff --git a/internal/api/context.go b/internal/api/context.go index 2179f1b19d..92dd1c4df5 100644 --- a/internal/api/context.go +++ b/internal/api/context.go @@ -16,23 +16,22 @@ var ( externalProviderTypeKey = ctxkey.New[string]("external_provider_type") externalProviderEmailOptionalKey = ctxkey.New[bool]("external_provider_allow_no_email") - tokenKey = ctxkey.New[*jwt.Token]("jwt") - inviteTokenKey = ctxkey.New[string]("invite_token") - signatureKey = ctxkey.New[string]("signature") - targetUserKey = ctxkey.New[*models.User]("target_user") - factorKey = ctxkey.New[*models.Factor]("factor") - sessionKey = ctxkey.New[*models.Session]("session") - externalReferrerKey = ctxkey.New[string]("external_referrer") - adminUserKey = ctxkey.New[*models.User]("admin_user") - oauthTokenKey = ctxkey.New[string]("oauth_token") // for OAuth1.0, also known as request token - oauthVerifierKey = ctxkey.New[string]("oauth_verifier") - ssoProviderKey = ctxkey.New[*models.SSOProvider]("sso_provider") - externalHostKey = ctxkey.New[*url.URL]("external_host") - oauthClientStateKey = ctxkey.New[uuid.UUID]("oauth_client_state_id") - flowStateContextKey = ctxkey.New[*models.FlowState]("flow_state") - scimRequestKey = ctxkey.New[*http.Request]("scim_request") - scimSSOProviderIDKey = ctxkey.New[uuid.UUID]("scim_sso_provider_id") - scimTokenPrefixKey = ctxkey.New[string]("scim_token_prefix") + tokenKey = ctxkey.New[*jwt.Token]("jwt") + inviteTokenKey = ctxkey.New[string]("invite_token") + signatureKey = ctxkey.New[string]("signature") + targetUserKey = ctxkey.New[*models.User]("target_user") + factorKey = ctxkey.New[*models.Factor]("factor") + sessionKey = ctxkey.New[*models.Session]("session") + externalReferrerKey = ctxkey.New[string]("external_referrer") + adminUserKey = ctxkey.New[*models.User]("admin_user") + oauthTokenKey = ctxkey.New[string]("oauth_token") // for OAuth1.0, also known as request token + oauthVerifierKey = ctxkey.New[string]("oauth_verifier") + ssoProviderKey = ctxkey.New[*models.SSOProvider]("sso_provider") + externalHostKey = ctxkey.New[*url.URL]("external_host") + oauthClientStateKey = ctxkey.New[uuid.UUID]("oauth_client_state_id") + flowStateContextKey = ctxkey.New[*models.FlowState]("flow_state") + scimRequestKey = ctxkey.New[*http.Request]("scim_request") + scimTokenKey = ctxkey.New[*models.SCIMToken]("scim_token") ) // withToken adds the JWT token to the context. diff --git a/internal/api/scim.go b/internal/api/scim.go index 7675f74b3f..dcf52fa67e 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -103,8 +103,7 @@ func newSCIMTokenValidator(db *storage.Connection) server.TokenValidator { if err != nil { return ctx, err } - ctx = scimTokenPrefixKey.WithValue(ctx, token.Prefix) - return scimSSOProviderIDKey.WithValue(ctx, token.SSOProviderID), nil + return scimTokenKey.WithValue(ctx, token), nil } } @@ -153,8 +152,8 @@ func (a *API) limitSCIMInvalidToken(validate server.TokenValidator, lmt *limiter func (a *API) limitSCIMProvider(lmt *limiter.Limiter) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - providerID, ok := scimSSOProviderIDKey.Lookup(r.Context()) - if !ok { + providerID, err := scimProviderID(r.Context()) + if err != nil { next.ServeHTTP(w, r) return } @@ -225,11 +224,11 @@ func scimTarget(ctx context.Context, id, version string) (providerID, resourceID } func scimProviderID(ctx context.Context) (uuid.UUID, error) { - providerID, ok := scimSSOProviderIDKey.Lookup(ctx) - if !ok || providerID == uuid.Nil { + token := scimTokenKey.Value(ctx) + if token == nil || token.SSOProviderID == uuid.Nil { return uuid.Nil, errMissingSSOProvider } - return providerID, nil + return token.SSOProviderID, nil } func scimRequest(ctx context.Context) (*http.Request, error) { @@ -311,7 +310,11 @@ func errSCIMTooManyRequests() error { } func scimActor(r *http.Request) *models.User { - return &models.User{Email: storage.NullString("scim:" + scimTokenPrefixKey.Value(r.Context()))} + prefix := "" + if token := scimTokenKey.Value(r.Context()); token != nil { + prefix = token.Prefix + } + return &models.User{Email: storage.NullString("scim:" + prefix)} } func (a *API) auditSCIM(tx *storage.Connection, r *http.Request, actor *models.User, action models.AuditAction, providerID uuid.UUID, traits map[string]any) error { diff --git a/internal/api/scim_admin_test.go b/internal/api/scim_admin_test.go index d0368d35fe..5281592180 100644 --- a/internal/api/scim_admin_test.go +++ b/internal/api/scim_admin_test.go @@ -126,14 +126,14 @@ func (ts *SCIMTokensTestSuite) TestTokenValidatorResolvesSSOProvider() { ctx, err := validate(context.Background(), created.Token) require.NoError(ts.T(), err) - providerID, ok := scimSSOProviderIDKey.Lookup(ctx) - require.True(ts.T(), ok) + providerID, err := scimProviderID(ctx) + require.NoError(ts.T(), err) require.Equal(ts.T(), ts.Provider.ID, providerID) + require.Equal(ts.T(), created.Prefix, scimTokenKey.Value(ctx).Prefix) ctx, err = validate(context.Background(), "scim_invalid") require.ErrorIs(ts.T(), err, server.ErrInvalidToken) - _, ok = scimSSOProviderIDKey.Lookup(ctx) - require.False(ts.T(), ok) + require.Nil(ts.T(), scimTokenKey.Value(ctx)) cancelled, cancel := context.WithCancel(context.Background()) cancel() From 9b821f7287134e790db1e50973f375462963b2ac Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 21:09:31 -0600 Subject: [PATCH 11/88] chore(scim): test that every authorization header shape is rate limited --- internal/api/scim_users_test.go | 26 ++++++++++++++++++++++++++ 1 file changed, 26 insertions(+) diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index 67b586279c..a28e81896a 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -801,3 +801,29 @@ func TestSCIMRateLimit(t *testing.T) { require.JSONEq(t, limited, w.Body.String()) } } + +func TestSCIMRateLimitEveryAuthorizationHeader(t *testing.T) { + api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { + config.RateLimitScim = 1 + }) + defer api.db.Close() + require.NoError(t, models.TruncateAll(api.db)) + + for i, header := range []string{"", "Bearer", "Bearer ", "bearer scim_invalid", "Bearer scim a", "Basic scim_invalid", "Bearer\tscim_invalid", "Bearer scim_invalid"} { + ip := fmt.Sprintf("203.0.113.%d", 100+i) + send := func() *httptest.ResponseRecorder { + r := httptest.NewRequest(http.MethodGet, "/scim/v2/Users", nil) + if header != "" { + r.Header.Set("Authorization", header) + } + r.Header.Set(api.config.RateLimitHeader, ip) + w := httptest.NewRecorder() + api.handler.ServeHTTP(w, r) + return w + } + for range 30 { + require.NotEqual(t, http.StatusTooManyRequests, send().Code, header) + } + require.Equal(t, http.StatusTooManyRequests, send().Code, header) + } +} From 4ee94207bdd5e3d3982082b9d712436d0cdec0f8 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 21:14:28 -0600 Subject: [PATCH 12/88] chore(scim): mount scim-go under /scim/v2 and reuse the auth bearer token parser --- internal/api/api.go | 22 ++------------- internal/api/router.go | 4 +-- internal/api/scim.go | 24 +++++----------- internal/api/scim_test.go | 23 +++++---------- internal/api/scim_users_test.go | 34 ++++++++++++++--------- internal/api/testdata/scim/not_found.json | 7 ----- 6 files changed, 39 insertions(+), 75 deletions(-) delete mode 100644 internal/api/testdata/scim/not_found.json diff --git a/internal/api/api.go b/internal/api/api.go index f396b0152e..8fc2ebabf8 100644 --- a/internal/api/api.go +++ b/internal/api/api.go @@ -483,26 +483,8 @@ func NewAPIWithVersion(globalConfig *conf.GlobalConfiguration, db *storage.Conne r.Route(scimBasePath, func(r *router) { r.Use(api.requireScimServerEnabled) r.Use(api.withSCIMRequest) - r.UseBypass(api.limitSCIMHandler(api.limiterOpts.SCIMIP, r.chi)) - r.NotFound(scimNotFound) - - r.Method(http.MethodGet, "/ServiceProviderConfig", api.scim) - r.Method(http.MethodGet, "/ResourceTypes", api.scim) - r.Method(http.MethodGet, "/ResourceTypes/{id}", api.scim) - r.Method(http.MethodGet, "/Schemas", api.scim) - r.Method(http.MethodGet, "/Schemas/{id}", api.scim) - r.Method(http.MethodGet, "/Users", api.scim) - r.Method(http.MethodPost, "/Users", api.scim) - r.Method(http.MethodGet, "/Users/{id}", api.scim) - r.Method(http.MethodPut, "/Users/{id}", api.scim) - r.Method(http.MethodPatch, "/Users/{id}", api.scim) - r.Method(http.MethodDelete, "/Users/{id}", api.scim) - r.Method(http.MethodGet, "/Groups", api.scim) - r.Method(http.MethodPost, "/Groups", api.scim) - r.Method(http.MethodGet, "/Groups/{id}", api.scim) - r.Method(http.MethodPut, "/Groups/{id}", api.scim) - r.Method(http.MethodPatch, "/Groups/{id}", api.scim) - r.Method(http.MethodDelete, "/Groups/{id}", api.scim) + r.UseBypass(api.limitSCIMByIP(api.limiterOpts.SCIMIP)) + r.Handle("/*", api.scim) }) }) diff --git a/internal/api/router.go b/internal/api/router.go index 97c047f33c..114b6aff6a 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -37,8 +37,8 @@ func (r *router) Delete(pattern string, fn apiHandler) { r.chi.Delete(pattern, handler(fn)) } -func (r *router) Method(method, pattern string, h http.Handler) { - r.chi.Method(method, pattern, h) +func (r *router) Handle(pattern string, h http.Handler) { + r.chi.Handle(pattern, h) } func (r *router) With(fn middlewareHandler) *router { diff --git a/internal/api/scim.go b/internal/api/scim.go index dcf52fa67e..8cb42b5abb 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -6,13 +6,13 @@ import ( "errors" "fmt" "net/http" + "path" "strconv" "strings" "time" "github.com/didip/tollbooth/v5" "github.com/didip/tollbooth/v5/limiter" - "github.com/go-chi/chi/v5" "github.com/gofrs/uuid" "github.com/supabase-community/scim-go/pkg/core" "github.com/supabase-community/scim-go/pkg/protocol" @@ -82,10 +82,6 @@ func scimBaseURL(config *conf.GlobalConfiguration) string { return strings.TrimRight(config.API.ExternalURL, "/") + scimBasePath } -func scimNotFound(w http.ResponseWriter, r *http.Request) error { - return protocol.SendError(w, scimerrors.ErrNotFound("Endpoint or resource does not exist")) -} - func scimTooManyRequests(w http.ResponseWriter, r *http.Request) error { return protocol.SendError(w, errSCIMTooManyRequests()) } @@ -111,10 +107,10 @@ func (a *API) withSCIMRequest(w http.ResponseWriter, req *http.Request) (context return scimRequestKey.WithValue(req.Context(), req), nil } -func (a *API) limitSCIMHandler(lmt *limiter.Limiter, routes chi.Routes) func(http.Handler) http.Handler { +func (a *API) limitSCIMByIP(lmt *limiter.Limiter) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if !(scimHasBearerToken(r) && scimValidatesToken(routes, r)) && a.performRateLimiting(lmt, r) != nil { + if a.scimSkipsTokenValidator(r) && a.performRateLimiting(lmt, r) != nil { handler(scimTooManyRequests)(w, r) return } @@ -123,17 +119,11 @@ func (a *API) limitSCIMHandler(lmt *limiter.Limiter, routes chi.Routes) func(htt } } -func scimHasBearerToken(r *http.Request) bool { - scheme, token, _ := strings.Cut(r.Header.Get("Authorization"), " ") - return strings.EqualFold(scheme, "Bearer") && token != "" -} - -func scimValidatesToken(routes chi.Routes, r *http.Request) bool { - path := r.URL.Path - if rctx := chi.RouteContext(r.Context()); rctx != nil && rctx.RoutePath != "" { - path = rctx.RoutePath +func (a *API) scimSkipsTokenValidator(r *http.Request) bool { + if _, err := a.extractBearerToken(r); err != nil { + return true } - return path != "/ServiceProviderConfig" && routes.Match(chi.NewRouteContext(), r.Method, path) + return r.URL.Path == scimBasePath+"/ServiceProviderConfig" || path.Clean(r.URL.Path) != r.URL.Path } func (a *API) limitSCIMInvalidToken(validate server.TokenValidator, lmt *limiter.Limiter) server.TokenValidator { diff --git a/internal/api/scim_test.go b/internal/api/scim_test.go index 3ddfb2b1e9..8d88d18166 100644 --- a/internal/api/scim_test.go +++ b/internal/api/scim_test.go @@ -233,17 +233,18 @@ func TestSCIM(t *testing.T) { api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusMethodNotAllowed, w.Code) - require.Equal(t, []string{http.MethodGet}, w.Header().Values("Allow")) + require.Equal(t, "GET, HEAD", w.Header().Get("Allow")) + require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) }) } } for _, tc := range []struct { method, path string - allow []string + allow string }{ - {http.MethodPut, scimUsersPath, []string{http.MethodGet, http.MethodPost}}, - {http.MethodPost, scimUsersPath + "/missing", []string{http.MethodGet, http.MethodPut, http.MethodPatch, http.MethodDelete}}, + {http.MethodPut, scimUsersPath, "GET, HEAD, POST"}, + {http.MethodPost, scimUsersPath + "/missing", "DELETE, GET, HEAD, PATCH, PUT"}, } { t.Run(tc.method+" "+tc.path, func(t *testing.T) { r := httptest.NewRequest(tc.method, tc.path, nil) @@ -253,7 +254,8 @@ func TestSCIM(t *testing.T) { api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusMethodNotAllowed, w.Code) - require.ElementsMatch(t, tc.allow, w.Header().Values("Allow")) + require.Equal(t, tc.allow, w.Header().Get("Allow")) + require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) }) } }) @@ -456,17 +458,6 @@ func TestSCIMServer(t *testing.T) { }) } }) - - t.Run("NotFound", func(t *testing.T) { - r := httptest.NewRequest(http.MethodGet, scimBasePath+"/Unknown", nil) - w := httptest.NewRecorder() - - require.NoError(t, scimNotFound(w, r)) - - require.Equal(t, http.StatusNotFound, w.Code) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) - require.JSONEq(t, scimFixture(t, "not_found.json"), w.Body.String()) - }) } func TestSCIMUserFields(t *testing.T) { diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index a28e81896a..55c59ad791 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -780,8 +780,9 @@ func TestSCIMRateLimit(t *testing.T) { status int }{ {http.MethodGet, "/scim/v2/ServiceProviderConfig", http.StatusOK}, - {http.MethodGet, "/scim/v2/Unknown", http.StatusNotFound}, - {http.MethodDelete, "/scim/v2/Users", http.StatusMethodNotAllowed}, + {http.MethodHead, "/scim/v2/ServiceProviderConfig", http.StatusOK}, + {http.MethodGet, "/scim/v2/Unknown", http.StatusUnauthorized}, + {http.MethodDelete, "/scim/v2/Users", http.StatusUnauthorized}, } { ip := fmt.Sprintf("203.0.113.%d", i+1) send := func() *httptest.ResponseRecorder { @@ -802,28 +803,35 @@ func TestSCIMRateLimit(t *testing.T) { } } -func TestSCIMRateLimitEveryAuthorizationHeader(t *testing.T) { +func TestSCIMRateLimitBoundsEveryUnauthenticatedRequest(t *testing.T) { api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { config.RateLimitScim = 1 }) defer api.db.Close() require.NoError(t, models.TruncateAll(api.db)) - for i, header := range []string{"", "Bearer", "Bearer ", "bearer scim_invalid", "Bearer scim a", "Basic scim_invalid", "Bearer\tscim_invalid", "Bearer scim_invalid"} { + type request struct{ path, authorization string } + requests := []request{} + for _, header := range []string{"", "Bearer", "Bearer ", "bearer scim_invalid", "Bearer scim a", "Basic scim_invalid", "Bearer\tscim_invalid", "Bearer scim_invalid"} { + requests = append(requests, request{"/scim/v2/Users", header}) + } + for _, path := range []string{"/scim/v2//Users", "/scim/v2/./Users", "/scim/v2/Users/", "/scim/v2/ServiceProviderConfig/"} { + requests = append(requests, request{path, "Bearer scim_invalid"}) + } + + for i, tc := range requests { ip := fmt.Sprintf("203.0.113.%d", 100+i) - send := func() *httptest.ResponseRecorder { - r := httptest.NewRequest(http.MethodGet, "/scim/v2/Users", nil) - if header != "" { - r.Header.Set("Authorization", header) + limited := false + for range 31 { + r := httptest.NewRequest(http.MethodGet, tc.path, nil) + if tc.authorization != "" { + r.Header.Set("Authorization", tc.authorization) } r.Header.Set(api.config.RateLimitHeader, ip) w := httptest.NewRecorder() api.handler.ServeHTTP(w, r) - return w - } - for range 30 { - require.NotEqual(t, http.StatusTooManyRequests, send().Code, header) + limited = limited || w.Code == http.StatusTooManyRequests } - require.Equal(t, http.StatusTooManyRequests, send().Code, header) + require.True(t, limited, "%q %q", tc.path, tc.authorization) } } diff --git a/internal/api/testdata/scim/not_found.json b/internal/api/testdata/scim/not_found.json deleted file mode 100644 index 4d241ba672..0000000000 --- a/internal/api/testdata/scim/not_found.json +++ /dev/null @@ -1,7 +0,0 @@ -{ - "schemas": [ - "urn:ietf:params:scim:api:messages:2.0:Error" - ], - "status": "404", - "detail": "Endpoint or resource does not exist" -} From 1200172dcafd28c3bc17563b6e83cae874a0f8d2 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 21:16:01 -0600 Subject: [PATCH 13/88] chore(scim): compute group member changes from the rows found and pass uuid slices to postgres --- internal/models/scim.go | 21 --------------------- internal/models/scim_group.go | 24 +++++++++++------------- internal/models/scim_user.go | 2 +- 3 files changed, 12 insertions(+), 35 deletions(-) diff --git a/internal/models/scim.go b/internal/models/scim.go index 5b55c34359..34f0433b56 100644 --- a/internal/models/scim.go +++ b/internal/models/scim.go @@ -176,27 +176,6 @@ func isCheckViolation(err error, constraint string) bool { return errors.As(err, &pgErr) && pgErr.Code == pgerrcode.CheckViolation && pgErr.ConstraintName == constraint } -func uuidArray(ids []uuid.UUID) string { - values := make([]string, len(ids)) - for i, id := range ids { - values[i] = id.String() - } - return "{" + strings.Join(values, ",") + "}" -} - -func dedupeUUIDs(ids []uuid.UUID) []uuid.UUID { - seen := make(map[uuid.UUID]struct{}, len(ids)) - unique := make([]uuid.UUID, 0, len(ids)) - for _, id := range ids { - if _, ok := seen[id]; ok { - continue - } - seen[id] = struct{}{} - unique = append(unique, id) - } - return unique -} - func differenceUUIDs(from, subtract []uuid.UUID) []uuid.UUID { exclude := make(map[uuid.UUID]struct{}, len(subtract)) for _, id := range subtract { diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go index bcd136bf82..24cbe605c4 100644 --- a/internal/models/scim_group.go +++ b/internal/models/scim_group.go @@ -113,7 +113,7 @@ func FindSCIMGroupMembers(tx *storage.Connection, providerID uuid.UUID, groupIDs } err := tx.RawQuery( fmt.Sprintf("SELECT m.group_id, m.scim_user_id, u.resource->>'userName' AS display FROM %q m JOIN %q u ON u.id = m.scim_user_id WHERE m.group_id = ANY(?::uuid[]) AND u.sso_provider_id = ? AND u.deleted_at IS NULL ORDER BY m.group_id, m.created_at, m.scim_user_id", (&SCIMGroupMember{}).TableName(), (&SCIMUser{}).TableName()), - uuidArray(groupIDs), providerID, + groupIDs, providerID, ).All(&members) if err != nil { return nil, errors.Wrap(err, "error finding SCIM group members") @@ -128,7 +128,7 @@ func FindSCIMGroupsForUsers(tx *storage.Connection, providerID uuid.UUID, scimUs } err := tx.RawQuery( fmt.Sprintf("SELECT m.group_id, m.scim_user_id, g.resource->>'displayName' AS display FROM %q m JOIN %q g ON g.id = m.group_id WHERE m.scim_user_id = ANY(?::uuid[]) AND g.sso_provider_id = ? ORDER BY m.scim_user_id, g.display_name COLLATE \"C\", g.id", (&SCIMGroupMember{}).TableName(), (&SCIMGroup{}).TableName()), - uuidArray(scimUserIDs), providerID, + scimUserIDs, providerID, ).All(&groups) if err != nil { return nil, errors.Wrap(err, "error finding SCIM groups for users") @@ -137,18 +137,16 @@ func FindSCIMGroupsForUsers(tx *storage.Connection, providerID uuid.UUID, scimUs } func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserIDs []uuid.UUID) (added, removed []uuid.UUID, err error) { - wanted := dedupeUUIDs(scimUserIDs) - - foundIDs := []uuid.UUID{} - if len(wanted) > 0 { + wanted := []uuid.UUID{} + if len(scimUserIDs) > 0 { if err := tx.RawQuery( fmt.Sprintf("SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL", (&SCIMUser{}).TableName()), - uuidArray(wanted), group.SSOProviderID, - ).All(&foundIDs); err != nil { + scimUserIDs, group.SSOProviderID, + ).All(&wanted); err != nil { return nil, nil, errors.Wrap(err, "error finding SCIM group members") } } - if missing := differenceUUIDs(wanted, foundIDs); len(missing) > 0 { + if missing := differenceUUIDs(scimUserIDs, wanted); len(missing) > 0 { return nil, nil, SCIMGroupMemberNotFoundError{IDs: missing} } @@ -170,7 +168,7 @@ func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserI if len(removed) > 0 { if err := tx.RawQuery( fmt.Sprintf("DELETE FROM %q WHERE group_id = ? AND scim_user_id = ANY(?::uuid[])", (&SCIMGroupMember{}).TableName()), - group.ID, uuidArray(removed), + group.ID, removed, ).Exec(); err != nil { return nil, nil, errors.Wrap(err, "error removing SCIM group members") } @@ -179,7 +177,7 @@ func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserI locked := []uuid.UUID{} if err := tx.RawQuery( fmt.Sprintf("SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL ORDER BY id FOR SHARE", (&SCIMUser{}).TableName()), - uuidArray(added), group.SSOProviderID, + added, group.SSOProviderID, ).All(&locked); err != nil { return nil, nil, errors.Wrap(err, "error locking SCIM group members") } @@ -188,7 +186,7 @@ func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserI } if err := tx.RawQuery( fmt.Sprintf("INSERT INTO %q (group_id, scim_user_id) SELECT ?, unnest(?::uuid[])", (&SCIMGroupMember{}).TableName()), - group.ID, uuidArray(locked), + group.ID, locked, ).Exec(); err != nil { return nil, nil, errors.Wrap(err, "error adding SCIM group members") } @@ -218,7 +216,7 @@ func RemoveSCIMUserFromGroups(tx *storage.Connection, scimUserID uuid.UUID) ([]u if len(groupIDs) > 0 { if err := tx.RawQuery( fmt.Sprintf("UPDATE %q SET updated_at = now() WHERE id = ANY(?::uuid[])", groups), - uuidArray(groupIDs), + groupIDs, ).Exec(); err != nil { return nil, errors.Wrap(err, "error updating SCIM groups") } diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index 633ccc8163..5773f52f7a 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -97,7 +97,7 @@ func FindSCIMUserLinks(tx *storage.Connection, ids []uuid.UUID) (map[uuid.UUID]u }{} if err := tx.RawQuery( fmt.Sprintf("SELECT id, user_id FROM %q WHERE id = ANY(?::uuid[]) AND user_id IS NOT NULL", (&SCIMUser{}).TableName()), - uuidArray(ids), + ids, ).All(&rows); err != nil { return nil, errors.Wrap(err, "error finding SCIM user links") } From 01b32c85e55e9ad0be051e919888935de89a9f89 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 21:18:06 -0600 Subject: [PATCH 14/88] chore(scim): read tokens and settings instead of writing them before provider delete --- internal/api/scim_admin.go | 38 ++++++++++++------------------ internal/models/scim_token.go | 11 --------- internal/models/scim_token_test.go | 31 ------------------------ 3 files changed, 15 insertions(+), 65 deletions(-) diff --git a/internal/api/scim_admin.go b/internal/api/scim_admin.go index 887c3adb1e..ef98437a39 100644 --- a/internal/api/scim_admin.go +++ b/internal/api/scim_admin.go @@ -91,43 +91,35 @@ func (a *API) isSCIMEnabled(db *storage.Connection, provider *models.SSOProvider } func (a *API) deprovisionSCIM(tx *storage.Connection, r *http.Request, provider *models.SSOProvider) error { - disabled, err := models.DisableSCIM(tx, provider.ID) + enabled, err := models.IsSCIMEnabled(tx, provider.ID) if err != nil { return err } - prefixes, err := a.revokeSCIMTokens(tx, r, provider) - if err != nil { - return err - } - if disabled && a.config.SSO.SCIM.Enabled { - if err := a.auditSCIM(tx, r, getAdminUser(r.Context()), models.SCIMDisabledAction, provider.ID, map[string]any{"token_prefixes": prefixes}); err != nil { - return err - } - } - banned, err := models.BanDeprovisionedSCIMUsers(tx, provider.ID, a.Now().Add(scimProviderDeletedBan)) - if err != nil || banned == 0 { - return err - } - return a.auditSCIM(tx, r, getAdminUser(r.Context()), models.SCIMUsersBannedAction, provider.ID, map[string]any{"banned_user_count": banned}) -} - -func (a *API) revokeSCIMTokens(tx *storage.Connection, r *http.Request, provider *models.SSOProvider) ([]string, error) { if err := models.LockSCIMTokens(tx, provider.ID); err != nil { - return nil, err + return err } - tokens, err := models.RevokeSCIMTokensBySSOProvider(tx, provider.ID) + tokens, err := models.FindActiveSCIMTokensBySSOProvider(tx, provider.ID) if err != nil { - return nil, err + return err } actor := getAdminUser(r.Context()) prefixes := make([]string, len(tokens)) for i := range tokens { prefixes[i] = tokens[i].Prefix if err := a.auditSCIM(tx, r, actor, models.SCIMTokenRevokedAction, provider.ID, map[string]any{"token_prefix": tokens[i].Prefix}); err != nil { - return nil, err + return err } } - return prefixes, nil + if enabled && a.config.SSO.SCIM.Enabled { + if err := a.auditSCIM(tx, r, actor, models.SCIMDisabledAction, provider.ID, map[string]any{"token_prefixes": prefixes}); err != nil { + return err + } + } + banned, err := models.BanDeprovisionedSCIMUsers(tx, provider.ID, a.Now().Add(scimProviderDeletedBan)) + if err != nil || banned == 0 { + return err + } + return a.auditSCIM(tx, r, actor, models.SCIMUsersBannedAction, provider.ID, map[string]any{"banned_user_count": banned}) } func (a *API) adminSCIMTokensCreate(w http.ResponseWriter, r *http.Request) error { diff --git a/internal/models/scim_token.go b/internal/models/scim_token.go index 701d8858c2..bf55bcc408 100644 --- a/internal/models/scim_token.go +++ b/internal/models/scim_token.go @@ -105,17 +105,6 @@ func LockSCIMTokens(tx *storage.Connection, providerID uuid.UUID) error { return nil } -func RevokeSCIMTokensBySSOProvider(tx *storage.Connection, providerID uuid.UUID) ([]SCIMToken, error) { - tokens := []SCIMToken{} - if err := tx.RawQuery( - fmt.Sprintf("WITH revoked AS (UPDATE %q SET revoked_at = now() WHERE sso_provider_id = ? AND "+activeSCIMTokenClause+" RETURNING *) SELECT * FROM revoked ORDER BY created_at ASC, id ASC", (&SCIMToken{}).TableName()), - providerID, - ).All(&tokens); err != nil { - return nil, errors.Wrap(err, "error revoking SCIM tokens") - } - return tokens, nil -} - func FindSCIMTokenByPrefix(tx *storage.Connection, providerID uuid.UUID, prefix string) (*SCIMToken, error) { tokens := []SCIMToken{} if err := tx.Q().Where("sso_provider_id = ? AND prefix = ?", providerID, prefix).Limit(2).All(&tokens); err != nil { diff --git a/internal/models/scim_token_test.go b/internal/models/scim_token_test.go index 5a295f37cd..387485ed49 100644 --- a/internal/models/scim_token_test.go +++ b/internal/models/scim_token_test.go @@ -266,34 +266,3 @@ func (ts *SCIMTokenTestSuite) TestFindActiveBySSOProvider() { require.NoError(ts.T(), err) require.Empty(ts.T(), tokens) } - -func (ts *SCIMTokenTestSuite) TestRevokeBySSOProvider() { - first, _ := ts.createToken(nil) - second, _ := ts.createToken(nil) - revoked, _ := ts.createToken(nil) - require.NoError(ts.T(), revoked.Revoke(ts.db)) - expired, _ := ts.createToken(nil) - ts.expire(expired) - other, _, err := CreateSCIMToken(ts.db, ts.createProvider(), nil) - require.NoError(ts.T(), err) - - tokens, err := RevokeSCIMTokensBySSOProvider(ts.db, ts.provider.ID) - require.NoError(ts.T(), err) - require.Len(ts.T(), tokens, 2) - require.Equal(ts.T(), []uuid.UUID{first.ID, second.ID}, []uuid.UUID{tokens[0].ID, tokens[1].ID}) - for _, token := range tokens { - require.NotNil(ts.T(), token.RevokedAt) - } - - active, err := FindActiveSCIMTokensBySSOProvider(ts.db, ts.provider.ID) - require.NoError(ts.T(), err) - require.Empty(ts.T(), active) - - tokens, err = RevokeSCIMTokensBySSOProvider(ts.db, ts.provider.ID) - require.NoError(ts.T(), err) - require.Empty(ts.T(), tokens) - - still, err := FindActiveSCIMTokensBySSOProvider(ts.db, other.SSOProviderID) - require.NoError(ts.T(), err) - require.Len(ts.T(), still, 1) -} From e34b7f3d5de35bc3381b28b28f06bf864ef1a8b4 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 21:20:23 -0600 Subject: [PATCH 15/88] chore(scim): fix scim-go lint findings and move linking locks and logout into models --- internal/api/api.go | 4 +--- internal/api/scim.go | 12 +++++------ internal/api/scim_admin.go | 41 +++++++++++++++++++++++------------- internal/api/scim_groups.go | 2 +- internal/api/scim_test.go | 2 +- internal/api/scim_users.go | 15 ++++--------- internal/models/linking.go | 26 +++++++++++++++++++++++ internal/models/scim_user.go | 33 ++++++----------------------- 8 files changed, 72 insertions(+), 63 deletions(-) diff --git a/internal/api/api.go b/internal/api/api.go index 8fc2ebabf8..abc6eac3d6 100644 --- a/internal/api/api.go +++ b/internal/api/api.go @@ -138,11 +138,9 @@ func NewAPIWithVersion(globalConfig *conf.GlobalConfiguration, db *storage.Conne api.oauthServer = oauthserver.NewServer(globalConfig, db, api.tokenService) } - api.scim = newSCIMServer(globalConfig, + api.scim = api.newSCIMServer( api.limitSCIMInvalidToken(newSCIMTokenValidator(db), api.limiterOpts.SCIMIP), api.limitSCIMProvider(api.limiterOpts.SCIM), - &scimUsers{api: api}, - &scimGroups{api: api}, ) if api.config.Password.HIBP.Enabled { diff --git a/internal/api/scim.go b/internal/api/scim.go index 8cb42b5abb..05974b71d1 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -34,7 +34,7 @@ const ( var errMissingSSOProvider = errors.New("scim: request has no SSO provider") -func newSCIMServer(config *conf.GlobalConfiguration, validate server.TokenValidator, limit func(http.Handler) http.Handler, users server.Repository[*core.User], groups server.Repository[*core.Group]) *server.Server { +func (a *API) newSCIMServer(validate server.TokenValidator, limit func(http.Handler) http.Handler) *server.Server { requireToken := server.RequireBearerToken(validate) authenticate := func(next http.Handler) http.Handler { if limit != nil { @@ -44,13 +44,13 @@ func newSCIMServer(config *conf.GlobalConfiguration, validate server.TokenValida } return server.New(scimBasePath, core.NewServiceProviderConfig().Filtering(protocol.DefaultLimits.MaxCount).Patching().Sorting().Versioning(), - server.WithBaseURL(scimBaseURL(config)), + server.WithBaseURL(scimBaseURL(a.config)), server.ErrorHandler(scimLogError), server.WithResource(server.NewResource[*core.User](scimResourceTypeUser, "/Users", core.SchemaUser, scimUserSchemas.Base().Attributes...). WithExtension(core.SchemaEnterpriseUser, scimUserSchemas.Extensions()[0].Attributes...). - WithRepository(users)), + WithRepository(&scimUsers{api: a})), server.WithResource(server.NewResource[*core.Group](scimResourceTypeGroup, "/Groups", core.SchemaGroup, scimGroupSchemas.Base().Attributes...). - WithRepository(groups)), + WithRepository(&scimGroups{api: a})), server.WithAuthentication(core.NewOAuthBearerToken().AsPrimary(), authenticate), ) } @@ -313,7 +313,7 @@ func (a *API) auditSCIM(tx *storage.Connection, r *http.Request, actor *models.U return models.NewAuditLogEntry(a.config.AuditLog, r, tx, actor, action, utilities.GetIPAddress(r), traits) } -func (a *API) auditSCIMMember(tx *storage.Connection, r *http.Request, actor *models.User, action models.AuditAction, providerID, groupID, scimUserID uuid.UUID, userID *uuid.UUID) error { +func scimMemberTraits(groupID, scimUserID uuid.UUID, userID *uuid.UUID) map[string]any { traits := map[string]any{ "scim_group_id": groupID, "scim_user_id": scimUserID, @@ -321,5 +321,5 @@ func (a *API) auditSCIMMember(tx *storage.Connection, r *http.Request, actor *mo if userID != nil { traits["user_id"] = *userID } - return a.auditSCIM(tx, r, actor, action, providerID, traits) + return traits } diff --git a/internal/api/scim_admin.go b/internal/api/scim_admin.go index ef98437a39..16e7a2daf6 100644 --- a/internal/api/scim_admin.go +++ b/internal/api/scim_admin.go @@ -13,7 +13,10 @@ import ( "github.com/supabase/auth/internal/utilities" ) -const scimProviderDeletedBan = 100 * 365 * 24 * time.Hour +const ( + scimProviderDeletedBan = 100 * 365 * 24 * time.Hour + scimTokenPrefixTrait = "token_prefix" +) type AdminSCIMTokenCreateParams struct { ExpiresAt *time.Time `json:"expires_at"` @@ -95,21 +98,11 @@ func (a *API) deprovisionSCIM(tx *storage.Connection, r *http.Request, provider if err != nil { return err } - if err := models.LockSCIMTokens(tx, provider.ID); err != nil { - return err - } - tokens, err := models.FindActiveSCIMTokensBySSOProvider(tx, provider.ID) + actor := getAdminUser(r.Context()) + prefixes, err := a.auditSCIMTokensRevoked(tx, r, actor, provider.ID) if err != nil { return err } - actor := getAdminUser(r.Context()) - prefixes := make([]string, len(tokens)) - for i := range tokens { - prefixes[i] = tokens[i].Prefix - if err := a.auditSCIM(tx, r, actor, models.SCIMTokenRevokedAction, provider.ID, map[string]any{"token_prefix": tokens[i].Prefix}); err != nil { - return err - } - } if enabled && a.config.SSO.SCIM.Enabled { if err := a.auditSCIM(tx, r, actor, models.SCIMDisabledAction, provider.ID, map[string]any{"token_prefixes": prefixes}); err != nil { return err @@ -122,6 +115,24 @@ func (a *API) deprovisionSCIM(tx *storage.Connection, r *http.Request, provider return a.auditSCIM(tx, r, actor, models.SCIMUsersBannedAction, provider.ID, map[string]any{"banned_user_count": banned}) } +func (a *API) auditSCIMTokensRevoked(tx *storage.Connection, r *http.Request, actor *models.User, providerID uuid.UUID) ([]string, error) { + if err := models.LockSCIMTokens(tx, providerID); err != nil { + return nil, err + } + tokens, err := models.FindActiveSCIMTokensBySSOProvider(tx, providerID) + if err != nil { + return nil, err + } + prefixes := make([]string, len(tokens)) + for i := range tokens { + prefixes[i] = tokens[i].Prefix + if err := a.auditSCIM(tx, r, actor, models.SCIMTokenRevokedAction, providerID, map[string]any{scimTokenPrefixTrait: tokens[i].Prefix}); err != nil { + return nil, err + } + } + return prefixes, nil +} + func (a *API) adminSCIMTokensCreate(w http.ResponseWriter, r *http.Request) error { ctx := r.Context() db := a.db.WithContext(ctx) @@ -149,7 +160,7 @@ func (a *API) adminSCIMTokensCreate(w http.ResponseWriter, r *http.Request) erro if token, plaintext, err = models.CreateSCIMToken(tx, provider, params.ExpiresAt); err != nil { return err } - return a.auditSCIM(tx, r, getAdminUser(ctx), models.SCIMTokenCreatedAction, provider.ID, map[string]any{"token_prefix": token.Prefix}) + return a.auditSCIM(tx, r, getAdminUser(ctx), models.SCIMTokenCreatedAction, provider.ID, map[string]any{scimTokenPrefixTrait: token.Prefix}) }); err != nil { if errors.Is(err, models.SCIMTokenExpiryError{}) { return apierrors.NewBadRequestError(apierrors.ErrorCodeValidationFailed, "expires_at must be in the future") @@ -196,7 +207,7 @@ func (a *API) adminSCIMTokensRevoke(w http.ResponseWriter, r *http.Request) erro if err = token.Revoke(tx); err != nil { return err } - return a.auditSCIM(tx, r, getAdminUser(ctx), models.SCIMTokenRevokedAction, provider.ID, map[string]any{"token_prefix": token.Prefix}) + return a.auditSCIM(tx, r, getAdminUser(ctx), models.SCIMTokenRevokedAction, provider.ID, map[string]any{scimTokenPrefixTrait: token.Prefix}) }); err != nil { if models.IsNotFoundError(err) { return apierrors.NewNotFoundError(apierrors.ErrorCodeSCIMTokenNotFound, "SCIM token not found") diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index d0e066afc2..58d4a961a7 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -224,7 +224,7 @@ func (s *scimGroups) auditMembers(tx *storage.Connection, r *http.Request, row * if linked, ok := links[id]; ok { userID = &linked } - if err := s.api.auditSCIMMember(tx, r, scimActor(r), change.action, row.SSOProviderID, row.ID, id, userID); err != nil { + if err := s.api.auditSCIM(tx, r, scimActor(r), change.action, row.SSOProviderID, scimMemberTraits(row.ID, id, userID)); err != nil { return err } } diff --git a/internal/api/scim_test.go b/internal/api/scim_test.go index 8d88d18166..4b8302086c 100644 --- a/internal/api/scim_test.go +++ b/internal/api/scim_test.go @@ -277,7 +277,7 @@ func newSCIMServerFor(externalURL string) *server.Server { } return ctx, nil } - return newSCIMServer(&conf.GlobalConfiguration{API: conf.APIConfiguration{ExternalURL: externalURL}}, validate, nil, nil, nil) + return (&API{config: &conf.GlobalConfiguration{API: conf.APIConfiguration{ExternalURL: externalURL}}}).newSCIMServer(validate, nil) } func scimServe(t *testing.T, srv *server.Server, method, path, body string, headers ...string) *httptest.ResponseRecorder { diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index 5d99f52d7d..3a8b3c8023 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -198,7 +198,7 @@ func (s *scimUsers) Delete(ctx context.Context, id, version string) error { return err } if row.UserID != nil { - if err := scimDeactivate(tx, *row.UserID); err != nil { + if err := models.LogoutSCIMUser(tx, *row.UserID); err != nil { return err } } @@ -284,7 +284,7 @@ func (s *scimUsers) sync(tx *storage.Connection, providerID uuid.UUID, old, row } } if old.Active && !row.Active { - return nil, scimDeactivate(tx, linked.ID) + return nil, models.LogoutSCIMUser(tx, linked.ID) } return nil, nil } @@ -299,7 +299,7 @@ func (s *scimUsers) linkNew(tx *storage.Connection, row *models.SCIMUser, user * created = linked } if !row.Active { - return created, scimDeactivate(tx, linked.ID) + return created, models.LogoutSCIMUser(tx, linked.ID) } return created, nil } @@ -413,7 +413,7 @@ func (a *API) removeSCIMUserFromGroups(tx *storage.Connection, r *http.Request, return err } for _, groupID := range groupIDs { - if err := a.auditSCIMMember(tx, r, actor, models.SCIMGroupMemberRemovedAction, row.SSOProviderID, groupID, row.ID, row.UserID); err != nil { + if err := a.auditSCIM(tx, r, actor, models.SCIMGroupMemberRemovedAction, row.SSOProviderID, scimMemberTraits(groupID, row.ID, row.UserID)); err != nil { return err } } @@ -495,13 +495,6 @@ func scimUserAuditAction(before, after *models.SCIMUser) models.AuditAction { return models.SCIMUserUpdatedAction } -func scimDeactivate(tx *storage.Connection, userID uuid.UUID) error { - if err := models.LockUserForSCIM(tx, userID); err != nil { - return err - } - return models.Logout(tx, userID) -} - func scimHookError(err error) error { var httpErr *apierrors.HTTPError if errors.As(err, &httpErr) && httpErr.HTTPStatus < http.StatusInternalServerError { diff --git a/internal/models/linking.go b/internal/models/linking.go index 84ccacb50c..e9a61a5d70 100644 --- a/internal/models/linking.go +++ b/internal/models/linking.go @@ -1,8 +1,10 @@ package models import ( + "slices" "strings" + "github.com/pkg/errors" "github.com/supabase/auth/internal/api/provider" "github.com/supabase/auth/internal/conf" "github.com/supabase/auth/internal/storage" @@ -219,3 +221,27 @@ func DetermineAccountLinking(tx *storage.Connection, config *conf.GlobalConfigur CandidateEmail: candidateEmail, }, nil } + +func LockAccountLinkingEmails(tx *storage.Connection, providerType string, emails []string) error { + keys := make([]string, 0, len(emails)) + for _, email := range emails { + if email != "" { + keys = append(keys, strings.ToLower(email)) + } + } + slices.Sort(keys) + for _, email := range slices.Compact(keys) { + if err := LockAccountLinking(tx, providerType, email); err != nil { + return err + } + } + return nil +} + +func LockAccountLinking(tx *storage.Connection, providerType, email string) error { + key := providerType + "|" + strings.ToLower(email) + if err := tx.RawQuery("SELECT pg_advisory_xact_lock(hashtextextended(?, 0))", key).Exec(); err != nil { + return errors.Wrap(err, "error locking account linking") + } + return nil +} diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index 5773f52f7a..49eb4353db 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -3,8 +3,6 @@ package models import ( "encoding/json" "fmt" - "slices" - "strings" "time" "github.com/gofrs/uuid" @@ -117,6 +115,13 @@ func LockUserForSCIM(tx *storage.Connection, userID uuid.UUID) error { return nil } +func LogoutSCIMUser(tx *storage.Connection, userID uuid.UUID) error { + if err := LockUserForSCIM(tx, userID); err != nil { + return err + } + return Logout(tx, userID) +} + func SoftDeleteSCIMUsersByUserID(tx *storage.Connection, userID uuid.UUID) ([]SCIMUser, error) { rows := []SCIMUser{} if err := tx.RawQuery( @@ -252,27 +257,3 @@ func RenameSCIMIdentity(tx *storage.Connection, userID uuid.UUID, provider, from } return nil } - -func LockAccountLinkingEmails(tx *storage.Connection, providerType string, emails []string) error { - keys := make([]string, 0, len(emails)) - for _, email := range emails { - if email != "" { - keys = append(keys, strings.ToLower(email)) - } - } - slices.Sort(keys) - for _, email := range slices.Compact(keys) { - if err := LockAccountLinking(tx, providerType, email); err != nil { - return err - } - } - return nil -} - -func LockAccountLinking(tx *storage.Connection, providerType, email string) error { - key := providerType + "|" + strings.ToLower(email) - if err := tx.RawQuery("SELECT pg_advisory_xact_lock(hashtextextended(?, 0))", key).Exec(); err != nil { - return errors.Wrap(err, "error locking account linking") - } - return nil -} From 25f7ec3fff5e39835fac215e20e90f8ca0593401 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 21:22:15 -0600 Subject: [PATCH 16/88] chore(scim): simplify the provider rate limit and scim error mapping --- internal/api/scim.go | 9 +-------- 1 file changed, 1 insertion(+), 8 deletions(-) diff --git a/internal/api/scim.go b/internal/api/scim.go index 05974b71d1..e6549ec135 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -142,12 +142,7 @@ func (a *API) limitSCIMInvalidToken(validate server.TokenValidator, lmt *limiter func (a *API) limitSCIMProvider(lmt *limiter.Limiter) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - providerID, err := scimProviderID(r.Context()) - if err != nil { - next.ServeHTTP(w, r) - return - } - if err := tollbooth.LimitByKeys(lmt, []string{providerID.String()}); err != nil { + if providerID, err := scimProviderID(r.Context()); err == nil && tollbooth.LimitByKeys(lmt, []string{providerID.String()}) != nil { handler(scimTooManyRequests)(w, r) return } @@ -257,8 +252,6 @@ func scimParseVersion(version string) (*time.Time, error) { func scimTranslate(err error) error { switch { - case err == nil: - return nil case models.IsNotFoundError(err): return errSCIMNotFound() case errors.Is(err, models.SCIMUserStaleError{}), errors.Is(err, models.SCIMGroupStaleError{}): From a98c4dea4821450615d08bc8b7cc438eaf3c9a86 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 21:25:53 -0600 Subject: [PATCH 17/88] chore(scim): rename SCIM helpers to say what they do --- internal/api/api.go | 2 +- internal/api/scim.go | 8 +-- internal/api/scim_groups.go | 28 +++++----- internal/api/scim_users.go | 90 ++++++++++++++++----------------- internal/api/scim_users_test.go | 4 +- internal/models/scim_user.go | 4 +- 6 files changed, 68 insertions(+), 68 deletions(-) diff --git a/internal/api/api.go b/internal/api/api.go index abc6eac3d6..2ed37c94cb 100644 --- a/internal/api/api.go +++ b/internal/api/api.go @@ -140,7 +140,7 @@ func NewAPIWithVersion(globalConfig *conf.GlobalConfiguration, db *storage.Conne api.scim = api.newSCIMServer( api.limitSCIMInvalidToken(newSCIMTokenValidator(db), api.limiterOpts.SCIMIP), - api.limitSCIMProvider(api.limiterOpts.SCIM), + api.limitSCIMByProvider(api.limiterOpts.SCIM), ) if api.config.Password.HIBP.Enabled { diff --git a/internal/api/scim.go b/internal/api/scim.go index e6549ec135..5c4405fc82 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -48,9 +48,9 @@ func (a *API) newSCIMServer(validate server.TokenValidator, limit func(http.Hand server.ErrorHandler(scimLogError), server.WithResource(server.NewResource[*core.User](scimResourceTypeUser, "/Users", core.SchemaUser, scimUserSchemas.Base().Attributes...). WithExtension(core.SchemaEnterpriseUser, scimUserSchemas.Extensions()[0].Attributes...). - WithRepository(&scimUsers{api: a})), + WithRepository(&scimUserRepository{api: a})), server.WithResource(server.NewResource[*core.Group](scimResourceTypeGroup, "/Groups", core.SchemaGroup, scimGroupSchemas.Base().Attributes...). - WithRepository(&scimGroups{api: a})), + WithRepository(&scimGroupRepository{api: a})), server.WithAuthentication(core.NewOAuthBearerToken().AsPrimary(), authenticate), ) } @@ -139,7 +139,7 @@ func (a *API) limitSCIMInvalidToken(validate server.TokenValidator, lmt *limiter } } -func (a *API) limitSCIMProvider(lmt *limiter.Limiter) func(http.Handler) http.Handler { +func (a *API) limitSCIMByProvider(lmt *limiter.Limiter) func(http.Handler) http.Handler { return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if providerID, err := scimProviderID(r.Context()); err == nil && tollbooth.LimitByKeys(lmt, []string{providerID.String()}) != nil { @@ -250,7 +250,7 @@ func scimParseVersion(version string) (*time.Time, error) { return &updatedAt, nil } -func scimTranslate(err error) error { +func scimError(err error) error { switch { case models.IsNotFoundError(err): return errSCIMNotFound() diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index 58d4a961a7..67f2db5dfe 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -12,11 +12,11 @@ import ( "github.com/supabase/auth/internal/storage" ) -type scimGroups struct { +type scimGroupRepository struct { api *API } -func (s *scimGroups) List(ctx context.Context, query *protocol.SearchRequest) ([]*core.Group, int, error) { +func (s *scimGroupRepository) List(ctx context.Context, query *protocol.SearchRequest) ([]*core.Group, int, error) { providerID, err := scimProviderID(ctx) if err != nil { return nil, 0, err @@ -37,7 +37,7 @@ func (s *scimGroups) List(ctx context.Context, query *protocol.SearchRequest) ([ return groups, total, nil } -func (s *scimGroups) Get(ctx context.Context, id string) (*core.Group, error) { +func (s *scimGroupRepository) Get(ctx context.Context, id string) (*core.Group, error) { providerID, resourceID, _, err := scimTarget(ctx, id, "") if err != nil { return nil, err @@ -45,12 +45,12 @@ func (s *scimGroups) Get(ctx context.Context, id string) (*core.Group, error) { db := s.api.db.WithContext(ctx) row, err := models.FindSCIMGroup(db, providerID, resourceID) if err != nil { - return nil, scimTranslate(err) + return nil, scimError(err) } return s.renderOne(db, providerID, row, protocol.ProjectionFrom(ctx)) } -func (s *scimGroups) Create(ctx context.Context, group *core.Group) (*core.Group, error) { +func (s *scimGroupRepository) Create(ctx context.Context, group *core.Group) (*core.Group, error) { providerID, err := scimProviderID(ctx) if err != nil { return nil, err @@ -61,7 +61,7 @@ func (s *scimGroups) Create(ctx context.Context, group *core.Group) (*core.Group }) } -func (s *scimGroups) Replace(ctx context.Context, group *core.Group) (*core.Group, error) { +func (s *scimGroupRepository) Replace(ctx context.Context, group *core.Group) (*core.Group, error) { providerID, id, updatedAt, err := scimTarget(ctx, group.ID, group.Meta.Version) if err != nil { return nil, err @@ -76,7 +76,7 @@ func (s *scimGroups) Replace(ctx context.Context, group *core.Group) (*core.Grou }) } -func (s *scimGroups) Delete(ctx context.Context, id, version string) error { +func (s *scimGroupRepository) Delete(ctx context.Context, id, version string) error { providerID, resourceID, updatedAt, err := scimTarget(ctx, id, version) if err != nil { return err @@ -85,7 +85,7 @@ func (s *scimGroups) Delete(ctx context.Context, id, version string) error { if err != nil { return err } - return scimTranslate(s.api.db.WithContext(ctx).Transaction(func(tx *storage.Connection) error { + return scimError(s.api.db.WithContext(ctx).Transaction(func(tx *storage.Connection) error { row, err := models.FindSCIMGroupForUpdate(tx, providerID, resourceID) if err != nil { return err @@ -104,7 +104,7 @@ func (s *scimGroups) Delete(ctx context.Context, id, version string) error { })) } -func (s *scimGroups) save(ctx context.Context, providerID uuid.UUID, action models.AuditAction, group *core.Group, write func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error)) (*core.Group, error) { +func (s *scimGroupRepository) save(ctx context.Context, providerID uuid.UUID, action models.AuditAction, group *core.Group, write func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error)) (*core.Group, error) { members, err := scimMemberIDs(group.Members) if err != nil { return nil, err @@ -144,12 +144,12 @@ func (s *scimGroups) save(ctx context.Context, providerID uuid.UUID, action mode return s.auditMembers(tx, r, row, added, removed) }) if err != nil { - return nil, scimTranslate(err) + return nil, scimError(err) } return s.renderOne(db, providerID, row, protocol.Projection{}) } -func (s *scimGroups) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMGroup, projection protocol.Projection) ([]*core.Group, error) { +func (s *scimGroupRepository) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMGroup, projection protocol.Projection) ([]*core.Group, error) { memberships := []models.SCIMGroupMembership{} if projection.Returns("members") { ids := make([]uuid.UUID, len(rows)) @@ -186,7 +186,7 @@ func (s *scimGroups) render(tx *storage.Connection, providerID uuid.UUID, rows [ return groups, nil } -func (s *scimGroups) renderOne(tx *storage.Connection, providerID uuid.UUID, row *models.SCIMGroup, projection protocol.Projection) (*core.Group, error) { +func (s *scimGroupRepository) renderOne(tx *storage.Connection, providerID uuid.UUID, row *models.SCIMGroup, projection protocol.Projection) (*core.Group, error) { groups, err := s.render(tx, providerID, []models.SCIMGroup{*row}, projection) if err != nil { return nil, err @@ -194,7 +194,7 @@ func (s *scimGroups) renderOne(tx *storage.Connection, providerID uuid.UUID, row return groups[0], nil } -func (s *scimGroups) audit(tx *storage.Connection, r *http.Request, action models.AuditAction, row *models.SCIMGroup) error { +func (s *scimGroupRepository) audit(tx *storage.Connection, r *http.Request, action models.AuditAction, row *models.SCIMGroup) error { var resource struct { DisplayName string `json:"displayName"` } @@ -207,7 +207,7 @@ func (s *scimGroups) audit(tx *storage.Connection, r *http.Request, action model }) } -func (s *scimGroups) auditMembers(tx *storage.Connection, r *http.Request, row *models.SCIMGroup, added, removed []uuid.UUID) error { +func (s *scimGroupRepository) auditMembers(tx *storage.Connection, r *http.Request, row *models.SCIMGroup, added, removed []uuid.UUID) error { links, err := models.FindSCIMUserLinks(tx, append(append([]uuid.UUID{}, added...), removed...)) if err != nil { return err diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index 3a8b3c8023..ea423c4fed 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -20,11 +20,11 @@ import ( "github.com/supabase/auth/internal/storage" ) -type scimUsers struct { +type scimUserRepository struct { api *API } -func (s *scimUsers) List(ctx context.Context, query *protocol.SearchRequest) ([]*core.User, int, error) { +func (s *scimUserRepository) List(ctx context.Context, query *protocol.SearchRequest) ([]*core.User, int, error) { providerID, err := scimProviderID(ctx) if err != nil { return nil, 0, err @@ -45,7 +45,7 @@ func (s *scimUsers) List(ctx context.Context, query *protocol.SearchRequest) ([] return users, total, nil } -func (s *scimUsers) Get(ctx context.Context, id string) (*core.User, error) { +func (s *scimUserRepository) Get(ctx context.Context, id string) (*core.User, error) { providerID, resourceID, _, err := scimTarget(ctx, id, "") if err != nil { return nil, err @@ -53,12 +53,12 @@ func (s *scimUsers) Get(ctx context.Context, id string) (*core.User, error) { db := s.api.db.WithContext(ctx) row, err := models.FindSCIMUser(db, providerID, resourceID) if err != nil { - return nil, scimTranslate(err) + return nil, scimError(err) } return s.renderOne(db, providerID, row, protocol.ProjectionFrom(ctx)) } -func (s *scimUsers) Create(ctx context.Context, user *core.User) (*core.User, error) { +func (s *scimUserRepository) Create(ctx context.Context, user *core.User) (*core.User, error) { providerID, err := scimProviderID(ctx) if err != nil { return nil, err @@ -67,10 +67,10 @@ func (s *scimUsers) Create(ctx context.Context, user *core.User) (*core.User, er if err != nil { return nil, err } - if err := scimValidateEmails(user); err != nil { + if err := scimValidatePrimaryEmail(user); err != nil { return nil, err } - if scimEmail(user) == "" { + if scimUserEmail(user) == "" { return nil, errSCIMEmailRequired() } r, err := scimRequest(ctx) @@ -78,33 +78,33 @@ func (s *scimUsers) Create(ctx context.Context, user *core.User) (*core.User, er return nil, err } db := s.api.db.WithContext(ctx) - if err := s.beforeCreate(r, db, providerID, user); err != nil { - return nil, scimTranslate(err) + if err := s.runBeforeUserCreatedHook(r, db, providerID, user); err != nil { + return nil, scimError(err) } var row *models.SCIMUser var created *models.User err = db.Transaction(func(tx *storage.Connection) error { - if terr := models.LockAccountLinking(tx, "sso:"+providerID.String(), scimEmail(user)); terr != nil { + if terr := models.LockAccountLinking(tx, "sso:"+providerID.String(), scimUserEmail(user)); terr != nil { return terr } var terr error if row, terr = models.CreateSCIMUser(tx, providerID, resource); terr != nil { return terr } - if created, terr = s.linkNew(tx, row, user); terr != nil { + if created, terr = s.provisionAuthUser(tx, row, user); terr != nil { return terr } return s.audit(tx, r, models.SCIMUserCreatedAction, row) }) if err != nil { - return nil, scimTranslate(err) + return nil, scimError(err) } - s.afterCreate(r, db, created) + s.runAfterUserCreatedHook(r, db, created) return s.renderOne(db, providerID, row, protocol.Projection{}) } -func (s *scimUsers) Replace(ctx context.Context, user *core.User) (*core.User, error) { +func (s *scimUserRepository) Replace(ctx context.Context, user *core.User) (*core.User, error) { providerID, id, updatedAt, err := scimTarget(ctx, user.ID, user.Meta.Version) if err != nil { return nil, err @@ -113,10 +113,10 @@ func (s *scimUsers) Replace(ctx context.Context, user *core.User) (*core.User, e if err != nil { return nil, err } - if err := scimValidateEmails(user); err != nil { + if err := scimValidatePrimaryEmail(user); err != nil { return nil, err } - email := scimEmail(user) + email := scimUserEmail(user) r, err := scimRequest(ctx) if err != nil { return nil, err @@ -124,14 +124,14 @@ func (s *scimUsers) Replace(ctx context.Context, user *core.User) (*core.User, e db := s.api.db.WithContext(ctx) existing, err := models.FindSCIMUser(db, providerID, id) if err != nil { - return nil, scimTranslate(err) + return nil, scimError(err) } if existing.UserID == nil { if email == "" { return nil, errSCIMEmailRequired() } - if err := s.beforeCreate(r, db, providerID, user); err != nil { - return nil, scimTranslate(err) + if err := s.runBeforeUserCreatedHook(r, db, providerID, user); err != nil { + return nil, scimError(err) } } @@ -161,19 +161,19 @@ func (s *scimUsers) Replace(ctx context.Context, user *core.User) (*core.User, e if row, terr = models.ReplaceSCIMUser(tx, providerID, id, resource, updatedAt); terr != nil { return terr } - if created, terr = s.sync(tx, providerID, old, row, user); terr != nil { + if created, terr = s.syncAuthUser(tx, providerID, old, row, user); terr != nil { return terr } return s.audit(tx, r, scimUserAuditAction(old, row), row) }) if err != nil { - return nil, scimTranslate(err) + return nil, scimError(err) } - s.afterCreate(r, db, created) + s.runAfterUserCreatedHook(r, db, created) return s.renderOne(db, providerID, row, protocol.Projection{}) } -func (s *scimUsers) Delete(ctx context.Context, id, version string) error { +func (s *scimUserRepository) Delete(ctx context.Context, id, version string) error { providerID, resourceID, updatedAt, err := scimTarget(ctx, id, version) if err != nil { return err @@ -185,9 +185,9 @@ func (s *scimUsers) Delete(ctx context.Context, id, version string) error { db := s.api.db.WithContext(ctx) existing, err := models.FindSCIMUser(db, providerID, resourceID) if err != nil { - return scimTranslate(err) + return scimError(err) } - return scimTranslate(db.Transaction(func(tx *storage.Connection) error { + return scimError(db.Transaction(func(tx *storage.Connection) error { if existing.UserID != nil { if err := models.LockUserForSCIM(tx, *existing.UserID); err != nil { return err @@ -209,7 +209,7 @@ func (s *scimUsers) Delete(ctx context.Context, id, version string) error { })) } -func (s *scimUsers) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMUser, projection protocol.Projection) ([]*core.User, error) { +func (s *scimUserRepository) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMUser, projection protocol.Projection) ([]*core.User, error) { memberships := []models.SCIMGroupMembership{} if projection.Returns("groups") { ids := make([]uuid.UUID, len(rows)) @@ -251,7 +251,7 @@ func (s *scimUsers) render(tx *storage.Connection, providerID uuid.UUID, rows [] return users, nil } -func (s *scimUsers) renderOne(tx *storage.Connection, providerID uuid.UUID, row *models.SCIMUser, projection protocol.Projection) (*core.User, error) { +func (s *scimUserRepository) renderOne(tx *storage.Connection, providerID uuid.UUID, row *models.SCIMUser, projection protocol.Projection) (*core.User, error) { users, err := s.render(tx, providerID, []models.SCIMUser{*row}, projection) if err != nil { return nil, err @@ -259,12 +259,12 @@ func (s *scimUsers) renderOne(tx *storage.Connection, providerID uuid.UUID, row return users[0], nil } -func (s *scimUsers) sync(tx *storage.Connection, providerID uuid.UUID, old, row *models.SCIMUser, user *core.User) (*models.User, error) { +func (s *scimUserRepository) syncAuthUser(tx *storage.Connection, providerID uuid.UUID, old, row *models.SCIMUser, user *core.User) (*models.User, error) { if old.UserID == nil { - if scimEmail(user) == "" { + if scimUserEmail(user) == "" { return nil, errSCIMEmailRequired() } - return s.linkNew(tx, row, user) + return s.provisionAuthUser(tx, row, user) } linked, err := models.FindUserByID(tx, *old.UserID) @@ -273,7 +273,7 @@ func (s *scimUsers) sync(tx *storage.Connection, providerID uuid.UUID, old, row } if from := scimUserName(old.Resource); from != user.UserName { data := map[string]any{"sub": user.UserName} - if email := scimEmail(user); email != "" { + if email := scimUserEmail(user); email != "" { data["email"] = email } err := models.RenameSCIMIdentity(tx, linked.ID, "sso:"+providerID.String(), from, user.UserName, data) @@ -289,8 +289,8 @@ func (s *scimUsers) sync(tx *storage.Connection, providerID uuid.UUID, old, row return nil, nil } -func (s *scimUsers) linkNew(tx *storage.Connection, row *models.SCIMUser, user *core.User) (*models.User, error) { - linked, isNew, err := s.link(tx, row, user) +func (s *scimUserRepository) provisionAuthUser(tx *storage.Connection, row *models.SCIMUser, user *core.User) (*models.User, error) { + linked, isNew, err := s.linkAuthUser(tx, row, user) if err != nil { return nil, err } @@ -304,9 +304,9 @@ func (s *scimUsers) linkNew(tx *storage.Connection, row *models.SCIMUser, user * return created, nil } -func (s *scimUsers) link(tx *storage.Connection, row *models.SCIMUser, user *core.User) (*models.User, bool, error) { +func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SCIMUser, user *core.User) (*models.User, bool, error) { providerType := "sso:" + row.SSOProviderID.String() - decision, err := s.decide(tx, providerType, user) + decision, err := s.decideAccountLinking(tx, providerType, user) if err != nil { return nil, false, err } @@ -341,12 +341,12 @@ func (s *scimUsers) link(tx *storage.Connection, row *models.SCIMUser, user *cor return linked, false, models.LinkSCIMUser(tx, row, linked.ID) } -func (s *scimUsers) beforeCreate(r *http.Request, db *storage.Connection, providerID uuid.UUID, user *core.User) error { +func (s *scimUserRepository) runBeforeUserCreatedHook(r *http.Request, db *storage.Connection, providerID uuid.UUID, user *core.User) error { if !s.api.hooksMgr.Enabled(v0hooks.BeforeUserCreated) { return nil } providerType := "sso:" + providerID.String() - decision, err := s.decide(db, providerType, user) + decision, err := s.decideAccountLinking(db, providerType, user) if err != nil || decision.Decision != models.CreateAccount { return err } @@ -357,7 +357,7 @@ func (s *scimUsers) beforeCreate(r *http.Request, db *storage.Connection, provid return scimHookError(s.api.triggerBeforeUserCreated(r, db, candidate)) } -func (s *scimUsers) afterCreate(r *http.Request, db *storage.Connection, user *models.User) { +func (s *scimUserRepository) runAfterUserCreatedHook(r *http.Request, db *storage.Connection, user *models.User) { if user == nil { return } @@ -366,12 +366,12 @@ func (s *scimUsers) afterCreate(r *http.Request, db *storage.Connection, user *m } } -func (s *scimUsers) decide(conn *storage.Connection, providerType string, user *core.User) (models.AccountLinkingResult, error) { - emails := []provider.Email{{Email: scimEmail(user), Verified: true, Primary: true}} +func (s *scimUserRepository) decideAccountLinking(conn *storage.Connection, providerType string, user *core.User) (models.AccountLinkingResult, error) { + emails := []provider.Email{{Email: scimUserEmail(user), Verified: true, Primary: true}} return models.DetermineAccountLinking(conn, s.api.config, emails, s.api.config.JWT.Aud, providerType, user.UserName) } -func (s *scimUsers) newUser(providerType string, decision models.AccountLinkingResult, user *core.User) (*models.User, error) { +func (s *scimUserRepository) newUser(providerType string, decision models.AccountLinkingResult, user *core.User) (*models.User, error) { params := &SignupParams{ Provider: providerType, Email: decision.CandidateEmail.Email, @@ -387,7 +387,7 @@ func (s *scimUsers) newUser(providerType string, decision models.AccountLinkingR return candidate, nil } -func (s *scimUsers) audit(tx *storage.Connection, r *http.Request, action models.AuditAction, row *models.SCIMUser) error { +func (s *scimUserRepository) audit(tx *storage.Connection, r *http.Request, action models.AuditAction, row *models.SCIMUser) error { return s.api.auditSCIM(tx, r, scimActor(r), action, row.SSOProviderID, scimUserTraits(row)) } @@ -424,7 +424,7 @@ func scimUserResource(user *core.User) ([]byte, error) { return scimEncode(user, "id", "meta", "password", "groups") } -func scimEmail(user *core.User) string { +func scimUserEmail(user *core.User) string { if email := scimPrimaryEmail(user.Emails); email != "" { return email } @@ -434,7 +434,7 @@ func scimEmail(user *core.User) string { return "" } -func scimValidateEmails(user *core.User) error { +func scimValidatePrimaryEmail(user *core.User) error { if email := scimPrimaryEmail(user.Emails); email != "" && !isEmailAddress(email) { return errSCIMEmailInvalid() } @@ -460,7 +460,7 @@ func scimPrimaryEmail(emails []core.Email) string { func scimIdentityData(user *core.User) map[string]any { return map[string]any{ "sub": user.UserName, - "email": scimEmail(user), + "email": scimUserEmail(user), "email_verified": true, } } diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index 55c59ad791..9b73a6a124 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -142,7 +142,7 @@ func (ts *SCIMUsersTestSuite) repository() (context.Context, server.Repository[* ctx, err := newSCIMTokenValidator(ts.API.db)(context.Background(), ts.TokenA) require.NoError(ts.T(), err) ctx = scimRequestKey.WithValue(ctx, httptest.NewRequest(http.MethodPost, "/scim/v2/Users", nil)) - return ctx, &scimUsers{api: ts.API} + return ctx, &scimUserRepository{api: ts.API} } func emails(value string) []core.Email { @@ -715,7 +715,7 @@ func (ts *SCIMUsersTestSuite) TestUnknownID() { } func (ts *SCIMUsersTestSuite) TestRequiresSSOProviderOnContext() { - users := &scimUsers{api: ts.API} + users := &scimUserRepository{api: ts.API} _, _, err := users.List(context.Background(), &protocol.SearchRequest{Count: 10}) require.Error(ts.T(), err) diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index 49eb4353db..b185644b8e 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -174,7 +174,7 @@ func LinkSCIMUser(tx *storage.Connection, user *SCIMUser, userID uuid.UUID) erro return nil } -func IsSCIMProvisioned(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { +func wasSCIMProvisioned(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { provisioned, err := tx.Q().Where("sso_provider_id = ? AND user_id = ?", providerID, userID).Exists(&SCIMUser{}) if err != nil { return false, errors.Wrap(err, "error finding SCIM user") @@ -191,7 +191,7 @@ func IsSCIMManaged(tx *storage.Connection, providerID, userID uuid.UUID) (bool, } func IsSCIMDeprovisioned(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { - provisioned, err := IsSCIMProvisioned(tx, providerID, userID) + provisioned, err := wasSCIMProvisioned(tx, providerID, userID) if err != nil { return false, err } From 56d1622df0de860eb90543c20dc729010a6d8470 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 21:27:10 -0600 Subject: [PATCH 18/88] chore(scim): order scim.go and scim_admin.go by constructors, methods, then functions --- internal/api/scim.go | 52 ++++++------ internal/api/scim_admin.go | 164 ++++++++++++++++++------------------- 2 files changed, 108 insertions(+), 108 deletions(-) diff --git a/internal/api/scim.go b/internal/api/scim.go index 5c4405fc82..d09a8b7ec4 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -34,6 +34,14 @@ const ( var errMissingSSOProvider = errors.New("scim: request has no SSO provider") +var ( + scimUserSchemas = core.Schemas{ + core.NewSchema(core.SchemaUser).With(core.UserAttributes()...), + core.NewSchema(core.SchemaEnterpriseUser).With(core.EnterpriseUserAttributes()...), + } + scimGroupSchemas = newSCIMGroupSchemas() +) + func (a *API) newSCIMServer(validate server.TokenValidator, limit func(http.Handler) http.Handler) *server.Server { requireToken := server.RequireBearerToken(validate) authenticate := func(next http.Handler) http.Handler { @@ -55,14 +63,6 @@ func (a *API) newSCIMServer(validate server.TokenValidator, limit func(http.Hand ) } -var ( - scimUserSchemas = core.Schemas{ - core.NewSchema(core.SchemaUser).With(core.UserAttributes()...), - core.NewSchema(core.SchemaEnterpriseUser).With(core.EnterpriseUserAttributes()...), - } - scimGroupSchemas = newSCIMGroupSchemas() -) - func newSCIMGroupSchemas() core.Schemas { attributes := core.GroupAttributes() for _, attribute := range attributes { @@ -78,18 +78,6 @@ func newSCIMGroupSchemas() core.Schemas { return core.Schemas{core.NewSchema(core.SchemaGroup).With(attributes...)} } -func scimBaseURL(config *conf.GlobalConfiguration) string { - return strings.TrimRight(config.API.ExternalURL, "/") + scimBasePath -} - -func scimTooManyRequests(w http.ResponseWriter, r *http.Request) error { - return protocol.SendError(w, errSCIMTooManyRequests()) -} - -func scimLogError(r *http.Request, err error) { - observability.GetLogEntry(r).Entry.WithError(err).Error("scim: request failed") -} - func newSCIMTokenValidator(db *storage.Connection) server.TokenValidator { return func(ctx context.Context, candidate string) (context.Context, error) { token, err := models.AuthenticateSCIMToken(db.WithContext(ctx), candidate) @@ -151,6 +139,24 @@ func (a *API) limitSCIMByProvider(lmt *limiter.Limiter) func(http.Handler) http. } } +func (a *API) auditSCIM(tx *storage.Connection, r *http.Request, actor *models.User, action models.AuditAction, providerID uuid.UUID, traits map[string]any) error { + traits["sso_provider_id"] = providerID + traits["outcome"] = "success" + return models.NewAuditLogEntry(a.config.AuditLog, r, tx, actor, action, utilities.GetIPAddress(r), traits) +} + +func scimBaseURL(config *conf.GlobalConfiguration) string { + return strings.TrimRight(config.API.ExternalURL, "/") + scimBasePath +} + +func scimTooManyRequests(w http.ResponseWriter, r *http.Request) error { + return protocol.SendError(w, errSCIMTooManyRequests()) +} + +func scimLogError(r *http.Request, err error) { + observability.GetLogEntry(r).Entry.WithError(err).Error("scim: request failed") +} + func scimSearch(query *protocol.SearchRequest, schemas core.Schemas, name string) (models.SCIMQuery, error) { search := models.SCIMQuery{Offset: query.Offset(), Limit: query.Count} if query.Filter != "" { @@ -300,12 +306,6 @@ func scimActor(r *http.Request) *models.User { return &models.User{Email: storage.NullString("scim:" + prefix)} } -func (a *API) auditSCIM(tx *storage.Connection, r *http.Request, actor *models.User, action models.AuditAction, providerID uuid.UUID, traits map[string]any) error { - traits["sso_provider_id"] = providerID - traits["outcome"] = "success" - return models.NewAuditLogEntry(a.config.AuditLog, r, tx, actor, action, utilities.GetIPAddress(r), traits) -} - func scimMemberTraits(groupID, scimUserID uuid.UUID, userID *uuid.UUID) map[string]any { traits := map[string]any{ "scim_group_id": groupID, diff --git a/internal/api/scim_admin.go b/internal/api/scim_admin.go index 16e7a2daf6..3f522e88db 100644 --- a/internal/api/scim_admin.go +++ b/internal/api/scim_admin.go @@ -51,88 +51,6 @@ func (a *API) adminSCIMDisable(w http.ResponseWriter, r *http.Request) error { return a.changeSCIMEnabled(w, r, models.DisableSCIM, models.SCIMDisabledAction, "disabling") } -func (a *API) changeSCIMEnabled(w http.ResponseWriter, r *http.Request, change func(*storage.Connection, uuid.UUID) (bool, error), action models.AuditAction, verb string) error { - ctx := r.Context() - db := a.db.WithContext(ctx) - provider := getSSOProvider(ctx) - - if err := db.Transaction(func(tx *storage.Connection) error { - changed, err := change(tx, provider.ID) - if err != nil || !changed { - return err - } - return a.auditSCIM(tx, r, getAdminUser(ctx), action, provider.ID, map[string]any{}) - }); err != nil { - return apierrors.NewInternalServerError("Error %s SCIM", verb).WithInternalError(err) - } - - return a.sendSCIMStatus(w, db, provider) -} - -func (a *API) sendSCIMStatus(w http.ResponseWriter, db *storage.Connection, provider *models.SSOProvider) error { - tokens, err := models.FindActiveSCIMTokensBySSOProvider(db, provider.ID) - if err != nil { - return apierrors.NewInternalServerError("Error finding SCIM tokens").WithInternalError(err) - } - enabled, err := a.isSCIMEnabled(db, provider) - if err != nil { - return apierrors.NewInternalServerError("Error finding SCIM settings").WithInternalError(err) - } - - return sendJSON(w, http.StatusOK, &AdminSCIMStatusResponse{ - Enabled: enabled, - BaseURL: scimBaseURL(a.config), - Tokens: tokens, - }) -} - -func (a *API) isSCIMEnabled(db *storage.Connection, provider *models.SSOProvider) (bool, error) { - if !a.config.SSO.SCIM.Enabled || !provider.IsEnabled() { - return false, nil - } - return models.IsSCIMEnabled(db, provider.ID) -} - -func (a *API) deprovisionSCIM(tx *storage.Connection, r *http.Request, provider *models.SSOProvider) error { - enabled, err := models.IsSCIMEnabled(tx, provider.ID) - if err != nil { - return err - } - actor := getAdminUser(r.Context()) - prefixes, err := a.auditSCIMTokensRevoked(tx, r, actor, provider.ID) - if err != nil { - return err - } - if enabled && a.config.SSO.SCIM.Enabled { - if err := a.auditSCIM(tx, r, actor, models.SCIMDisabledAction, provider.ID, map[string]any{"token_prefixes": prefixes}); err != nil { - return err - } - } - banned, err := models.BanDeprovisionedSCIMUsers(tx, provider.ID, a.Now().Add(scimProviderDeletedBan)) - if err != nil || banned == 0 { - return err - } - return a.auditSCIM(tx, r, actor, models.SCIMUsersBannedAction, provider.ID, map[string]any{"banned_user_count": banned}) -} - -func (a *API) auditSCIMTokensRevoked(tx *storage.Connection, r *http.Request, actor *models.User, providerID uuid.UUID) ([]string, error) { - if err := models.LockSCIMTokens(tx, providerID); err != nil { - return nil, err - } - tokens, err := models.FindActiveSCIMTokensBySSOProvider(tx, providerID) - if err != nil { - return nil, err - } - prefixes := make([]string, len(tokens)) - for i := range tokens { - prefixes[i] = tokens[i].Prefix - if err := a.auditSCIM(tx, r, actor, models.SCIMTokenRevokedAction, providerID, map[string]any{scimTokenPrefixTrait: tokens[i].Prefix}); err != nil { - return nil, err - } - } - return prefixes, nil -} - func (a *API) adminSCIMTokensCreate(w http.ResponseWriter, r *http.Request) error { ctx := r.Context() db := a.db.WithContext(ctx) @@ -217,3 +135,85 @@ func (a *API) adminSCIMTokensRevoke(w http.ResponseWriter, r *http.Request) erro return sendJSON(w, http.StatusOK, token) } + +func (a *API) changeSCIMEnabled(w http.ResponseWriter, r *http.Request, change func(*storage.Connection, uuid.UUID) (bool, error), action models.AuditAction, verb string) error { + ctx := r.Context() + db := a.db.WithContext(ctx) + provider := getSSOProvider(ctx) + + if err := db.Transaction(func(tx *storage.Connection) error { + changed, err := change(tx, provider.ID) + if err != nil || !changed { + return err + } + return a.auditSCIM(tx, r, getAdminUser(ctx), action, provider.ID, map[string]any{}) + }); err != nil { + return apierrors.NewInternalServerError("Error %s SCIM", verb).WithInternalError(err) + } + + return a.sendSCIMStatus(w, db, provider) +} + +func (a *API) sendSCIMStatus(w http.ResponseWriter, db *storage.Connection, provider *models.SSOProvider) error { + tokens, err := models.FindActiveSCIMTokensBySSOProvider(db, provider.ID) + if err != nil { + return apierrors.NewInternalServerError("Error finding SCIM tokens").WithInternalError(err) + } + enabled, err := a.isSCIMEnabled(db, provider) + if err != nil { + return apierrors.NewInternalServerError("Error finding SCIM settings").WithInternalError(err) + } + + return sendJSON(w, http.StatusOK, &AdminSCIMStatusResponse{ + Enabled: enabled, + BaseURL: scimBaseURL(a.config), + Tokens: tokens, + }) +} + +func (a *API) isSCIMEnabled(db *storage.Connection, provider *models.SSOProvider) (bool, error) { + if !a.config.SSO.SCIM.Enabled || !provider.IsEnabled() { + return false, nil + } + return models.IsSCIMEnabled(db, provider.ID) +} + +func (a *API) deprovisionSCIM(tx *storage.Connection, r *http.Request, provider *models.SSOProvider) error { + enabled, err := models.IsSCIMEnabled(tx, provider.ID) + if err != nil { + return err + } + actor := getAdminUser(r.Context()) + prefixes, err := a.auditSCIMTokensRevoked(tx, r, actor, provider.ID) + if err != nil { + return err + } + if enabled && a.config.SSO.SCIM.Enabled { + if err := a.auditSCIM(tx, r, actor, models.SCIMDisabledAction, provider.ID, map[string]any{"token_prefixes": prefixes}); err != nil { + return err + } + } + banned, err := models.BanDeprovisionedSCIMUsers(tx, provider.ID, a.Now().Add(scimProviderDeletedBan)) + if err != nil || banned == 0 { + return err + } + return a.auditSCIM(tx, r, actor, models.SCIMUsersBannedAction, provider.ID, map[string]any{"banned_user_count": banned}) +} + +func (a *API) auditSCIMTokensRevoked(tx *storage.Connection, r *http.Request, actor *models.User, providerID uuid.UUID) ([]string, error) { + if err := models.LockSCIMTokens(tx, providerID); err != nil { + return nil, err + } + tokens, err := models.FindActiveSCIMTokensBySSOProvider(tx, providerID) + if err != nil { + return nil, err + } + prefixes := make([]string, len(tokens)) + for i := range tokens { + prefixes[i] = tokens[i].Prefix + if err := a.auditSCIM(tx, r, actor, models.SCIMTokenRevokedAction, providerID, map[string]any{scimTokenPrefixTrait: tokens[i].Prefix}); err != nil { + return nil, err + } + } + return prefixes, nil +} From 39e2f9d69a7051b4d7d0d041afbe9db83a0a2b20 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 21:32:20 -0600 Subject: [PATCH 19/88] chore(scim): test that /Me, /Bulk and .search return 501 --- internal/api/scim_users_test.go | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index 9b73a6a124..80537f260b 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -235,6 +235,21 @@ func (ts *SCIMUsersTestSuite) TestOktaContentTypesAndReactivate() { } } +func (ts *SCIMUsersTestSuite) TestUnsupportedEndpointsReturnNotImplemented() { + for _, tc := range []struct{ method, path string }{ + {http.MethodGet, "/Me"}, + {http.MethodPost, "/Bulk"}, + {http.MethodPost, "/.search"}, + {http.MethodPost, "/Users/.search"}, + {http.MethodPost, "/Groups/.search"}, + } { + w, _ := ts.do(ts.TokenA, tc.method, tc.path, "{}") + require.Equal(ts.T(), http.StatusNotImplemented, w.Code, tc.path) + require.Equal(ts.T(), protocol.MediaType, w.Header().Get("Content-Type"), tc.path) + require.Contains(ts.T(), w.Body.String(), protocol.SchemaError, tc.path) + } +} + func (ts *SCIMUsersTestSuite) TestUniquenessWithinProvider() { ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) From 4a1aee8b8ca6bf6c69a808b10fa6561a7b5c3490 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 21:33:35 -0600 Subject: [PATCH 20/88] docs(scim): say unknown routes with a valid token count against the provider limit --- README.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/README.md b/README.md index 21cfd47f02..d1bf238a5c 100644 --- a/README.md +++ b/README.md @@ -218,7 +218,7 @@ Mounts the SCIM 2.0 routes at `/scim/v2` and the SCIM admin routes at `/admin/ss `GOTRUE_RATE_LIMIT_SCIM` - `number` -Requests per 5 minutes to `/scim/v2`, with a burst of 30. Requests with a valid SCIM token are limited per SSO provider. Requests without a valid token, and requests to `/ServiceProviderConfig` or to an unknown route, are limited per IP. Defaults to 3000. +Requests per 5 minutes to `/scim/v2`, with a burst of 30. Requests with a valid SCIM token are limited per SSO provider. Requests without a valid token, and requests to `/ServiceProviderConfig`, are limited per IP. Defaults to 3000. `GOTRUE_PASSWORD_MIN_LENGTH` - `int` From 524098ed905751a8a2fa00fe6f792c0fffc80c59 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 21:55:22 -0600 Subject: [PATCH 21/88] chore(scim): share admin jwt and audit query test helpers --- internal/api/scim_admin_test.go | 10 +++------- internal/api/scim_provider_delete_test.go | 8 ++------ internal/api/scim_users_test.go | 17 ++++++++++++++--- 3 files changed, 19 insertions(+), 16 deletions(-) diff --git a/internal/api/scim_admin_test.go b/internal/api/scim_admin_test.go index 5281592180..ec8dba0433 100644 --- a/internal/api/scim_admin_test.go +++ b/internal/api/scim_admin_test.go @@ -41,9 +41,7 @@ func (ts *SCIMTokensTestSuite) SetupTest() { require.NoError(ts.T(), models.TruncateAll(ts.API.db)) ts.API.config.SSO.SCIM.Enabled = true - token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, &AccessTokenClaims{Role: "supabase_admin"}).SignedString([]byte(ts.Config.JWT.Secret)) - require.NoError(ts.T(), err) - ts.AdminJWT = token + ts.AdminJWT = adminJWT(ts.T(), ts.Config.JWT.Secret) ts.Provider = ts.createProvider() } @@ -449,8 +447,7 @@ func (ts *SCIMTokensTestSuite) TestConcurrentEnableAndDisable() { } func (ts *SCIMTokensTestSuite) scimActions(provider *models.SSOProvider) []string { - entries := []models.AuditLogEntry{} - require.NoError(ts.T(), ts.API.db.Q().Where("payload->>'log_type' = ? AND payload->'traits'->>'sso_provider_id' = ?", "scim", provider.ID.String()).Order("created_at asc").All(&entries)) + entries := queryAuditEntries(ts.T(), ts.API.db, "payload->>'log_type' = ? AND payload->'traits'->>'sso_provider_id' = ?", "scim", provider.ID.String()) actions := []string{} for _, entry := range entries { actions = append(actions, entry.Payload["action"].(string)) @@ -508,8 +505,7 @@ func (ts *SCIMTokensTestSuite) TestStatusForDisabledProvider() { type scimTokenEvent struct{ action, prefix string } func (ts *SCIMTokensTestSuite) tokenEvents() []scimTokenEvent { - entries := []models.AuditLogEntry{} - require.NoError(ts.T(), ts.API.db.Q().Where("payload->>'log_type' = ?", "scim").Order("created_at asc").All(&entries)) + entries := queryAuditEntries(ts.T(), ts.API.db, "payload->>'log_type' = ?", "scim") events := []scimTokenEvent{} for _, entry := range entries { diff --git a/internal/api/scim_provider_delete_test.go b/internal/api/scim_provider_delete_test.go index b1b37489e2..52263779fc 100644 --- a/internal/api/scim_provider_delete_test.go +++ b/internal/api/scim_provider_delete_test.go @@ -6,7 +6,6 @@ import ( "time" "github.com/gofrs/uuid" - jwt "github.com/golang-jwt/jwt/v5" "github.com/stretchr/testify/require" "github.com/supabase/auth/internal/api/provider" "github.com/supabase/auth/internal/models" @@ -18,8 +17,7 @@ func scimUser(name string) string { } func (ts *SCIMUsersTestSuite) deleteProvider(p *models.SSOProvider) { - token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, &AccessTokenClaims{Role: "supabase_admin"}).SignedString([]byte(ts.API.config.JWT.Secret)) - require.NoError(ts.T(), err) + token := adminJWT(ts.T(), ts.API.config.JWT.Secret) r := httptest.NewRequest(http.MethodDelete, "/admin/sso/providers/"+p.ID.String(), nil) r.Header.Set("Authorization", "Bearer "+token) w := httptest.NewRecorder() @@ -40,9 +38,7 @@ func (ts *SCIMUsersTestSuite) countRows(model any, where string, args ...any) in } func (ts *SCIMUsersTestSuite) auditActions(action models.AuditAction) []models.AuditLogEntry { - entries := []models.AuditLogEntry{} - require.NoError(ts.T(), ts.API.db.Q().Where("payload->>'action' = ?", string(action)).All(&entries)) - return entries + return queryAuditEntries(ts.T(), ts.API.db, "payload->>'action' = ?", string(action)) } func (ts *SCIMUsersTestSuite) TestProviderDeleteBansDeprovisionedUsers() { diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index 80537f260b..a86e964d35 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -16,6 +16,7 @@ import ( "time" "github.com/gofrs/uuid" + jwt "github.com/golang-jwt/jwt/v5" "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" "github.com/supabase-community/scim-go/pkg/core" @@ -101,6 +102,18 @@ func createSSOProviderWithSCIMToken(t require.TestingT, db *storage.Connection) return provider, token } +func adminJWT(t require.TestingT, secret string) string { + token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, &AccessTokenClaims{Role: "supabase_admin"}).SignedString([]byte(secret)) + require.NoError(t, err) + return token +} + +func queryAuditEntries(t require.TestingT, db *storage.Connection, where string, args ...any) []models.AuditLogEntry { + entries := []models.AuditLogEntry{} + require.NoError(t, db.Q().Where(where, args...).Order("created_at asc").All(&entries)) + return entries +} + func (ts *SCIMUsersTestSuite) do(token, method, path, body string) (*httptest.ResponseRecorder, map[string]any) { return ts.doAs(protocol.MediaType, token, method, path, body) } @@ -378,9 +391,7 @@ func (ts *SCIMUsersTestSuite) whileLocked(lock, finish func(tx *storage.Connecti } func (ts *SCIMUsersTestSuite) scimAuditEntries() []models.AuditLogEntry { - entries := []models.AuditLogEntry{} - require.NoError(ts.T(), ts.API.db.Q().Where("payload->>'log_type' = ?", "scim").Order("created_at asc").All(&entries)) - return entries + return queryAuditEntries(ts.T(), ts.API.db, "payload->>'log_type' = ?", "scim") } func (ts *SCIMUsersTestSuite) TestAuditLog() { From f940242587fef07d98e18ebe0935fe32e4ea9e77 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:06:22 -0600 Subject: [PATCH 22/88] chore(scim): collapse SCIM deprovisioned check into a single query --- internal/models/scim_user.go | 32 +++++++++++++++----------------- 1 file changed, 15 insertions(+), 17 deletions(-) diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index b185644b8e..f0b6a812e4 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -174,14 +174,6 @@ func LinkSCIMUser(tx *storage.Connection, user *SCIMUser, userID uuid.UUID) erro return nil } -func wasSCIMProvisioned(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { - provisioned, err := tx.Q().Where("sso_provider_id = ? AND user_id = ?", providerID, userID).Exists(&SCIMUser{}) - if err != nil { - return false, errors.Wrap(err, "error finding SCIM user") - } - return provisioned, nil -} - func IsSCIMManaged(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { managed, err := tx.Q().Where("sso_provider_id = ? AND user_id = ? AND deleted_at IS NULL", providerID, userID).Exists(&SCIMUser{}) if err != nil { @@ -191,18 +183,24 @@ func IsSCIMManaged(tx *storage.Connection, providerID, userID uuid.UUID) (bool, } func IsSCIMDeprovisioned(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { - provisioned, err := wasSCIMProvisioned(tx, providerID, userID) - if err != nil { - return false, err + result := struct { + AnyRow bool `db:"any_row"` + Live bool `db:"live"` + }{} + if err := tx.RawQuery( + fmt.Sprintf( + "SELECT EXISTS(SELECT 1 FROM %q WHERE sso_provider_id = ? AND user_id = ?) AS any_row, "+ + "EXISTS(SELECT 1 FROM %q WHERE sso_provider_id = ? AND user_id = ? AND deleted_at IS NULL AND active) AS live", + (&SCIMUser{}).TableName(), (&SCIMUser{}).TableName(), + ), + providerID, userID, providerID, userID, + ).First(&result); err != nil { + return false, errors.Wrap(err, "error finding SCIM user") } - if !provisioned { + if !result.AnyRow { return false, nil } - live, err := tx.Q().Where("sso_provider_id = ? AND user_id = ? AND deleted_at IS NULL AND active", providerID, userID).Exists(&SCIMUser{}) - if err != nil { - return false, errors.Wrap(err, "error finding live SCIM user") - } - return !live, nil + return !result.Live, nil } func IsSCIMUserDeprovisionedForUpdate(tx *storage.Connection, userID uuid.UUID) (bool, error) { From 41da9b6f352a410e553399f036c1a8d9139007af Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:16:56 -0600 Subject: [PATCH 23/88] chore(scim): pass a SCIMTarget to the SCIM row writers --- internal/api/scim.go | 19 ++++---- internal/api/scim_groups.go | 20 ++++---- internal/api/scim_users.go | 23 ++++----- internal/models/scim.go | 77 ++++++++++++++++++++++-------- internal/models/scim_group.go | 26 ++++------ internal/models/scim_group_test.go | 18 +++---- internal/models/scim_user.go | 26 ++++------ 7 files changed, 117 insertions(+), 92 deletions(-) diff --git a/internal/api/scim.go b/internal/api/scim.go index d09a8b7ec4..e07862430b 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -201,17 +201,20 @@ func scimEncode(resource core.Resource, drop ...string) ([]byte, error) { return json.Marshal(fields) } -func scimTarget(ctx context.Context, id, version string) (providerID, resourceID uuid.UUID, updatedAt *time.Time, err error) { - if providerID, err = scimProviderID(ctx); err != nil { - return uuid.Nil, uuid.Nil, nil, err +func scimTarget(ctx context.Context, id, version string) (models.SCIMTarget, error) { + providerID, err := scimProviderID(ctx) + if err != nil { + return models.SCIMTarget{}, err } - if resourceID, err = uuid.FromString(id); err != nil { - return uuid.Nil, uuid.Nil, nil, errSCIMNotFound() + resourceID, err := uuid.FromString(id) + if err != nil { + return models.SCIMTarget{}, errSCIMNotFound() } - if updatedAt, err = scimParseVersion(version); err != nil { - return uuid.Nil, uuid.Nil, nil, err + updatedAt, err := scimParseVersion(version) + if err != nil { + return models.SCIMTarget{}, err } - return providerID, resourceID, updatedAt, nil + return models.SCIMTarget{ProviderID: providerID, ID: resourceID, UpdatedAt: updatedAt}, nil } func scimProviderID(ctx context.Context) (uuid.UUID, error) { diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index 67f2db5dfe..85bccbd67a 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -38,16 +38,16 @@ func (s *scimGroupRepository) List(ctx context.Context, query *protocol.SearchRe } func (s *scimGroupRepository) Get(ctx context.Context, id string) (*core.Group, error) { - providerID, resourceID, _, err := scimTarget(ctx, id, "") + target, err := scimTarget(ctx, id, "") if err != nil { return nil, err } db := s.api.db.WithContext(ctx) - row, err := models.FindSCIMGroup(db, providerID, resourceID) + row, err := models.FindSCIMGroup(db, target.ProviderID, target.ID) if err != nil { return nil, scimError(err) } - return s.renderOne(db, providerID, row, protocol.ProjectionFrom(ctx)) + return s.renderOne(db, target.ProviderID, row, protocol.ProjectionFrom(ctx)) } func (s *scimGroupRepository) Create(ctx context.Context, group *core.Group) (*core.Group, error) { @@ -62,22 +62,22 @@ func (s *scimGroupRepository) Create(ctx context.Context, group *core.Group) (*c } func (s *scimGroupRepository) Replace(ctx context.Context, group *core.Group) (*core.Group, error) { - providerID, id, updatedAt, err := scimTarget(ctx, group.ID, group.Meta.Version) + target, err := scimTarget(ctx, group.ID, group.Meta.Version) if err != nil { return nil, err } - return s.save(ctx, providerID, models.SCIMGroupUpdatedAction, group, func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error) { - unchanged, err := models.FindUnchangedSCIMGroup(tx, providerID, id, resource, updatedAt) + return s.save(ctx, target.ProviderID, models.SCIMGroupUpdatedAction, group, func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error) { + unchanged, err := models.FindUnchangedSCIMGroup(tx, target, resource) if err != nil || unchanged != nil { return unchanged, false, err } - row, err := models.ReplaceSCIMGroup(tx, providerID, id, resource, updatedAt) + row, err := models.ReplaceSCIMGroup(tx, target, resource) return row, true, err }) } func (s *scimGroupRepository) Delete(ctx context.Context, id, version string) error { - providerID, resourceID, updatedAt, err := scimTarget(ctx, id, version) + target, err := scimTarget(ctx, id, version) if err != nil { return err } @@ -86,7 +86,7 @@ func (s *scimGroupRepository) Delete(ctx context.Context, id, version string) er return err } return scimError(s.api.db.WithContext(ctx).Transaction(func(tx *storage.Connection) error { - row, err := models.FindSCIMGroupForUpdate(tx, providerID, resourceID) + row, err := models.FindSCIMGroupForUpdate(tx, target.ProviderID, target.ID) if err != nil { return err } @@ -94,7 +94,7 @@ func (s *scimGroupRepository) Delete(ctx context.Context, id, version string) er if err != nil { return err } - if row, err = models.DeleteSCIMGroup(tx, providerID, resourceID, updatedAt); err != nil { + if row, err = models.DeleteSCIMGroup(tx, target); err != nil { return err } if err := s.auditMembers(tx, r, row, nil, removed); err != nil { diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index ea423c4fed..e8cc9ee6db 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -46,16 +46,16 @@ func (s *scimUserRepository) List(ctx context.Context, query *protocol.SearchReq } func (s *scimUserRepository) Get(ctx context.Context, id string) (*core.User, error) { - providerID, resourceID, _, err := scimTarget(ctx, id, "") + target, err := scimTarget(ctx, id, "") if err != nil { return nil, err } db := s.api.db.WithContext(ctx) - row, err := models.FindSCIMUser(db, providerID, resourceID) + row, err := models.FindSCIMUser(db, target.ProviderID, target.ID) if err != nil { return nil, scimError(err) } - return s.renderOne(db, providerID, row, protocol.ProjectionFrom(ctx)) + return s.renderOne(db, target.ProviderID, row, protocol.ProjectionFrom(ctx)) } func (s *scimUserRepository) Create(ctx context.Context, user *core.User) (*core.User, error) { @@ -105,10 +105,11 @@ func (s *scimUserRepository) Create(ctx context.Context, user *core.User) (*core } func (s *scimUserRepository) Replace(ctx context.Context, user *core.User) (*core.User, error) { - providerID, id, updatedAt, err := scimTarget(ctx, user.ID, user.Meta.Version) + target, err := scimTarget(ctx, user.ID, user.Meta.Version) if err != nil { return nil, err } + providerID := target.ProviderID resource, err := scimUserResource(user) if err != nil { return nil, err @@ -122,7 +123,7 @@ func (s *scimUserRepository) Replace(ctx context.Context, user *core.User) (*cor return nil, err } db := s.api.db.WithContext(ctx) - existing, err := models.FindSCIMUser(db, providerID, id) + existing, err := models.FindSCIMUser(db, providerID, target.ID) if err != nil { return nil, scimError(err) } @@ -146,7 +147,7 @@ func (s *scimUserRepository) Replace(ctx context.Context, user *core.User) (*cor return terr } } - old, terr := models.FindSCIMUserForUpdate(tx, providerID, id) + old, terr := models.FindSCIMUserForUpdate(tx, providerID, target.ID) if terr != nil { return terr } @@ -154,11 +155,11 @@ func (s *scimUserRepository) Replace(ctx context.Context, user *core.User) (*cor if terr := models.LockUserForSCIM(tx, *old.UserID); terr != nil { return terr } - if row, terr = models.FindUnchangedSCIMUser(tx, providerID, id, resource, updatedAt); terr != nil || row != nil { + if row, terr = models.FindUnchangedSCIMUser(tx, target, resource); terr != nil || row != nil { return terr } } - if row, terr = models.ReplaceSCIMUser(tx, providerID, id, resource, updatedAt); terr != nil { + if row, terr = models.ReplaceSCIMUser(tx, target, resource); terr != nil { return terr } if created, terr = s.syncAuthUser(tx, providerID, old, row, user); terr != nil { @@ -174,7 +175,7 @@ func (s *scimUserRepository) Replace(ctx context.Context, user *core.User) (*cor } func (s *scimUserRepository) Delete(ctx context.Context, id, version string) error { - providerID, resourceID, updatedAt, err := scimTarget(ctx, id, version) + target, err := scimTarget(ctx, id, version) if err != nil { return err } @@ -183,7 +184,7 @@ func (s *scimUserRepository) Delete(ctx context.Context, id, version string) err return err } db := s.api.db.WithContext(ctx) - existing, err := models.FindSCIMUser(db, providerID, resourceID) + existing, err := models.FindSCIMUser(db, target.ProviderID, target.ID) if err != nil { return scimError(err) } @@ -193,7 +194,7 @@ func (s *scimUserRepository) Delete(ctx context.Context, id, version string) err return err } } - row, err := models.DeleteSCIMUser(tx, providerID, resourceID, updatedAt) + row, err := models.DeleteSCIMUser(tx, target) if err != nil { return err } diff --git a/internal/models/scim.go b/internal/models/scim.go index 34f0433b56..373b99749a 100644 --- a/internal/models/scim.go +++ b/internal/models/scim.go @@ -39,6 +39,12 @@ type SCIMQuery struct { Limit int } +type SCIMTarget struct { + ProviderID uuid.UUID + ID uuid.UUID + UpdatedAt *time.Time +} + type scimTable struct { name string label string @@ -80,30 +86,26 @@ func createSCIMRow[T any](tx *storage.Connection, table scimTable, providerID uu return row, nil } -func findSCIMRow[T any](tx *storage.Connection, table scimTable, providerID, id uuid.UUID, lock string) (*T, error) { - where := "id = ? AND sso_provider_id = ?" - if table.live != "" { - where += " AND " + table.live +func findSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarget, forUpdate bool) (*T, error) { + lock := "" + if forUpdate { + lock = " FOR UPDATE" } row := new(T) if err := tx.RawQuery( - fmt.Sprintf("SELECT %s FROM %q WHERE %s%s", table.columns, table.name, where, lock), - id, providerID, + fmt.Sprintf("SELECT %s FROM %q WHERE %s%s", table.columns, table.name, table.targetClause(), lock), + target.ID, target.ProviderID, ).First(row); err != nil { return nil, table.error(err, "finding") } return row, nil } -func findUnchangedSCIMRow[T any](tx *storage.Connection, table scimTable, providerID, id uuid.UUID, resource []byte, updatedAt *time.Time) (*T, error) { - where := "id = ? AND sso_provider_id = ? AND resource = ?::jsonb AND (?::timestamptz IS NULL OR updated_at = ?)" - if table.live != "" { - where += " AND " + table.live - } +func findUnchangedSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarget, resource []byte) (*T, error) { row := new(T) if err := tx.RawQuery( - fmt.Sprintf("SELECT %s FROM %q WHERE %s FOR UPDATE", table.columns, table.name, where), - id, providerID, string(resource), updatedAt, updatedAt, + fmt.Sprintf("SELECT %s FROM %q WHERE %s AND resource = ?::jsonb AND (?::timestamptz IS NULL OR updated_at = ?) FOR UPDATE", table.columns, table.name, table.targetClause()), + target.ID, target.ProviderID, string(resource), target.UpdatedAt, target.UpdatedAt, ).First(row); err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, nil @@ -113,14 +115,15 @@ func findUnchangedSCIMRow[T any](tx *storage.Connection, table scimTable, provid return row, nil } -func scimWriteError[T any](tx *storage.Connection, table scimTable, err error, providerID, id uuid.UUID, updatedAt *time.Time, verb string) error { - if errors.Is(err, sql.ErrNoRows) && updatedAt != nil { - if _, findErr := findSCIMRow[T](tx, table, providerID, id, ""); findErr != nil { - return findErr - } - return table.stale +func replaceSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarget, resource []byte) (*T, error) { + row := new(T) + if err := tx.RawQuery( + fmt.Sprintf("UPDATE %q SET resource = ?::jsonb, updated_at = now() WHERE %s AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING %s", table.name, table.targetClause(), table.columns), + string(resource), target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, + ).First(row); err != nil { + return nil, table.writeError(tx, target, err, "replacing") } - return table.error(err, verb) + return row, nil } func (t scimTable) where(providerID uuid.UUID, filter SCIMFilter) (string, []any) { @@ -140,6 +143,40 @@ func (t scimTable) where(providerID uuid.UUID, filter SCIMFilter) (string, []any return strings.Join(clauses, " AND "), args } +func (t scimTable) targetClause() string { + if t.live == "" { + return "id = ? AND sso_provider_id = ?" + } + return "id = ? AND sso_provider_id = ? AND " + t.live +} + +func (t scimTable) exists(tx *storage.Connection, target SCIMTarget) (bool, error) { + result := struct { + Exists bool `db:"exists"` + }{} + if err := tx.RawQuery( + fmt.Sprintf("SELECT EXISTS(SELECT 1 FROM %q WHERE %s) AS exists", t.name, t.targetClause()), + target.ID, target.ProviderID, + ).First(&result); err != nil { + return false, errors.Wrapf(err, "error finding %s", t.label) + } + return result.Exists, nil +} + +func (t scimTable) writeError(tx *storage.Connection, target SCIMTarget, err error, verb string) error { + if !errors.Is(err, sql.ErrNoRows) || target.UpdatedAt == nil { + return t.error(err, verb) + } + exists, findErr := t.exists(tx, target) + if findErr != nil { + return findErr + } + if !exists { + return t.notFound + } + return t.stale +} + func (t scimTable) orderBy(order SCIMOrder) string { direction := "ASC" if order.Descending { diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go index 24cbe605c4..03f71fa4e2 100644 --- a/internal/models/scim_group.go +++ b/internal/models/scim_group.go @@ -56,31 +56,23 @@ func CreateSCIMGroup(tx *storage.Connection, providerID uuid.UUID, resource []by } func FindSCIMGroup(tx *storage.Connection, providerID, id uuid.UUID) (*SCIMGroup, error) { - return findSCIMRow[SCIMGroup](tx, scimGroupsTable, providerID, id, "") + return findSCIMRow[SCIMGroup](tx, scimGroupsTable, SCIMTarget{ProviderID: providerID, ID: id}, false) } func FindSCIMGroupForUpdate(tx *storage.Connection, providerID, id uuid.UUID) (*SCIMGroup, error) { - return findSCIMRow[SCIMGroup](tx, scimGroupsTable, providerID, id, " FOR UPDATE") + return findSCIMRow[SCIMGroup](tx, scimGroupsTable, SCIMTarget{ProviderID: providerID, ID: id}, true) } func FindSCIMGroups(tx *storage.Connection, providerID uuid.UUID, query SCIMQuery) ([]SCIMGroup, int, error) { return findSCIMPage[SCIMGroup](tx, scimGroupsTable, providerID, query) } -func ReplaceSCIMGroup(tx *storage.Connection, providerID, id uuid.UUID, resource []byte, updatedAt *time.Time) (*SCIMGroup, error) { - group := &SCIMGroup{} - err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET resource = ?::jsonb, updated_at = now() WHERE id = ? AND sso_provider_id = ? AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimGroupColumns, group.TableName()), - string(resource), id, providerID, updatedAt, updatedAt, - ).First(group) - if err != nil { - return nil, scimWriteError[SCIMGroup](tx, scimGroupsTable, err, providerID, id, updatedAt, "replacing") - } - return group, nil +func ReplaceSCIMGroup(tx *storage.Connection, target SCIMTarget, resource []byte) (*SCIMGroup, error) { + return replaceSCIMRow[SCIMGroup](tx, scimGroupsTable, target, resource) } -func FindUnchangedSCIMGroup(tx *storage.Connection, providerID, id uuid.UUID, resource []byte, updatedAt *time.Time) (*SCIMGroup, error) { - return findUnchangedSCIMRow[SCIMGroup](tx, scimGroupsTable, providerID, id, resource, updatedAt) +func FindUnchangedSCIMGroup(tx *storage.Connection, target SCIMTarget, resource []byte) (*SCIMGroup, error) { + return findUnchangedSCIMRow[SCIMGroup](tx, scimGroupsTable, target, resource) } func TouchSCIMGroup(tx *storage.Connection, group *SCIMGroup) (*SCIMGroup, error) { @@ -94,14 +86,14 @@ func TouchSCIMGroup(tx *storage.Connection, group *SCIMGroup) (*SCIMGroup, error return touched, nil } -func DeleteSCIMGroup(tx *storage.Connection, providerID, id uuid.UUID, updatedAt *time.Time) (*SCIMGroup, error) { +func DeleteSCIMGroup(tx *storage.Connection, target SCIMTarget) (*SCIMGroup, error) { group := &SCIMGroup{} err := tx.RawQuery( fmt.Sprintf("DELETE FROM %q WHERE id = ? AND sso_provider_id = ? AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimGroupColumns, group.TableName()), - id, providerID, updatedAt, updatedAt, + target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(group) if err != nil { - return nil, scimWriteError[SCIMGroup](tx, scimGroupsTable, err, providerID, id, updatedAt, "deleting") + return nil, scimGroupsTable.writeError(tx, target, err, "deleting") } return group, nil } diff --git a/internal/models/scim_group_test.go b/internal/models/scim_group_test.go index 1fbf03dcec..815a9ec82d 100644 --- a/internal/models/scim_group_test.go +++ b/internal/models/scim_group_test.go @@ -117,14 +117,14 @@ func (ts *SCIMGroupTestSuite) TestFindGroupsFiltersAndSorts() { func (ts *SCIMGroupTestSuite) TestReplaceChecksVersion() { group := ts.createGroup(ts.provider.ID, "Engineering") - replaced, err := ReplaceSCIMGroup(ts.db, ts.provider.ID, group.ID, []byte(`{"displayName":"Platform"}`), &group.UpdatedAt) + replaced, err := ReplaceSCIMGroup(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: group.ID, UpdatedAt: &group.UpdatedAt}, []byte(`{"displayName":"Platform"}`)) require.NoError(ts.T(), err) require.Equal(ts.T(), "platform", replaced.DisplayName) - _, err = ReplaceSCIMGroup(ts.db, ts.provider.ID, group.ID, []byte(`{"displayName":"Stale"}`), &group.UpdatedAt) + _, err = ReplaceSCIMGroup(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: group.ID, UpdatedAt: &group.UpdatedAt}, []byte(`{"displayName":"Stale"}`)) require.ErrorIs(ts.T(), err, SCIMGroupStaleError{}) - _, err = ReplaceSCIMGroup(ts.db, ts.createProvider().ID, group.ID, []byte(`{"displayName":"Other"}`), nil) + _, err = ReplaceSCIMGroup(ts.db, SCIMTarget{ProviderID: ts.createProvider().ID, ID: group.ID}, []byte(`{"displayName":"Other"}`)) require.ErrorIs(ts.T(), err, SCIMGroupNotFoundError{}) } @@ -134,7 +134,7 @@ func (ts *SCIMGroupTestSuite) TestDeleteRemovesMembers() { _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{user.ID}) require.NoError(ts.T(), err) - _, err = DeleteSCIMGroup(ts.db, ts.provider.ID, group.ID, nil) + _, err = DeleteSCIMGroup(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: group.ID}) require.NoError(ts.T(), err) _, err = FindSCIMGroup(ts.db, ts.provider.ID, group.ID) @@ -143,7 +143,7 @@ func (ts *SCIMGroupTestSuite) TestDeleteRemovesMembers() { require.NoError(ts.T(), err) require.Zero(ts.T(), count) - _, err = DeleteSCIMGroup(ts.db, ts.provider.ID, group.ID, nil) + _, err = DeleteSCIMGroup(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: group.ID}) require.ErrorIs(ts.T(), err, SCIMGroupNotFoundError{}) } @@ -191,7 +191,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersRejectsOtherProviderUsers() { func (ts *SCIMGroupTestSuite) TestReplaceMembersRejectsDeletedUsers() { group := ts.createGroup(ts.provider.ID, "Engineering") alice := ts.createUser(ts.provider.ID, "alice") - _, err := DeleteSCIMUser(ts.db, ts.provider.ID, alice.ID, nil) + _, err := DeleteSCIMUser(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) require.NoError(ts.T(), err) _, _, err = ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) @@ -207,7 +207,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersWaitsForConcurrentUserDelete() { deleting := ts.beginTx() defer func() { _ = deleting.TX.Rollback() }() - _, err := DeleteSCIMUser(deleting, ts.provider.ID, alice.ID, nil) + _, err := DeleteSCIMUser(deleting, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) require.NoError(ts.T(), err) _, err = RemoveSCIMUserFromGroups(deleting, alice.ID) require.NoError(ts.T(), err) @@ -237,7 +237,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersDoesNotLockExistingMembers() { deleting := ts.beginTx() defer func() { _ = deleting.TX.Rollback() }() - _, err = DeleteSCIMUser(deleting, ts.provider.ID, alice.ID, nil) + _, err = DeleteSCIMUser(deleting, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) require.NoError(ts.T(), err) result := make(chan error, 1) @@ -275,7 +275,7 @@ func (ts *SCIMGroupTestSuite) TestFindMembersHidesDeletedUsers() { _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) require.NoError(ts.T(), err) - _, err = DeleteSCIMUser(ts.db, ts.provider.ID, alice.ID, nil) + _, err = DeleteSCIMUser(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) require.NoError(ts.T(), err) members, err := FindSCIMGroupMembers(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index f0b6a812e4..794662dd81 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -45,41 +45,33 @@ func CreateSCIMUser(tx *storage.Connection, providerID uuid.UUID, resource []byt } func FindSCIMUser(tx *storage.Connection, providerID, id uuid.UUID) (*SCIMUser, error) { - return findSCIMRow[SCIMUser](tx, scimUsersTable, providerID, id, "") + return findSCIMRow[SCIMUser](tx, scimUsersTable, SCIMTarget{ProviderID: providerID, ID: id}, false) } func FindSCIMUserForUpdate(tx *storage.Connection, providerID, id uuid.UUID) (*SCIMUser, error) { - return findSCIMRow[SCIMUser](tx, scimUsersTable, providerID, id, " FOR UPDATE") + return findSCIMRow[SCIMUser](tx, scimUsersTable, SCIMTarget{ProviderID: providerID, ID: id}, true) } -func FindUnchangedSCIMUser(tx *storage.Connection, providerID, id uuid.UUID, resource []byte, updatedAt *time.Time) (*SCIMUser, error) { - return findUnchangedSCIMRow[SCIMUser](tx, scimUsersTable, providerID, id, resource, updatedAt) +func FindUnchangedSCIMUser(tx *storage.Connection, target SCIMTarget, resource []byte) (*SCIMUser, error) { + return findUnchangedSCIMRow[SCIMUser](tx, scimUsersTable, target, resource) } func FindSCIMUsers(tx *storage.Connection, providerID uuid.UUID, query SCIMQuery) ([]SCIMUser, int, error) { return findSCIMPage[SCIMUser](tx, scimUsersTable, providerID, query) } -func ReplaceSCIMUser(tx *storage.Connection, providerID, id uuid.UUID, resource []byte, updatedAt *time.Time) (*SCIMUser, error) { - user := &SCIMUser{} - err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET resource = ?::jsonb, updated_at = now() WHERE id = ? AND sso_provider_id = ? AND deleted_at IS NULL AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimUserColumns, user.TableName()), - string(resource), id, providerID, updatedAt, updatedAt, - ).First(user) - if err != nil { - return nil, scimWriteError[SCIMUser](tx, scimUsersTable, err, providerID, id, updatedAt, "replacing") - } - return user, nil +func ReplaceSCIMUser(tx *storage.Connection, target SCIMTarget, resource []byte) (*SCIMUser, error) { + return replaceSCIMRow[SCIMUser](tx, scimUsersTable, target, resource) } -func DeleteSCIMUser(tx *storage.Connection, providerID, id uuid.UUID, updatedAt *time.Time) (*SCIMUser, error) { +func DeleteSCIMUser(tx *storage.Connection, target SCIMTarget) (*SCIMUser, error) { user := &SCIMUser{} err := tx.RawQuery( fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE id = ? AND sso_provider_id = ? AND deleted_at IS NULL AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimUserColumns, user.TableName()), - id, providerID, updatedAt, updatedAt, + target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(user) if err != nil { - return nil, scimWriteError[SCIMUser](tx, scimUsersTable, err, providerID, id, updatedAt, "deleting") + return nil, scimUsersTable.writeError(tx, target, err, "deleting") } return user, nil } From a9255d75568c308fc8ecfa76400732a54e5d731b Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:18:03 -0600 Subject: [PATCH 24/88] chore(scim): pass a SCIMIdentityRename to RenameSCIMIdentity --- internal/api/scim_users.go | 8 +++++++- internal/models/scim_user.go | 16 ++++++++++++---- 2 files changed, 19 insertions(+), 5 deletions(-) diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index e8cc9ee6db..d31f21306d 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -277,7 +277,13 @@ func (s *scimUserRepository) syncAuthUser(tx *storage.Connection, providerID uui if email := scimUserEmail(user); email != "" { data["email"] = email } - err := models.RenameSCIMIdentity(tx, linked.ID, "sso:"+providerID.String(), from, user.UserName, data) + err := models.RenameSCIMIdentity(tx, models.SCIMIdentityRename{ + UserID: linked.ID, + Provider: "sso:" + providerID.String(), + From: from, + To: user.UserName, + Data: data, + }) if errors.Is(err, models.SCIMIdentityNotFoundError{}) { logrus.WithField("user_id", linked.ID).WithField("sso_provider_id", providerID).Warn("scim: SCIM identity not found, rename skipped") } else if err != nil { diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index 794662dd81..988f985833 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -29,6 +29,14 @@ func (SCIMUser) TableName() string { return "scim_users" } +type SCIMIdentityRename struct { + UserID uuid.UUID + Provider string + From string + To string + Data map[string]any +} + var scimUsersTable = scimTable{ name: SCIMUser{}.TableName(), label: "SCIM user", @@ -220,21 +228,21 @@ func IsSCIMUserDeprovisionedForUpdate(tx *storage.Connection, userID uuid.UUID) return len(rows) > 0, nil } -func RenameSCIMIdentity(tx *storage.Connection, userID uuid.UUID, provider, from, to string, data map[string]any) error { - encoded, err := json.Marshal(data) +func RenameSCIMIdentity(tx *storage.Connection, rename SCIMIdentityRename) error { + encoded, err := json.Marshal(rename.Data) if err != nil { return errors.Wrap(err, "error encoding identity data") } table := (&Identity{}).TableName() if err := tx.RawQuery( fmt.Sprintf("DELETE FROM %[1]q WHERE user_id = ? AND provider = ? AND provider_id <> ? AND (lower(provider_id) = lower(?) OR provider_id = ?) AND EXISTS (SELECT 1 FROM %[1]q WHERE user_id = ? AND provider = ? AND provider_id = ?)", table), - userID, provider, from, from, to, userID, provider, from, + rename.UserID, rename.Provider, rename.From, rename.From, rename.To, rename.UserID, rename.Provider, rename.From, ).Exec(); err != nil { return errors.Wrap(err, "error removing stale SCIM identities") } count, err := tx.RawQuery( fmt.Sprintf("UPDATE %q SET provider_id = ?, identity_data = identity_data || ?::jsonb, updated_at = now() WHERE user_id = ? AND provider = ? AND provider_id = ?", table), - to, string(encoded), userID, provider, from, + rename.To, string(encoded), rename.UserID, rename.Provider, rename.From, ).ExecWithCount() if err != nil { if isUniqueViolation(err) { From 2f8bd371131546b1105cfcc45b404cbed9d33dbb Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:19:17 -0600 Subject: [PATCH 25/88] chore(scim): name scimTable fields and table references consistently --- internal/models/scim.go | 38 +++++++++++++++--------------- internal/models/scim_group.go | 34 +++++++++++++-------------- internal/models/scim_settings.go | 4 ++-- internal/models/scim_token.go | 2 +- internal/models/scim_user.go | 40 ++++++++++++++++---------------- 5 files changed, 59 insertions(+), 59 deletions(-) diff --git a/internal/models/scim.go b/internal/models/scim.go index 373b99749a..69665b9d34 100644 --- a/internal/models/scim.go +++ b/internal/models/scim.go @@ -46,14 +46,14 @@ type SCIMTarget struct { } type scimTable struct { - name string - label string - columns string - nameCol string - live string - notFound error - stale error - conflict error + name string + label string + columns string + nameColumn string + liveClause string + notFound error + stale error + conflict error } func findSCIMPage[T any](tx *storage.Connection, table scimTable, providerID uuid.UUID, query SCIMQuery) ([]T, int, error) { @@ -81,7 +81,7 @@ func createSCIMRow[T any](tx *storage.Connection, table scimTable, providerID uu fmt.Sprintf("INSERT INTO %q (id, sso_provider_id, resource) VALUES (?, ?, ?::jsonb) RETURNING "+table.columns, table.name), uuid.Must(uuid.NewV4()), providerID, string(resource), ).First(row); err != nil { - return nil, table.error(err, "creating") + return nil, table.wrapError(err, "creating") } return row, nil } @@ -96,7 +96,7 @@ func findSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarg fmt.Sprintf("SELECT %s FROM %q WHERE %s%s", table.columns, table.name, table.targetClause(), lock), target.ID, target.ProviderID, ).First(row); err != nil { - return nil, table.error(err, "finding") + return nil, table.wrapError(err, "finding") } return row, nil } @@ -110,7 +110,7 @@ func findUnchangedSCIMRow[T any](tx *storage.Connection, table scimTable, target if errors.Is(err, sql.ErrNoRows) { return nil, nil } - return nil, table.error(err, "finding") + return nil, table.wrapError(err, "finding") } return row, nil } @@ -129,11 +129,11 @@ func replaceSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMT func (t scimTable) where(providerID uuid.UUID, filter SCIMFilter) (string, []any) { clauses := []string{"sso_provider_id = ?"} args := []any{providerID} - if t.live != "" { - clauses = append(clauses, t.live) + if t.liveClause != "" { + clauses = append(clauses, t.liveClause) } if filter.Name != nil { - clauses = append(clauses, t.nameCol+` COLLATE "C" = lower(?)`) + clauses = append(clauses, t.nameColumn+` COLLATE "C" = lower(?)`) args = append(args, *filter.Name) } if filter.ExternalID != nil { @@ -144,10 +144,10 @@ func (t scimTable) where(providerID uuid.UUID, filter SCIMFilter) (string, []any } func (t scimTable) targetClause() string { - if t.live == "" { + if t.liveClause == "" { return "id = ? AND sso_provider_id = ?" } - return "id = ? AND sso_provider_id = ? AND " + t.live + return "id = ? AND sso_provider_id = ? AND " + t.liveClause } func (t scimTable) exists(tx *storage.Connection, target SCIMTarget) (bool, error) { @@ -165,7 +165,7 @@ func (t scimTable) exists(tx *storage.Connection, target SCIMTarget) (bool, erro func (t scimTable) writeError(tx *storage.Connection, target SCIMTarget, err error, verb string) error { if !errors.Is(err, sql.ErrNoRows) || target.UpdatedAt == nil { - return t.error(err, verb) + return t.wrapError(err, verb) } exists, findErr := t.exists(tx, target) if findErr != nil { @@ -186,14 +186,14 @@ func (t scimTable) orderBy(order SCIMOrder) string { case SCIMSortByID: return "id " + direction case SCIMSortByName: - return t.nameCol + ` COLLATE "C" ` + direction + ", id " + direction + return t.nameColumn + ` COLLATE "C" ` + direction + ", id " + direction case SCIMSortByUpdatedAt: return "updated_at " + direction + ", id " + direction } return "created_at " + direction + ", id " + direction } -func (t scimTable) error(err error, verb string) error { +func (t scimTable) wrapError(err error, verb string) error { switch { case errors.Is(err, sql.ErrNoRows): return t.notFound diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go index 03f71fa4e2..7760f7988e 100644 --- a/internal/models/scim_group.go +++ b/internal/models/scim_group.go @@ -42,13 +42,13 @@ type SCIMGroupMembership struct { } var scimGroupsTable = scimTable{ - name: SCIMGroup{}.TableName(), - label: "SCIM group", - columns: scimGroupColumns, - nameCol: "display_name", - notFound: SCIMGroupNotFoundError{}, - stale: SCIMGroupStaleError{}, - conflict: SCIMGroupConflictError{}, + name: SCIMGroup{}.TableName(), + label: "SCIM group", + columns: scimGroupColumns, + nameColumn: "display_name", + notFound: SCIMGroupNotFoundError{}, + stale: SCIMGroupStaleError{}, + conflict: SCIMGroupConflictError{}, } func CreateSCIMGroup(tx *storage.Connection, providerID uuid.UUID, resource []byte) (*SCIMGroup, error) { @@ -78,7 +78,7 @@ func FindUnchangedSCIMGroup(tx *storage.Connection, target SCIMTarget, resource func TouchSCIMGroup(tx *storage.Connection, group *SCIMGroup) (*SCIMGroup, error) { touched := &SCIMGroup{} if err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET updated_at = now() WHERE id = ? RETURNING "+scimGroupColumns, group.TableName()), + fmt.Sprintf("UPDATE %q SET updated_at = now() WHERE id = ? RETURNING "+scimGroupColumns, scimGroupsTable.name), group.ID, ).First(touched); err != nil { return nil, errors.Wrap(err, "error updating SCIM group") @@ -89,7 +89,7 @@ func TouchSCIMGroup(tx *storage.Connection, group *SCIMGroup) (*SCIMGroup, error func DeleteSCIMGroup(tx *storage.Connection, target SCIMTarget) (*SCIMGroup, error) { group := &SCIMGroup{} err := tx.RawQuery( - fmt.Sprintf("DELETE FROM %q WHERE id = ? AND sso_provider_id = ? AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimGroupColumns, group.TableName()), + fmt.Sprintf("DELETE FROM %q WHERE id = ? AND sso_provider_id = ? AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimGroupColumns, scimGroupsTable.name), target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(group) if err != nil { @@ -104,7 +104,7 @@ func FindSCIMGroupMembers(tx *storage.Connection, providerID uuid.UUID, groupIDs return members, nil } err := tx.RawQuery( - fmt.Sprintf("SELECT m.group_id, m.scim_user_id, u.resource->>'userName' AS display FROM %q m JOIN %q u ON u.id = m.scim_user_id WHERE m.group_id = ANY(?::uuid[]) AND u.sso_provider_id = ? AND u.deleted_at IS NULL ORDER BY m.group_id, m.created_at, m.scim_user_id", (&SCIMGroupMember{}).TableName(), (&SCIMUser{}).TableName()), + fmt.Sprintf("SELECT m.group_id, m.scim_user_id, u.resource->>'userName' AS display FROM %q m JOIN %q u ON u.id = m.scim_user_id WHERE m.group_id = ANY(?::uuid[]) AND u.sso_provider_id = ? AND u.deleted_at IS NULL ORDER BY m.group_id, m.created_at, m.scim_user_id", SCIMGroupMember{}.TableName(), scimUsersTable.name), groupIDs, providerID, ).All(&members) if err != nil { @@ -119,7 +119,7 @@ func FindSCIMGroupsForUsers(tx *storage.Connection, providerID uuid.UUID, scimUs return groups, nil } err := tx.RawQuery( - fmt.Sprintf("SELECT m.group_id, m.scim_user_id, g.resource->>'displayName' AS display FROM %q m JOIN %q g ON g.id = m.group_id WHERE m.scim_user_id = ANY(?::uuid[]) AND g.sso_provider_id = ? ORDER BY m.scim_user_id, g.display_name COLLATE \"C\", g.id", (&SCIMGroupMember{}).TableName(), (&SCIMGroup{}).TableName()), + fmt.Sprintf("SELECT m.group_id, m.scim_user_id, g.resource->>'displayName' AS display FROM %q m JOIN %q g ON g.id = m.group_id WHERE m.scim_user_id = ANY(?::uuid[]) AND g.sso_provider_id = ? ORDER BY m.scim_user_id, g.display_name COLLATE \"C\", g.id", SCIMGroupMember{}.TableName(), scimGroupsTable.name), scimUserIDs, providerID, ).All(&groups) if err != nil { @@ -132,7 +132,7 @@ func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserI wanted := []uuid.UUID{} if len(scimUserIDs) > 0 { if err := tx.RawQuery( - fmt.Sprintf("SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL", (&SCIMUser{}).TableName()), + fmt.Sprintf("SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL", scimUsersTable.name), scimUserIDs, group.SSOProviderID, ).All(&wanted); err != nil { return nil, nil, errors.Wrap(err, "error finding SCIM group members") @@ -144,7 +144,7 @@ func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserI current := []SCIMGroupMember{} if err := tx.RawQuery( - fmt.Sprintf("SELECT group_id, scim_user_id, created_at FROM %q WHERE group_id = ?", (&SCIMGroupMember{}).TableName()), + fmt.Sprintf("SELECT group_id, scim_user_id, created_at FROM %q WHERE group_id = ?", SCIMGroupMember{}.TableName()), group.ID, ).All(¤t); err != nil { return nil, nil, errors.Wrap(err, "error finding SCIM group members") @@ -159,7 +159,7 @@ func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserI if len(removed) > 0 { if err := tx.RawQuery( - fmt.Sprintf("DELETE FROM %q WHERE group_id = ? AND scim_user_id = ANY(?::uuid[])", (&SCIMGroupMember{}).TableName()), + fmt.Sprintf("DELETE FROM %q WHERE group_id = ? AND scim_user_id = ANY(?::uuid[])", SCIMGroupMember{}.TableName()), group.ID, removed, ).Exec(); err != nil { return nil, nil, errors.Wrap(err, "error removing SCIM group members") @@ -168,7 +168,7 @@ func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserI if len(added) > 0 { locked := []uuid.UUID{} if err := tx.RawQuery( - fmt.Sprintf("SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL ORDER BY id FOR SHARE", (&SCIMUser{}).TableName()), + fmt.Sprintf("SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL ORDER BY id FOR SHARE", scimUsersTable.name), added, group.SSOProviderID, ).All(&locked); err != nil { return nil, nil, errors.Wrap(err, "error locking SCIM group members") @@ -177,7 +177,7 @@ func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserI return nil, nil, SCIMGroupMemberNotFoundError{IDs: missing} } if err := tx.RawQuery( - fmt.Sprintf("INSERT INTO %q (group_id, scim_user_id) SELECT ?, unnest(?::uuid[])", (&SCIMGroupMember{}).TableName()), + fmt.Sprintf("INSERT INTO %q (group_id, scim_user_id) SELECT ?, unnest(?::uuid[])", SCIMGroupMember{}.TableName()), group.ID, locked, ).Exec(); err != nil { return nil, nil, errors.Wrap(err, "error adding SCIM group members") @@ -187,7 +187,7 @@ func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserI } func RemoveSCIMUserFromGroups(tx *storage.Connection, scimUserID uuid.UUID) ([]uuid.UUID, error) { - groups, members := (&SCIMGroup{}).TableName(), (&SCIMGroupMember{}).TableName() + groups, members := scimGroupsTable.name, SCIMGroupMember{}.TableName() if err := tx.RawQuery( fmt.Sprintf("SELECT id FROM %q WHERE id IN (SELECT group_id FROM %q WHERE scim_user_id = ?) ORDER BY id FOR UPDATE", groups, members), scimUserID, diff --git a/internal/models/scim_settings.go b/internal/models/scim_settings.go index c47deb0c8e..e211e42f6f 100644 --- a/internal/models/scim_settings.go +++ b/internal/models/scim_settings.go @@ -28,7 +28,7 @@ func (s *SCIMSettings) AfterFind(*pop.Connection) error { } func EnableSCIM(tx *storage.Connection, providerID uuid.UUID) (bool, error) { - table := (&SCIMSettings{}).TableName() + table := SCIMSettings{}.TableName() rows := []SCIMSettings{} if err := tx.RawQuery( fmt.Sprintf("INSERT INTO %[1]q (sso_provider_id, enabled) VALUES (?, true) ON CONFLICT (sso_provider_id) DO UPDATE SET enabled = true, updated_at = now() WHERE %[1]q.enabled = false RETURNING *", table), @@ -42,7 +42,7 @@ func EnableSCIM(tx *storage.Connection, providerID uuid.UUID) (bool, error) { func DisableSCIM(tx *storage.Connection, providerID uuid.UUID) (bool, error) { rows := []SCIMSettings{} if err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET enabled = false, updated_at = now() WHERE sso_provider_id = ? AND enabled RETURNING *", (&SCIMSettings{}).TableName()), + fmt.Sprintf("UPDATE %q SET enabled = false, updated_at = now() WHERE sso_provider_id = ? AND enabled RETURNING *", SCIMSettings{}.TableName()), providerID, ).All(&rows); err != nil { return false, errors.Wrap(err, "error disabling SCIM") diff --git a/internal/models/scim_token.go b/internal/models/scim_token.go index bf55bcc408..7c13e7aa3e 100644 --- a/internal/models/scim_token.go +++ b/internal/models/scim_token.go @@ -154,7 +154,7 @@ func AuthenticateSCIMToken(tx *storage.Connection, plaintext string) (*SCIMToken ) SELECT * FROM touched UNION ALL -SELECT * FROM authenticated WHERE NOT EXISTS (SELECT 1 FROM touched)`, token.TableName(), (&SSOProvider{}).TableName(), (&SCIMSettings{}).TableName()), +SELECT * FROM authenticated WHERE NOT EXISTS (SELECT 1 FROM touched)`, token.TableName(), SSOProvider{}.TableName(), SCIMSettings{}.TableName()), HashSCIMToken(plaintext), ).First(token) if err != nil { diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index 988f985833..9806fb86d0 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -38,14 +38,14 @@ type SCIMIdentityRename struct { } var scimUsersTable = scimTable{ - name: SCIMUser{}.TableName(), - label: "SCIM user", - columns: scimUserColumns, - nameCol: "user_name", - live: "deleted_at IS NULL", - notFound: SCIMUserNotFoundError{}, - stale: SCIMUserStaleError{}, - conflict: SCIMUserConflictError{}, + name: SCIMUser{}.TableName(), + label: "SCIM user", + columns: scimUserColumns, + nameColumn: "user_name", + liveClause: "deleted_at IS NULL", + notFound: SCIMUserNotFoundError{}, + stale: SCIMUserStaleError{}, + conflict: SCIMUserConflictError{}, } func CreateSCIMUser(tx *storage.Connection, providerID uuid.UUID, resource []byte) (*SCIMUser, error) { @@ -75,7 +75,7 @@ func ReplaceSCIMUser(tx *storage.Connection, target SCIMTarget, resource []byte) func DeleteSCIMUser(tx *storage.Connection, target SCIMTarget) (*SCIMUser, error) { user := &SCIMUser{} err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE id = ? AND sso_provider_id = ? AND deleted_at IS NULL AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimUserColumns, user.TableName()), + fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE id = ? AND sso_provider_id = ? AND deleted_at IS NULL AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimUserColumns, scimUsersTable.name), target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(user) if err != nil { @@ -94,7 +94,7 @@ func FindSCIMUserLinks(tx *storage.Connection, ids []uuid.UUID) (map[uuid.UUID]u UserID uuid.UUID `db:"user_id"` }{} if err := tx.RawQuery( - fmt.Sprintf("SELECT id, user_id FROM %q WHERE id = ANY(?::uuid[]) AND user_id IS NOT NULL", (&SCIMUser{}).TableName()), + fmt.Sprintf("SELECT id, user_id FROM %q WHERE id = ANY(?::uuid[]) AND user_id IS NOT NULL", scimUsersTable.name), ids, ).All(&rows); err != nil { return nil, errors.Wrap(err, "error finding SCIM user links") @@ -107,7 +107,7 @@ func FindSCIMUserLinks(tx *storage.Connection, ids []uuid.UUID) (map[uuid.UUID]u func LockUserForSCIM(tx *storage.Connection, userID uuid.UUID) error { if err := tx.RawQuery( - fmt.Sprintf("SELECT id FROM %q WHERE id = ? FOR UPDATE", (&User{}).TableName()), + fmt.Sprintf("SELECT id FROM %q WHERE id = ? FOR UPDATE", User{}.TableName()), userID, ).Exec(); err != nil { return errors.Wrap(err, "error locking user") @@ -125,7 +125,7 @@ func LogoutSCIMUser(tx *storage.Connection, userID uuid.UUID) error { func SoftDeleteSCIMUsersByUserID(tx *storage.Connection, userID uuid.UUID) ([]SCIMUser, error) { rows := []SCIMUser{} if err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE user_id = ? AND deleted_at IS NULL RETURNING "+scimUserColumns, (&SCIMUser{}).TableName()), + fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE user_id = ? AND deleted_at IS NULL RETURNING "+scimUserColumns, scimUsersTable.name), userID, ).All(&rows); err != nil { return nil, errors.Wrap(err, "error deleting SCIM users by user id") @@ -134,7 +134,7 @@ func SoftDeleteSCIMUsersByUserID(tx *storage.Connection, userID uuid.UUID) ([]SC } func BanDeprovisionedSCIMUsers(tx *storage.Connection, providerID uuid.UUID, until time.Time) (int, error) { - users, scimUsers := (&User{}).TableName(), (&SCIMUser{}).TableName() + users, scimUsers := User{}.TableName(), scimUsersTable.name count, err := tx.RawQuery( fmt.Sprintf( "UPDATE %[1]q u SET banned_until = ?, updated_at = now() "+ @@ -165,7 +165,7 @@ func LinkSCIMUser(tx *storage.Connection, user *SCIMUser, userID uuid.UUID) erro } if err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET user_id = ? WHERE id = ?", user.TableName()), + fmt.Sprintf("UPDATE %q SET user_id = ? WHERE id = ?", scimUsersTable.name), userID, user.ID, ).Exec(); err != nil { return errors.Wrap(err, "error linking SCIM user") @@ -189,9 +189,9 @@ func IsSCIMDeprovisioned(tx *storage.Connection, providerID, userID uuid.UUID) ( }{} if err := tx.RawQuery( fmt.Sprintf( - "SELECT EXISTS(SELECT 1 FROM %q WHERE sso_provider_id = ? AND user_id = ?) AS any_row, "+ - "EXISTS(SELECT 1 FROM %q WHERE sso_provider_id = ? AND user_id = ? AND deleted_at IS NULL AND active) AS live", - (&SCIMUser{}).TableName(), (&SCIMUser{}).TableName(), + "SELECT EXISTS(SELECT 1 FROM %[1]q WHERE sso_provider_id = ? AND user_id = ?) AS any_row, "+ + "EXISTS(SELECT 1 FROM %[1]q WHERE sso_provider_id = ? AND user_id = ? AND deleted_at IS NULL AND active) AS live", + scimUsersTable.name, ), providerID, userID, providerID, userID, ).First(&result); err != nil { @@ -205,7 +205,7 @@ func IsSCIMDeprovisioned(tx *storage.Connection, providerID, userID uuid.UUID) ( func IsSCIMUserDeprovisionedForUpdate(tx *storage.Connection, userID uuid.UUID) (bool, error) { if err := tx.RawQuery( - fmt.Sprintf("SELECT id FROM %q WHERE id = ? FOR NO KEY UPDATE", (&User{}).TableName()), + fmt.Sprintf("SELECT id FROM %q WHERE id = ? FOR NO KEY UPDATE", User{}.TableName()), userID, ).Exec(); err != nil { return false, errors.Wrap(err, "error locking user") @@ -215,7 +215,7 @@ func IsSCIMUserDeprovisionedForUpdate(tx *storage.Connection, userID uuid.UUID) DeletedAt *time.Time `db:"deleted_at"` }{} if err := tx.RawQuery( - fmt.Sprintf("SELECT active, deleted_at FROM %q WHERE user_id = ?", (&SCIMUser{}).TableName()), + fmt.Sprintf("SELECT active, deleted_at FROM %q WHERE user_id = ?", scimUsersTable.name), userID, ).All(&rows); err != nil { return false, errors.Wrap(err, "error finding SCIM users") @@ -233,7 +233,7 @@ func RenameSCIMIdentity(tx *storage.Connection, rename SCIMIdentityRename) error if err != nil { return errors.Wrap(err, "error encoding identity data") } - table := (&Identity{}).TableName() + table := Identity{}.TableName() if err := tx.RawQuery( fmt.Sprintf("DELETE FROM %[1]q WHERE user_id = ? AND provider = ? AND provider_id <> ? AND (lower(provider_id) = lower(?) OR provider_id = ?) AND EXISTS (SELECT 1 FROM %[1]q WHERE user_id = ? AND provider = ? AND provider_id = ?)", table), rename.UserID, rename.Provider, rename.From, rename.From, rename.To, rename.UserID, rename.Provider, rename.From, From 982e8e1fcb09a9bcdcb6a308f4ad2e057813b967 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:20:26 -0600 Subject: [PATCH 26/88] chore(scim): split ReplaceSCIMGroupMembers into lookup and write steps --- internal/models/scim_group.go | 125 ++++++++++++++++++++-------------- 1 file changed, 75 insertions(+), 50 deletions(-) diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go index 7760f7988e..7fcb0e21a4 100644 --- a/internal/models/scim_group.go +++ b/internal/models/scim_group.go @@ -129,59 +129,21 @@ func FindSCIMGroupsForUsers(tx *storage.Connection, providerID uuid.UUID, scimUs } func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserIDs []uuid.UUID) (added, removed []uuid.UUID, err error) { - wanted := []uuid.UUID{} - if len(scimUserIDs) > 0 { - if err := tx.RawQuery( - fmt.Sprintf("SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL", scimUsersTable.name), - scimUserIDs, group.SSOProviderID, - ).All(&wanted); err != nil { - return nil, nil, errors.Wrap(err, "error finding SCIM group members") - } + wanted, err := findLiveSCIMUserIDs(tx, group.SSOProviderID, scimUserIDs, false) + if err != nil { + return nil, nil, err } - if missing := differenceUUIDs(scimUserIDs, wanted); len(missing) > 0 { - return nil, nil, SCIMGroupMemberNotFoundError{IDs: missing} + current, err := findSCIMGroupMemberIDs(tx, group.ID) + if err != nil { + return nil, nil, err } - - current := []SCIMGroupMember{} - if err := tx.RawQuery( - fmt.Sprintf("SELECT group_id, scim_user_id, created_at FROM %q WHERE group_id = ?", SCIMGroupMember{}.TableName()), - group.ID, - ).All(¤t); err != nil { - return nil, nil, errors.Wrap(err, "error finding SCIM group members") + removed = differenceUUIDs(current, wanted) + added = differenceUUIDs(wanted, current) + if err := removeSCIMGroupMembers(tx, group.ID, removed); err != nil { + return nil, nil, err } - currentIDs := make([]uuid.UUID, len(current)) - for i := range current { - currentIDs[i] = current[i].SCIMUserID - } - - removed = differenceUUIDs(currentIDs, wanted) - added = differenceUUIDs(wanted, currentIDs) - - if len(removed) > 0 { - if err := tx.RawQuery( - fmt.Sprintf("DELETE FROM %q WHERE group_id = ? AND scim_user_id = ANY(?::uuid[])", SCIMGroupMember{}.TableName()), - group.ID, removed, - ).Exec(); err != nil { - return nil, nil, errors.Wrap(err, "error removing SCIM group members") - } - } - if len(added) > 0 { - locked := []uuid.UUID{} - if err := tx.RawQuery( - fmt.Sprintf("SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL ORDER BY id FOR SHARE", scimUsersTable.name), - added, group.SSOProviderID, - ).All(&locked); err != nil { - return nil, nil, errors.Wrap(err, "error locking SCIM group members") - } - if missing := differenceUUIDs(added, locked); len(missing) > 0 { - return nil, nil, SCIMGroupMemberNotFoundError{IDs: missing} - } - if err := tx.RawQuery( - fmt.Sprintf("INSERT INTO %q (group_id, scim_user_id) SELECT ?, unnest(?::uuid[])", SCIMGroupMember{}.TableName()), - group.ID, locked, - ).Exec(); err != nil { - return nil, nil, errors.Wrap(err, "error adding SCIM group members") - } + if err := addSCIMGroupMembers(tx, group, added); err != nil { + return nil, nil, err } return added, removed, nil } @@ -215,3 +177,66 @@ func RemoveSCIMUserFromGroups(tx *storage.Connection, scimUserID uuid.UUID) ([]u } return groupIDs, nil } + +func findLiveSCIMUserIDs(tx *storage.Connection, providerID uuid.UUID, ids []uuid.UUID, lock bool) ([]uuid.UUID, error) { + found := []uuid.UUID{} + if len(ids) == 0 { + return found, nil + } + query, message := "SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL", "error finding SCIM group members" + if lock { + query, message = query+" ORDER BY id FOR SHARE", "error locking SCIM group members" + } + if err := tx.RawQuery(fmt.Sprintf(query, scimUsersTable.name), ids, providerID).All(&found); err != nil { + return nil, errors.Wrap(err, message) + } + if missing := differenceUUIDs(ids, found); len(missing) > 0 { + return nil, SCIMGroupMemberNotFoundError{IDs: missing} + } + return found, nil +} + +func findSCIMGroupMemberIDs(tx *storage.Connection, groupID uuid.UUID) ([]uuid.UUID, error) { + members := []SCIMGroupMember{} + if err := tx.RawQuery( + fmt.Sprintf("SELECT group_id, scim_user_id, created_at FROM %q WHERE group_id = ?", SCIMGroupMember{}.TableName()), + groupID, + ).All(&members); err != nil { + return nil, errors.Wrap(err, "error finding SCIM group members") + } + ids := make([]uuid.UUID, len(members)) + for i := range members { + ids[i] = members[i].SCIMUserID + } + return ids, nil +} + +func removeSCIMGroupMembers(tx *storage.Connection, groupID uuid.UUID, scimUserIDs []uuid.UUID) error { + if len(scimUserIDs) == 0 { + return nil + } + if err := tx.RawQuery( + fmt.Sprintf("DELETE FROM %q WHERE group_id = ? AND scim_user_id = ANY(?::uuid[])", SCIMGroupMember{}.TableName()), + groupID, scimUserIDs, + ).Exec(); err != nil { + return errors.Wrap(err, "error removing SCIM group members") + } + return nil +} + +func addSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserIDs []uuid.UUID) error { + if len(scimUserIDs) == 0 { + return nil + } + locked, err := findLiveSCIMUserIDs(tx, group.SSOProviderID, scimUserIDs, true) + if err != nil { + return err + } + if err := tx.RawQuery( + fmt.Sprintf("INSERT INTO %q (group_id, scim_user_id) SELECT ?, unnest(?::uuid[])", SCIMGroupMember{}.TableName()), + group.ID, locked, + ).Exec(); err != nil { + return errors.Wrap(err, "error adding SCIM group members") + } + return nil +} From 86f088bb2cc5f3bd78cfee313287131cae27685d Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:21:57 -0600 Subject: [PATCH 27/88] chore(scim): share provider and search setup between Users and Groups List --- internal/api/scim.go | 31 +++++++++++++++++++++++++++++++ internal/api/scim_groups.go | 24 ++++++------------------ internal/api/scim_users.go | 34 +++++++++++----------------------- 3 files changed, 48 insertions(+), 41 deletions(-) diff --git a/internal/api/scim.go b/internal/api/scim.go index e07862430b..3ebad7ad03 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -32,6 +32,13 @@ const ( scimResourceTypeGroup = "Group" ) +type scimLister[Row, Resource any] struct { + schemas core.Schemas + name string + find func(*storage.Connection, uuid.UUID, models.SCIMQuery) ([]Row, int, error) + render func(*storage.Connection, uuid.UUID, []Row, protocol.Projection) ([]Resource, error) +} + var errMissingSSOProvider = errors.New("scim: request has no SSO provider") var ( @@ -42,6 +49,26 @@ var ( scimGroupSchemas = newSCIMGroupSchemas() ) +func (l scimLister[Row, Resource]) list(ctx context.Context, db *storage.Connection, query *protocol.SearchRequest) ([]Resource, int, error) { + providerID, err := scimProviderID(ctx) + if err != nil { + return nil, 0, err + } + search, err := scimSearch(query, l.schemas, l.name) + if err != nil { + return nil, 0, err + } + rows, total, err := l.find(db, providerID, search) + if err != nil { + return nil, 0, err + } + resources, err := l.render(db, providerID, rows, protocol.ProjectionFrom(ctx)) + if err != nil { + return nil, 0, err + } + return resources, total, nil +} + func (a *API) newSCIMServer(validate server.TokenValidator, limit func(http.Handler) http.Handler) *server.Server { requireToken := server.RequireBearerToken(validate) authenticate := func(next http.Handler) http.Handler { @@ -225,6 +252,10 @@ func scimProviderID(ctx context.Context) (uuid.UUID, error) { return token.SSOProviderID, nil } +func scimProviderType(providerID uuid.UUID) string { + return "sso:" + providerID.String() +} + func scimRequest(ctx context.Context) (*http.Request, error) { r := scimRequestKey.Value(ctx) if r == nil { diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index 85bccbd67a..300bfdb9fd 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -17,24 +17,12 @@ type scimGroupRepository struct { } func (s *scimGroupRepository) List(ctx context.Context, query *protocol.SearchRequest) ([]*core.Group, int, error) { - providerID, err := scimProviderID(ctx) - if err != nil { - return nil, 0, err - } - search, err := scimSearch(query, scimGroupSchemas, "displayName") - if err != nil { - return nil, 0, err - } - db := s.api.db.WithContext(ctx) - rows, total, err := models.FindSCIMGroups(db, providerID, search) - if err != nil { - return nil, 0, err - } - groups, err := s.render(db, providerID, rows, protocol.ProjectionFrom(ctx)) - if err != nil { - return nil, 0, err - } - return groups, total, nil + return scimLister[models.SCIMGroup, *core.Group]{ + schemas: scimGroupSchemas, + name: "displayName", + find: models.FindSCIMGroups, + render: s.render, + }.list(ctx, s.api.db.WithContext(ctx), query) } func (s *scimGroupRepository) Get(ctx context.Context, id string) (*core.Group, error) { diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index d31f21306d..ff4f90a1cd 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -25,24 +25,12 @@ type scimUserRepository struct { } func (s *scimUserRepository) List(ctx context.Context, query *protocol.SearchRequest) ([]*core.User, int, error) { - providerID, err := scimProviderID(ctx) - if err != nil { - return nil, 0, err - } - search, err := scimSearch(query, scimUserSchemas, "userName") - if err != nil { - return nil, 0, err - } - db := s.api.db.WithContext(ctx) - rows, total, err := models.FindSCIMUsers(db, providerID, search) - if err != nil { - return nil, 0, err - } - users, err := s.render(db, providerID, rows, protocol.ProjectionFrom(ctx)) - if err != nil { - return nil, 0, err - } - return users, total, nil + return scimLister[models.SCIMUser, *core.User]{ + schemas: scimUserSchemas, + name: "userName", + find: models.FindSCIMUsers, + render: s.render, + }.list(ctx, s.api.db.WithContext(ctx), query) } func (s *scimUserRepository) Get(ctx context.Context, id string) (*core.User, error) { @@ -85,7 +73,7 @@ func (s *scimUserRepository) Create(ctx context.Context, user *core.User) (*core var row *models.SCIMUser var created *models.User err = db.Transaction(func(tx *storage.Connection) error { - if terr := models.LockAccountLinking(tx, "sso:"+providerID.String(), scimUserEmail(user)); terr != nil { + if terr := models.LockAccountLinking(tx, scimProviderType(providerID), scimUserEmail(user)); terr != nil { return terr } var terr error @@ -139,7 +127,7 @@ func (s *scimUserRepository) Replace(ctx context.Context, user *core.User) (*cor var row *models.SCIMUser var created *models.User err = db.Transaction(func(tx *storage.Connection) error { - if terr := models.LockAccountLinking(tx, "sso:"+providerID.String(), email); terr != nil { + if terr := models.LockAccountLinking(tx, scimProviderType(providerID), email); terr != nil { return terr } if existing.UserID != nil { @@ -279,7 +267,7 @@ func (s *scimUserRepository) syncAuthUser(tx *storage.Connection, providerID uui } err := models.RenameSCIMIdentity(tx, models.SCIMIdentityRename{ UserID: linked.ID, - Provider: "sso:" + providerID.String(), + Provider: scimProviderType(providerID), From: from, To: user.UserName, Data: data, @@ -312,7 +300,7 @@ func (s *scimUserRepository) provisionAuthUser(tx *storage.Connection, row *mode } func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SCIMUser, user *core.User) (*models.User, bool, error) { - providerType := "sso:" + row.SSOProviderID.String() + providerType := scimProviderType(row.SSOProviderID) decision, err := s.decideAccountLinking(tx, providerType, user) if err != nil { return nil, false, err @@ -352,7 +340,7 @@ func (s *scimUserRepository) runBeforeUserCreatedHook(r *http.Request, db *stora if !s.api.hooksMgr.Enabled(v0hooks.BeforeUserCreated) { return nil } - providerType := "sso:" + providerID.String() + providerType := scimProviderType(providerID) decision, err := s.decideAccountLinking(db, providerType, user) if err != nil || decision.Decision != models.CreateAccount { return err From feb886d8629e2ac28c9b368310afeeec1ed39703 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:24:26 -0600 Subject: [PATCH 28/88] chore(scim): split SCIM user Create, Replace and Delete into steps --- internal/api/scim_users.go | 322 +++++++++++++++++++++++-------------- 1 file changed, 199 insertions(+), 123 deletions(-) diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index ff4f90a1cd..b3a4bce2c7 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -24,6 +24,13 @@ type scimUserRepository struct { api *API } +type scimUserChange struct { + r *http.Request + target models.SCIMTarget + resource []byte + user *core.User +} + func (s *scimUserRepository) List(ctx context.Context, query *protocol.SearchRequest) ([]*core.User, int, error) { return scimLister[models.SCIMUser, *core.User]{ schemas: scimUserSchemas, @@ -55,35 +62,22 @@ func (s *scimUserRepository) Create(ctx context.Context, user *core.User) (*core if err != nil { return nil, err } - if err := scimValidatePrimaryEmail(user); err != nil { - return nil, err - } - if scimUserEmail(user) == "" { - return nil, errSCIMEmailRequired() - } r, err := scimRequest(ctx) if err != nil { return nil, err } db := s.api.db.WithContext(ctx) - if err := s.runBeforeUserCreatedHook(r, db, providerID, user); err != nil { - return nil, scimError(err) + if err := s.beforeProvision(r, db, providerID, user); err != nil { + return nil, err } + change := scimUserChange{r: r, target: models.SCIMTarget{ProviderID: providerID}, resource: resource, user: user} var row *models.SCIMUser var created *models.User err = db.Transaction(func(tx *storage.Connection) error { - if terr := models.LockAccountLinking(tx, scimProviderType(providerID), scimUserEmail(user)); terr != nil { - return terr - } var terr error - if row, terr = models.CreateSCIMUser(tx, providerID, resource); terr != nil { - return terr - } - if created, terr = s.provisionAuthUser(tx, row, user); terr != nil { - return terr - } - return s.audit(tx, r, models.SCIMUserCreatedAction, row) + row, created, terr = s.create(tx, change) + return terr }) if err != nil { return nil, scimError(err) @@ -97,69 +91,38 @@ func (s *scimUserRepository) Replace(ctx context.Context, user *core.User) (*cor if err != nil { return nil, err } - providerID := target.ProviderID resource, err := scimUserResource(user) if err != nil { return nil, err } - if err := scimValidatePrimaryEmail(user); err != nil { - return nil, err - } - email := scimUserEmail(user) r, err := scimRequest(ctx) if err != nil { return nil, err } db := s.api.db.WithContext(ctx) - existing, err := models.FindSCIMUser(db, providerID, target.ID) + existing, err := models.FindSCIMUser(db, target.ProviderID, target.ID) if err != nil { return nil, scimError(err) } if existing.UserID == nil { - if email == "" { - return nil, errSCIMEmailRequired() - } - if err := s.runBeforeUserCreatedHook(r, db, providerID, user); err != nil { - return nil, scimError(err) + if err := s.beforeProvision(r, db, target.ProviderID, user); err != nil { + return nil, err } } + change := scimUserChange{r: r, target: target, resource: resource, user: user} var row *models.SCIMUser var created *models.User err = db.Transaction(func(tx *storage.Connection) error { - if terr := models.LockAccountLinking(tx, scimProviderType(providerID), email); terr != nil { - return terr - } - if existing.UserID != nil { - if terr := models.LockUserForSCIM(tx, *existing.UserID); terr != nil { - return terr - } - } - old, terr := models.FindSCIMUserForUpdate(tx, providerID, target.ID) - if terr != nil { - return terr - } - if old.UserID != nil { - if terr := models.LockUserForSCIM(tx, *old.UserID); terr != nil { - return terr - } - if row, terr = models.FindUnchangedSCIMUser(tx, target, resource); terr != nil || row != nil { - return terr - } - } - if row, terr = models.ReplaceSCIMUser(tx, target, resource); terr != nil { - return terr - } - if created, terr = s.syncAuthUser(tx, providerID, old, row, user); terr != nil { - return terr - } - return s.audit(tx, r, scimUserAuditAction(old, row), row) + var terr error + row, created, terr = s.replace(tx, change, existing) + return terr }) if err != nil { return nil, scimError(err) } s.runAfterUserCreatedHook(r, db, created) - return s.renderOne(db, providerID, row, protocol.Projection{}) + return s.renderOne(db, target.ProviderID, row, protocol.Projection{}) } func (s *scimUserRepository) Delete(ctx context.Context, id, version string) error { @@ -177,50 +140,16 @@ func (s *scimUserRepository) Delete(ctx context.Context, id, version string) err return scimError(err) } return scimError(db.Transaction(func(tx *storage.Connection) error { - if existing.UserID != nil { - if err := models.LockUserForSCIM(tx, *existing.UserID); err != nil { - return err - } - } - row, err := models.DeleteSCIMUser(tx, target) - if err != nil { - return err - } - if row.UserID != nil { - if err := models.LogoutSCIMUser(tx, *row.UserID); err != nil { - return err - } - } - if err := s.api.removeSCIMUserFromGroups(tx, r, scimActor(r), row); err != nil { - return err - } - return s.audit(tx, r, models.SCIMUserDeletedAction, row) + return s.delete(tx, r, target, existing) })) } func (s *scimUserRepository) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMUser, projection protocol.Projection) ([]*core.User, error) { - memberships := []models.SCIMGroupMembership{} - if projection.Returns("groups") { - ids := make([]uuid.UUID, len(rows)) - for i, row := range rows { - ids[i] = row.ID - } - var err error - if memberships, err = models.FindSCIMGroupsForUsers(tx, providerID, ids); err != nil { - return nil, err - } + groups, err := s.groupMemberships(tx, providerID, rows, projection) + if err != nil { + return nil, err } base := scimBaseURL(s.api.config) - groups := map[uuid.UUID][]core.GroupMembership{} - for _, m := range memberships { - groups[m.SCIMUserID] = append(groups[m.SCIMUserID], core.GroupMembership{ - Value: m.GroupID.String(), - Ref: base + "/Groups/" + m.GroupID.String(), - Display: m.Display, - Type: "direct", - }) - } - users := make([]*core.User, 0, len(rows)) for _, row := range rows { user := &core.User{} @@ -240,6 +169,31 @@ func (s *scimUserRepository) render(tx *storage.Connection, providerID uuid.UUID return users, nil } +func (s *scimUserRepository) groupMemberships(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMUser, projection protocol.Projection) (map[uuid.UUID][]core.GroupMembership, error) { + groups := map[uuid.UUID][]core.GroupMembership{} + if !projection.Returns("groups") { + return groups, nil + } + ids := make([]uuid.UUID, len(rows)) + for i, row := range rows { + ids[i] = row.ID + } + memberships, err := models.FindSCIMGroupsForUsers(tx, providerID, ids) + if err != nil { + return nil, err + } + base := scimBaseURL(s.api.config) + for _, m := range memberships { + groups[m.SCIMUserID] = append(groups[m.SCIMUserID], core.GroupMembership{ + Value: m.GroupID.String(), + Ref: base + "/Groups/" + m.GroupID.String(), + Display: m.Display, + Type: "direct", + }) + } + return groups, nil +} + func (s *scimUserRepository) renderOne(tx *storage.Connection, providerID uuid.UUID, row *models.SCIMUser, projection protocol.Projection) (*core.User, error) { users, err := s.render(tx, providerID, []models.SCIMUser{*row}, projection) if err != nil { @@ -248,35 +202,96 @@ func (s *scimUserRepository) renderOne(tx *storage.Connection, providerID uuid.U return users[0], nil } -func (s *scimUserRepository) syncAuthUser(tx *storage.Connection, providerID uuid.UUID, old, row *models.SCIMUser, user *core.User) (*models.User, error) { +func (s *scimUserRepository) create(tx *storage.Connection, change scimUserChange) (*models.SCIMUser, *models.User, error) { + if err := models.LockAccountLinking(tx, scimProviderType(change.target.ProviderID), scimUserEmail(change.user)); err != nil { + return nil, nil, err + } + row, err := models.CreateSCIMUser(tx, change.target.ProviderID, change.resource) + if err != nil { + return nil, nil, err + } + created, err := s.provisionAuthUser(tx, row, change.user) + if err != nil { + return nil, nil, err + } + return row, created, s.audit(tx, change.r, models.SCIMUserCreatedAction, row) +} + +func (s *scimUserRepository) replace(tx *storage.Connection, change scimUserChange, existing *models.SCIMUser) (*models.SCIMUser, *models.User, error) { + old, err := s.lockForReplace(tx, change.target, scimUserEmail(change.user), existing) + if err != nil { + return nil, nil, err + } + row, changed, err := s.replaceRow(tx, change.target, old, change.resource) + if err != nil || !changed { + return row, nil, err + } + created, err := s.syncAuthUser(tx, change, old, row) + if err != nil { + return nil, nil, err + } + return row, created, s.audit(tx, change.r, scimUserAuditAction(old, row), row) +} + +func (s *scimUserRepository) delete(tx *storage.Connection, r *http.Request, target models.SCIMTarget, existing *models.SCIMUser) error { + if err := lockSCIMLinkedUser(tx, existing.UserID); err != nil { + return err + } + row, err := models.DeleteSCIMUser(tx, target) + if err != nil { + return err + } + if err := logoutSCIMLinkedUser(tx, row.UserID); err != nil { + return err + } + if err := s.api.removeSCIMUserFromGroups(tx, r, scimActor(r), row); err != nil { + return err + } + return s.audit(tx, r, models.SCIMUserDeletedAction, row) +} + +func (s *scimUserRepository) lockForReplace(tx *storage.Connection, target models.SCIMTarget, email string, existing *models.SCIMUser) (*models.SCIMUser, error) { + if err := models.LockAccountLinking(tx, scimProviderType(target.ProviderID), email); err != nil { + return nil, err + } + if err := lockSCIMLinkedUser(tx, existing.UserID); err != nil { + return nil, err + } + old, err := models.FindSCIMUserForUpdate(tx, target.ProviderID, target.ID) + if err != nil { + return nil, err + } + if err := lockSCIMLinkedUser(tx, old.UserID); err != nil { + return nil, err + } + return old, nil +} + +func (s *scimUserRepository) replaceRow(tx *storage.Connection, target models.SCIMTarget, old *models.SCIMUser, resource []byte) (*models.SCIMUser, bool, error) { + if old.UserID != nil { + unchanged, err := models.FindUnchangedSCIMUser(tx, target, resource) + if err != nil || unchanged != nil { + return unchanged, false, err + } + } + row, err := models.ReplaceSCIMUser(tx, target, resource) + return row, err == nil, err +} + +func (s *scimUserRepository) syncAuthUser(tx *storage.Connection, change scimUserChange, old, row *models.SCIMUser) (*models.User, error) { if old.UserID == nil { - if scimUserEmail(user) == "" { + if scimUserEmail(change.user) == "" { return nil, errSCIMEmailRequired() } - return s.provisionAuthUser(tx, row, user) + return s.provisionAuthUser(tx, row, change.user) } linked, err := models.FindUserByID(tx, *old.UserID) if err != nil { return nil, err } - if from := scimUserName(old.Resource); from != user.UserName { - data := map[string]any{"sub": user.UserName} - if email := scimUserEmail(user); email != "" { - data["email"] = email - } - err := models.RenameSCIMIdentity(tx, models.SCIMIdentityRename{ - UserID: linked.ID, - Provider: scimProviderType(providerID), - From: from, - To: user.UserName, - Data: data, - }) - if errors.Is(err, models.SCIMIdentityNotFoundError{}) { - logrus.WithField("user_id", linked.ID).WithField("sso_provider_id", providerID).Warn("scim: SCIM identity not found, rename skipped") - } else if err != nil { - return nil, err - } + if err := s.renameIdentity(tx, change.target.ProviderID, old, change.user); err != nil { + return nil, err } if old.Active && !row.Active { return nil, models.LogoutSCIMUser(tx, linked.ID) @@ -317,13 +332,7 @@ func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SC return nil, false, err } case models.CreateAccount: - if linked, err = s.newUser(providerType, decision, user); err != nil { - return nil, false, err - } - if linked, err = s.api.signupNewUser(tx, linked); err != nil { - return nil, false, err - } - if _, err = s.api.createNewIdentity(tx, linked, providerType, scimIdentityData(user)); err != nil { + if linked, err = s.createAuthUser(tx, providerType, decision, user); err != nil { return nil, false, err } return linked, true, models.LinkSCIMUser(tx, row, linked.ID) @@ -336,6 +345,28 @@ func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SC return linked, false, models.LinkSCIMUser(tx, row, linked.ID) } +func (s *scimUserRepository) createAuthUser(tx *storage.Connection, providerType string, decision models.AccountLinkingResult, user *core.User) (*models.User, error) { + candidate, err := s.newUser(providerType, decision, user) + if err != nil { + return nil, err + } + created, err := s.api.signupNewUser(tx, candidate) + if err != nil { + return nil, err + } + if _, err := s.api.createNewIdentity(tx, created, providerType, scimIdentityData(user)); err != nil { + return nil, err + } + return created, nil +} + +func (s *scimUserRepository) beforeProvision(r *http.Request, db *storage.Connection, providerID uuid.UUID, user *core.User) error { + if scimUserEmail(user) == "" { + return errSCIMEmailRequired() + } + return scimError(s.runBeforeUserCreatedHook(r, db, providerID, user)) +} + func (s *scimUserRepository) runBeforeUserCreatedHook(r *http.Request, db *storage.Connection, providerID uuid.UUID, user *core.User) error { if !s.api.hooksMgr.Enabled(v0hooks.BeforeUserCreated) { return nil @@ -382,6 +413,30 @@ func (s *scimUserRepository) newUser(providerType string, decision models.Accoun return candidate, nil } +func (s *scimUserRepository) renameIdentity(tx *storage.Connection, providerID uuid.UUID, old *models.SCIMUser, user *core.User) error { + userID := *old.UserID + from := scimUserName(old.Resource) + if from == user.UserName { + return nil + } + data := map[string]any{"sub": user.UserName} + if email := scimUserEmail(user); email != "" { + data["email"] = email + } + err := models.RenameSCIMIdentity(tx, models.SCIMIdentityRename{ + UserID: userID, + Provider: scimProviderType(providerID), + From: from, + To: user.UserName, + Data: data, + }) + if errors.Is(err, models.SCIMIdentityNotFoundError{}) { + logrus.WithField("user_id", userID).WithField("sso_provider_id", providerID).Warn("scim: SCIM identity not found, rename skipped") + return nil + } + return err +} + func (s *scimUserRepository) audit(tx *storage.Connection, r *http.Request, action models.AuditAction, row *models.SCIMUser) error { return s.api.auditSCIM(tx, r, scimActor(r), action, row.SSOProviderID, scimUserTraits(row)) } @@ -416,7 +471,28 @@ func (a *API) removeSCIMUserFromGroups(tx *storage.Connection, r *http.Request, } func scimUserResource(user *core.User) ([]byte, error) { - return scimEncode(user, "id", "meta", "password", "groups") + resource, err := scimEncode(user, "id", "meta", "password", "groups") + if err != nil { + return nil, err + } + if err := scimValidatePrimaryEmail(user); err != nil { + return nil, err + } + return resource, nil +} + +func lockSCIMLinkedUser(tx *storage.Connection, userID *uuid.UUID) error { + if userID == nil { + return nil + } + return models.LockUserForSCIM(tx, *userID) +} + +func logoutSCIMLinkedUser(tx *storage.Connection, userID *uuid.UUID) error { + if userID == nil { + return nil + } + return models.LogoutSCIMUser(tx, *userID) } func scimUserEmail(user *core.User) string { From 3b16e99d91c3c7e764486be46fbf4fb2836420d0 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:25:54 -0600 Subject: [PATCH 29/88] chore(scim): flatten SCIM group save --- internal/api/scim_groups.go | 97 ++++++++++++++++++++++--------------- 1 file changed, 59 insertions(+), 38 deletions(-) diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index 300bfdb9fd..33cd69ce57 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -16,6 +16,12 @@ type scimGroupRepository struct { api *API } +type scimGroupChange struct { + r *http.Request + action models.AuditAction + members []uuid.UUID +} + func (s *scimGroupRepository) List(ctx context.Context, query *protocol.SearchRequest) ([]*core.Group, int, error) { return scimLister[models.SCIMGroup, *core.Group]{ schemas: scimGroupSchemas, @@ -43,7 +49,7 @@ func (s *scimGroupRepository) Create(ctx context.Context, group *core.Group) (*c if err != nil { return nil, err } - return s.save(ctx, providerID, models.SCIMGroupCreatedAction, group, func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error) { + return s.save(ctx, models.SCIMGroupCreatedAction, group, func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error) { row, err := models.CreateSCIMGroup(tx, providerID, resource) return row, true, err }) @@ -54,7 +60,7 @@ func (s *scimGroupRepository) Replace(ctx context.Context, group *core.Group) (* if err != nil { return nil, err } - return s.save(ctx, target.ProviderID, models.SCIMGroupUpdatedAction, group, func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error) { + return s.save(ctx, models.SCIMGroupUpdatedAction, group, func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error) { unchanged, err := models.FindUnchangedSCIMGroup(tx, target, resource) if err != nil || unchanged != nil { return unchanged, false, err @@ -92,7 +98,7 @@ func (s *scimGroupRepository) Delete(ctx context.Context, id, version string) er })) } -func (s *scimGroupRepository) save(ctx context.Context, providerID uuid.UUID, action models.AuditAction, group *core.Group, write func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error)) (*core.Group, error) { +func (s *scimGroupRepository) save(ctx context.Context, action models.AuditAction, group *core.Group, write func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error)) (*core.Group, error) { members, err := scimMemberIDs(group.Members) if err != nil { return nil, err @@ -105,6 +111,7 @@ func (s *scimGroupRepository) save(ctx context.Context, providerID uuid.UUID, ac if err != nil { return nil, err } + change := scimGroupChange{r: r, action: action, members: members} db := s.api.db.WithContext(ctx) var row *models.SCIMGroup err = db.Transaction(func(tx *storage.Connection) error { @@ -115,50 +122,40 @@ func (s *scimGroupRepository) save(ctx context.Context, providerID uuid.UUID, ac if row, changed, terr = write(tx, resource); terr != nil { return terr } - added, removed, terr := models.ReplaceSCIMGroupMembers(tx, row, members) - if terr != nil { - return terr - } - if !changed { - if len(added) == 0 && len(removed) == 0 { - return nil - } - if row, terr = models.TouchSCIMGroup(tx, row); terr != nil { - return terr - } - } else if terr := s.audit(tx, r, action, row); terr != nil { - return terr - } - return s.auditMembers(tx, r, row, added, removed) + row, terr = s.applyMembers(tx, change, row, changed) + return terr }) if err != nil { return nil, scimError(err) } - return s.renderOne(db, providerID, row, protocol.Projection{}) + return s.renderOne(db, row.SSOProviderID, row, protocol.Projection{}) } -func (s *scimGroupRepository) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMGroup, projection protocol.Projection) ([]*core.Group, error) { - memberships := []models.SCIMGroupMembership{} - if projection.Returns("members") { - ids := make([]uuid.UUID, len(rows)) - for i, row := range rows { - ids[i] = row.ID - } - var err error - if memberships, err = models.FindSCIMGroupMembers(tx, providerID, ids); err != nil { - return nil, err - } +func (s *scimGroupRepository) applyMembers(tx *storage.Connection, change scimGroupChange, row *models.SCIMGroup, changed bool) (*models.SCIMGroup, error) { + added, removed, err := models.ReplaceSCIMGroupMembers(tx, row, change.members) + if err != nil { + return nil, err } - base := scimBaseURL(s.api.config) - members := map[uuid.UUID][]core.Member{} - for _, m := range memberships { - members[m.GroupID] = append(members[m.GroupID], core.Member{ - Value: m.SCIMUserID.String(), - Ref: base + "/Users/" + m.SCIMUserID.String(), - Type: scimResourceTypeUser, - }) + switch { + case changed: + err = s.audit(tx, change.r, change.action, row) + case len(added) > 0 || len(removed) > 0: + row, err = models.TouchSCIMGroup(tx, row) + default: + return row, nil } + if err != nil { + return nil, err + } + return row, s.auditMembers(tx, change.r, row, added, removed) +} +func (s *scimGroupRepository) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMGroup, projection protocol.Projection) ([]*core.Group, error) { + members, err := s.members(tx, providerID, rows, projection) + if err != nil { + return nil, err + } + base := scimBaseURL(s.api.config) groups := make([]*core.Group, 0, len(rows)) for _, row := range rows { group := &core.Group{} @@ -174,6 +171,30 @@ func (s *scimGroupRepository) render(tx *storage.Connection, providerID uuid.UUI return groups, nil } +func (s *scimGroupRepository) members(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMGroup, projection protocol.Projection) (map[uuid.UUID][]core.Member, error) { + members := map[uuid.UUID][]core.Member{} + if !projection.Returns("members") { + return members, nil + } + ids := make([]uuid.UUID, len(rows)) + for i, row := range rows { + ids[i] = row.ID + } + memberships, err := models.FindSCIMGroupMembers(tx, providerID, ids) + if err != nil { + return nil, err + } + base := scimBaseURL(s.api.config) + for _, m := range memberships { + members[m.GroupID] = append(members[m.GroupID], core.Member{ + Value: m.SCIMUserID.String(), + Ref: base + "/Users/" + m.SCIMUserID.String(), + Type: scimResourceTypeUser, + }) + } + return members, nil +} + func (s *scimGroupRepository) renderOne(tx *storage.Connection, providerID uuid.UUID, row *models.SCIMGroup, projection protocol.Projection) (*core.Group, error) { groups, err := s.render(tx, providerID, []models.SCIMGroup{*row}, projection) if err != nil { From 4710b8b009deb30d2eeb404551b1fa34a5a7db08 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:26:56 -0600 Subject: [PATCH 30/88] chore(scim): split SCIM admin token create and revoke --- internal/api/scim_admin.go | 56 ++++++++++++++++++++++++-------------- 1 file changed, 35 insertions(+), 21 deletions(-) diff --git a/internal/api/scim_admin.go b/internal/api/scim_admin.go index 3f522e88db..986e1f21da 100644 --- a/internal/api/scim_admin.go +++ b/internal/api/scim_admin.go @@ -56,14 +56,9 @@ func (a *API) adminSCIMTokensCreate(w http.ResponseWriter, r *http.Request) erro db := a.db.WithContext(ctx) provider := getSSOProvider(ctx) - params := &AdminSCIMTokenCreateParams{} - if body, err := utilities.GetBodyBytes(r); err != nil || len(body) > 0 { - if err := retrieveRequestParams(r, params); err != nil { - return err - } - } - if params.ExpiresAt != nil && !params.ExpiresAt.After(a.Now()) { - return apierrors.NewBadRequestError(apierrors.ErrorCodeValidationFailed, "expires_at must be in the future") + params, err := a.scimTokenCreateParams(r) + if err != nil { + return err } var ( @@ -112,20 +107,9 @@ func (a *API) adminSCIMTokensRevoke(w http.ResponseWriter, r *http.Request) erro var token *models.SCIMToken if err := db.Transaction(func(tx *storage.Connection) error { - if err := models.LockSCIMTokens(tx, provider.ID); err != nil { - return err - } var err error - if token, err = models.FindSCIMTokenByPrefix(tx, provider.ID, chi.URLParam(r, "prefix")); err != nil { - return err - } - if token.IsRevoked() { - return nil - } - if err = token.Revoke(tx); err != nil { - return err - } - return a.auditSCIM(tx, r, getAdminUser(ctx), models.SCIMTokenRevokedAction, provider.ID, map[string]any{scimTokenPrefixTrait: token.Prefix}) + token, err = a.revokeSCIMToken(tx, r, provider.ID, chi.URLParam(r, "prefix")) + return err }); err != nil { if models.IsNotFoundError(err) { return apierrors.NewNotFoundError(apierrors.ErrorCodeSCIMTokenNotFound, "SCIM token not found") @@ -136,6 +120,36 @@ func (a *API) adminSCIMTokensRevoke(w http.ResponseWriter, r *http.Request) erro return sendJSON(w, http.StatusOK, token) } +func (a *API) scimTokenCreateParams(r *http.Request) (*AdminSCIMTokenCreateParams, error) { + params := &AdminSCIMTokenCreateParams{} + if body, err := utilities.GetBodyBytes(r); err != nil || len(body) > 0 { + if err := retrieveRequestParams(r, params); err != nil { + return nil, err + } + } + if params.ExpiresAt != nil && !params.ExpiresAt.After(a.Now()) { + return nil, apierrors.NewBadRequestError(apierrors.ErrorCodeValidationFailed, "expires_at must be in the future") + } + return params, nil +} + +func (a *API) revokeSCIMToken(tx *storage.Connection, r *http.Request, providerID uuid.UUID, prefix string) (*models.SCIMToken, error) { + if err := models.LockSCIMTokens(tx, providerID); err != nil { + return nil, err + } + token, err := models.FindSCIMTokenByPrefix(tx, providerID, prefix) + if err != nil { + return nil, err + } + if token.IsRevoked() { + return token, nil + } + if err := token.Revoke(tx); err != nil { + return nil, err + } + return token, a.auditSCIM(tx, r, getAdminUser(r.Context()), models.SCIMTokenRevokedAction, providerID, map[string]any{scimTokenPrefixTrait: token.Prefix}) +} + func (a *API) changeSCIMEnabled(w http.ResponseWriter, r *http.Request, change func(*storage.Connection, uuid.UUID) (bool, error), action models.AuditAction, verb string) error { ctx := r.Context() db := a.db.WithContext(ctx) From f6218b24f6b700de2eb20f67d81be4e519baa7c5 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:28:14 -0600 Subject: [PATCH 31/88] chore(scim): use the API clock and handle errors SCIM ignored --- internal/api/scim.go | 4 ++-- internal/api/scim_groups.go | 3 ++- internal/api/scim_users.go | 19 ++++++++++--------- 3 files changed, 14 insertions(+), 12 deletions(-) diff --git a/internal/api/scim.go b/internal/api/scim.go index 3ebad7ad03..fd47862584 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -39,9 +39,9 @@ type scimLister[Row, Resource any] struct { render func(*storage.Connection, uuid.UUID, []Row, protocol.Projection) ([]Resource, error) } -var errMissingSSOProvider = errors.New("scim: request has no SSO provider") - var ( + errMissingSSOProvider = errors.New("scim: request has no SSO provider") + scimUserSchemas = core.Schemas{ core.NewSchema(core.SchemaUser).With(core.UserAttributes()...), core.NewSchema(core.SchemaEnterpriseUser).With(core.EnterpriseUserAttributes()...), diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index 33cd69ce57..d7269337b9 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "net/http" + "slices" "github.com/gofrs/uuid" "github.com/supabase-community/scim-go/pkg/core" @@ -217,7 +218,7 @@ func (s *scimGroupRepository) audit(tx *storage.Connection, r *http.Request, act } func (s *scimGroupRepository) auditMembers(tx *storage.Connection, r *http.Request, row *models.SCIMGroup, added, removed []uuid.UUID) error { - links, err := models.FindSCIMUserLinks(tx, append(append([]uuid.UUID{}, added...), removed...)) + links, err := models.FindSCIMUserLinks(tx, slices.Concat(added, removed)) if err != nil { return err } diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index b3a4bce2c7..b266df5bf1 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -5,7 +5,6 @@ import ( "encoding/json" "errors" "net/http" - "time" "github.com/badoux/checkmail" "github.com/gofrs/uuid" @@ -408,22 +407,22 @@ func (s *scimUserRepository) newUser(providerType string, decision models.Accoun if err != nil { return nil, err } - now := time.Now() + now := s.api.Now() candidate.EmailConfirmedAt = &now return candidate, nil } func (s *scimUserRepository) renameIdentity(tx *storage.Connection, providerID uuid.UUID, old *models.SCIMUser, user *core.User) error { userID := *old.UserID - from := scimUserName(old.Resource) - if from == user.UserName { - return nil + from, err := scimUserName(old.Resource) + if err != nil || from == user.UserName { + return err } data := map[string]any{"sub": user.UserName} if email := scimUserEmail(user); email != "" { data["email"] = email } - err := models.RenameSCIMIdentity(tx, models.SCIMIdentityRename{ + err = models.RenameSCIMIdentity(tx, models.SCIMIdentityRename{ UserID: userID, Provider: scimProviderType(providerID), From: from, @@ -536,12 +535,14 @@ func scimIdentityData(user *core.User) map[string]any { } } -func scimUserName(resource []byte) string { +func scimUserName(resource []byte) (string, error) { var r struct { UserName string `json:"userName"` } - _ = json.Unmarshal(resource, &r) - return r.UserName + if err := json.Unmarshal(resource, &r); err != nil { + return "", err + } + return r.UserName, nil } func scimUserTraits(row *models.SCIMUser) map[string]any { From b47e3b4c385053bf419a321af2946a8b38509909 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:35:11 -0600 Subject: [PATCH 32/88] chore(scim): log SCIM user warnings through the request logger --- internal/api/scim_users.go | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index b266df5bf1..6f46f03c71 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -8,7 +8,6 @@ import ( "github.com/badoux/checkmail" "github.com/gofrs/uuid" - "github.com/sirupsen/logrus" "github.com/supabase-community/scim-go/pkg/core" "github.com/supabase-community/scim-go/pkg/protocol" "github.com/supabase-community/scim-go/pkg/scimerrors" @@ -16,6 +15,7 @@ import ( "github.com/supabase/auth/internal/api/provider" "github.com/supabase/auth/internal/hooks/v0hooks" "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/observability" "github.com/supabase/auth/internal/storage" ) @@ -289,7 +289,7 @@ func (s *scimUserRepository) syncAuthUser(tx *storage.Connection, change scimUse if err != nil { return nil, err } - if err := s.renameIdentity(tx, change.target.ProviderID, old, change.user); err != nil { + if err := s.renameIdentity(tx, change, old); err != nil { return nil, err } if old.Active && !row.Active { @@ -387,7 +387,7 @@ func (s *scimUserRepository) runAfterUserCreatedHook(r *http.Request, db *storag return } if err := s.api.triggerAfterUserCreated(r, db, user); err != nil { - logrus.WithError(err).WithField("user_id", user.ID).Error("scim: after user created hook failed") + observability.GetLogEntry(r).Entry.WithError(err).WithField("user_id", user.ID).Error("scim: after user created hook failed") } } @@ -412,8 +412,8 @@ func (s *scimUserRepository) newUser(providerType string, decision models.Accoun return candidate, nil } -func (s *scimUserRepository) renameIdentity(tx *storage.Connection, providerID uuid.UUID, old *models.SCIMUser, user *core.User) error { - userID := *old.UserID +func (s *scimUserRepository) renameIdentity(tx *storage.Connection, change scimUserChange, old *models.SCIMUser) error { + providerID, user, userID := change.target.ProviderID, change.user, *old.UserID from, err := scimUserName(old.Resource) if err != nil || from == user.UserName { return err @@ -430,7 +430,7 @@ func (s *scimUserRepository) renameIdentity(tx *storage.Connection, providerID u Data: data, }) if errors.Is(err, models.SCIMIdentityNotFoundError{}) { - logrus.WithField("user_id", userID).WithField("sso_provider_id", providerID).Warn("scim: SCIM identity not found, rename skipped") + observability.GetLogEntry(change.r).Entry.WithField("user_id", userID).WithField("sso_provider_id", providerID).Warn("scim: SCIM identity not found, rename skipped") return nil } return err From 7fb5885f9d8c7d8e6133664da54c25598e5a2d73 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:39:39 -0600 Subject: [PATCH 33/88] chore(scim): name filter conditions, use field names and unexport token helpers --- internal/api/scim_filter.go | 6 ++++-- internal/api/scim_groups.go | 4 ++-- internal/models/scim_token.go | 20 ++++++++++---------- internal/models/scim_token_test.go | 4 ++-- 4 files changed, 18 insertions(+), 16 deletions(-) diff --git a/internal/api/scim_filter.go b/internal/api/scim_filter.go index 3129bdcfc4..ef1c2e0f31 100644 --- a/internal/api/scim_filter.go +++ b/internal/api/scim_filter.go @@ -14,8 +14,10 @@ type scimEqFilter struct { } func (f scimEqFilter) Compare(attribute *protocol.Attribute, op filter.Operator, value any) (models.SCIMFilter, error) { - text, ok := value.(string) - if op != filter.OpEquals || attribute.Parent != nil || !ok { + text, isString := value.(string) + isEquals := op == filter.OpEquals + isTopLevel := attribute.Parent == nil + if !isEquals || !isTopLevel || !isString { return f.unsupported() } switch attribute.Definition.Name { diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index d7269337b9..35fbeac8ed 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -226,8 +226,8 @@ func (s *scimGroupRepository) auditMembers(tx *storage.Connection, r *http.Reque action models.AuditAction ids []uuid.UUID }{ - {models.SCIMGroupMemberAddedAction, added}, - {models.SCIMGroupMemberRemovedAction, removed}, + {action: models.SCIMGroupMemberAddedAction, ids: added}, + {action: models.SCIMGroupMemberRemovedAction, ids: removed}, } { for _, id := range change.ids { var userID *uuid.UUID diff --git a/internal/models/scim_token.go b/internal/models/scim_token.go index 7c13e7aa3e..78bafdc7ab 100644 --- a/internal/models/scim_token.go +++ b/internal/models/scim_token.go @@ -15,10 +15,10 @@ import ( ) const ( - SCIMTokenMarker = "scim_" + scimTokenMarker = "scim_" scimTokenBytes = 20 - scimTokenPrefixLength = len(SCIMTokenMarker) + 7 + scimTokenPrefixLength = len(scimTokenMarker) + 7 ) type SCIMToken struct { @@ -50,11 +50,6 @@ func (t *SCIMToken) IsRevoked() bool { return t.RevokedAt != nil } -func HashSCIMToken(token string) string { - sum := sha256.Sum256([]byte(token)) - return hex.EncodeToString(sum[:]) -} - func CreateSCIMToken(tx *storage.Connection, provider *SSOProvider, expiresAt *time.Time) (*SCIMToken, string, error) { plaintext, err := generateSCIMToken() if err != nil { @@ -64,7 +59,7 @@ func CreateSCIMToken(tx *storage.Connection, provider *SSOProvider, expiresAt *t token := &SCIMToken{ ID: uuid.Must(uuid.NewV4()), SSOProviderID: provider.ID, - TokenHash: HashSCIMToken(plaintext), + TokenHash: hashSCIMToken(plaintext), Prefix: plaintext[:scimTokenPrefixLength], ExpiresAt: expiresAt, } @@ -155,7 +150,7 @@ func AuthenticateSCIMToken(tx *storage.Connection, plaintext string) (*SCIMToken SELECT * FROM touched UNION ALL SELECT * FROM authenticated WHERE NOT EXISTS (SELECT 1 FROM touched)`, token.TableName(), SSOProvider{}.TableName(), SCIMSettings{}.TableName()), - HashSCIMToken(plaintext), + hashSCIMToken(plaintext), ).First(token) if err != nil { if errors.Is(err, sql.ErrNoRows) { @@ -171,5 +166,10 @@ func generateSCIMToken() (string, error) { if _, err := rand.Read(b); err != nil { return "", err } - return SCIMTokenMarker + hex.EncodeToString(b), nil + return scimTokenMarker + hex.EncodeToString(b), nil +} + +func hashSCIMToken(token string) string { + sum := sha256.Sum256([]byte(token)) + return hex.EncodeToString(sum[:]) } diff --git a/internal/models/scim_token_test.go b/internal/models/scim_token_test.go index 387485ed49..175fcee92e 100644 --- a/internal/models/scim_token_test.go +++ b/internal/models/scim_token_test.go @@ -57,7 +57,7 @@ func (ts *SCIMTokenTestSuite) TestCreate() { require.Regexp(ts.T(), regexp.MustCompile(`^scim_[0-9a-f]{40}$`), plaintext) require.Equal(ts.T(), plaintext[:12], token.Prefix) - require.Equal(ts.T(), HashSCIMToken(plaintext), token.TokenHash) + require.Equal(ts.T(), hashSCIMToken(plaintext), token.TokenHash) require.NotEqual(ts.T(), plaintext, token.TokenHash) require.Equal(ts.T(), ts.provider.ID, token.SSOProviderID) require.False(ts.T(), token.CreatedAt.IsZero()) @@ -145,7 +145,7 @@ func (ts *SCIMTokenTestSuite) TestFindByPrefixAmbiguous() { duplicate := &SCIMToken{ ID: uuid.Must(uuid.NewV4()), SSOProviderID: ts.provider.ID, - TokenHash: HashSCIMToken("duplicate"), + TokenHash: hashSCIMToken("duplicate"), Prefix: token.Prefix, } require.NoError(ts.T(), ts.db.RawQuery( From 915423073a8ed65a810e8ce15d6f39f945e285c9 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:40:57 -0600 Subject: [PATCH 34/88] chore(scim): stop asserting in test goroutines and use matching testify helpers --- internal/api/scim_link_test.go | 11 +++++------ internal/api/scim_users_test.go | 20 ++++++++++++-------- internal/models/scim_token_test.go | 3 +-- 3 files changed, 18 insertions(+), 16 deletions(-) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index 5cd799d404..f5b4e53457 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -1,7 +1,6 @@ package api import ( - "errors" "net/http" "net/http/httptest" "strconv" @@ -72,7 +71,7 @@ func (ts *SCIMUsersTestSuite) TestCreateDoesNotLinkOutsideProvider() { require.NotEqual(ts.T(), password.ID, user.ID) require.NotEqual(ts.T(), other.ID, user.ID) - require.Len(ts.T(), ts.identities(password), 0) + require.Empty(ts.T(), ts.identities(password)) require.Len(ts.T(), ts.identities(other), 1) } @@ -131,7 +130,7 @@ func (ts *SCIMUsersTestSuite) TestPasswordUserWithSameEmailIsNeverLinked() { reloaded, err := models.FindUserByID(ts.API.db, password.ID) require.NoError(ts.T(), err) require.False(ts.T(), reloaded.IsSSOUser) - require.Len(ts.T(), ts.identities(reloaded), 0) + require.Empty(ts.T(), ts.identities(reloaded)) } func (ts *SCIMUsersTestSuite) TestLinkAccountKeepsUserSSO() { @@ -146,7 +145,7 @@ func (ts *SCIMUsersTestSuite) TestLinkAccountKeepsUserSSO() { require.Len(ts.T(), ts.identities(user), 2) require.True(ts.T(), user.IsSSOUser) require.Equal(ts.T(), http.StatusUnprocessableEntity, ts.passkeyRegistrationOptions(user)) - require.Len(ts.T(), ts.identities(password), 0) + require.Empty(ts.T(), ts.identities(password)) } func (ts *SCIMUsersTestSuite) TestOldEmailCannotSignInAfterEmailChange() { @@ -507,7 +506,7 @@ func (ts *SCIMUsersTestSuite) issueSession(conn *storage.Connection, user *model func (ts *SCIMUsersTestSuite) requireBanned(err error) { var httpErr *apierrors.HTTPError - require.True(ts.T(), errors.As(err, &httpErr), err) + require.ErrorAs(ts.T(), err, &httpErr) require.Equal(ts.T(), http.StatusForbidden, httpErr.HTTPStatus) require.Equal(ts.T(), apierrors.ErrorCodeUserBanned, httpErr.ErrorCode) } @@ -758,7 +757,7 @@ func (ts *SCIMUsersTestSuite) TestCreateConcurrentSameEmailLinksToOneUser() { count, err := ts.API.db.Q().Where("email = ?", "race@example.com").Count(&models.User{}) require.NoError(ts.T(), err) - require.EqualValues(ts.T(), 1, count) + require.Equal(ts.T(), 1, count) } func (ts *SCIMUsersTestSuite) rename(id, userName string) (int, string) { diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index a86e964d35..d9ce102a5d 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -119,6 +119,16 @@ func (ts *SCIMUsersTestSuite) do(token, method, path, body string) (*httptest.Re } func (ts *SCIMUsersTestSuite) doAs(contentType, token, method, path, body string, headers ...string) (*httptest.ResponseRecorder, map[string]any) { + w := ts.serve(contentType, token, method, path, body, headers...) + + var decoded map[string]any + if w.Body.Len() > 0 { + require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &decoded), w.Body.String()) + } + return w, decoded +} + +func (ts *SCIMUsersTestSuite) serve(contentType, token, method, path, body string, headers ...string) *httptest.ResponseRecorder { r := httptest.NewRequest(method, "/scim/v2"+path, strings.NewReader(body)) r.Header.Set("Authorization", "Bearer "+token) r.Header.Set("Content-Type", contentType) @@ -127,12 +137,7 @@ func (ts *SCIMUsersTestSuite) doAs(contentType, token, method, path, body string } w := httptest.NewRecorder() ts.API.handler.ServeHTTP(w, r) - - var decoded map[string]any - if w.Body.Len() > 0 { - require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &decoded), w.Body.String()) - } - return w, decoded + return w } func (ts *SCIMUsersTestSuite) create(token, body string) string { @@ -378,8 +383,7 @@ func (ts *SCIMUsersTestSuite) whileLocked(lock, finish func(tx *storage.Connecti code := make(chan int, 1) go func() { - w, _ := ts.do(ts.TokenA, method, path, body) - code <- w.Code + code <- ts.serve(protocol.MediaType, ts.TokenA, method, path, body).Code }() select { case c := <-code: diff --git a/internal/models/scim_token_test.go b/internal/models/scim_token_test.go index 175fcee92e..1dbae59042 100644 --- a/internal/models/scim_token_test.go +++ b/internal/models/scim_token_test.go @@ -1,7 +1,6 @@ package models import ( - "regexp" "testing" "time" @@ -55,7 +54,7 @@ func (ts *SCIMTokenTestSuite) createToken(expiresAt *time.Time) (*SCIMToken, str func (ts *SCIMTokenTestSuite) TestCreate() { token, plaintext := ts.createToken(nil) - require.Regexp(ts.T(), regexp.MustCompile(`^scim_[0-9a-f]{40}$`), plaintext) + require.Regexp(ts.T(), `^scim_[0-9a-f]{40}$`, plaintext) require.Equal(ts.T(), plaintext[:12], token.Prefix) require.Equal(ts.T(), hashSCIMToken(plaintext), token.TokenHash) require.NotEqual(ts.T(), plaintext, token.TokenHash) From bc50decb3f53fbda50614c6d752f0603122dfe62 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:49:09 -0600 Subject: [PATCH 35/88] chore: update scim-go to v0.8.1 --- go.mod | 2 +- go.sum | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/go.mod b/go.mod index 5579bd69ad..1c790ba850 100644 --- a/go.mod +++ b/go.mod @@ -163,7 +163,7 @@ require ( github.com/spf13/cobra v1.8.1 github.com/standard-webhooks/standard-webhooks/libraries v0.0.0-20240303152453-e0e82adf1721 github.com/stretchr/testify v1.12.1 - github.com/supabase-community/scim-go v0.8.0 + github.com/supabase-community/scim-go v0.8.1 github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 github.com/xeipuuv/gojsonschema v1.2.0 go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.64.0 diff --git a/go.sum b/go.sum index b0f8ea7ef3..e78d212301 100644 --- a/go.sum +++ b/go.sum @@ -498,8 +498,8 @@ github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= -github.com/supabase-community/scim-go v0.8.0 h1:eeLUw37qFBerMY86JLfMS6mV1yFVxUj9EFleEGpoXEw= -github.com/supabase-community/scim-go v0.8.0/go.mod h1:oEMij9JuKtAl0wl0jeyIHDKqHvPpUmTXegKeCBKyxXw= +github.com/supabase-community/scim-go v0.8.1 h1:VCgPiHrHYMjd0d/QcVQ0P0TB6W0CER2edArdqlLVz1A= +github.com/supabase-community/scim-go v0.8.1/go.mod h1:oEMij9JuKtAl0wl0jeyIHDKqHvPpUmTXegKeCBKyxXw= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 h1:VDuRtwen5Z7QQ5ctuHUse4wAv/JozkKZkdic5vUV4Lg= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869/go.mod h1:eHX5nlSMSnyPjUrbYzeqrA8snCe2SKyfizKjU3dkfOw= github.com/supranational/blst v0.3.16-0.20250831170142-f48500c1fdbe h1:nbdqkIGOGfUAD54q1s2YBcBz/WcsxCO9HUQ4aGV5hUw= From 1d0c825079fcf9e4ca11b86118f226d9d5f4fe30 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 22:54:24 -0600 Subject: [PATCH 36/88] chore(scim): test that versioned writes to missing SCIM rows return not found --- internal/models/scim_group_test.go | 18 ++++++++++++++++++ 1 file changed, 18 insertions(+) diff --git a/internal/models/scim_group_test.go b/internal/models/scim_group_test.go index 815a9ec82d..d7152e90a0 100644 --- a/internal/models/scim_group_test.go +++ b/internal/models/scim_group_test.go @@ -147,6 +147,24 @@ func (ts *SCIMGroupTestSuite) TestDeleteRemovesMembers() { require.ErrorIs(ts.T(), err, SCIMGroupNotFoundError{}) } +func (ts *SCIMGroupTestSuite) TestVersionedWriteToMissingRowIsNotFound() { + group := ts.createGroup(ts.provider.ID, "Engineering") + alice := ts.createUser(ts.provider.ID, "alice") + + _, err := ReplaceSCIMGroup(ts.db, SCIMTarget{ProviderID: ts.createProvider().ID, ID: group.ID, UpdatedAt: &group.UpdatedAt}, []byte(`{"displayName":"Other"}`)) + require.ErrorIs(ts.T(), err, SCIMGroupNotFoundError{}) + + _, err = DeleteSCIMGroup(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: group.ID}) + require.NoError(ts.T(), err) + _, err = DeleteSCIMGroup(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: group.ID, UpdatedAt: &group.UpdatedAt}) + require.ErrorIs(ts.T(), err, SCIMGroupNotFoundError{}) + + _, err = DeleteSCIMUser(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) + require.NoError(ts.T(), err) + _, err = DeleteSCIMUser(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID, UpdatedAt: &alice.UpdatedAt}) + require.ErrorIs(ts.T(), err, SCIMUserNotFoundError{}) +} + func (ts *SCIMGroupTestSuite) TestReplaceMembersDiffs() { group := ts.createGroup(ts.provider.ID, "Engineering") alice := ts.createUser(ts.provider.ID, "Alice") From 16baa7c79ca3b44c8a4c7ce9b0502cf35ecd37f0 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 23:31:24 -0600 Subject: [PATCH 37/88] chore: bump scim-go to v0.8.2 --- go.mod | 2 +- go.sum | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/go.mod b/go.mod index 1c790ba850..9ac25c3fce 100644 --- a/go.mod +++ b/go.mod @@ -163,7 +163,7 @@ require ( github.com/spf13/cobra v1.8.1 github.com/standard-webhooks/standard-webhooks/libraries v0.0.0-20240303152453-e0e82adf1721 github.com/stretchr/testify v1.12.1 - github.com/supabase-community/scim-go v0.8.1 + github.com/supabase-community/scim-go v0.8.2 github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 github.com/xeipuuv/gojsonschema v1.2.0 go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.64.0 diff --git a/go.sum b/go.sum index e78d212301..d49be69ca5 100644 --- a/go.sum +++ b/go.sum @@ -498,8 +498,8 @@ github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= -github.com/supabase-community/scim-go v0.8.1 h1:VCgPiHrHYMjd0d/QcVQ0P0TB6W0CER2edArdqlLVz1A= -github.com/supabase-community/scim-go v0.8.1/go.mod h1:oEMij9JuKtAl0wl0jeyIHDKqHvPpUmTXegKeCBKyxXw= +github.com/supabase-community/scim-go v0.8.2 h1:MwyFTe88+z9RV5MxYLRBvHr2qhPqEvseTRf413ufo10= +github.com/supabase-community/scim-go v0.8.2/go.mod h1:oEMij9JuKtAl0wl0jeyIHDKqHvPpUmTXegKeCBKyxXw= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 h1:VDuRtwen5Z7QQ5ctuHUse4wAv/JozkKZkdic5vUV4Lg= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869/go.mod h1:eHX5nlSMSnyPjUrbYzeqrA8snCe2SKyfizKjU3dkfOw= github.com/supranational/blst v0.3.16-0.20250831170142-f48500c1fdbe h1:nbdqkIGOGfUAD54q1s2YBcBz/WcsxCO9HUQ4aGV5hUw= From 990ebe725200979246ad48372bb12175c7a307c1 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 23:43:22 -0600 Subject: [PATCH 38/88] chore: bump scim-go to v0.8.3 --- go.mod | 2 +- go.sum | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/go.mod b/go.mod index 9ac25c3fce..ef7832e378 100644 --- a/go.mod +++ b/go.mod @@ -163,7 +163,7 @@ require ( github.com/spf13/cobra v1.8.1 github.com/standard-webhooks/standard-webhooks/libraries v0.0.0-20240303152453-e0e82adf1721 github.com/stretchr/testify v1.12.1 - github.com/supabase-community/scim-go v0.8.2 + github.com/supabase-community/scim-go v0.8.3 github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 github.com/xeipuuv/gojsonschema v1.2.0 go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.64.0 diff --git a/go.sum b/go.sum index d49be69ca5..103ab4cadb 100644 --- a/go.sum +++ b/go.sum @@ -498,8 +498,8 @@ github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= -github.com/supabase-community/scim-go v0.8.2 h1:MwyFTe88+z9RV5MxYLRBvHr2qhPqEvseTRf413ufo10= -github.com/supabase-community/scim-go v0.8.2/go.mod h1:oEMij9JuKtAl0wl0jeyIHDKqHvPpUmTXegKeCBKyxXw= +github.com/supabase-community/scim-go v0.8.3 h1:LZFBpi3IEeRpJbY8Ji6HrcjEiurV79JPGj/1J8NOytg= +github.com/supabase-community/scim-go v0.8.3/go.mod h1:oEMij9JuKtAl0wl0jeyIHDKqHvPpUmTXegKeCBKyxXw= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 h1:VDuRtwen5Z7QQ5ctuHUse4wAv/JozkKZkdic5vUV4Lg= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869/go.mod h1:eHX5nlSMSnyPjUrbYzeqrA8snCe2SKyfizKjU3dkfOw= github.com/supranational/blst v0.3.16-0.20250831170142-f48500c1fdbe h1:nbdqkIGOGfUAD54q1s2YBcBz/WcsxCO9HUQ4aGV5hUw= From f93e6bed671c64a84bd7006eef3ccc71dc418b85 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 23:46:53 -0600 Subject: [PATCH 39/88] chore(scim): encode SCIM group attributes without rendering members --- internal/api/scim_groups.go | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index 35fbeac8ed..a0fb69e56d 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -104,7 +104,9 @@ func (s *scimGroupRepository) save(ctx context.Context, action models.AuditActio if err != nil { return nil, err } - resource, err := scimEncode(group, "id", "meta", "members") + attributes := *group + attributes.Members = nil + resource, err := scimEncode(&attributes, "id", "meta", "members") if err != nil { return nil, err } From 189c0c917151ca9bd734b5f90b8487ef04732ffe Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 23:48:32 -0600 Subject: [PATCH 40/88] chore(scim): validate only added SCIM group members --- internal/models/scim_group.go | 38 ++++++++++++++--------------------- 1 file changed, 15 insertions(+), 23 deletions(-) diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go index 7fcb0e21a4..1bf98bc7bb 100644 --- a/internal/models/scim_group.go +++ b/internal/models/scim_group.go @@ -129,20 +129,16 @@ func FindSCIMGroupsForUsers(tx *storage.Connection, providerID uuid.UUID, scimUs } func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserIDs []uuid.UUID) (added, removed []uuid.UUID, err error) { - wanted, err := findLiveSCIMUserIDs(tx, group.SSOProviderID, scimUserIDs, false) - if err != nil { - return nil, nil, err - } current, err := findSCIMGroupMemberIDs(tx, group.ID) if err != nil { return nil, nil, err } - removed = differenceUUIDs(current, wanted) - added = differenceUUIDs(wanted, current) + removed = differenceUUIDs(current, scimUserIDs) if err := removeSCIMGroupMembers(tx, group.ID, removed); err != nil { return nil, nil, err } - if err := addSCIMGroupMembers(tx, group, added); err != nil { + added, err = addSCIMGroupMembers(tx, group, differenceUUIDs(scimUserIDs, current)) + if err != nil { return nil, nil, err } return added, removed, nil @@ -178,17 +174,13 @@ func RemoveSCIMUserFromGroups(tx *storage.Connection, scimUserID uuid.UUID) ([]u return groupIDs, nil } -func findLiveSCIMUserIDs(tx *storage.Connection, providerID uuid.UUID, ids []uuid.UUID, lock bool) ([]uuid.UUID, error) { +func lockLiveSCIMUserIDs(tx *storage.Connection, providerID uuid.UUID, ids []uuid.UUID) ([]uuid.UUID, error) { found := []uuid.UUID{} - if len(ids) == 0 { - return found, nil - } - query, message := "SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL", "error finding SCIM group members" - if lock { - query, message = query+" ORDER BY id FOR SHARE", "error locking SCIM group members" - } - if err := tx.RawQuery(fmt.Sprintf(query, scimUsersTable.name), ids, providerID).All(&found); err != nil { - return nil, errors.Wrap(err, message) + if err := tx.RawQuery( + fmt.Sprintf("SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL ORDER BY id FOR SHARE", scimUsersTable.name), + ids, providerID, + ).All(&found); err != nil { + return nil, errors.Wrap(err, "error locking SCIM group members") } if missing := differenceUUIDs(ids, found); len(missing) > 0 { return nil, SCIMGroupMemberNotFoundError{IDs: missing} @@ -224,19 +216,19 @@ func removeSCIMGroupMembers(tx *storage.Connection, groupID uuid.UUID, scimUserI return nil } -func addSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserIDs []uuid.UUID) error { +func addSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserIDs []uuid.UUID) ([]uuid.UUID, error) { if len(scimUserIDs) == 0 { - return nil + return []uuid.UUID{}, nil } - locked, err := findLiveSCIMUserIDs(tx, group.SSOProviderID, scimUserIDs, true) + locked, err := lockLiveSCIMUserIDs(tx, group.SSOProviderID, scimUserIDs) if err != nil { - return err + return nil, err } if err := tx.RawQuery( fmt.Sprintf("INSERT INTO %q (group_id, scim_user_id) SELECT ?, unnest(?::uuid[])", SCIMGroupMember{}.TableName()), group.ID, locked, ).Exec(); err != nil { - return errors.Wrap(err, "error adding SCIM group members") + return nil, errors.Wrap(err, "error adding SCIM group members") } - return nil + return locked, nil } From 58c07f90e0d28a2519d84ac6d201d3a48e97fe33 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 23:50:31 -0600 Subject: [PATCH 41/88] chore(scim): stop selecting unused userName for SCIM group members --- internal/models/scim_group.go | 2 +- internal/models/scim_group_test.go | 3 +-- 2 files changed, 2 insertions(+), 3 deletions(-) diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go index 1bf98bc7bb..294b86c358 100644 --- a/internal/models/scim_group.go +++ b/internal/models/scim_group.go @@ -104,7 +104,7 @@ func FindSCIMGroupMembers(tx *storage.Connection, providerID uuid.UUID, groupIDs return members, nil } err := tx.RawQuery( - fmt.Sprintf("SELECT m.group_id, m.scim_user_id, u.resource->>'userName' AS display FROM %q m JOIN %q u ON u.id = m.scim_user_id WHERE m.group_id = ANY(?::uuid[]) AND u.sso_provider_id = ? AND u.deleted_at IS NULL ORDER BY m.group_id, m.created_at, m.scim_user_id", SCIMGroupMember{}.TableName(), scimUsersTable.name), + fmt.Sprintf("SELECT m.group_id, m.scim_user_id FROM %q m JOIN %q u ON u.id = m.scim_user_id WHERE m.group_id = ANY(?::uuid[]) AND u.sso_provider_id = ? AND u.deleted_at IS NULL ORDER BY m.group_id, m.created_at, m.scim_user_id", SCIMGroupMember{}.TableName(), scimUsersTable.name), groupIDs, providerID, ).All(&members) if err != nil { diff --git a/internal/models/scim_group_test.go b/internal/models/scim_group_test.go index d7152e90a0..539b6f1087 100644 --- a/internal/models/scim_group_test.go +++ b/internal/models/scim_group_test.go @@ -184,8 +184,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersDiffs() { members, err := FindSCIMGroupMembers(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) require.NoError(ts.T(), err) require.Len(ts.T(), members, 2) - displays := []string{members[0].Display, members[1].Display} - require.ElementsMatch(ts.T(), []string{"bob", "carol"}, displays) + require.ElementsMatch(ts.T(), []uuid.UUID{bob.ID, carol.ID}, []uuid.UUID{members[0].SCIMUserID, members[1].SCIMUserID}) added, removed, err = ReplaceSCIMGroupMembers(ts.db, group, nil) require.NoError(ts.T(), err) From ddf83b16a4137b7417e166c7f0eb164b752e4489 Mon Sep 17 00:00:00 2001 From: mo khan Date: Wed, 30 Sep 2026 23:53:35 -0600 Subject: [PATCH 42/88] chore(scim): test that SCIM group replace validates only added members --- internal/models/scim_group_test.go | 18 ++++++++++++++++++ 1 file changed, 18 insertions(+) diff --git a/internal/models/scim_group_test.go b/internal/models/scim_group_test.go index 539b6f1087..b6e8feceb2 100644 --- a/internal/models/scim_group_test.go +++ b/internal/models/scim_group_test.go @@ -300,6 +300,24 @@ func (ts *SCIMGroupTestSuite) TestFindMembersHidesDeletedUsers() { require.Empty(ts.T(), members) } +func (ts *SCIMGroupTestSuite) TestReplaceMembersValidatesOnlyAddedMembers() { + group := ts.createGroup(ts.provider.ID, "Engineering") + alice := ts.createUser(ts.provider.ID, "alice") + bob := ts.createUser(ts.provider.ID, "bob") + _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) + require.NoError(ts.T(), err) + _, err = DeleteSCIMUser(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) + require.NoError(ts.T(), err) + + added, removed, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, bob.ID}) + require.NoError(ts.T(), err) + require.Equal(ts.T(), []uuid.UUID{bob.ID}, added) + require.Empty(ts.T(), removed) + + _, _, err = ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, bob.ID, uuid.Nil}) + require.Equal(ts.T(), SCIMGroupMemberNotFoundError{IDs: []uuid.UUID{uuid.Nil}}, err) +} + func (ts *SCIMGroupTestSuite) TestFindGroupsForUsers() { engineering := ts.createGroup(ts.provider.ID, "Engineering") admins := ts.createGroup(ts.provider.ID, "Admins") From 29ac50e0bb976468bce6de961894215be4810425 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:19:29 -0600 Subject: [PATCH 43/88] fix(scim): sync the linked user email when the SCIM primary email changes --- internal/api/scim_link_test.go | 90 +++++++++++++++++++++++++--------- internal/api/scim_users.go | 23 +++++++++ internal/models/scim_user.go | 34 +++++++++++++ 3 files changed, 124 insertions(+), 23 deletions(-) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index f5b4e53457..6480c95ba8 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -150,34 +150,78 @@ func (ts *SCIMUsersTestSuite) TestLinkAccountKeepsUserSSO() { func (ts *SCIMUsersTestSuite) TestOldEmailCannotSignInAfterEmailChange() { id := ts.create(ts.TokenA, oktaUser) - w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, strings.Replace(oktaUser, `"value": "alice@example.com"`, `"value": "alice.smith@example.com"`, 1)) + w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, ts.withEmail("alice.smith@example.com")) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) user := ts.linkedUser(id) - require.Equal(ts.T(), "alice@example.com", user.GetEmail()) - - for _, req := range []struct{ path, body string }{ - {"/recover", `{"email":"alice@example.com"}`}, - {"/otp", `{"email":"alice@example.com","create_user":false}`}, - {"/magiclink", `{"email":"alice@example.com"}`}, - {"/token?grant_type=password", `{"email":"alice@example.com","password":"hunter2hunter2"}`}, - } { - r := httptest.NewRequest(http.MethodPost, req.path, strings.NewReader(req.body)) - r.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - ts.API.handler.ServeHTTP(w, r) - require.NotContains(ts.T(), w.Body.String(), "access_token", req.path) - - reloaded := ts.linkedUser(id) - require.Nil(ts.T(), reloaded.RecoverySentAt, req.path) - require.Empty(ts.T(), reloaded.RecoveryToken, req.path) - require.Empty(ts.T(), reloaded.ConfirmationToken, req.path) - count, err := ts.API.db.Q().Where("user_id = ?", user.ID).Count(&models.OneTimeToken{}) - require.NoError(ts.T(), err) - require.Zero(ts.T(), count, req.path) - require.Zero(ts.T(), ts.sessions(user), req.path) + require.Equal(ts.T(), "alice.smith@example.com", user.GetEmail()) + + for _, email := range []string{"alice@example.com", "alice.smith@example.com"} { + for _, req := range []struct{ path, body string }{ + {"/recover", `{"email":"` + email + `"}`}, + {"/otp", `{"email":"` + email + `","create_user":false}`}, + {"/magiclink", `{"email":"` + email + `"}`}, + {"/token?grant_type=password", `{"email":"` + email + `","password":"hunter2hunter2"}`}, + } { + r := httptest.NewRequest(http.MethodPost, req.path, strings.NewReader(req.body)) + r.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + require.NotContains(ts.T(), w.Body.String(), "access_token", req.path) + + reloaded := ts.linkedUser(id) + require.Nil(ts.T(), reloaded.RecoverySentAt, req.path) + require.Empty(ts.T(), reloaded.RecoveryToken, req.path) + require.Empty(ts.T(), reloaded.ConfirmationToken, req.path) + count, err := ts.API.db.Q().Where("user_id = ?", user.ID).Count(&models.OneTimeToken{}) + require.NoError(ts.T(), err) + require.Zero(ts.T(), count, req.path) + require.Zero(ts.T(), ts.sessions(user), req.path) + } } } +func (ts *SCIMUsersTestSuite) withEmail(email string) string { + return strings.Replace(oktaUser, `"value": "alice@example.com"`, `"value": "`+email+`"`, 1) +} + +func (ts *SCIMUsersTestSuite) TestReplaceChangesEmail() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + + w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, ts.withEmail("Alice.Smith@example.com")) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + + reloaded := ts.linkedUser(id) + require.Equal(ts.T(), "alice.smith@example.com", reloaded.GetEmail()) + require.Equal(ts.T(), "Alice.Smith@example.com", reloaded.UserMetaData["email"]) + identity, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", "sso:"+ts.A.ID.String()) + require.NoError(ts.T(), err) + require.Equal(ts.T(), "Alice.Smith@example.com", identity.IdentityData["email"]) + + signedIn, err := ts.samlLogin(ts.A, "saml-name-id", "alice.smith@example.com") + require.NoError(ts.T(), err) + require.Equal(ts.T(), user.ID, signedIn.ID) +} + +func (ts *SCIMUsersTestSuite) TestReplaceRejectsEmailTakenInProvider() { + id := ts.create(ts.TokenA, oktaUser) + ts.ssoUser(ts.A, "bob", "bob@example.com") + + w, body := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, ts.withEmail("Bob@example.com")) + require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) + require.Equal(ts.T(), "uniqueness", body["scimType"]) + require.Equal(ts.T(), "alice@example.com", ts.linkedUser(id).GetEmail()) +} + +func (ts *SCIMUsersTestSuite) TestReplaceAllowsEmailTakenInAnotherProvider() { + id := ts.create(ts.TokenA, oktaUser) + ts.ssoUser(ts.B, "bob", "bob@example.com") + + w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, ts.withEmail("bob@example.com")) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), "bob@example.com", ts.linkedUser(id).GetEmail()) +} + func (ts *SCIMUsersTestSuite) TestRemovingEmailsKeepsUserEmail() { id := ts.create(ts.TokenA, strings.Replace(oktaUser, `"userName": "Alice@Example.com"`, `"userName": "alice.smith"`, 1)) user := ts.linkedUser(id) diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index 6f46f03c71..f7fe9e5250 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "net/http" + "strings" "github.com/badoux/checkmail" "github.com/gofrs/uuid" @@ -292,6 +293,9 @@ func (s *scimUserRepository) syncAuthUser(tx *storage.Connection, change scimUse if err := s.renameIdentity(tx, change, old); err != nil { return nil, err } + if err := s.changeEmail(tx, change, linked); err != nil { + return nil, err + } if old.Active && !row.Active { return nil, models.LogoutSCIMUser(tx, linked.ID) } @@ -436,6 +440,25 @@ func (s *scimUserRepository) renameIdentity(tx *storage.Connection, change scimU return err } +func (s *scimUserRepository) changeEmail(tx *storage.Connection, change scimUserChange, linked *models.User) error { + email := scimUserEmail(change.user) + if email == "" || strings.EqualFold(email, linked.GetEmail()) { + return nil + } + if err := models.ChangeSCIMIdentityEmail(tx, models.SCIMIdentityEmailChange{ + UserID: linked.ID, + Provider: scimProviderType(change.target.ProviderID), + Subject: change.user.UserName, + Email: email, + }); err != nil { + return err + } + if err := linked.SetEmail(tx, strings.ToLower(email)); err != nil { + return err + } + return linked.UpdateUserMetaData(tx, map[string]any{"email": email}) +} + func (s *scimUserRepository) audit(tx *storage.Connection, r *http.Request, action models.AuditAction, row *models.SCIMUser) error { return s.api.auditSCIM(tx, r, scimActor(r), action, row.SSOProviderID, scimUserTraits(row)) } diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index 9806fb86d0..e6dd766545 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -37,6 +37,13 @@ type SCIMIdentityRename struct { Data map[string]any } +type SCIMIdentityEmailChange struct { + UserID uuid.UUID + Provider string + Subject string + Email string +} + var scimUsersTable = scimTable{ name: SCIMUser{}.TableName(), label: "SCIM user", @@ -255,3 +262,30 @@ func RenameSCIMIdentity(tx *storage.Connection, rename SCIMIdentityRename) error } return nil } + +func ChangeSCIMIdentityEmail(tx *storage.Connection, change SCIMIdentityEmailChange) error { + table := Identity{}.TableName() + taken := struct { + Exists bool `db:"exists"` + }{} + if err := tx.RawQuery( + fmt.Sprintf("SELECT EXISTS(SELECT 1 FROM %q WHERE provider = ? AND email = lower(?) AND user_id <> ?) AS exists", table), + change.Provider, change.Email, change.UserID, + ).First(&taken); err != nil { + return errors.Wrap(err, "error finding SCIM identity email") + } + if taken.Exists { + return SCIMUserConflictError{} + } + encoded, err := json.Marshal(map[string]any{"email": change.Email}) + if err != nil { + return errors.Wrap(err, "error encoding identity data") + } + if err := tx.RawQuery( + fmt.Sprintf("UPDATE %q SET identity_data = identity_data || ?::jsonb, updated_at = now() WHERE user_id = ? AND provider = ? AND provider_id = ?", table), + string(encoded), change.UserID, change.Provider, change.Subject, + ).Exec(); err != nil { + return errors.Wrap(err, "error changing SCIM identity email") + } + return nil +} From 044eb48ff3df1c24d86d1ac71b2e9f64e2c88fdd Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:22:19 -0600 Subject: [PATCH 44/88] chore(scim): test that a SCIM rename with a new email keeps the user's own identity --- internal/api/scim_link_test.go | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index 6480c95ba8..84ccdfda19 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -203,6 +203,25 @@ func (ts *SCIMUsersTestSuite) TestReplaceChangesEmail() { require.Equal(ts.T(), user.ID, signedIn.ID) } +func (ts *SCIMUsersTestSuite) TestReplaceRenamesAndChangesEmail() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + + body := strings.Replace(ts.withEmail("alice.smith@example.com"), `"userName": "Alice@Example.com"`, `"userName": "alice.smith@example.com"`, 1) + w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, body) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + + require.Equal(ts.T(), "alice.smith@example.com", ts.linkedUser(id).GetEmail()) + identities := ts.identities(user) + require.Len(ts.T(), identities, 1) + require.Equal(ts.T(), "alice.smith@example.com", identities[0].ProviderID) + require.Equal(ts.T(), "alice.smith@example.com", identities[0].IdentityData["email"]) + + signedIn, err := ts.samlLogin(ts.A, "saml-name-id", "alice.smith@example.com") + require.NoError(ts.T(), err) + require.Equal(ts.T(), user.ID, signedIn.ID) +} + func (ts *SCIMUsersTestSuite) TestReplaceRejectsEmailTakenInProvider() { id := ts.create(ts.TokenA, oktaUser) ts.ssoUser(ts.A, "bob", "bob@example.com") From 355471d72b8cfd683c39d03721af21fa5c92a8ce Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:29:14 -0600 Subject: [PATCH 45/88] chore(deps): bump scim-go to v0.8.4 --- go.mod | 2 +- go.sum | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/go.mod b/go.mod index ef7832e378..b3e282de34 100644 --- a/go.mod +++ b/go.mod @@ -163,7 +163,7 @@ require ( github.com/spf13/cobra v1.8.1 github.com/standard-webhooks/standard-webhooks/libraries v0.0.0-20240303152453-e0e82adf1721 github.com/stretchr/testify v1.12.1 - github.com/supabase-community/scim-go v0.8.3 + github.com/supabase-community/scim-go v0.8.4 github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 github.com/xeipuuv/gojsonschema v1.2.0 go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.64.0 diff --git a/go.sum b/go.sum index 103ab4cadb..359c873f60 100644 --- a/go.sum +++ b/go.sum @@ -498,8 +498,8 @@ github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= -github.com/supabase-community/scim-go v0.8.3 h1:LZFBpi3IEeRpJbY8Ji6HrcjEiurV79JPGj/1J8NOytg= -github.com/supabase-community/scim-go v0.8.3/go.mod h1:oEMij9JuKtAl0wl0jeyIHDKqHvPpUmTXegKeCBKyxXw= +github.com/supabase-community/scim-go v0.8.4 h1:9FBo8r9c2u3rL//ZXa8YbTQod5MS4/2l6UjQjoBj9JI= +github.com/supabase-community/scim-go v0.8.4/go.mod h1:oEMij9JuKtAl0wl0jeyIHDKqHvPpUmTXegKeCBKyxXw= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 h1:VDuRtwen5Z7QQ5ctuHUse4wAv/JozkKZkdic5vUV4Lg= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869/go.mod h1:eHX5nlSMSnyPjUrbYzeqrA8snCe2SKyfizKjU3dkfOw= github.com/supranational/blst v0.3.16-0.20250831170142-f48500c1fdbe h1:nbdqkIGOGfUAD54q1s2YBcBz/WcsxCO9HUQ4aGV5hUw= From e64538c221184b860b42a99be46bde75b334d67d Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:38:23 -0600 Subject: [PATCH 46/88] fix(scim): refuse to link a SCIM user to a non-SSO account --- internal/api/scim_link_test.go | 26 ++++++++++++++++++++++++++ internal/api/scim_users.go | 13 +++++++++++++ 2 files changed, 39 insertions(+) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index 84ccdfda19..b98ea38e8b 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -133,6 +133,32 @@ func (ts *SCIMUsersTestSuite) TestPasswordUserWithSameEmailIsNeverLinked() { require.Empty(ts.T(), ts.identities(reloaded)) } +func (ts *SCIMUsersTestSuite) TestNonSSOUserWithSSOIdentityEmailIsNeverLinked() { + ts.requireNonSSOUserNeverLinked("saml-name-id") +} + +func (ts *SCIMUsersTestSuite) TestNonSSOUserWithSSOIdentitySubjectIsNeverLinked() { + ts.requireNonSSOUserNeverLinked("Alice@Example.com") +} + +func (ts *SCIMUsersTestSuite) requireNonSSOUserNeverLinked(sub string) { + password, err := models.NewUser("", "alice@example.com", "", ts.API.config.JWT.Aud, nil) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Create(password)) + identity, err := models.NewIdentity(password, "sso:"+ts.A.ID.String(), map[string]any{"sub": sub, "email": "alice@example.com", "email_verified": true}) + require.NoError(ts.T(), err) + require.NoError(ts.T(), ts.API.db.Create(identity)) + + w, body := ts.do(ts.TokenA, http.MethodPost, "/Users", oktaUser) + + require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) + require.Equal(ts.T(), "uniqueness", body["scimType"]) + reloaded := ts.reloadUser(password.ID) + require.False(ts.T(), reloaded.IsSSOUser) + require.Len(ts.T(), ts.identities(reloaded), 1) + require.Zero(ts.T(), ts.countRows(&models.SCIMUser{}, "user_id = ?", password.ID)) +} + func (ts *SCIMUsersTestSuite) TestLinkAccountKeepsUserSSO() { password, err := models.NewUser("", "alice@example.com", "", ts.API.config.JWT.Aud, nil) require.NoError(ts.T(), err) diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index f7fe9e5250..47fa99823c 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -327,7 +327,13 @@ func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SC linked := decision.User switch decision.Decision { case models.AccountExists: + if err = scimCanLink(linked); err != nil { + return nil, false, err + } case models.LinkAccount: + if err = scimCanLink(linked); err != nil { + return nil, false, err + } if _, err = s.api.createNewIdentity(tx, linked, providerType, scimIdentityData(user)); err != nil { return nil, false, err } @@ -348,6 +354,13 @@ func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SC return linked, false, models.LinkSCIMUser(tx, row, linked.ID) } +func scimCanLink(linked *models.User) error { + if !linked.IsSSOUser { + return scimerrors.ErrUniqueness("user is not an SSO user") + } + return nil +} + func (s *scimUserRepository) createAuthUser(tx *storage.Connection, providerType string, decision models.AccountLinkingResult, user *core.User) (*models.User, error) { candidate, err := s.newUser(providerType, decision, user) if err != nil { From 5b79f40b8ab274231e3b8f3c2fc0f87b9737174a Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:40:17 -0600 Subject: [PATCH 47/88] fix(scim): refuse to relink a SCIM user to an account the provider deleted --- internal/api/scim_link_test.go | 42 +++++++++++++++++------ internal/api/scim_provider_delete_test.go | 3 +- internal/api/scim_users.go | 13 +++++-- internal/api/scim_users_test.go | 5 +-- internal/models/scim_user.go | 8 +++++ 5 files changed, 54 insertions(+), 17 deletions(-) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index b98ea38e8b..8544660b75 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -484,10 +484,10 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedAfterDelete() { _, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") require.Error(ts.T(), err) - relinked := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) - user, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") - require.NoError(ts.T(), err) - require.Equal(ts.T(), relinked.ID, user.ID) + w, _ = ts.do(ts.TokenA, http.MethodPost, "/Users", oktaUser) + require.Equal(ts.T(), http.StatusConflict, w.Code) + _, err = ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") + require.Error(ts.T(), err) } func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedWhenCreatedInactive() { @@ -616,9 +616,9 @@ func (ts *SCIMUsersTestSuite) TestSessionRefusedWhileDeprovisioned() { require.Equal(ts.T(), http.StatusNoContent, w.Code) ts.requireBanned(ts.issueSession(ts.API.db, user)) - relinked := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) - require.Equal(ts.T(), user.ID, relinked.ID) - require.NoError(ts.T(), ts.issueSession(ts.API.db, relinked)) + w, _ = ts.do(ts.TokenA, http.MethodPost, "/Users", oktaUser) + require.Equal(ts.T(), http.StatusConflict, w.Code) + ts.requireBanned(ts.issueSession(ts.API.db, user)) } func (ts *SCIMUsersTestSuite) TestSessionAllowedWhileSCIMFlagOff() { @@ -794,16 +794,36 @@ func (ts *SCIMUsersTestSuite) TestDeleteLogsOutWithoutBanning() { require.Equal(ts.T(), http.StatusBadRequest, ts.refresh(refreshToken)) } -func (ts *SCIMUsersTestSuite) TestCreateDoesNotBanAfterDelete() { +func (ts *SCIMUsersTestSuite) TestCreateRefusesUserDeletedByProvider() { + for _, body := range []string{oktaUser, strings.Replace(oktaUser, `"userName": "Alice@Example.com"`, `"userName": "alice.new@example.com"`, 1)} { + ts.SetupTest() + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + + w, got := ts.do(ts.TokenA, http.MethodPost, "/Users", body) + + require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) + require.Equal(ts.T(), "uniqueness", got["scimType"]) + require.False(ts.T(), ts.reloadUser(user.ID).IsBanned()) + require.Zero(ts.T(), ts.countRows(&models.SCIMUser{}, "user_id = ? AND deleted_at IS NULL", user.ID)) + require.Len(ts.T(), ts.identities(user), 1) + require.Equal(ts.T(), 1, ts.users("alice@example.com")) + } +} + +func (ts *SCIMUsersTestSuite) TestCreateAfterAdminDeletesProviderDeletedUser() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") require.Equal(ts.T(), http.StatusNoContent, w.Code) + require.NoError(ts.T(), ts.API.db.Destroy(user)) - relinked := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + created := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) - require.Equal(ts.T(), user.ID, relinked.ID) - require.False(ts.T(), relinked.IsBanned()) + require.NotEqual(ts.T(), user.ID, created.ID) + require.True(ts.T(), created.IsSSOUser) } func (ts *SCIMUsersTestSuite) TestCreateConcurrentSameEmailLinksToOneUser() { diff --git a/internal/api/scim_provider_delete_test.go b/internal/api/scim_provider_delete_test.go index 52263779fc..c428b05d37 100644 --- a/internal/api/scim_provider_delete_test.go +++ b/internal/api/scim_provider_delete_test.go @@ -57,7 +57,8 @@ func (ts *SCIMUsersTestSuite) TestProviderDeleteBansDeprovisionedUsers() { recreated := ts.linkedUser(recreatedID) w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Users/"+recreatedID, "") require.Equal(ts.T(), http.StatusNoContent, w.Code) - require.Equal(ts.T(), recreated.ID, ts.linkedUser(ts.create(ts.TokenA, scimUser("recreated"))).ID) + relinkedID := ts.create(ts.TokenA, scimUser("relinked")) + require.NoError(ts.T(), ts.API.db.RawQuery("UPDATE scim_users SET user_id = ? WHERE id = ?", recreated.ID, relinkedID).Exec()) ts.createGroup(ts.TokenA, groupWith("A", "", deactivatedID)) diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index 47fa99823c..98738ad9e8 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -327,11 +327,11 @@ func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SC linked := decision.User switch decision.Decision { case models.AccountExists: - if err = scimCanLink(linked); err != nil { + if err = scimCanLink(tx, row.SSOProviderID, linked); err != nil { return nil, false, err } case models.LinkAccount: - if err = scimCanLink(linked); err != nil { + if err = scimCanLink(tx, row.SSOProviderID, linked); err != nil { return nil, false, err } if _, err = s.api.createNewIdentity(tx, linked, providerType, scimIdentityData(user)); err != nil { @@ -354,10 +354,17 @@ func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SC return linked, false, models.LinkSCIMUser(tx, row, linked.ID) } -func scimCanLink(linked *models.User) error { +func scimCanLink(tx *storage.Connection, providerID uuid.UUID, linked *models.User) error { if !linked.IsSSOUser { return scimerrors.ErrUniqueness("user is not an SSO user") } + deleted, err := models.IsSCIMDeleted(tx, providerID, linked.ID) + if err != nil { + return err + } + if deleted { + return scimerrors.ErrUniqueness("user was deleted by this provider") + } return nil } diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index d9ce102a5d..d582b6f0a3 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -229,8 +229,9 @@ func (ts *SCIMUsersTestSuite) TestOktaLifecycle() { require.EqualValues(ts.T(), 0, ts.list(ts.TokenA, filter)["totalResults"], filter) } - reprovisioned := ts.create(ts.TokenA, oktaUser) - require.NotEqual(ts.T(), id, reprovisioned) + w, body := ts.do(ts.TokenA, http.MethodPost, "/Users", oktaUser) + require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) + require.Equal(ts.T(), "uniqueness", body["scimType"]) require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&stored)) require.NotNil(ts.T(), stored.DeletedAt) } diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index e6dd766545..4678912df7 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -189,6 +189,14 @@ func IsSCIMManaged(tx *storage.Connection, providerID, userID uuid.UUID) (bool, return managed, nil } +func IsSCIMDeleted(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { + deleted, err := tx.Q().Where("sso_provider_id = ? AND user_id = ? AND deleted_at IS NOT NULL", providerID, userID).Exists(&SCIMUser{}) + if err != nil { + return false, errors.Wrap(err, "error finding deleted SCIM user") + } + return deleted, nil +} + func IsSCIMDeprovisioned(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { result := struct { AnyRow bool `db:"any_row"` From 7b0dcfc50d0251ea6dadb0558f5f3d15d7b01696 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:42:25 -0600 Subject: [PATCH 48/88] fix(scim): revoke SCIM tokens when SCIM is disabled --- internal/api/scim_admin.go | 62 +++++++++++++++++++++------------ internal/api/scim_admin_test.go | 52 ++++++++++++++++++++------- 2 files changed, 80 insertions(+), 34 deletions(-) diff --git a/internal/api/scim_admin.go b/internal/api/scim_admin.go index 986e1f21da..69cf738f77 100644 --- a/internal/api/scim_admin.go +++ b/internal/api/scim_admin.go @@ -44,11 +44,44 @@ func (a *API) adminSCIMGet(w http.ResponseWriter, r *http.Request) error { } func (a *API) adminSCIMEnable(w http.ResponseWriter, r *http.Request) error { - return a.changeSCIMEnabled(w, r, models.EnableSCIM, models.SCIMEnabledAction, "enabling") + ctx := r.Context() + db := a.db.WithContext(ctx) + provider := getSSOProvider(ctx) + + if err := db.Transaction(func(tx *storage.Connection) error { + changed, err := models.EnableSCIM(tx, provider.ID) + if err != nil || !changed { + return err + } + return a.auditSCIM(tx, r, getAdminUser(ctx), models.SCIMEnabledAction, provider.ID, map[string]any{}) + }); err != nil { + return apierrors.NewInternalServerError("Error enabling SCIM").WithInternalError(err) + } + + return a.sendSCIMStatus(w, db, provider) } func (a *API) adminSCIMDisable(w http.ResponseWriter, r *http.Request) error { - return a.changeSCIMEnabled(w, r, models.DisableSCIM, models.SCIMDisabledAction, "disabling") + ctx := r.Context() + db := a.db.WithContext(ctx) + provider := getSSOProvider(ctx) + actor := getAdminUser(ctx) + + if err := db.Transaction(func(tx *storage.Connection) error { + changed, err := models.DisableSCIM(tx, provider.ID) + if err != nil { + return err + } + prefixes, err := a.revokeSCIMTokens(tx, r, actor, provider.ID) + if err != nil || !changed { + return err + } + return a.auditSCIM(tx, r, actor, models.SCIMDisabledAction, provider.ID, map[string]any{"token_prefixes": prefixes}) + }); err != nil { + return apierrors.NewInternalServerError("Error disabling SCIM").WithInternalError(err) + } + + return a.sendSCIMStatus(w, db, provider) } func (a *API) adminSCIMTokensCreate(w http.ResponseWriter, r *http.Request) error { @@ -150,24 +183,6 @@ func (a *API) revokeSCIMToken(tx *storage.Connection, r *http.Request, providerI return token, a.auditSCIM(tx, r, getAdminUser(r.Context()), models.SCIMTokenRevokedAction, providerID, map[string]any{scimTokenPrefixTrait: token.Prefix}) } -func (a *API) changeSCIMEnabled(w http.ResponseWriter, r *http.Request, change func(*storage.Connection, uuid.UUID) (bool, error), action models.AuditAction, verb string) error { - ctx := r.Context() - db := a.db.WithContext(ctx) - provider := getSSOProvider(ctx) - - if err := db.Transaction(func(tx *storage.Connection) error { - changed, err := change(tx, provider.ID) - if err != nil || !changed { - return err - } - return a.auditSCIM(tx, r, getAdminUser(ctx), action, provider.ID, map[string]any{}) - }); err != nil { - return apierrors.NewInternalServerError("Error %s SCIM", verb).WithInternalError(err) - } - - return a.sendSCIMStatus(w, db, provider) -} - func (a *API) sendSCIMStatus(w http.ResponseWriter, db *storage.Connection, provider *models.SSOProvider) error { tokens, err := models.FindActiveSCIMTokensBySSOProvider(db, provider.ID) if err != nil { @@ -198,7 +213,7 @@ func (a *API) deprovisionSCIM(tx *storage.Connection, r *http.Request, provider return err } actor := getAdminUser(r.Context()) - prefixes, err := a.auditSCIMTokensRevoked(tx, r, actor, provider.ID) + prefixes, err := a.revokeSCIMTokens(tx, r, actor, provider.ID) if err != nil { return err } @@ -214,7 +229,7 @@ func (a *API) deprovisionSCIM(tx *storage.Connection, r *http.Request, provider return a.auditSCIM(tx, r, actor, models.SCIMUsersBannedAction, provider.ID, map[string]any{"banned_user_count": banned}) } -func (a *API) auditSCIMTokensRevoked(tx *storage.Connection, r *http.Request, actor *models.User, providerID uuid.UUID) ([]string, error) { +func (a *API) revokeSCIMTokens(tx *storage.Connection, r *http.Request, actor *models.User, providerID uuid.UUID) ([]string, error) { if err := models.LockSCIMTokens(tx, providerID); err != nil { return nil, err } @@ -225,6 +240,9 @@ func (a *API) auditSCIMTokensRevoked(tx *storage.Connection, r *http.Request, ac prefixes := make([]string, len(tokens)) for i := range tokens { prefixes[i] = tokens[i].Prefix + if err := tokens[i].Revoke(tx); err != nil { + return nil, err + } if err := a.auditSCIM(tx, r, actor, models.SCIMTokenRevokedAction, providerID, map[string]any{scimTokenPrefixTrait: tokens[i].Prefix}); err != nil { return nil, err } diff --git a/internal/api/scim_admin_test.go b/internal/api/scim_admin_test.go index ec8dba0433..d6ca29e36f 100644 --- a/internal/api/scim_admin_test.go +++ b/internal/api/scim_admin_test.go @@ -341,20 +341,37 @@ func (ts *SCIMTokensTestSuite) TestEnableWithZeroTokens() { require.True(ts.T(), ts.status(http.MethodGet, provider).Enabled) } -func (ts *SCIMTokensTestSuite) TestEnableAndDisableLeaveTokensUnchanged() { - ts.create(ts.Provider, map[string]any{}) +func (ts *SCIMTokensTestSuite) TestDisableRevokesActiveTokens() { + active := ts.create(ts.Provider, map[string]any{}) expiring := ts.create(ts.Provider, map[string]any{"expires_at": time.Now().Add(time.Hour)}) - ts.revoke(ts.create(ts.Provider, map[string]any{}).Prefix) - before, err := models.FindSCIMTokensBySSOProvider(ts.API.db, ts.Provider.ID) + revoked := ts.create(ts.Provider, map[string]any{}) + ts.revoke(revoked.Prefix) + otherProvider := createSSOProvider(ts.T(), ts.API.db) + ts.status(http.MethodPost, otherProvider) + other := ts.create(otherProvider, map[string]any{}) + before, err := models.FindSCIMTokenByPrefix(ts.API.db, ts.Provider.ID, revoked.Prefix) require.NoError(ts.T(), err) - require.NotNil(ts.T(), expiring.ExpiresAt) - for _, method := range []string{http.MethodDelete, http.MethodDelete, http.MethodPost, http.MethodPost} { + ts.status(http.MethodDelete, ts.Provider) + disabled, err := models.FindSCIMTokensBySSOProvider(ts.API.db, ts.Provider.ID) + require.NoError(ts.T(), err) + + for _, method := range []string{http.MethodDelete, http.MethodPost, http.MethodPost} { ts.status(method, ts.Provider) after, err := models.FindSCIMTokensBySSOProvider(ts.API.db, ts.Provider.ID) require.NoError(ts.T(), err) - require.Equal(ts.T(), before, after, method) + require.Equal(ts.T(), disabled, after, method) + } + require.Len(ts.T(), disabled, 3) + for _, token := range disabled { + require.True(ts.T(), token.IsRevoked(), token.Prefix) } + after, err := models.FindSCIMTokenByPrefix(ts.API.db, ts.Provider.ID, revoked.Prefix) + require.NoError(ts.T(), err) + require.Equal(ts.T(), before.RevokedAt, after.RevokedAt) + require.Equal(ts.T(), http.StatusUnauthorized, ts.scimRequest(active.Token).Code) + require.Equal(ts.T(), http.StatusUnauthorized, ts.scimRequest(expiring.Token).Code) + require.Equal(ts.T(), http.StatusOK, ts.scimRequest(other.Token).Code) } func (ts *SCIMTokensTestSuite) TestDisableStopsAuthenticatedRequests() { @@ -362,7 +379,7 @@ func (ts *SCIMTokensTestSuite) TestDisableStopsAuthenticatedRequests() { status := ts.status(http.MethodDelete, ts.Provider) require.False(ts.T(), status.Enabled) - require.Len(ts.T(), status.Tokens, 1) + require.Empty(ts.T(), status.Tokens) for _, path := range []string{"/scim/v2/Users", "/scim/v2/Groups", "/scim/v2/Schemas", "/scim/v2/ResourceTypes", "/scim/v2/ServiceProviderConfig"} { r := httptest.NewRequest(http.MethodGet, path, nil) @@ -378,7 +395,7 @@ func (ts *SCIMTokensTestSuite) TestDisableStopsAuthenticatedRequests() { } } -func (ts *SCIMTokensTestSuite) TestReenableRestoresExistingTokens() { +func (ts *SCIMTokensTestSuite) TestReenableKeepsOldTokensRevoked() { token := ts.create(ts.Provider, map[string]any{}) require.Equal(ts.T(), http.StatusOK, ts.scimRequest(token.Token).Code) @@ -387,7 +404,18 @@ func (ts *SCIMTokensTestSuite) TestReenableRestoresExistingTokens() { status := ts.status(http.MethodPost, ts.Provider) require.True(ts.T(), status.Enabled) - require.Equal(ts.T(), http.StatusOK, ts.scimRequest(token.Token).Code) + require.Empty(ts.T(), status.Tokens) + require.Equal(ts.T(), http.StatusUnauthorized, ts.scimRequest(token.Token).Code) + replacement := ts.create(ts.Provider, map[string]any{}) + require.Equal(ts.T(), http.StatusOK, ts.scimRequest(replacement.Token).Code) + + require.Equal(ts.T(), []scimTokenEvent{ + {string(models.SCIMTokenCreatedAction), token.Prefix}, + {string(models.SCIMTokenRevokedAction), token.Prefix}, + {string(models.SCIMDisabledAction), token.Prefix}, + {string(models.SCIMEnabledAction), ""}, + {string(models.SCIMTokenCreatedAction), replacement.Prefix}, + }, ts.tokenEvents()) } func (ts *SCIMTokensTestSuite) TestMintAndRevokeWhileDisabled() { @@ -422,7 +450,7 @@ func (ts *SCIMTokensTestSuite) TestStatusIndependentOfTokens() { ts.status(http.MethodDelete, ts.Provider) status := ts.status(http.MethodGet, ts.Provider) require.False(ts.T(), status.Enabled) - require.Len(ts.T(), status.Tokens, 1) + require.Empty(ts.T(), status.Tokens) } func (ts *SCIMTokensTestSuite) TestConcurrentEnableAndDisable() { @@ -514,7 +542,7 @@ func (ts *SCIMTokensTestSuite) tokenEvents() []scimTokenEvent { require.Equal(ts.T(), ts.Provider.ID.String(), traits["sso_provider_id"]) require.Equal(ts.T(), "success", traits["outcome"]) prefix, _ := traits["token_prefix"].(string) - if prefixes, ok := traits["token_prefixes"].([]any); ok { + if prefixes, ok := traits["token_prefixes"].([]any); ok && len(prefixes) > 0 { require.Len(ts.T(), prefixes, 1) prefix = prefixes[0].(string) } From 8e0a2169fd153504ae7d1dc7d77de4928ea418e3 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:46:39 -0600 Subject: [PATCH 49/88] fix(scim): clear pending one-time tokens when the SCIM primary email changes --- internal/api/scim_link_test.go | 14 ++++++++++++++ internal/api/scim_users.go | 3 +++ 2 files changed, 17 insertions(+) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index 8544660b75..1b99225869 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -229,6 +229,20 @@ func (ts *SCIMUsersTestSuite) TestReplaceChangesEmail() { require.Equal(ts.T(), user.ID, signedIn.ID) } +func (ts *SCIMUsersTestSuite) TestReplaceChangingEmailClearsPendingTokens() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + user.RecoveryToken = "recovery-token-hash" + require.NoError(ts.T(), ts.API.db.UpdateOnly(user, "recovery_token")) + require.NoError(ts.T(), models.CreateOneTimeToken(ts.API.db, user.ID, "alice@example.com", "recovery-token-hash", models.RecoveryToken, time.Hour, true)) + + w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, ts.withEmail("alice.smith@example.com")) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + + require.Empty(ts.T(), ts.reloadUser(user.ID).RecoveryToken) + require.Zero(ts.T(), ts.countRows(&models.OneTimeToken{}, "user_id = ?", user.ID)) +} + func (ts *SCIMUsersTestSuite) TestReplaceRenamesAndChangesEmail() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index 98738ad9e8..fe439487c7 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -476,6 +476,9 @@ func (s *scimUserRepository) changeEmail(tx *storage.Connection, change scimUser if err := linked.SetEmail(tx, strings.ToLower(email)); err != nil { return err } + if err := linked.ClearAllPendingTokens(tx); err != nil { + return err + } return linked.UpdateUserMetaData(tx, map[string]any{"email": email}) } From 09df1c0cdbec321db25ce755bc7e0252b0c661ed Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:47:38 -0600 Subject: [PATCH 50/88] fix(scim): check for a provider-deleted user under the SCIM user lock --- internal/api/scim.go | 2 ++ internal/api/scim_users.go | 13 +++---------- internal/models/errors.go | 6 ++++++ internal/models/scim_user.go | 15 +++++++-------- 4 files changed, 18 insertions(+), 18 deletions(-) diff --git a/internal/api/scim.go b/internal/api/scim.go index fd47862584..57c35d3f88 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -304,6 +304,8 @@ func scimError(err error) error { return scimerrors.ErrUniqueness(`"userName" and "externalId" must be unique`) case errors.Is(err, models.SCIMUserLinkedError{}): return scimerrors.ErrUniqueness("user is already provisioned by this provider") + case errors.Is(err, models.SCIMUserDeletedError{}): + return scimerrors.ErrUniqueness("user was deleted by this provider") } return err } diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index fe439487c7..f3c40263b2 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -327,11 +327,11 @@ func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SC linked := decision.User switch decision.Decision { case models.AccountExists: - if err = scimCanLink(tx, row.SSOProviderID, linked); err != nil { + if err = scimCanLink(linked); err != nil { return nil, false, err } case models.LinkAccount: - if err = scimCanLink(tx, row.SSOProviderID, linked); err != nil { + if err = scimCanLink(linked); err != nil { return nil, false, err } if _, err = s.api.createNewIdentity(tx, linked, providerType, scimIdentityData(user)); err != nil { @@ -354,17 +354,10 @@ func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SC return linked, false, models.LinkSCIMUser(tx, row, linked.ID) } -func scimCanLink(tx *storage.Connection, providerID uuid.UUID, linked *models.User) error { +func scimCanLink(linked *models.User) error { if !linked.IsSSOUser { return scimerrors.ErrUniqueness("user is not an SSO user") } - deleted, err := models.IsSCIMDeleted(tx, providerID, linked.ID) - if err != nil { - return err - } - if deleted { - return scimerrors.ErrUniqueness("user was deleted by this provider") - } return nil } diff --git a/internal/models/errors.go b/internal/models/errors.go index 04176e11a5..67716c51b7 100644 --- a/internal/models/errors.go +++ b/internal/models/errors.go @@ -260,6 +260,12 @@ func (e SCIMUserLinkedError) Error() string { return "user is already linked to a SCIM user in this provider" } +type SCIMUserDeletedError struct{} + +func (e SCIMUserDeletedError) Error() string { + return "user was deleted by this provider" +} + type SCIMGroupNotFoundError struct{} func (e SCIMGroupNotFoundError) Error() string { diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index 4678912df7..6b7fa53f5f 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -170,6 +170,13 @@ func LinkSCIMUser(tx *storage.Connection, user *SCIMUser, userID uuid.UUID) erro if linked { return SCIMUserLinkedError{} } + deleted, err := tx.Q().Where("sso_provider_id = ? AND user_id = ? AND deleted_at IS NOT NULL", user.SSOProviderID, userID).Exists(&SCIMUser{}) + if err != nil { + return errors.Wrap(err, "error finding deleted SCIM user") + } + if deleted { + return SCIMUserDeletedError{} + } if err := tx.RawQuery( fmt.Sprintf("UPDATE %q SET user_id = ? WHERE id = ?", scimUsersTable.name), @@ -189,14 +196,6 @@ func IsSCIMManaged(tx *storage.Connection, providerID, userID uuid.UUID) (bool, return managed, nil } -func IsSCIMDeleted(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { - deleted, err := tx.Q().Where("sso_provider_id = ? AND user_id = ? AND deleted_at IS NOT NULL", providerID, userID).Exists(&SCIMUser{}) - if err != nil { - return false, errors.Wrap(err, "error finding deleted SCIM user") - } - return deleted, nil -} - func IsSCIMDeprovisioned(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { result := struct { AnyRow bool `db:"any_row"` From 58faafa1d273e127594744b0950ba7a977f4bc5e Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:48:30 -0600 Subject: [PATCH 51/88] chore(scim): test that admin user delete lets SCIM create a fresh account --- internal/api/scim_link_test.go | 24 ++++++++++++++++-------- 1 file changed, 16 insertions(+), 8 deletions(-) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index 1b99225869..6c37c1db66 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -828,16 +828,24 @@ func (ts *SCIMUsersTestSuite) TestCreateRefusesUserDeletedByProvider() { } func (ts *SCIMUsersTestSuite) TestCreateAfterAdminDeletesProviderDeletedUser() { - id := ts.create(ts.TokenA, oktaUser) - user := ts.linkedUser(id) - w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") - require.Equal(ts.T(), http.StatusNoContent, w.Code) - require.NoError(ts.T(), ts.API.db.Destroy(user)) + for _, soft := range []bool{false, true} { + ts.SetupTest() + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + r := httptest.NewRequest(http.MethodDelete, "/admin/users/"+user.ID.String(), strings.NewReader(`{"should_soft_delete":`+strconv.FormatBool(soft)+`}`)) + r.Header.Set("Authorization", "Bearer "+adminJWT(ts.T(), ts.API.config.JWT.Secret)) + r.Header.Set("Content-Type", "application/json") + w = httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - created := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + created := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) - require.NotEqual(ts.T(), user.ID, created.ID) - require.True(ts.T(), created.IsSSOUser) + require.NotEqual(ts.T(), user.ID, created.ID, soft) + require.True(ts.T(), created.IsSSOUser, soft) + } } func (ts *SCIMUsersTestSuite) TestCreateConcurrentSameEmailLinksToOneUser() { From d697dc7f259539c0b727495ec3a7d6df5445e6fc Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:55:21 -0600 Subject: [PATCH 52/88] chore(scim): check live and deleted SCIM links in one query --- internal/models/scim_user.go | 22 ++++++++++++++-------- 1 file changed, 14 insertions(+), 8 deletions(-) diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index 6b7fa53f5f..7234db84f9 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -163,18 +163,24 @@ func LinkSCIMUser(tx *storage.Connection, user *SCIMUser, userID uuid.UUID) erro return err } - linked, err := tx.Q().Where("sso_provider_id = ? AND user_id = ? AND deleted_at IS NULL", user.SSOProviderID, userID).Exists(&SCIMUser{}) - if err != nil { + existing := struct { + Live bool `db:"live"` + Deleted bool `db:"deleted"` + }{} + if err := tx.RawQuery( + fmt.Sprintf( + "SELECT EXISTS(SELECT 1 FROM %[1]q WHERE sso_provider_id = ? AND user_id = ? AND deleted_at IS NULL) AS live, "+ + "EXISTS(SELECT 1 FROM %[1]q WHERE sso_provider_id = ? AND user_id = ? AND deleted_at IS NOT NULL) AS deleted", + scimUsersTable.name, + ), + user.SSOProviderID, userID, user.SSOProviderID, userID, + ).First(&existing); err != nil { return errors.Wrap(err, "error finding linked SCIM user") } - if linked { + if existing.Live { return SCIMUserLinkedError{} } - deleted, err := tx.Q().Where("sso_provider_id = ? AND user_id = ? AND deleted_at IS NOT NULL", user.SSOProviderID, userID).Exists(&SCIMUser{}) - if err != nil { - return errors.Wrap(err, "error finding deleted SCIM user") - } - if deleted { + if existing.Deleted { return SCIMUserDeletedError{} } From 1339c3b15c4cbdbbe6262f2ecf456729a24f4a37 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:55:35 -0600 Subject: [PATCH 53/88] chore(scim): drop comment from concurrent SCIM create test --- internal/api/scim_link_test.go | 3 --- 1 file changed, 3 deletions(-) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index 6c37c1db66..1a53b0e293 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -873,9 +873,6 @@ func (ts *SCIMUsersTestSuite) TestCreateConcurrentSameEmailLinksToOneUser() { close(start) wg.Wait() - // The lock serializes the two creates: whichever commits first creates the - // user, the other observes that account under the same provider and is - // rejected as already linked -- never silently creating a second user. created := 0 for _, code := range codes { if code == http.StatusCreated { From 9a93754e5386f1bbd3c0b3c9b29971d6ec50eca0 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:58:05 -0600 Subject: [PATCH 54/88] chore(scim): split SCIM rate limiting, errors, linking and cleanup into their own files --- internal/api/scim.go | 95 ----------------- internal/api/scim_errors.go | 62 +++++++++++ internal/api/scim_ratelimit.go | 61 +++++++++++ internal/api/scim_user_cleanup.go | 38 +++++++ internal/api/scim_user_linking.go | 142 +++++++++++++++++++++++++ internal/api/scim_users.go | 168 ------------------------------ 6 files changed, 303 insertions(+), 263 deletions(-) create mode 100644 internal/api/scim_errors.go create mode 100644 internal/api/scim_ratelimit.go create mode 100644 internal/api/scim_user_cleanup.go create mode 100644 internal/api/scim_user_linking.go diff --git a/internal/api/scim.go b/internal/api/scim.go index 57c35d3f88..0de124edc0 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -6,13 +6,10 @@ import ( "errors" "fmt" "net/http" - "path" "strconv" "strings" "time" - "github.com/didip/tollbooth/v5" - "github.com/didip/tollbooth/v5/limiter" "github.com/gofrs/uuid" "github.com/supabase-community/scim-go/pkg/core" "github.com/supabase-community/scim-go/pkg/protocol" @@ -122,50 +119,6 @@ func (a *API) withSCIMRequest(w http.ResponseWriter, req *http.Request) (context return scimRequestKey.WithValue(req.Context(), req), nil } -func (a *API) limitSCIMByIP(lmt *limiter.Limiter) func(http.Handler) http.Handler { - return func(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if a.scimSkipsTokenValidator(r) && a.performRateLimiting(lmt, r) != nil { - handler(scimTooManyRequests)(w, r) - return - } - next.ServeHTTP(w, r) - }) - } -} - -func (a *API) scimSkipsTokenValidator(r *http.Request) bool { - if _, err := a.extractBearerToken(r); err != nil { - return true - } - return r.URL.Path == scimBasePath+"/ServiceProviderConfig" || path.Clean(r.URL.Path) != r.URL.Path -} - -func (a *API) limitSCIMInvalidToken(validate server.TokenValidator, lmt *limiter.Limiter) server.TokenValidator { - return func(ctx context.Context, candidate string) (context.Context, error) { - next, err := validate(ctx, candidate) - if !errors.Is(err, server.ErrInvalidToken) { - return next, err - } - if r := scimRequestKey.Value(ctx); r != nil && a.performRateLimiting(lmt, r) != nil { - return ctx, errSCIMTooManyRequests() - } - return next, err - } -} - -func (a *API) limitSCIMByProvider(lmt *limiter.Limiter) func(http.Handler) http.Handler { - return func(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if providerID, err := scimProviderID(r.Context()); err == nil && tollbooth.LimitByKeys(lmt, []string{providerID.String()}) != nil { - handler(scimTooManyRequests)(w, r) - return - } - next.ServeHTTP(w, r) - }) - } -} - func (a *API) auditSCIM(tx *storage.Connection, r *http.Request, actor *models.User, action models.AuditAction, providerID uuid.UUID, traits map[string]any) error { traits["sso_provider_id"] = providerID traits["outcome"] = "success" @@ -176,10 +129,6 @@ func scimBaseURL(config *conf.GlobalConfiguration) string { return strings.TrimRight(config.API.ExternalURL, "/") + scimBasePath } -func scimTooManyRequests(w http.ResponseWriter, r *http.Request) error { - return protocol.SendError(w, errSCIMTooManyRequests()) -} - func scimLogError(r *http.Request, err error) { observability.GetLogEntry(r).Entry.WithError(err).Error("scim: request failed") } @@ -290,50 +239,6 @@ func scimParseVersion(version string) (*time.Time, error) { return &updatedAt, nil } -func scimError(err error) error { - switch { - case models.IsNotFoundError(err): - return errSCIMNotFound() - case errors.Is(err, models.SCIMUserStaleError{}), errors.Is(err, models.SCIMGroupStaleError{}): - return errSCIMStale() - case errors.Is(err, models.SCIMGroupConflictError{}): - return scimerrors.ErrUniqueness(`"externalId" must be unique`) - case errors.As(err, &models.SCIMGroupMemberNotFoundError{}): - return errSCIMMemberNotFound() - case errors.Is(err, models.SCIMUserConflictError{}): - return scimerrors.ErrUniqueness(`"userName" and "externalId" must be unique`) - case errors.Is(err, models.SCIMUserLinkedError{}): - return scimerrors.ErrUniqueness("user is already provisioned by this provider") - case errors.Is(err, models.SCIMUserDeletedError{}): - return scimerrors.ErrUniqueness("user was deleted by this provider") - } - return err -} - -func errSCIMNotFound() error { - return scimerrors.ErrNotFound("Resource not found") -} - -func errSCIMStale() error { - return scimerrors.ErrPreconditionFailed("resource has changed on the server") -} - -func errSCIMMemberNotFound() error { - return scimerrors.ErrInvalidValue(`"members.value" must reference a User in this provider`) -} - -func errSCIMEmailRequired() error { - return scimerrors.ErrInvalidValue(`"emails" or an email address "userName" is required`) -} - -func errSCIMEmailInvalid() error { - return scimerrors.ErrInvalidValue(`"emails" value must be an email address`) -} - -func errSCIMTooManyRequests() error { - return scimerrors.NewError(http.StatusTooManyRequests, "", "Request rate limit reached") -} - func scimActor(r *http.Request) *models.User { prefix := "" if token := scimTokenKey.Value(r.Context()); token != nil { diff --git a/internal/api/scim_errors.go b/internal/api/scim_errors.go new file mode 100644 index 0000000000..2b9e7cdc67 --- /dev/null +++ b/internal/api/scim_errors.go @@ -0,0 +1,62 @@ +package api + +import ( + "net/http" + + "github.com/pkg/errors" + "github.com/supabase-community/scim-go/pkg/scimerrors" + "github.com/supabase/auth/internal/api/apierrors" + "github.com/supabase/auth/internal/models" +) + +func scimError(err error) error { + switch { + case models.IsNotFoundError(err): + return errSCIMNotFound() + case errors.Is(err, models.SCIMUserStaleError{}), errors.Is(err, models.SCIMGroupStaleError{}): + return errSCIMStale() + case errors.Is(err, models.SCIMGroupConflictError{}): + return scimerrors.ErrUniqueness(`"externalId" must be unique`) + case errors.As(err, &models.SCIMGroupMemberNotFoundError{}): + return errSCIMMemberNotFound() + case errors.Is(err, models.SCIMUserConflictError{}): + return scimerrors.ErrUniqueness(`"userName" and "externalId" must be unique`) + case errors.Is(err, models.SCIMUserLinkedError{}): + return scimerrors.ErrUniqueness("user is already provisioned by this provider") + case errors.Is(err, models.SCIMUserDeletedError{}): + return scimerrors.ErrUniqueness("user was deleted by this provider") + } + return err +} + +func errSCIMNotFound() error { + return scimerrors.ErrNotFound("Resource not found") +} + +func errSCIMStale() error { + return scimerrors.ErrPreconditionFailed("resource has changed on the server") +} + +func errSCIMMemberNotFound() error { + return scimerrors.ErrInvalidValue(`"members.value" must reference a User in this provider`) +} + +func errSCIMEmailRequired() error { + return scimerrors.ErrInvalidValue(`"emails" or an email address "userName" is required`) +} + +func errSCIMEmailInvalid() error { + return scimerrors.ErrInvalidValue(`"emails" value must be an email address`) +} + +func errSCIMTooManyRequests() error { + return scimerrors.NewError(http.StatusTooManyRequests, "", "Request rate limit reached") +} + +func scimHookError(err error) error { + var httpErr *apierrors.HTTPError + if errors.As(err, &httpErr) && httpErr.HTTPStatus < http.StatusInternalServerError { + return scimerrors.NewError(httpErr.HTTPStatus, "", httpErr.Message) + } + return err +} diff --git a/internal/api/scim_ratelimit.go b/internal/api/scim_ratelimit.go new file mode 100644 index 0000000000..98ed74b487 --- /dev/null +++ b/internal/api/scim_ratelimit.go @@ -0,0 +1,61 @@ +package api + +import ( + "context" + "net/http" + "path" + + "github.com/didip/tollbooth/v5" + "github.com/didip/tollbooth/v5/limiter" + "github.com/pkg/errors" + "github.com/supabase-community/scim-go/pkg/protocol" + "github.com/supabase-community/scim-go/pkg/server" +) + +func (a *API) limitSCIMByIP(lmt *limiter.Limiter) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if a.scimSkipsTokenValidator(r) && a.performRateLimiting(lmt, r) != nil { + handler(scimTooManyRequests)(w, r) + return + } + next.ServeHTTP(w, r) + }) + } +} + +func (a *API) scimSkipsTokenValidator(r *http.Request) bool { + if _, err := a.extractBearerToken(r); err != nil { + return true + } + return r.URL.Path == scimBasePath+"/ServiceProviderConfig" || path.Clean(r.URL.Path) != r.URL.Path +} + +func (a *API) limitSCIMInvalidToken(validate server.TokenValidator, lmt *limiter.Limiter) server.TokenValidator { + return func(ctx context.Context, candidate string) (context.Context, error) { + next, err := validate(ctx, candidate) + if !errors.Is(err, server.ErrInvalidToken) { + return next, err + } + if r := scimRequestKey.Value(ctx); r != nil && a.performRateLimiting(lmt, r) != nil { + return ctx, errSCIMTooManyRequests() + } + return next, err + } +} + +func (a *API) limitSCIMByProvider(lmt *limiter.Limiter) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if providerID, err := scimProviderID(r.Context()); err == nil && tollbooth.LimitByKeys(lmt, []string{providerID.String()}) != nil { + handler(scimTooManyRequests)(w, r) + return + } + next.ServeHTTP(w, r) + }) + } +} + +func scimTooManyRequests(w http.ResponseWriter, r *http.Request) error { + return protocol.SendError(w, errSCIMTooManyRequests()) +} diff --git a/internal/api/scim_user_cleanup.go b/internal/api/scim_user_cleanup.go new file mode 100644 index 0000000000..bff85bbe9b --- /dev/null +++ b/internal/api/scim_user_cleanup.go @@ -0,0 +1,38 @@ +package api + +import ( + "net/http" + + "github.com/gofrs/uuid" + "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/storage" +) + +func (a *API) deleteSCIMUsers(tx *storage.Connection, r *http.Request, actor *models.User, userID uuid.UUID) error { + rows, err := models.SoftDeleteSCIMUsersByUserID(tx, userID) + if err != nil { + return err + } + for i := range rows { + if err := a.removeSCIMUserFromGroups(tx, r, actor, &rows[i]); err != nil { + return err + } + if err := a.auditSCIM(tx, r, actor, models.SCIMUserDeletedAction, rows[i].SSOProviderID, scimUserTraits(&rows[i])); err != nil { + return err + } + } + return nil +} + +func (a *API) removeSCIMUserFromGroups(tx *storage.Connection, r *http.Request, actor *models.User, row *models.SCIMUser) error { + groupIDs, err := models.RemoveSCIMUserFromGroups(tx, row.ID) + if err != nil { + return err + } + for _, groupID := range groupIDs { + if err := a.auditSCIM(tx, r, actor, models.SCIMGroupMemberRemovedAction, row.SSOProviderID, scimMemberTraits(groupID, row.ID, row.UserID)); err != nil { + return err + } + } + return nil +} diff --git a/internal/api/scim_user_linking.go b/internal/api/scim_user_linking.go new file mode 100644 index 0000000000..45a290524f --- /dev/null +++ b/internal/api/scim_user_linking.go @@ -0,0 +1,142 @@ +package api + +import ( + "net/http" + + "github.com/gofrs/uuid" + "github.com/supabase-community/scim-go/pkg/core" + "github.com/supabase-community/scim-go/pkg/scimerrors" + "github.com/supabase/auth/internal/api/apierrors" + "github.com/supabase/auth/internal/api/provider" + "github.com/supabase/auth/internal/hooks/v0hooks" + "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/observability" + "github.com/supabase/auth/internal/storage" +) + +func (s *scimUserRepository) provisionAuthUser(tx *storage.Connection, row *models.SCIMUser, user *core.User) (*models.User, error) { + linked, isNew, err := s.linkAuthUser(tx, row, user) + if err != nil { + return nil, err + } + var created *models.User + if isNew { + created = linked + } + if !row.Active { + return created, models.LogoutSCIMUser(tx, linked.ID) + } + return created, nil +} + +func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SCIMUser, user *core.User) (*models.User, bool, error) { + providerType := scimProviderType(row.SSOProviderID) + decision, err := s.decideAccountLinking(tx, providerType, user) + if err != nil { + return nil, false, err + } + + linked := decision.User + switch decision.Decision { + case models.AccountExists: + if err = scimCanLink(linked); err != nil { + return nil, false, err + } + case models.LinkAccount: + if err = scimCanLink(linked); err != nil { + return nil, false, err + } + if _, err = s.api.createNewIdentity(tx, linked, providerType, scimIdentityData(user)); err != nil { + return nil, false, err + } + if err = linked.UpdateAppMetaDataProviders(tx); err != nil { + return nil, false, err + } + case models.CreateAccount: + if linked, err = s.createAuthUser(tx, providerType, decision, user); err != nil { + return nil, false, err + } + return linked, true, models.LinkSCIMUser(tx, row, linked.ID) + case models.MultipleAccounts: + return nil, false, scimerrors.ErrUniqueness("multiple users share this email in the SSO provider") + default: + return nil, false, apierrors.NewInternalServerError("Unknown automatic linking decision: %v", decision.Decision) + } + + return linked, false, models.LinkSCIMUser(tx, row, linked.ID) +} + +func scimCanLink(linked *models.User) error { + if !linked.IsSSOUser { + return scimerrors.ErrUniqueness("user is not an SSO user") + } + return nil +} + +func (s *scimUserRepository) createAuthUser(tx *storage.Connection, providerType string, decision models.AccountLinkingResult, user *core.User) (*models.User, error) { + candidate, err := s.newUser(providerType, decision, user) + if err != nil { + return nil, err + } + created, err := s.api.signupNewUser(tx, candidate) + if err != nil { + return nil, err + } + if _, err := s.api.createNewIdentity(tx, created, providerType, scimIdentityData(user)); err != nil { + return nil, err + } + return created, nil +} + +func (s *scimUserRepository) beforeProvision(r *http.Request, db *storage.Connection, providerID uuid.UUID, user *core.User) error { + if scimUserEmail(user) == "" { + return errSCIMEmailRequired() + } + return scimError(s.runBeforeUserCreatedHook(r, db, providerID, user)) +} + +func (s *scimUserRepository) runBeforeUserCreatedHook(r *http.Request, db *storage.Connection, providerID uuid.UUID, user *core.User) error { + if !s.api.hooksMgr.Enabled(v0hooks.BeforeUserCreated) { + return nil + } + providerType := scimProviderType(providerID) + decision, err := s.decideAccountLinking(db, providerType, user) + if err != nil || decision.Decision != models.CreateAccount { + return err + } + candidate, err := s.newUser(providerType, decision, user) + if err != nil { + return err + } + return scimHookError(s.api.triggerBeforeUserCreated(r, db, candidate)) +} + +func (s *scimUserRepository) runAfterUserCreatedHook(r *http.Request, db *storage.Connection, user *models.User) { + if user == nil { + return + } + if err := s.api.triggerAfterUserCreated(r, db, user); err != nil { + observability.GetLogEntry(r).Entry.WithError(err).WithField("user_id", user.ID).Error("scim: after user created hook failed") + } +} + +func (s *scimUserRepository) decideAccountLinking(conn *storage.Connection, providerType string, user *core.User) (models.AccountLinkingResult, error) { + emails := []provider.Email{{Email: scimUserEmail(user), Verified: true, Primary: true}} + return models.DetermineAccountLinking(conn, s.api.config, emails, s.api.config.JWT.Aud, providerType, user.UserName) +} + +func (s *scimUserRepository) newUser(providerType string, decision models.AccountLinkingResult, user *core.User) (*models.User, error) { + params := &SignupParams{ + Provider: providerType, + Email: decision.CandidateEmail.Email, + Aud: s.api.config.JWT.Aud, + Data: scimIdentityData(user), + } + candidate, err := params.ToUserModel(true) + if err != nil { + return nil, err + } + now := s.api.Now() + candidate.EmailConfirmedAt = &now + return candidate, nil +} diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index f3c40263b2..8b5d6e9f80 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -11,10 +11,6 @@ import ( "github.com/gofrs/uuid" "github.com/supabase-community/scim-go/pkg/core" "github.com/supabase-community/scim-go/pkg/protocol" - "github.com/supabase-community/scim-go/pkg/scimerrors" - "github.com/supabase/auth/internal/api/apierrors" - "github.com/supabase/auth/internal/api/provider" - "github.com/supabase/auth/internal/hooks/v0hooks" "github.com/supabase/auth/internal/models" "github.com/supabase/auth/internal/observability" "github.com/supabase/auth/internal/storage" @@ -302,133 +298,6 @@ func (s *scimUserRepository) syncAuthUser(tx *storage.Connection, change scimUse return nil, nil } -func (s *scimUserRepository) provisionAuthUser(tx *storage.Connection, row *models.SCIMUser, user *core.User) (*models.User, error) { - linked, isNew, err := s.linkAuthUser(tx, row, user) - if err != nil { - return nil, err - } - var created *models.User - if isNew { - created = linked - } - if !row.Active { - return created, models.LogoutSCIMUser(tx, linked.ID) - } - return created, nil -} - -func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SCIMUser, user *core.User) (*models.User, bool, error) { - providerType := scimProviderType(row.SSOProviderID) - decision, err := s.decideAccountLinking(tx, providerType, user) - if err != nil { - return nil, false, err - } - - linked := decision.User - switch decision.Decision { - case models.AccountExists: - if err = scimCanLink(linked); err != nil { - return nil, false, err - } - case models.LinkAccount: - if err = scimCanLink(linked); err != nil { - return nil, false, err - } - if _, err = s.api.createNewIdentity(tx, linked, providerType, scimIdentityData(user)); err != nil { - return nil, false, err - } - if err = linked.UpdateAppMetaDataProviders(tx); err != nil { - return nil, false, err - } - case models.CreateAccount: - if linked, err = s.createAuthUser(tx, providerType, decision, user); err != nil { - return nil, false, err - } - return linked, true, models.LinkSCIMUser(tx, row, linked.ID) - case models.MultipleAccounts: - return nil, false, scimerrors.ErrUniqueness("multiple users share this email in the SSO provider") - default: - return nil, false, apierrors.NewInternalServerError("Unknown automatic linking decision: %v", decision.Decision) - } - - return linked, false, models.LinkSCIMUser(tx, row, linked.ID) -} - -func scimCanLink(linked *models.User) error { - if !linked.IsSSOUser { - return scimerrors.ErrUniqueness("user is not an SSO user") - } - return nil -} - -func (s *scimUserRepository) createAuthUser(tx *storage.Connection, providerType string, decision models.AccountLinkingResult, user *core.User) (*models.User, error) { - candidate, err := s.newUser(providerType, decision, user) - if err != nil { - return nil, err - } - created, err := s.api.signupNewUser(tx, candidate) - if err != nil { - return nil, err - } - if _, err := s.api.createNewIdentity(tx, created, providerType, scimIdentityData(user)); err != nil { - return nil, err - } - return created, nil -} - -func (s *scimUserRepository) beforeProvision(r *http.Request, db *storage.Connection, providerID uuid.UUID, user *core.User) error { - if scimUserEmail(user) == "" { - return errSCIMEmailRequired() - } - return scimError(s.runBeforeUserCreatedHook(r, db, providerID, user)) -} - -func (s *scimUserRepository) runBeforeUserCreatedHook(r *http.Request, db *storage.Connection, providerID uuid.UUID, user *core.User) error { - if !s.api.hooksMgr.Enabled(v0hooks.BeforeUserCreated) { - return nil - } - providerType := scimProviderType(providerID) - decision, err := s.decideAccountLinking(db, providerType, user) - if err != nil || decision.Decision != models.CreateAccount { - return err - } - candidate, err := s.newUser(providerType, decision, user) - if err != nil { - return err - } - return scimHookError(s.api.triggerBeforeUserCreated(r, db, candidate)) -} - -func (s *scimUserRepository) runAfterUserCreatedHook(r *http.Request, db *storage.Connection, user *models.User) { - if user == nil { - return - } - if err := s.api.triggerAfterUserCreated(r, db, user); err != nil { - observability.GetLogEntry(r).Entry.WithError(err).WithField("user_id", user.ID).Error("scim: after user created hook failed") - } -} - -func (s *scimUserRepository) decideAccountLinking(conn *storage.Connection, providerType string, user *core.User) (models.AccountLinkingResult, error) { - emails := []provider.Email{{Email: scimUserEmail(user), Verified: true, Primary: true}} - return models.DetermineAccountLinking(conn, s.api.config, emails, s.api.config.JWT.Aud, providerType, user.UserName) -} - -func (s *scimUserRepository) newUser(providerType string, decision models.AccountLinkingResult, user *core.User) (*models.User, error) { - params := &SignupParams{ - Provider: providerType, - Email: decision.CandidateEmail.Email, - Aud: s.api.config.JWT.Aud, - Data: scimIdentityData(user), - } - candidate, err := params.ToUserModel(true) - if err != nil { - return nil, err - } - now := s.api.Now() - candidate.EmailConfirmedAt = &now - return candidate, nil -} - func (s *scimUserRepository) renameIdentity(tx *storage.Connection, change scimUserChange, old *models.SCIMUser) error { providerID, user, userID := change.target.ProviderID, change.user, *old.UserID from, err := scimUserName(old.Resource) @@ -479,35 +348,6 @@ func (s *scimUserRepository) audit(tx *storage.Connection, r *http.Request, acti return s.api.auditSCIM(tx, r, scimActor(r), action, row.SSOProviderID, scimUserTraits(row)) } -func (a *API) deleteSCIMUsers(tx *storage.Connection, r *http.Request, actor *models.User, userID uuid.UUID) error { - rows, err := models.SoftDeleteSCIMUsersByUserID(tx, userID) - if err != nil { - return err - } - for i := range rows { - if err := a.removeSCIMUserFromGroups(tx, r, actor, &rows[i]); err != nil { - return err - } - if err := a.auditSCIM(tx, r, actor, models.SCIMUserDeletedAction, rows[i].SSOProviderID, scimUserTraits(&rows[i])); err != nil { - return err - } - } - return nil -} - -func (a *API) removeSCIMUserFromGroups(tx *storage.Connection, r *http.Request, actor *models.User, row *models.SCIMUser) error { - groupIDs, err := models.RemoveSCIMUserFromGroups(tx, row.ID) - if err != nil { - return err - } - for _, groupID := range groupIDs { - if err := a.auditSCIM(tx, r, actor, models.SCIMGroupMemberRemovedAction, row.SSOProviderID, scimMemberTraits(groupID, row.ID, row.UserID)); err != nil { - return err - } - } - return nil -} - func scimUserResource(user *core.User) ([]byte, error) { resource, err := scimEncode(user, "id", "meta", "password", "groups") if err != nil { @@ -605,11 +445,3 @@ func scimUserAuditAction(before, after *models.SCIMUser) models.AuditAction { } return models.SCIMUserUpdatedAction } - -func scimHookError(err error) error { - var httpErr *apierrors.HTTPError - if errors.As(err, &httpErr) && httpErr.HTTPStatus < http.StatusInternalServerError { - return scimerrors.NewError(httpErr.HTTPStatus, "", httpErr.Message) - } - return err -} From 55f88c592604bf58a9665470fb34d6a2527d3cdb Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:58:53 -0600 Subject: [PATCH 55/88] chore(scim): share the SCIM user create and replace write path --- internal/api/scim_users.go | 43 ++++++++++++++++++-------------------- 1 file changed, 20 insertions(+), 23 deletions(-) diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index 8b5d6e9f80..fef0e2d127 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -68,18 +68,7 @@ func (s *scimUserRepository) Create(ctx context.Context, user *core.User) (*core } change := scimUserChange{r: r, target: models.SCIMTarget{ProviderID: providerID}, resource: resource, user: user} - var row *models.SCIMUser - var created *models.User - err = db.Transaction(func(tx *storage.Connection) error { - var terr error - row, created, terr = s.create(tx, change) - return terr - }) - if err != nil { - return nil, scimError(err) - } - s.runAfterUserCreatedHook(r, db, created) - return s.renderOne(db, providerID, row, protocol.Projection{}) + return s.save(db, change, s.create) } func (s *scimUserRepository) Replace(ctx context.Context, user *core.User) (*core.User, error) { @@ -107,18 +96,9 @@ func (s *scimUserRepository) Replace(ctx context.Context, user *core.User) (*cor } change := scimUserChange{r: r, target: target, resource: resource, user: user} - var row *models.SCIMUser - var created *models.User - err = db.Transaction(func(tx *storage.Connection) error { - var terr error - row, created, terr = s.replace(tx, change, existing) - return terr + return s.save(db, change, func(tx *storage.Connection, change scimUserChange) (*models.SCIMUser, *models.User, error) { + return s.replace(tx, change, existing) }) - if err != nil { - return nil, scimError(err) - } - s.runAfterUserCreatedHook(r, db, created) - return s.renderOne(db, target.ProviderID, row, protocol.Projection{}) } func (s *scimUserRepository) Delete(ctx context.Context, id, version string) error { @@ -140,6 +120,23 @@ func (s *scimUserRepository) Delete(ctx context.Context, id, version string) err })) } +func (s *scimUserRepository) save(db *storage.Connection, change scimUserChange, write func(*storage.Connection, scimUserChange) (*models.SCIMUser, *models.User, error)) (*core.User, error) { + var ( + row *models.SCIMUser + created *models.User + ) + err := db.Transaction(func(tx *storage.Connection) error { + var terr error + row, created, terr = write(tx, change) + return terr + }) + if err != nil { + return nil, scimError(err) + } + s.runAfterUserCreatedHook(change.r, db, created) + return s.renderOne(db, change.target.ProviderID, row, protocol.Projection{}) +} + func (s *scimUserRepository) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMUser, projection protocol.Projection) ([]*core.User, error) { groups, err := s.groupMemberships(tx, providerID, rows, projection) if err != nil { From 8a98cf9479e7f48b07e22f9312f7c70fcd8071c8 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 00:59:28 -0600 Subject: [PATCH 56/88] chore(scim): share the SCIM optimistic version clause --- internal/models/scim.go | 6 ++++-- internal/models/scim_group.go | 2 +- internal/models/scim_user.go | 2 +- 3 files changed, 6 insertions(+), 4 deletions(-) diff --git a/internal/models/scim.go b/internal/models/scim.go index 69665b9d34..7d0dd19dc8 100644 --- a/internal/models/scim.go +++ b/internal/models/scim.go @@ -13,6 +13,8 @@ import ( "github.com/supabase/auth/internal/storage" ) +const scimVersionClause = "(?::timestamptz IS NULL OR updated_at = ?)" + type SCIMFilter struct { Name *string ExternalID *string @@ -104,7 +106,7 @@ func findSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarg func findUnchangedSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarget, resource []byte) (*T, error) { row := new(T) if err := tx.RawQuery( - fmt.Sprintf("SELECT %s FROM %q WHERE %s AND resource = ?::jsonb AND (?::timestamptz IS NULL OR updated_at = ?) FOR UPDATE", table.columns, table.name, table.targetClause()), + fmt.Sprintf("SELECT %s FROM %q WHERE %s AND resource = ?::jsonb AND "+scimVersionClause+" FOR UPDATE", table.columns, table.name, table.targetClause()), target.ID, target.ProviderID, string(resource), target.UpdatedAt, target.UpdatedAt, ).First(row); err != nil { if errors.Is(err, sql.ErrNoRows) { @@ -118,7 +120,7 @@ func findUnchangedSCIMRow[T any](tx *storage.Connection, table scimTable, target func replaceSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarget, resource []byte) (*T, error) { row := new(T) if err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET resource = ?::jsonb, updated_at = now() WHERE %s AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING %s", table.name, table.targetClause(), table.columns), + fmt.Sprintf("UPDATE %q SET resource = ?::jsonb, updated_at = now() WHERE %s AND "+scimVersionClause+" RETURNING %s", table.name, table.targetClause(), table.columns), string(resource), target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(row); err != nil { return nil, table.writeError(tx, target, err, "replacing") diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go index 294b86c358..979009947c 100644 --- a/internal/models/scim_group.go +++ b/internal/models/scim_group.go @@ -89,7 +89,7 @@ func TouchSCIMGroup(tx *storage.Connection, group *SCIMGroup) (*SCIMGroup, error func DeleteSCIMGroup(tx *storage.Connection, target SCIMTarget) (*SCIMGroup, error) { group := &SCIMGroup{} err := tx.RawQuery( - fmt.Sprintf("DELETE FROM %q WHERE id = ? AND sso_provider_id = ? AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimGroupColumns, scimGroupsTable.name), + fmt.Sprintf("DELETE FROM %q WHERE id = ? AND sso_provider_id = ? AND "+scimVersionClause+" RETURNING "+scimGroupColumns, scimGroupsTable.name), target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(group) if err != nil { diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index 7234db84f9..15e9f4cd08 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -82,7 +82,7 @@ func ReplaceSCIMUser(tx *storage.Connection, target SCIMTarget, resource []byte) func DeleteSCIMUser(tx *storage.Connection, target SCIMTarget) (*SCIMUser, error) { user := &SCIMUser{} err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE id = ? AND sso_provider_id = ? AND deleted_at IS NULL AND (?::timestamptz IS NULL OR updated_at = ?) RETURNING "+scimUserColumns, scimUsersTable.name), + fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE id = ? AND sso_provider_id = ? AND deleted_at IS NULL AND "+scimVersionClause+" RETURNING "+scimUserColumns, scimUsersTable.name), target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(user) if err != nil { From 6035ed2341b350f0dd283f5c487e3ff5c2c97ec1 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 01:01:29 -0600 Subject: [PATCH 57/88] chore(scim): share the SCIM disabled audit event --- internal/api/scim_admin.go | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/internal/api/scim_admin.go b/internal/api/scim_admin.go index 69cf738f77..8e2a6164df 100644 --- a/internal/api/scim_admin.go +++ b/internal/api/scim_admin.go @@ -76,7 +76,7 @@ func (a *API) adminSCIMDisable(w http.ResponseWriter, r *http.Request) error { if err != nil || !changed { return err } - return a.auditSCIM(tx, r, actor, models.SCIMDisabledAction, provider.ID, map[string]any{"token_prefixes": prefixes}) + return a.auditSCIMDisabled(tx, r, provider.ID, prefixes) }); err != nil { return apierrors.NewInternalServerError("Error disabling SCIM").WithInternalError(err) } @@ -218,7 +218,7 @@ func (a *API) deprovisionSCIM(tx *storage.Connection, r *http.Request, provider return err } if enabled && a.config.SSO.SCIM.Enabled { - if err := a.auditSCIM(tx, r, actor, models.SCIMDisabledAction, provider.ID, map[string]any{"token_prefixes": prefixes}); err != nil { + if err := a.auditSCIMDisabled(tx, r, provider.ID, prefixes); err != nil { return err } } @@ -249,3 +249,7 @@ func (a *API) revokeSCIMTokens(tx *storage.Connection, r *http.Request, actor *m } return prefixes, nil } + +func (a *API) auditSCIMDisabled(tx *storage.Connection, r *http.Request, providerID uuid.UUID, prefixes []string) error { + return a.auditSCIM(tx, r, getAdminUser(r.Context()), models.SCIMDisabledAction, providerID, map[string]any{"token_prefixes": prefixes}) +} From 50b2f2029be982a47f61e8243b61243b2b808334 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 01:04:17 -0600 Subject: [PATCH 58/88] chore(scim): move the shared SCIM test suite and helpers into scim_test.go --- internal/api/scim_test.go | 102 ++++++++++++++++++++++++++++++++ internal/api/scim_users_test.go | 102 -------------------------------- 2 files changed, 102 insertions(+), 102 deletions(-) diff --git a/internal/api/scim_test.go b/internal/api/scim_test.go index 4b8302086c..ee998180d8 100644 --- a/internal/api/scim_test.go +++ b/internal/api/scim_test.go @@ -12,11 +12,14 @@ import ( "testing" "time" + jwt "github.com/golang-jwt/jwt/v5" "github.com/pkg/errors" "github.com/sirupsen/logrus" logrustest "github.com/sirupsen/logrus/hooks/test" "github.com/stretchr/testify/require" + "github.com/stretchr/testify/suite" scimCore "github.com/supabase-community/scim-go/pkg/core" + "github.com/supabase-community/scim-go/pkg/protocol" scimProtocol "github.com/supabase-community/scim-go/pkg/protocol" "github.com/supabase-community/scim-go/pkg/server" "github.com/supabase/auth/internal/conf" @@ -492,3 +495,102 @@ func TestSCIMUserFields(t *testing.T) { require.Equal(t, "alice", resource["userName"]) }) } + +type SCIMUsersTestSuite struct { + suite.Suite + API *API + TokenA string + TokenB string + A *models.SSOProvider + B *models.SSOProvider +} + +func TestSCIMUsers(t *testing.T) { + api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { + config.RateLimitScim = 1_000_000 + }) + defer api.db.Close() + + suite.Run(t, &SCIMUsersTestSuite{API: api}) +} + +func (ts *SCIMUsersTestSuite) SetupTest() { + require.NoError(ts.T(), models.TruncateAll(ts.API.db)) + ts.A, ts.TokenA = ts.provider() + ts.B, ts.TokenB = ts.provider() +} + +func (ts *SCIMUsersTestSuite) provider() (*models.SSOProvider, string) { + return createSSOProviderWithSCIMToken(ts.T(), ts.API.db) +} + +func (ts *SCIMUsersTestSuite) do(token, method, path, body string) (*httptest.ResponseRecorder, map[string]any) { + return ts.doAs(protocol.MediaType, token, method, path, body) +} + +func (ts *SCIMUsersTestSuite) doAs(contentType, token, method, path, body string, headers ...string) (*httptest.ResponseRecorder, map[string]any) { + w := ts.serve(contentType, token, method, path, body, headers...) + + var decoded map[string]any + if w.Body.Len() > 0 { + require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &decoded), w.Body.String()) + } + return w, decoded +} + +func (ts *SCIMUsersTestSuite) serve(contentType, token, method, path, body string, headers ...string) *httptest.ResponseRecorder { + r := httptest.NewRequest(method, "/scim/v2"+path, strings.NewReader(body)) + r.Header.Set("Authorization", "Bearer "+token) + r.Header.Set("Content-Type", contentType) + for i := 0; i+1 < len(headers); i += 2 { + r.Header.Set(headers[i], headers[i+1]) + } + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + return w +} + +func setupSCIMAPI(t *testing.T, tweak func(*conf.GlobalConfiguration)) (*API, *conf.GlobalConfiguration) { + api, config, err := setupAPIForTestWithCallback(func(config *conf.GlobalConfiguration, conn *storage.Connection) { + if config != nil { + config.SSO.SCIM.Enabled = true + if tweak != nil { + tweak(config) + } + } + }) + require.NoError(t, err) + return api, config +} + +func createSSOProvider(t require.TestingT, db *storage.Connection) *models.SSOProvider { + provider := &models.SSOProvider{} + require.NoError(t, db.Create(provider)) + return provider +} + +func createSCIMEnabledProvider(t require.TestingT, db *storage.Connection) *models.SSOProvider { + provider := createSSOProvider(t, db) + _, err := models.EnableSCIM(db, provider.ID) + require.NoError(t, err) + return provider +} + +func createSSOProviderWithSCIMToken(t require.TestingT, db *storage.Connection) (*models.SSOProvider, string) { + provider := createSCIMEnabledProvider(t, db) + _, token, err := models.CreateSCIMToken(db, provider, nil) + require.NoError(t, err) + return provider, token +} + +func adminJWT(t require.TestingT, secret string) string { + token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, &AccessTokenClaims{Role: "supabase_admin"}).SignedString([]byte(secret)) + require.NoError(t, err) + return token +} + +func queryAuditEntries(t require.TestingT, db *storage.Connection, where string, args ...any) []models.AuditLogEntry { + entries := []models.AuditLogEntry{} + require.NoError(t, db.Q().Where(where, args...).Order("created_at asc").All(&entries)) + return entries +} diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index d582b6f0a3..21a94346e9 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -2,7 +2,6 @@ package api import ( "context" - "encoding/json" "fmt" "maps" "net/http" @@ -16,9 +15,7 @@ import ( "time" "github.com/gofrs/uuid" - jwt "github.com/golang-jwt/jwt/v5" "github.com/stretchr/testify/require" - "github.com/stretchr/testify/suite" "github.com/supabase-community/scim-go/pkg/core" "github.com/supabase-community/scim-go/pkg/protocol" "github.com/supabase-community/scim-go/pkg/scimerrors" @@ -41,105 +38,6 @@ const oktaUser = `{ "active": true }` -type SCIMUsersTestSuite struct { - suite.Suite - API *API - TokenA string - TokenB string - A *models.SSOProvider - B *models.SSOProvider -} - -func TestSCIMUsers(t *testing.T) { - api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { - config.RateLimitScim = 1_000_000 - }) - defer api.db.Close() - - suite.Run(t, &SCIMUsersTestSuite{API: api}) -} - -func (ts *SCIMUsersTestSuite) SetupTest() { - require.NoError(ts.T(), models.TruncateAll(ts.API.db)) - ts.A, ts.TokenA = ts.provider() - ts.B, ts.TokenB = ts.provider() -} - -func (ts *SCIMUsersTestSuite) provider() (*models.SSOProvider, string) { - return createSSOProviderWithSCIMToken(ts.T(), ts.API.db) -} - -func setupSCIMAPI(t *testing.T, tweak func(*conf.GlobalConfiguration)) (*API, *conf.GlobalConfiguration) { - api, config, err := setupAPIForTestWithCallback(func(config *conf.GlobalConfiguration, conn *storage.Connection) { - if config != nil { - config.SSO.SCIM.Enabled = true - if tweak != nil { - tweak(config) - } - } - }) - require.NoError(t, err) - return api, config -} - -func createSSOProvider(t require.TestingT, db *storage.Connection) *models.SSOProvider { - provider := &models.SSOProvider{} - require.NoError(t, db.Create(provider)) - return provider -} - -func createSCIMEnabledProvider(t require.TestingT, db *storage.Connection) *models.SSOProvider { - provider := createSSOProvider(t, db) - _, err := models.EnableSCIM(db, provider.ID) - require.NoError(t, err) - return provider -} - -func createSSOProviderWithSCIMToken(t require.TestingT, db *storage.Connection) (*models.SSOProvider, string) { - provider := createSCIMEnabledProvider(t, db) - _, token, err := models.CreateSCIMToken(db, provider, nil) - require.NoError(t, err) - return provider, token -} - -func adminJWT(t require.TestingT, secret string) string { - token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, &AccessTokenClaims{Role: "supabase_admin"}).SignedString([]byte(secret)) - require.NoError(t, err) - return token -} - -func queryAuditEntries(t require.TestingT, db *storage.Connection, where string, args ...any) []models.AuditLogEntry { - entries := []models.AuditLogEntry{} - require.NoError(t, db.Q().Where(where, args...).Order("created_at asc").All(&entries)) - return entries -} - -func (ts *SCIMUsersTestSuite) do(token, method, path, body string) (*httptest.ResponseRecorder, map[string]any) { - return ts.doAs(protocol.MediaType, token, method, path, body) -} - -func (ts *SCIMUsersTestSuite) doAs(contentType, token, method, path, body string, headers ...string) (*httptest.ResponseRecorder, map[string]any) { - w := ts.serve(contentType, token, method, path, body, headers...) - - var decoded map[string]any - if w.Body.Len() > 0 { - require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &decoded), w.Body.String()) - } - return w, decoded -} - -func (ts *SCIMUsersTestSuite) serve(contentType, token, method, path, body string, headers ...string) *httptest.ResponseRecorder { - r := httptest.NewRequest(method, "/scim/v2"+path, strings.NewReader(body)) - r.Header.Set("Authorization", "Bearer "+token) - r.Header.Set("Content-Type", contentType) - for i := 0; i+1 < len(headers); i += 2 { - r.Header.Set(headers[i], headers[i+1]) - } - w := httptest.NewRecorder() - ts.API.handler.ServeHTTP(w, r) - return w -} - func (ts *SCIMUsersTestSuite) create(token, body string) string { w, created := ts.do(token, http.MethodPost, "/Users", body) require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) From 6fcad5e3d59dc17ba3a4b124151b41c54fe42b25 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 01:05:14 -0600 Subject: [PATCH 59/88] chore(scim): move SCIM tests next to the source they cover --- internal/api/scim_groups_test.go | 60 --------- internal/api/scim_ratelimit_test.go | 126 +++++++++++++++++++ internal/api/scim_test.go | 33 ----- internal/api/scim_users_test.go | 184 +++++++++++++--------------- 4 files changed, 208 insertions(+), 195 deletions(-) create mode 100644 internal/api/scim_ratelimit_test.go diff --git a/internal/api/scim_groups_test.go b/internal/api/scim_groups_test.go index 8eaadc6dab..0d251e6a29 100644 --- a/internal/api/scim_groups_test.go +++ b/internal/api/scim_groups_test.go @@ -494,66 +494,6 @@ func (ts *SCIMUsersTestSuite) TestGroupsUnknownID() { } } -func (ts *SCIMUsersTestSuite) TestUsersGroupsAttribute() { - alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) - bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) - ops := ts.createGroup(ts.TokenA, groupWith("Ops", "g-2", alice)) - eng := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) - - groupsOf := func(user map[string]any) []map[string]any { - found := []map[string]any{} - groups, _ := user["groups"].([]any) - for _, group := range groups { - found = append(found, group.(map[string]any)) - } - return found - } - storedResource := func(id string) string { - var stored models.SCIMUser - require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&stored)) - return string(stored.Resource) - } - - w, got := ts.do(ts.TokenA, http.MethodGet, "/Users/"+alice, "") - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - require.Equal(ts.T(), []map[string]any{ - {"value": eng, "$ref": "http://localhost:9999/scim/v2/Groups/" + eng, "display": "Engineering", "type": "direct"}, - {"value": ops, "$ref": "http://localhost:9999/scim/v2/Groups/" + ops, "display": "Ops", "type": "direct"}, - }, groupsOf(got)) - - w, got = ts.do(ts.TokenA, http.MethodGet, "/Users/"+bob, "") - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - require.NotContains(ts.T(), got, "groups") - - listed := ts.list(ts.TokenA, `userName eq "alice@example.com"`) - require.Len(ts.T(), groupsOf(listed["Resources"].([]any)[0].(map[string]any)), 2) - - claimed := `[{"value":"` + ops + `","display":"Forged"}]` - w, created := ts.do(ts.TokenA, http.MethodPost, "/Users", `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"carol@example.com","emails":[{"primary":true,"value":"carol@example.com"}],"groups":`+claimed+`}`) - require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) - require.NotContains(ts.T(), created, "groups") - require.NotContains(ts.T(), storedResource(created["id"].(string)), "groups") - - w, replaced := ts.do(ts.TokenA, http.MethodPut, "/Users/"+bob, `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"bob@example.com","groups":`+claimed+`}`) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - require.NotContains(ts.T(), replaced, "groups") - require.NotContains(ts.T(), storedResource(bob), "groups") - - w, patched := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+alice, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ - {"op":"replace","path":"displayName","value":"Alice"} - ]}`) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - require.Len(ts.T(), groupsOf(patched), 2) - require.NotContains(ts.T(), storedResource(alice), "groups") - - w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Groups/"+ops, "") - require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) - w, got = ts.do(ts.TokenA, http.MethodGet, "/Users/"+alice, "") - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - require.Len(ts.T(), groupsOf(got), 1) - require.Equal(ts.T(), eng, groupsOf(got)[0]["value"]) -} - func (ts *SCIMUsersTestSuite) TestGroupsAuditLog() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) diff --git a/internal/api/scim_ratelimit_test.go b/internal/api/scim_ratelimit_test.go new file mode 100644 index 0000000000..5108888236 --- /dev/null +++ b/internal/api/scim_ratelimit_test.go @@ -0,0 +1,126 @@ +package api + +import ( + "fmt" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" + "github.com/supabase-community/scim-go/pkg/protocol" + "github.com/supabase/auth/internal/conf" + "github.com/supabase/auth/internal/models" +) + +func TestSCIMRateLimit(t *testing.T) { + api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { + config.RateLimitScim = 1 + }) + defer api.db.Close() + require.NoError(t, models.TruncateAll(api.db)) + + token := func() string { + _, token := createSSOProviderWithSCIMToken(t, api.db) + return token + } + tokenA, tokenB := token(), token() + + get := func(token, ip string) *httptest.ResponseRecorder { + r := httptest.NewRequest(http.MethodGet, "/scim/v2/Users", nil) + if token != "" { + r.Header.Set("Authorization", "Bearer "+token) + } + r.Header.Set(api.config.RateLimitHeader, ip) + w := httptest.NewRecorder() + api.handler.ServeHTTP(w, r) + return w + } + + limited := `{"schemas":["urn:ietf:params:scim:api:messages:2.0:Error"],"detail":"Request rate limit reached","status":"429"}` + const ip = "192.0.2.1" + + for range 30 { + w := get(tokenA, ip) + require.Equal(t, http.StatusOK, w.Code, w.Body.String()) + } + w := get(tokenA, ip) + require.Equal(t, http.StatusTooManyRequests, w.Code) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) + require.JSONEq(t, limited, w.Body.String()) + w = get(tokenB, ip) + require.Equal(t, http.StatusOK, w.Code, w.Body.String()) + + for range 30 { + w := get("scim_invalid", ip) + require.Equal(t, http.StatusUnauthorized, w.Code, w.Body.String()) + } + for _, token := range []string{"scim_invalid", ""} { + w := get(token, ip) + require.Equal(t, http.StatusTooManyRequests, w.Code, w.Body.String()) + require.JSONEq(t, limited, w.Body.String()) + } + w = get(tokenB, ip) + require.Equal(t, http.StatusOK, w.Code, w.Body.String()) + w = get("scim_invalid", "198.51.100.1") + require.Equal(t, http.StatusUnauthorized, w.Code, w.Body.String()) + + for i, tc := range []struct { + method, path string + status int + }{ + {http.MethodGet, "/scim/v2/ServiceProviderConfig", http.StatusOK}, + {http.MethodHead, "/scim/v2/ServiceProviderConfig", http.StatusOK}, + {http.MethodGet, "/scim/v2/Unknown", http.StatusUnauthorized}, + {http.MethodDelete, "/scim/v2/Users", http.StatusUnauthorized}, + } { + ip := fmt.Sprintf("203.0.113.%d", i+1) + send := func() *httptest.ResponseRecorder { + r := httptest.NewRequest(tc.method, tc.path, nil) + r.Header.Set("Authorization", "Bearer scim_invalid") + r.Header.Set(api.config.RateLimitHeader, ip) + w := httptest.NewRecorder() + api.handler.ServeHTTP(w, r) + return w + } + for range 30 { + w := send() + require.Equal(t, tc.status, w.Code, tc.method+" "+tc.path) + } + w := send() + require.Equal(t, http.StatusTooManyRequests, w.Code, tc.method+" "+tc.path) + require.JSONEq(t, limited, w.Body.String()) + } +} + +func TestSCIMRateLimitBoundsEveryUnauthenticatedRequest(t *testing.T) { + api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { + config.RateLimitScim = 1 + }) + defer api.db.Close() + require.NoError(t, models.TruncateAll(api.db)) + + type request struct{ path, authorization string } + requests := []request{} + for _, header := range []string{"", "Bearer", "Bearer ", "bearer scim_invalid", "Bearer scim a", "Basic scim_invalid", "Bearer\tscim_invalid", "Bearer scim_invalid"} { + requests = append(requests, request{"/scim/v2/Users", header}) + } + for _, path := range []string{"/scim/v2//Users", "/scim/v2/./Users", "/scim/v2/Users/", "/scim/v2/ServiceProviderConfig/"} { + requests = append(requests, request{path, "Bearer scim_invalid"}) + } + + for i, tc := range requests { + ip := fmt.Sprintf("203.0.113.%d", 100+i) + limited := false + for range 31 { + r := httptest.NewRequest(http.MethodGet, tc.path, nil) + if tc.authorization != "" { + r.Header.Set("Authorization", tc.authorization) + } + r.Header.Set(api.config.RateLimitHeader, ip) + w := httptest.NewRecorder() + api.handler.ServeHTTP(w, r) + limited = limited || w.Code == http.StatusTooManyRequests + } + require.True(t, limited, "%q %q", tc.path, tc.authorization) + } +} diff --git a/internal/api/scim_test.go b/internal/api/scim_test.go index ee998180d8..6ee9215c25 100644 --- a/internal/api/scim_test.go +++ b/internal/api/scim_test.go @@ -463,39 +463,6 @@ func TestSCIMServer(t *testing.T) { }) } -func TestSCIMUserFields(t *testing.T) { - t.Run("has no email without emails", func(t *testing.T) { - require.Empty(t, scimPrimaryEmail(nil)) - }) - - t.Run("prefers the primary email", func(t *testing.T) { - require.Equal(t, "home@example.com", scimPrimaryEmail([]scimCore.Email{ - {Value: "work@example.com"}, - {Value: "home@example.com", Primary: new(true)}, - })) - }) - - t.Run("falls back to the first email", func(t *testing.T) { - require.Equal(t, "work@example.com", scimPrimaryEmail([]scimCore.Email{{Value: "work@example.com"}, {Value: "home@example.com"}})) - }) - - t.Run("drops id, meta and password from the resource", func(t *testing.T) { - user := &scimCore.User{UserName: "alice", Password: "secret"} - user.ID = "abc" - user.Meta = scimCore.Meta{Version: `W/"1"`} - - encoded, err := scimUserResource(user) - require.NoError(t, err) - - resource := map[string]any{} - require.NoError(t, json.Unmarshal(encoded, &resource)) - require.NotContains(t, resource, "id") - require.NotContains(t, resource, "meta") - require.NotContains(t, resource, "password") - require.Equal(t, "alice", resource["userName"]) - }) -} - type SCIMUsersTestSuite struct { suite.Suite API *API diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index 21a94346e9..95d6e54d0e 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -2,7 +2,7 @@ package api import ( "context" - "fmt" + "encoding/json" "maps" "net/http" "net/http/httptest" @@ -17,10 +17,10 @@ import ( "github.com/gofrs/uuid" "github.com/stretchr/testify/require" "github.com/supabase-community/scim-go/pkg/core" + scimCore "github.com/supabase-community/scim-go/pkg/core" "github.com/supabase-community/scim-go/pkg/protocol" "github.com/supabase-community/scim-go/pkg/scimerrors" "github.com/supabase-community/scim-go/pkg/server" - "github.com/supabase/auth/internal/conf" "github.com/supabase/auth/internal/models" "github.com/supabase/auth/internal/storage" ) @@ -652,115 +652,95 @@ func (ts *SCIMUsersTestSuite) TestRequiresSSOProviderOnContext() { require.Error(ts.T(), err) } -func TestSCIMRateLimit(t *testing.T) { - api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { - config.RateLimitScim = 1 +func TestSCIMUserFields(t *testing.T) { + t.Run("has no email without emails", func(t *testing.T) { + require.Empty(t, scimPrimaryEmail(nil)) }) - defer api.db.Close() - require.NoError(t, models.TruncateAll(api.db)) - token := func() string { - _, token := createSSOProviderWithSCIMToken(t, api.db) - return token - } - tokenA, tokenB := token(), token() + t.Run("prefers the primary email", func(t *testing.T) { + require.Equal(t, "home@example.com", scimPrimaryEmail([]scimCore.Email{ + {Value: "work@example.com"}, + {Value: "home@example.com", Primary: new(true)}, + })) + }) - get := func(token, ip string) *httptest.ResponseRecorder { - r := httptest.NewRequest(http.MethodGet, "/scim/v2/Users", nil) - if token != "" { - r.Header.Set("Authorization", "Bearer "+token) - } - r.Header.Set(api.config.RateLimitHeader, ip) - w := httptest.NewRecorder() - api.handler.ServeHTTP(w, r) - return w - } + t.Run("falls back to the first email", func(t *testing.T) { + require.Equal(t, "work@example.com", scimPrimaryEmail([]scimCore.Email{{Value: "work@example.com"}, {Value: "home@example.com"}})) + }) - limited := `{"schemas":["urn:ietf:params:scim:api:messages:2.0:Error"],"detail":"Request rate limit reached","status":"429"}` - const ip = "192.0.2.1" + t.Run("drops id, meta and password from the resource", func(t *testing.T) { + user := &scimCore.User{UserName: "alice", Password: "secret"} + user.ID = "abc" + user.Meta = scimCore.Meta{Version: `W/"1"`} - for range 30 { - w := get(tokenA, ip) - require.Equal(t, http.StatusOK, w.Code, w.Body.String()) - } - w := get(tokenA, ip) - require.Equal(t, http.StatusTooManyRequests, w.Code) - require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) - require.JSONEq(t, limited, w.Body.String()) - w = get(tokenB, ip) - require.Equal(t, http.StatusOK, w.Code, w.Body.String()) - - for range 30 { - w := get("scim_invalid", ip) - require.Equal(t, http.StatusUnauthorized, w.Code, w.Body.String()) - } - for _, token := range []string{"scim_invalid", ""} { - w := get(token, ip) - require.Equal(t, http.StatusTooManyRequests, w.Code, w.Body.String()) - require.JSONEq(t, limited, w.Body.String()) - } - w = get(tokenB, ip) - require.Equal(t, http.StatusOK, w.Code, w.Body.String()) - w = get("scim_invalid", "198.51.100.1") - require.Equal(t, http.StatusUnauthorized, w.Code, w.Body.String()) - - for i, tc := range []struct { - method, path string - status int - }{ - {http.MethodGet, "/scim/v2/ServiceProviderConfig", http.StatusOK}, - {http.MethodHead, "/scim/v2/ServiceProviderConfig", http.StatusOK}, - {http.MethodGet, "/scim/v2/Unknown", http.StatusUnauthorized}, - {http.MethodDelete, "/scim/v2/Users", http.StatusUnauthorized}, - } { - ip := fmt.Sprintf("203.0.113.%d", i+1) - send := func() *httptest.ResponseRecorder { - r := httptest.NewRequest(tc.method, tc.path, nil) - r.Header.Set("Authorization", "Bearer scim_invalid") - r.Header.Set(api.config.RateLimitHeader, ip) - w := httptest.NewRecorder() - api.handler.ServeHTTP(w, r) - return w - } - for range 30 { - w := send() - require.Equal(t, tc.status, w.Code, tc.method+" "+tc.path) - } - w := send() - require.Equal(t, http.StatusTooManyRequests, w.Code, tc.method+" "+tc.path) - require.JSONEq(t, limited, w.Body.String()) - } -} + encoded, err := scimUserResource(user) + require.NoError(t, err) -func TestSCIMRateLimitBoundsEveryUnauthenticatedRequest(t *testing.T) { - api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { - config.RateLimitScim = 1 + resource := map[string]any{} + require.NoError(t, json.Unmarshal(encoded, &resource)) + require.NotContains(t, resource, "id") + require.NotContains(t, resource, "meta") + require.NotContains(t, resource, "password") + require.Equal(t, "alice", resource["userName"]) }) - defer api.db.Close() - require.NoError(t, models.TruncateAll(api.db)) +} - type request struct{ path, authorization string } - requests := []request{} - for _, header := range []string{"", "Bearer", "Bearer ", "bearer scim_invalid", "Bearer scim a", "Basic scim_invalid", "Bearer\tscim_invalid", "Bearer scim_invalid"} { - requests = append(requests, request{"/scim/v2/Users", header}) +func (ts *SCIMUsersTestSuite) TestUsersGroupsAttribute() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + ops := ts.createGroup(ts.TokenA, groupWith("Ops", "g-2", alice)) + eng := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) + + groupsOf := func(user map[string]any) []map[string]any { + found := []map[string]any{} + groups, _ := user["groups"].([]any) + for _, group := range groups { + found = append(found, group.(map[string]any)) + } + return found } - for _, path := range []string{"/scim/v2//Users", "/scim/v2/./Users", "/scim/v2/Users/", "/scim/v2/ServiceProviderConfig/"} { - requests = append(requests, request{path, "Bearer scim_invalid"}) + storedResource := func(id string) string { + var stored models.SCIMUser + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&stored)) + return string(stored.Resource) } - for i, tc := range requests { - ip := fmt.Sprintf("203.0.113.%d", 100+i) - limited := false - for range 31 { - r := httptest.NewRequest(http.MethodGet, tc.path, nil) - if tc.authorization != "" { - r.Header.Set("Authorization", tc.authorization) - } - r.Header.Set(api.config.RateLimitHeader, ip) - w := httptest.NewRecorder() - api.handler.ServeHTTP(w, r) - limited = limited || w.Code == http.StatusTooManyRequests - } - require.True(t, limited, "%q %q", tc.path, tc.authorization) - } + w, got := ts.do(ts.TokenA, http.MethodGet, "/Users/"+alice, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), []map[string]any{ + {"value": eng, "$ref": "http://localhost:9999/scim/v2/Groups/" + eng, "display": "Engineering", "type": "direct"}, + {"value": ops, "$ref": "http://localhost:9999/scim/v2/Groups/" + ops, "display": "Ops", "type": "direct"}, + }, groupsOf(got)) + + w, got = ts.do(ts.TokenA, http.MethodGet, "/Users/"+bob, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.NotContains(ts.T(), got, "groups") + + listed := ts.list(ts.TokenA, `userName eq "alice@example.com"`) + require.Len(ts.T(), groupsOf(listed["Resources"].([]any)[0].(map[string]any)), 2) + + claimed := `[{"value":"` + ops + `","display":"Forged"}]` + w, created := ts.do(ts.TokenA, http.MethodPost, "/Users", `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"carol@example.com","emails":[{"primary":true,"value":"carol@example.com"}],"groups":`+claimed+`}`) + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) + require.NotContains(ts.T(), created, "groups") + require.NotContains(ts.T(), storedResource(created["id"].(string)), "groups") + + w, replaced := ts.do(ts.TokenA, http.MethodPut, "/Users/"+bob, `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"bob@example.com","groups":`+claimed+`}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.NotContains(ts.T(), replaced, "groups") + require.NotContains(ts.T(), storedResource(bob), "groups") + + w, patched := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+alice, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ + {"op":"replace","path":"displayName","value":"Alice"} + ]}`) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Len(ts.T(), groupsOf(patched), 2) + require.NotContains(ts.T(), storedResource(alice), "groups") + + w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Groups/"+ops, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + w, got = ts.do(ts.TokenA, http.MethodGet, "/Users/"+alice, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Len(ts.T(), groupsOf(got), 1) + require.Equal(ts.T(), eng, groupsOf(got)[0]["value"]) } From 3ff636037cffb0a76ca78b430e11b020699602b9 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 01:05:52 -0600 Subject: [PATCH 60/88] chore(scim): rename the shared SCIM test suite to SCIMTestSuite --- internal/api/scim_groups_test.go | 50 +++---- internal/api/scim_isolation_test.go | 10 +- internal/api/scim_link_test.go | 152 +++++++++++----------- internal/api/scim_okta_spec_test.go | 8 +- internal/api/scim_provider_delete_test.go | 26 ++-- internal/api/scim_test.go | 16 +-- internal/api/scim_users_test.go | 56 ++++---- 7 files changed, 159 insertions(+), 159 deletions(-) diff --git a/internal/api/scim_groups_test.go b/internal/api/scim_groups_test.go index 0d251e6a29..ffae9c9e39 100644 --- a/internal/api/scim_groups_test.go +++ b/internal/api/scim_groups_test.go @@ -22,13 +22,13 @@ func groupWith(displayName, externalID string, memberIDs ...string) string { return `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:Group"],"displayName":"` + displayName + `","externalId":"` + externalID + `","members":[` + strings.Join(members, ",") + `]}` } -func (ts *SCIMUsersTestSuite) createGroup(token, body string) string { +func (ts *SCIMTestSuite) createGroup(token, body string) string { w, created := ts.do(token, http.MethodPost, "/Groups", body) require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) return created["id"].(string) } -func (ts *SCIMUsersTestSuite) listGroups(token, filter string) map[string]any { +func (ts *SCIMTestSuite) listGroups(token, filter string) map[string]any { w, body := ts.do(token, http.MethodGet, "/Groups?"+url.Values{"filter": {filter}}.Encode(), "") require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) return body @@ -43,7 +43,7 @@ func memberValues(group map[string]any) []string { return values } -func (ts *SCIMUsersTestSuite) TestGroupsLifecycle() { +func (ts *SCIMTestSuite) TestGroupsLifecycle() { alice := ts.create(ts.TokenA, userWith("Alice@Example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) @@ -109,7 +109,7 @@ func (ts *SCIMUsersTestSuite) TestGroupsLifecycle() { require.Equal(ts.T(), http.StatusOK, w.Code) } -func (ts *SCIMUsersTestSuite) TestGroupsWithoutMembers() { +func (ts *SCIMTestSuite) TestGroupsWithoutMembers() { w, created := ts.do(ts.TokenA, http.MethodPost, "/Groups", `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:Group"],"displayName":"Empty"}`) require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) require.NotContains(ts.T(), created, "members") @@ -118,7 +118,7 @@ func (ts *SCIMUsersTestSuite) TestGroupsWithoutMembers() { require.EqualValues(ts.T(), 2, ts.listGroups(ts.TokenA, `displayName eq "Empty"`)["totalResults"]) } -func (ts *SCIMUsersTestSuite) TestGroupsRejectInvalidMembers() { +func (ts *SCIMTestSuite) TestGroupsRejectInvalidMembers() { outsider := ts.create(ts.TokenB, userWith("mallory@example.com", "m-1")) deleted := ts.create(ts.TokenA, userWith("gone@example.com", "g-1")) w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+deleted, "") @@ -146,7 +146,7 @@ func (ts *SCIMUsersTestSuite) TestGroupsRejectInvalidMembers() { } } -func (ts *SCIMUsersTestSuite) TestGroupsMemberTypeIsCaseInsensitive() { +func (ts *SCIMTestSuite) TestGroupsMemberTypeIsCaseInsensitive() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) for _, kind := range []string{"user", "USER", "User"} { @@ -157,7 +157,7 @@ func (ts *SCIMUsersTestSuite) TestGroupsMemberTypeIsCaseInsensitive() { } } -func (ts *SCIMUsersTestSuite) TestPatchReplaceMembers() { +func (ts *SCIMTestSuite) TestPatchReplaceMembers() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) carol := ts.create(ts.TokenA, userWith("carol@example.com", "c-1")) @@ -181,7 +181,7 @@ func (ts *SCIMUsersTestSuite) TestPatchReplaceMembers() { }, events) } -func (ts *SCIMUsersTestSuite) TestGroupMemberEventsCarryUserID() { +func (ts *SCIMTestSuite) TestGroupMemberEventsCarryUserID() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) userIDs := map[string]string{} @@ -223,7 +223,7 @@ func (ts *SCIMUsersTestSuite) TestGroupMemberEventsCarryUserID() { }, events) } -func (ts *SCIMUsersTestSuite) TestExcludedMembersKeepsWrites() { +func (ts *SCIMTestSuite) TestExcludedMembersKeepsWrites() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) carol := ts.create(ts.TokenA, userWith("carol@example.com", "c-1")) @@ -253,7 +253,7 @@ func (ts *SCIMUsersTestSuite) TestExcludedMembersKeepsWrites() { require.Equal(ts.T(), "Platform", got["displayName"]) } -func (ts *SCIMUsersTestSuite) TestPatchRemoveAbsentMember() { +func (ts *SCIMTestSuite) TestPatchRemoveAbsentMember() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice)) @@ -265,7 +265,7 @@ func (ts *SCIMUsersTestSuite) TestPatchRemoveAbsentMember() { require.Len(ts.T(), ts.auditActions(models.SCIMGroupMemberRemovedAction), before) } -func (ts *SCIMUsersTestSuite) TestPatchRejectsRemoveWithValue() { +func (ts *SCIMTestSuite) TestPatchRejectsRemoveWithValue() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, bob)) @@ -293,7 +293,7 @@ func (ts *SCIMUsersTestSuite) TestPatchRejectsRemoveWithValue() { require.Equal(ts.T(), []string{alice}, memberValues(got)) } -func (ts *SCIMUsersTestSuite) TestGroupsExternalIDUniqueWithinProvider() { +func (ts *SCIMTestSuite) TestGroupsExternalIDUniqueWithinProvider() { ts.createGroup(ts.TokenA, groupWith("A", "g-1")) w, body := ts.do(ts.TokenA, http.MethodPost, "/Groups", groupWith("B", "g-1")) @@ -303,7 +303,7 @@ func (ts *SCIMUsersTestSuite) TestGroupsExternalIDUniqueWithinProvider() { ts.createGroup(ts.TokenB, groupWith("A", "g-1")) } -func (ts *SCIMUsersTestSuite) TestGroupsETagAndIfMatch() { +func (ts *SCIMTestSuite) TestGroupsETagAndIfMatch() { w, created := ts.do(ts.TokenA, http.MethodPost, "/Groups", groupWith("Engineering", "g-1")) require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) id := created["id"].(string) @@ -328,7 +328,7 @@ func (ts *SCIMUsersTestSuite) TestGroupsETagAndIfMatch() { require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) } -func (ts *SCIMUsersTestSuite) TestIdenticalPutChecksIfMatch() { +func (ts *SCIMTestSuite) TestIdenticalPutChecksIfMatch() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) for path, body := range map[string]string{ "/Users/" + alice: userWith("alice@example.com", "a-2"), @@ -353,7 +353,7 @@ func (ts *SCIMUsersTestSuite) TestIdenticalPutChecksIfMatch() { } } -func (ts *SCIMUsersTestSuite) TestGroupsRemoveDeletedMembers() { +func (ts *SCIMTestSuite) TestGroupsRemoveDeletedMembers() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) eng := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice, bob)) @@ -386,7 +386,7 @@ func (ts *SCIMUsersTestSuite) TestGroupsRemoveDeletedMembers() { require.ElementsMatch(ts.T(), []string{eng, ops}, removed) } -func (ts *SCIMUsersTestSuite) TestGroupsVersionChangesWhenMemberDeleted() { +func (ts *SCIMTestSuite) TestGroupsVersionChangesWhenMemberDeleted() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) id := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) w, _ := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") @@ -401,7 +401,7 @@ func (ts *SCIMUsersTestSuite) TestGroupsVersionChangesWhenMemberDeleted() { require.Equal(ts.T(), http.StatusPreconditionFailed, w.Code, w.Body.String()) } -func (ts *SCIMUsersTestSuite) TestUserDeleteWaitsForGroupWrite() { +func (ts *SCIMTestSuite) TestUserDeleteWaitsForGroupWrite() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) id := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice, bob)) @@ -427,7 +427,7 @@ func (ts *SCIMUsersTestSuite) TestUserDeleteWaitsForGroupWrite() { require.Equal(ts.T(), []string{bob}, memberValues(got)) } -func (ts *SCIMUsersTestSuite) TestGroupsKeepDeactivatedMembers() { +func (ts *SCIMTestSuite) TestGroupsKeepDeactivatedMembers() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) id := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) @@ -454,7 +454,7 @@ func (ts *SCIMUsersTestSuite) TestGroupsKeepDeactivatedMembers() { } } -func (ts *SCIMUsersTestSuite) TestGroupsSortAndPaginate() { +func (ts *SCIMTestSuite) TestGroupsSortAndPaginate() { ts.createGroup(ts.TokenA, groupWith("beta", "g-2")) ts.createGroup(ts.TokenA, groupWith("Alpha", "g-1")) ts.createGroup(ts.TokenA, groupWith("gamma", "g-3")) @@ -473,7 +473,7 @@ func (ts *SCIMUsersTestSuite) TestGroupsSortAndPaginate() { require.Equal(ts.T(), "invalidValue", body["scimType"]) } -func (ts *SCIMUsersTestSuite) TestGroupsUnsupportedFilters() { +func (ts *SCIMTestSuite) TestGroupsUnsupportedFilters() { for _, filter := range []string{ `displayName co "eng"`, `members.value eq "00000000-0000-0000-0000-000000000000"`, @@ -487,14 +487,14 @@ func (ts *SCIMUsersTestSuite) TestGroupsUnsupportedFilters() { } } -func (ts *SCIMUsersTestSuite) TestGroupsUnknownID() { +func (ts *SCIMTestSuite) TestGroupsUnknownID() { for _, id := range []string{"not-a-uuid", "00000000-0000-0000-0000-000000000000"} { w, _ := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") require.Equal(ts.T(), http.StatusNotFound, w.Code, id) } } -func (ts *SCIMUsersTestSuite) TestGroupsAuditLog() { +func (ts *SCIMTestSuite) TestGroupsAuditLog() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) before := len(ts.scimAuditEntries()) @@ -551,7 +551,7 @@ func (ts *SCIMUsersTestSuite) TestGroupsAuditLog() { }, events) } -func (ts *SCIMUsersTestSuite) TestGroupsPushReplay() { +func (ts *SCIMTestSuite) TestGroupsPushReplay() { bjensen := ts.create(ts.TokenA, userWith("bjensen@example.com", "bjensen")) jsmith := ts.create(ts.TokenA, userWith("jsmith@example.com", "701984")) @@ -621,7 +621,7 @@ func (ts *SCIMUsersTestSuite) TestGroupsPushReplay() { }, events) } -func (ts *SCIMUsersTestSuite) TestGroupsPatchReplay() { +func (ts *SCIMTestSuite) TestGroupsPatchReplay() { bjensen := ts.create(ts.TokenA, userWith("bjensen@example.com", "bjensen")) jsmith := ts.create(ts.TokenA, userWith("jsmith@example.com", "701984")) babs := ts.create(ts.TokenA, userWith("babs@jensen.org", "babs")) @@ -644,7 +644,7 @@ func (ts *SCIMUsersTestSuite) TestGroupsPatchReplay() { require.Equal(ts.T(), len(expected), played) } -func (ts *SCIMUsersTestSuite) requireGroup(step, group, displayName string, members []string) { +func (ts *SCIMTestSuite) requireGroup(step, group, displayName string, members []string) { w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+group, "") require.Equal(ts.T(), http.StatusOK, w.Code, step) require.Equal(ts.T(), displayName, got["displayName"], step) diff --git a/internal/api/scim_isolation_test.go b/internal/api/scim_isolation_test.go index 2ab39830b0..7a5fa3fdf7 100644 --- a/internal/api/scim_isolation_test.go +++ b/internal/api/scim_isolation_test.go @@ -10,7 +10,7 @@ import ( "github.com/supabase/auth/internal/models" ) -func (ts *SCIMUsersTestSuite) TestTenantIsolation() { +func (ts *SCIMTestSuite) TestTenantIsolation() { idA := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) idB := ts.create(ts.TokenB, userWith("bob@example.com", "b-1")) @@ -36,7 +36,7 @@ func (ts *SCIMUsersTestSuite) TestTenantIsolation() { require.Equal(ts.T(), true, got["active"]) } -func (ts *SCIMUsersTestSuite) TestTenantIsolationWithSameEmail() { +func (ts *SCIMTestSuite) TestTenantIsolationWithSameEmail() { body := userWith("shared@example.com", "shared-1") ids := map[string]string{ts.TokenA: ts.create(ts.TokenA, body), ts.TokenB: ts.create(ts.TokenB, body)} require.NotEqual(ts.T(), ids[ts.TokenA], ids[ts.TokenB]) @@ -59,7 +59,7 @@ func (ts *SCIMUsersTestSuite) TestTenantIsolationWithSameEmail() { require.Equal(ts.T(), true, got["active"]) } -func (ts *SCIMUsersTestSuite) TestTombstonedUsersInvisibleToBothProviders() { +func (ts *SCIMTestSuite) TestTombstonedUsersInvisibleToBothProviders() { id := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) group := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", id)) w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") @@ -78,7 +78,7 @@ func (ts *SCIMUsersTestSuite) TestTombstonedUsersInvisibleToBothProviders() { require.Empty(ts.T(), memberValues(got)) } -func (ts *SCIMUsersTestSuite) TestRevokedAndExpiredTokensRefusedEverywhere() { +func (ts *SCIMTestSuite) TestRevokedAndExpiredTokensRefusedEverywhere() { user := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) group := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", user)) @@ -147,7 +147,7 @@ func (ts *SCIMUsersTestSuite) TestRevokedAndExpiredTokensRefusedEverywhere() { require.Equal(ts.T(), http.StatusOK, w.Code) } -func (ts *SCIMUsersTestSuite) TestGroupsTenantIsolation() { +func (ts *SCIMTestSuite) TestGroupsTenantIsolation() { aliceA := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) groupA := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", aliceA)) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index 1a53b0e293..21866e40e9 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -19,7 +19,7 @@ import ( "github.com/supabase/auth/internal/storage" ) -func (ts *SCIMUsersTestSuite) ssoUser(provider *models.SSOProvider, sub, email string) *models.User { +func (ts *SCIMTestSuite) ssoUser(provider *models.SSOProvider, sub, email string) *models.User { user, err := models.NewUser("", email, "", ts.API.config.JWT.Aud, nil) require.NoError(ts.T(), err) user.IsSSOUser = true @@ -30,7 +30,7 @@ func (ts *SCIMUsersTestSuite) ssoUser(provider *models.SSOProvider, sub, email s return user } -func (ts *SCIMUsersTestSuite) linkedUser(id string) *models.User { +func (ts *SCIMTestSuite) linkedUser(id string) *models.User { var row models.SCIMUser require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&row)) require.NotNil(ts.T(), row.UserID) @@ -39,13 +39,13 @@ func (ts *SCIMUsersTestSuite) linkedUser(id string) *models.User { return user } -func (ts *SCIMUsersTestSuite) identities(user *models.User) []*models.Identity { +func (ts *SCIMTestSuite) identities(user *models.User) []*models.Identity { identities, err := models.FindIdentitiesByUserID(ts.API.db, user.ID) require.NoError(ts.T(), err) return identities } -func (ts *SCIMUsersTestSuite) TestCreateProvisionsSSOUser() { +func (ts *SCIMTestSuite) TestCreateProvisionsSSOUser() { user := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) require.True(ts.T(), user.IsSSOUser) @@ -61,7 +61,7 @@ func (ts *SCIMUsersTestSuite) TestCreateProvisionsSSOUser() { require.Equal(ts.T(), "Alice@Example.com", identities[0].ProviderID) } -func (ts *SCIMUsersTestSuite) TestCreateDoesNotLinkOutsideProvider() { +func (ts *SCIMTestSuite) TestCreateDoesNotLinkOutsideProvider() { password, err := models.NewUser("", "alice@example.com", "", ts.API.config.JWT.Aud, nil) require.NoError(ts.T(), err) require.NoError(ts.T(), ts.API.db.Create(password)) @@ -75,7 +75,7 @@ func (ts *SCIMUsersTestSuite) TestCreateDoesNotLinkOutsideProvider() { require.Len(ts.T(), ts.identities(other), 1) } -func (ts *SCIMUsersTestSuite) TestCreateReusesSAMLIdentity() { +func (ts *SCIMTestSuite) TestCreateReusesSAMLIdentity() { existing := ts.ssoUser(ts.A, "Alice@Example.com", "alice@example.com") user := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) @@ -84,7 +84,7 @@ func (ts *SCIMUsersTestSuite) TestCreateReusesSAMLIdentity() { require.Len(ts.T(), ts.identities(user), 1) } -func (ts *SCIMUsersTestSuite) TestCreateLinksByEmailWithinProvider() { +func (ts *SCIMTestSuite) TestCreateLinksByEmailWithinProvider() { existing := ts.ssoUser(ts.A, "saml-name-id", "alice@example.com") user := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) @@ -93,7 +93,7 @@ func (ts *SCIMUsersTestSuite) TestCreateLinksByEmailWithinProvider() { require.Len(ts.T(), ts.identities(user), 2) } -func (ts *SCIMUsersTestSuite) passkeyRegistrationOptions(user *models.User) int { +func (ts *SCIMTestSuite) passkeyRegistrationOptions(user *models.User) int { passkey, webauthn := ts.API.config.Passkey, ts.API.config.WebAuthn defer func() { ts.API.config.Passkey, ts.API.config.WebAuthn = passkey, webauthn }() ts.API.config.Passkey.Enabled = true @@ -117,7 +117,7 @@ func (ts *SCIMUsersTestSuite) passkeyRegistrationOptions(user *models.User) int return w.Code } -func (ts *SCIMUsersTestSuite) TestPasswordUserWithSameEmailIsNeverLinked() { +func (ts *SCIMTestSuite) TestPasswordUserWithSameEmailIsNeverLinked() { password, err := models.NewUser("", "alice@example.com", "", ts.API.config.JWT.Aud, nil) require.NoError(ts.T(), err) require.NoError(ts.T(), ts.API.db.Create(password)) @@ -133,15 +133,15 @@ func (ts *SCIMUsersTestSuite) TestPasswordUserWithSameEmailIsNeverLinked() { require.Empty(ts.T(), ts.identities(reloaded)) } -func (ts *SCIMUsersTestSuite) TestNonSSOUserWithSSOIdentityEmailIsNeverLinked() { +func (ts *SCIMTestSuite) TestNonSSOUserWithSSOIdentityEmailIsNeverLinked() { ts.requireNonSSOUserNeverLinked("saml-name-id") } -func (ts *SCIMUsersTestSuite) TestNonSSOUserWithSSOIdentitySubjectIsNeverLinked() { +func (ts *SCIMTestSuite) TestNonSSOUserWithSSOIdentitySubjectIsNeverLinked() { ts.requireNonSSOUserNeverLinked("Alice@Example.com") } -func (ts *SCIMUsersTestSuite) requireNonSSOUserNeverLinked(sub string) { +func (ts *SCIMTestSuite) requireNonSSOUserNeverLinked(sub string) { password, err := models.NewUser("", "alice@example.com", "", ts.API.config.JWT.Aud, nil) require.NoError(ts.T(), err) require.NoError(ts.T(), ts.API.db.Create(password)) @@ -159,7 +159,7 @@ func (ts *SCIMUsersTestSuite) requireNonSSOUserNeverLinked(sub string) { require.Zero(ts.T(), ts.countRows(&models.SCIMUser{}, "user_id = ?", password.ID)) } -func (ts *SCIMUsersTestSuite) TestLinkAccountKeepsUserSSO() { +func (ts *SCIMTestSuite) TestLinkAccountKeepsUserSSO() { password, err := models.NewUser("", "alice@example.com", "", ts.API.config.JWT.Aud, nil) require.NoError(ts.T(), err) require.NoError(ts.T(), ts.API.db.Create(password)) @@ -174,7 +174,7 @@ func (ts *SCIMUsersTestSuite) TestLinkAccountKeepsUserSSO() { require.Empty(ts.T(), ts.identities(password)) } -func (ts *SCIMUsersTestSuite) TestOldEmailCannotSignInAfterEmailChange() { +func (ts *SCIMTestSuite) TestOldEmailCannotSignInAfterEmailChange() { id := ts.create(ts.TokenA, oktaUser) w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, ts.withEmail("alice.smith@example.com")) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) @@ -206,11 +206,11 @@ func (ts *SCIMUsersTestSuite) TestOldEmailCannotSignInAfterEmailChange() { } } -func (ts *SCIMUsersTestSuite) withEmail(email string) string { +func (ts *SCIMTestSuite) withEmail(email string) string { return strings.Replace(oktaUser, `"value": "alice@example.com"`, `"value": "`+email+`"`, 1) } -func (ts *SCIMUsersTestSuite) TestReplaceChangesEmail() { +func (ts *SCIMTestSuite) TestReplaceChangesEmail() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) @@ -229,7 +229,7 @@ func (ts *SCIMUsersTestSuite) TestReplaceChangesEmail() { require.Equal(ts.T(), user.ID, signedIn.ID) } -func (ts *SCIMUsersTestSuite) TestReplaceChangingEmailClearsPendingTokens() { +func (ts *SCIMTestSuite) TestReplaceChangingEmailClearsPendingTokens() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) user.RecoveryToken = "recovery-token-hash" @@ -243,7 +243,7 @@ func (ts *SCIMUsersTestSuite) TestReplaceChangingEmailClearsPendingTokens() { require.Zero(ts.T(), ts.countRows(&models.OneTimeToken{}, "user_id = ?", user.ID)) } -func (ts *SCIMUsersTestSuite) TestReplaceRenamesAndChangesEmail() { +func (ts *SCIMTestSuite) TestReplaceRenamesAndChangesEmail() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) @@ -262,7 +262,7 @@ func (ts *SCIMUsersTestSuite) TestReplaceRenamesAndChangesEmail() { require.Equal(ts.T(), user.ID, signedIn.ID) } -func (ts *SCIMUsersTestSuite) TestReplaceRejectsEmailTakenInProvider() { +func (ts *SCIMTestSuite) TestReplaceRejectsEmailTakenInProvider() { id := ts.create(ts.TokenA, oktaUser) ts.ssoUser(ts.A, "bob", "bob@example.com") @@ -272,7 +272,7 @@ func (ts *SCIMUsersTestSuite) TestReplaceRejectsEmailTakenInProvider() { require.Equal(ts.T(), "alice@example.com", ts.linkedUser(id).GetEmail()) } -func (ts *SCIMUsersTestSuite) TestReplaceAllowsEmailTakenInAnotherProvider() { +func (ts *SCIMTestSuite) TestReplaceAllowsEmailTakenInAnotherProvider() { id := ts.create(ts.TokenA, oktaUser) ts.ssoUser(ts.B, "bob", "bob@example.com") @@ -281,7 +281,7 @@ func (ts *SCIMUsersTestSuite) TestReplaceAllowsEmailTakenInAnotherProvider() { require.Equal(ts.T(), "bob@example.com", ts.linkedUser(id).GetEmail()) } -func (ts *SCIMUsersTestSuite) TestRemovingEmailsKeepsUserEmail() { +func (ts *SCIMTestSuite) TestRemovingEmailsKeepsUserEmail() { id := ts.create(ts.TokenA, strings.Replace(oktaUser, `"userName": "Alice@Example.com"`, `"userName": "alice.smith"`, 1)) user := ts.linkedUser(id) @@ -300,7 +300,7 @@ func (ts *SCIMUsersTestSuite) TestRemovingEmailsKeepsUserEmail() { } } -func (ts *SCIMUsersTestSuite) TestCreateInactiveLogsOutWithoutBanning() { +func (ts *SCIMTestSuite) TestCreateInactiveLogsOutWithoutBanning() { existing := ts.ssoUser(ts.A, "Alice@Example.com", "alice@example.com") ts.session(existing) @@ -310,7 +310,7 @@ func (ts *SCIMUsersTestSuite) TestCreateInactiveLogsOutWithoutBanning() { require.Zero(ts.T(), ts.sessions(user)) } -func (ts *SCIMUsersTestSuite) TestCreateRejectsSharedUser() { +func (ts *SCIMTestSuite) TestCreateRejectsSharedUser() { ts.create(ts.TokenA, oktaUser) w, body := ts.do(ts.TokenA, http.MethodPost, "/Users", strings.Replace(oktaUser, `"userName": "Alice@Example.com"`, `"userName": "alice.smith"`, 1)) @@ -319,13 +319,13 @@ func (ts *SCIMUsersTestSuite) TestCreateRejectsSharedUser() { require.EqualValues(ts.T(), 0, ts.list(ts.TokenA, `userName eq "alice.smith"`)["totalResults"]) } -func (ts *SCIMUsersTestSuite) TestCreateRequiresEmail() { +func (ts *SCIMTestSuite) TestCreateRequiresEmail() { w, body := ts.do(ts.TokenA, http.MethodPost, "/Users", `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"alice"}`) require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) require.Equal(ts.T(), "invalidValue", body["scimType"]) } -func (ts *SCIMUsersTestSuite) TestCreateFallsBackToEmailUserName() { +func (ts *SCIMTestSuite) TestCreateFallsBackToEmailUserName() { id := ts.create(ts.TokenA, `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"Alice@Example.com"}`) user := ts.linkedUser(id) require.Equal(ts.T(), "alice@example.com", user.GetEmail()) @@ -335,7 +335,7 @@ func (ts *SCIMUsersTestSuite) TestCreateFallsBackToEmailUserName() { require.Equal(ts.T(), user.ID, signedIn.ID) } -func (ts *SCIMUsersTestSuite) TestRejectsInvalidEmailsValue() { +func (ts *SCIMTestSuite) TestRejectsInvalidEmailsValue() { invalid := strings.Replace(oktaUser, `"value": "alice@example.com"`, `"value": "not-an-email"`, 1) w, body := ts.do(ts.TokenA, http.MethodPost, "/Users", invalid) require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) @@ -355,14 +355,14 @@ func (ts *SCIMUsersTestSuite) TestRejectsInvalidEmailsValue() { require.Equal(ts.T(), "alice@example.com", ts.linkedUser(id).GetEmail()) } -func (ts *SCIMUsersTestSuite) TestCreateKeepsAdminBan() { +func (ts *SCIMTestSuite) TestCreateKeepsAdminBan() { existing := ts.ssoUser(ts.A, "Alice@Example.com", "alice@example.com") require.NoError(ts.T(), existing.Ban(ts.API.db, time.Hour)) require.True(ts.T(), ts.linkedUser(ts.create(ts.TokenA, oktaUser)).IsBanned()) } -func (ts *SCIMUsersTestSuite) TestCreateLeavesNoUserOnConflict() { +func (ts *SCIMTestSuite) TestCreateLeavesNoUserOnConflict() { ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) w, _ := ts.do(ts.TokenA, http.MethodPost, "/Users", userWith("bob@example.com", "a-1")) @@ -372,19 +372,19 @@ func (ts *SCIMUsersTestSuite) TestCreateLeavesNoUserOnConflict() { require.Zero(ts.T(), count) } -func (ts *SCIMUsersTestSuite) session(user *models.User) { +func (ts *SCIMTestSuite) session(user *models.User) { session, err := models.NewSession(user.ID, nil) require.NoError(ts.T(), err) require.NoError(ts.T(), ts.API.db.Create(session)) } -func (ts *SCIMUsersTestSuite) refreshToken(user *models.User) string { +func (ts *SCIMTestSuite) refreshToken(user *models.User) string { token, err := models.GrantAuthenticatedUser(ts.API.db, user, models.GrantParams{}) require.NoError(ts.T(), err) return token.Token } -func (ts *SCIMUsersTestSuite) refresh(token string) int { +func (ts *SCIMTestSuite) refresh(token string) int { r := httptest.NewRequest(http.MethodPost, "/token?grant_type=refresh_token", strings.NewReader(`{"refresh_token":"`+token+`"}`)) r.Header.Set("Content-Type", "application/json") w := httptest.NewRecorder() @@ -392,13 +392,13 @@ func (ts *SCIMUsersTestSuite) refresh(token string) int { return w.Code } -func (ts *SCIMUsersTestSuite) sessions(user *models.User) int { +func (ts *SCIMTestSuite) sessions(user *models.User) int { count, err := ts.API.db.Q().Where("user_id = ?", user.ID).Count(&models.Session{}) require.NoError(ts.T(), err) return count } -func (ts *SCIMUsersTestSuite) samlLogin(ssoProvider *models.SSOProvider, sub, email string) (*models.User, error) { +func (ts *SCIMTestSuite) samlLogin(ssoProvider *models.SSOProvider, sub, email string) (*models.User, error) { userData := &provider.UserProvidedData{ Metadata: &provider.Claims{ Subject: sub, @@ -414,7 +414,7 @@ func (ts *SCIMUsersTestSuite) samlLogin(ssoProvider *models.SSOProvider, sub, em return ts.samlLoginWith(ssoProvider, userData) } -func (ts *SCIMUsersTestSuite) samlLoginWith(ssoProvider *models.SSOProvider, userData *provider.UserProvidedData) (*models.User, error) { +func (ts *SCIMTestSuite) samlLoginWith(ssoProvider *models.SSOProvider, userData *provider.UserProvidedData) (*models.User, error) { r := httptest.NewRequest(http.MethodPost, "/sso/saml/acs", nil) var user *models.User @@ -426,7 +426,7 @@ func (ts *SCIMUsersTestSuite) samlLoginWith(ssoProvider *models.SSOProvider, use return user, err } -func (ts *SCIMUsersTestSuite) TestSAMLLoginLocksVerifiedEmailWithoutMetadataEmail() { +func (ts *SCIMTestSuite) TestSAMLLoginLocksVerifiedEmailWithoutMetadataEmail() { userData := &provider.UserProvidedData{ Metadata: &provider.Claims{Subject: "saml-name-id", EmailVerified: true}, Emails: []provider.Email{{Email: "Alice@Example.com", Primary: true, Verified: true}}, @@ -452,7 +452,7 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginLocksVerifiedEmailWithoutMetadataEmai require.NoError(ts.T(), <-done) } -func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowedForActiveSCIMUser() { +func (ts *SCIMTestSuite) TestSAMLLoginAllowedForActiveSCIMUser() { id := ts.create(ts.TokenA, oktaUser) linked := ts.linkedUser(id) @@ -462,7 +462,7 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowedForActiveSCIMUser() { require.Equal(ts.T(), linked.ID, user.ID) } -func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowedForDeprovisionedUserWhileSCIMFlagOff() { +func (ts *SCIMTestSuite) TestSAMLLoginAllowedForDeprovisionedUserWhileSCIMFlagOff() { id := ts.create(ts.TokenA, oktaUser) linked := ts.linkedUser(id) ts.setActive(id, false) @@ -474,7 +474,7 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowedForDeprovisionedUserWhileSCIMF require.Equal(ts.T(), linked.ID, user.ID) } -func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedWhilePATCHedInactive() { +func (ts *SCIMTestSuite) TestSAMLLoginBlockedWhilePATCHedInactive() { id := ts.create(ts.TokenA, oktaUser) ts.setActive(id, false) @@ -489,7 +489,7 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedWhilePATCHedInactive() { require.Equal(ts.T(), linked.ID, user.ID) } -func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedAfterDelete() { +func (ts *SCIMTestSuite) TestSAMLLoginBlockedAfterDelete() { id := ts.create(ts.TokenA, oktaUser) w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") @@ -504,7 +504,7 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedAfterDelete() { require.Error(ts.T(), err) } -func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedWhenCreatedInactive() { +func (ts *SCIMTestSuite) TestSAMLLoginBlockedWhenCreatedInactive() { id := ts.create(ts.TokenA, strings.Replace(oktaUser, `"active": true`, `"active": false`, 1)) ts.linkedUser(id) @@ -512,7 +512,7 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedWhenCreatedInactive() { require.Error(ts.T(), err) } -func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowedWhenActiveOmitted() { +func (ts *SCIMTestSuite) TestSAMLLoginAllowedWhenActiveOmitted() { id := ts.create(ts.TokenA, `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"Alice@Example.com","emails":[{"primary":true,"value":"alice@example.com"}]}`) linked := ts.linkedUser(id) @@ -521,7 +521,7 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowedWhenActiveOmitted() { require.Equal(ts.T(), linked.ID, user.ID) } -func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowedWithoutSCIMRow() { +func (ts *SCIMTestSuite) TestSAMLLoginAllowedWithoutSCIMRow() { existing := ts.ssoUser(ts.A, "jit-user", "jit@example.com") user, err := ts.samlLogin(ts.A, "jit-user", "jit@example.com") @@ -530,7 +530,7 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowedWithoutSCIMRow() { require.Equal(ts.T(), existing.ID, user.ID) } -func (ts *SCIMUsersTestSuite) TestSAMLLoginLinksDivergedNameIDToSCIMUser() { +func (ts *SCIMTestSuite) TestSAMLLoginLinksDivergedNameIDToSCIMUser() { linked := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) for range 2 { @@ -551,7 +551,7 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginLinksDivergedNameIDToSCIMUser() { require.ElementsMatch(ts.T(), []string{"Alice@Example.com", "saml-name-id"}, providerIDs) } -func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedForDivergedNameIDWhileInactive() { +func (ts *SCIMTestSuite) TestSAMLLoginBlockedForDivergedNameIDWhileInactive() { id := ts.create(ts.TokenA, oktaUser) ts.setActive(id, false) linked := ts.linkedUser(id) @@ -561,13 +561,13 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginBlockedForDivergedNameIDWhileInactive require.Len(ts.T(), ts.identities(linked), 1) } -func (ts *SCIMUsersTestSuite) users(email string) int { +func (ts *SCIMTestSuite) users(email string) int { count, err := ts.API.db.Q().Where("email = ?", email).Count(&models.User{}) require.NoError(ts.T(), err) return count } -func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowsJITWithoutSCIMToken() { +func (ts *SCIMTestSuite) TestSAMLLoginAllowsJITWithoutSCIMToken() { provider := createSSOProvider(ts.T(), ts.API.db) user, err := ts.samlLogin(provider, "jit-user", "jit@example.com") @@ -576,7 +576,7 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowsJITWithoutSCIMToken() { require.Equal(ts.T(), "jit@example.com", user.GetEmail()) } -func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowsJITWhileSCIMFlagOff() { +func (ts *SCIMTestSuite) TestSAMLLoginAllowsJITWhileSCIMFlagOff() { ts.API.config.SSO.SCIM.Enabled = false defer func() { ts.API.config.SSO.SCIM.Enabled = true }() @@ -586,7 +586,7 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginAllowsJITWhileSCIMFlagOff() { require.Equal(ts.T(), "jit@example.com", user.GetEmail()) } -func (ts *SCIMUsersTestSuite) TestSAMLLoginNotBlockedByOtherProvider() { +func (ts *SCIMTestSuite) TestSAMLLoginNotBlockedByOtherProvider() { id := ts.create(ts.TokenB, oktaUser) w, _ := ts.do(ts.TokenB, http.MethodPatch, "/Users/"+id, `{ "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], @@ -601,20 +601,20 @@ func (ts *SCIMUsersTestSuite) TestSAMLLoginNotBlockedByOtherProvider() { require.Equal(ts.T(), existing.ID, user.ID) } -func (ts *SCIMUsersTestSuite) issueSession(conn *storage.Connection, user *models.User) error { +func (ts *SCIMTestSuite) issueSession(conn *storage.Connection, user *models.User) error { r := httptest.NewRequest(http.MethodPost, "/token", nil) _, err := ts.API.tokenService.IssueRefreshToken(r, http.Header{}, conn, user, models.OAuth, models.GrantParams{}) return err } -func (ts *SCIMUsersTestSuite) requireBanned(err error) { +func (ts *SCIMTestSuite) requireBanned(err error) { var httpErr *apierrors.HTTPError require.ErrorAs(ts.T(), err, &httpErr) require.Equal(ts.T(), http.StatusForbidden, httpErr.HTTPStatus) require.Equal(ts.T(), apierrors.ErrorCodeUserBanned, httpErr.ErrorCode) } -func (ts *SCIMUsersTestSuite) TestSessionRefusedWhileDeprovisioned() { +func (ts *SCIMTestSuite) TestSessionRefusedWhileDeprovisioned() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) require.NoError(ts.T(), ts.issueSession(ts.API.db, user)) @@ -635,7 +635,7 @@ func (ts *SCIMUsersTestSuite) TestSessionRefusedWhileDeprovisioned() { ts.requireBanned(ts.issueSession(ts.API.db, user)) } -func (ts *SCIMUsersTestSuite) TestSessionAllowedWhileSCIMFlagOff() { +func (ts *SCIMTestSuite) TestSessionAllowedWhileSCIMFlagOff() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) ts.setActive(id, false) @@ -645,7 +645,7 @@ func (ts *SCIMUsersTestSuite) TestSessionAllowedWhileSCIMFlagOff() { require.NoError(ts.T(), ts.issueSession(ts.API.db, user)) } -func (ts *SCIMUsersTestSuite) TestSessionRefusedForLinkedOAuthIdentityWhileDeprovisioned() { +func (ts *SCIMTestSuite) TestSessionRefusedForLinkedOAuthIdentityWhileDeprovisioned() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) identity, err := models.NewIdentity(user, "google", map[string]any{"sub": "google-sub", "email": "alice@example.com"}) @@ -669,7 +669,7 @@ func (ts *SCIMUsersTestSuite) TestSessionRefusedForLinkedOAuthIdentityWhileDepro require.Zero(ts.T(), ts.sessions(user)) } -func (ts *SCIMUsersTestSuite) TestSessionWaitsForConcurrentDeactivation() { +func (ts *SCIMTestSuite) TestSessionWaitsForConcurrentDeactivation() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) locked, release := make(chan struct{}), make(chan struct{}) @@ -703,7 +703,7 @@ func (ts *SCIMUsersTestSuite) TestSessionWaitsForConcurrentDeactivation() { require.Zero(ts.T(), ts.sessions(user)) } -func (ts *SCIMUsersTestSuite) TestWritesWaitForAdminUserDelete() { +func (ts *SCIMTestSuite) TestWritesWaitForAdminUserDelete() { for _, method := range []string{http.MethodDelete, http.MethodPut} { body := userWith(strings.ToLower(method)+"@example.com", method) id := ts.create(ts.TokenA, body) @@ -721,16 +721,16 @@ func (ts *SCIMUsersTestSuite) TestWritesWaitForAdminUserDelete() { } } -func (ts *SCIMUsersTestSuite) TestSessionAllowedForSSOUserWithoutSCIMRow() { +func (ts *SCIMTestSuite) TestSessionAllowedForSSOUserWithoutSCIMRow() { user := ts.ssoUser(ts.A, "saml-sub", "carol@example.com") require.NoError(ts.T(), ts.issueSession(ts.API.db, user)) } -func (ts *SCIMUsersTestSuite) setActive(id string, active bool) { +func (ts *SCIMTestSuite) setActive(id string, active bool) { ts.setActiveAs(ts.TokenA, id, active) } -func (ts *SCIMUsersTestSuite) setActiveAs(token, id string, active bool) { +func (ts *SCIMTestSuite) setActiveAs(token, id string, active bool) { w, _ := ts.do(token, http.MethodPatch, "/Users/"+id, `{ "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], "Operations": [{"op": "replace", "value": {"active": `+strconv.FormatBool(active)+`}}] @@ -738,7 +738,7 @@ func (ts *SCIMUsersTestSuite) setActiveAs(token, id string, active bool) { require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) } -func (ts *SCIMUsersTestSuite) TestReplaceDeactivatesAndReactivates() { +func (ts *SCIMTestSuite) TestReplaceDeactivatesAndReactivates() { id := ts.create(ts.TokenA, oktaUser) require.Equal(ts.T(), http.StatusOK, ts.refresh(ts.refreshToken(ts.linkedUser(id)))) refreshToken := ts.refreshToken(ts.linkedUser(id)) @@ -760,7 +760,7 @@ func (ts *SCIMUsersTestSuite) TestReplaceDeactivatesAndReactivates() { require.Equal(ts.T(), http.StatusBadRequest, ts.refresh(refreshToken)) } -func (ts *SCIMUsersTestSuite) TestPutInactiveRevokesSessions() { +func (ts *SCIMTestSuite) TestPutInactiveRevokesSessions() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) refreshToken := ts.refreshToken(user) @@ -773,7 +773,7 @@ func (ts *SCIMUsersTestSuite) TestPutInactiveRevokesSessions() { require.Len(ts.T(), ts.auditActions(models.SCIMUserDeactivatedAction), 1) } -func (ts *SCIMUsersTestSuite) TestReplaceKeepsAdminBanWhenActiveDoesNotChange() { +func (ts *SCIMTestSuite) TestReplaceKeepsAdminBanWhenActiveDoesNotChange() { id := ts.create(ts.TokenA, oktaUser) require.NoError(ts.T(), ts.linkedUser(id).Ban(ts.API.db, time.Hour)) @@ -781,7 +781,7 @@ func (ts *SCIMUsersTestSuite) TestReplaceKeepsAdminBanWhenActiveDoesNotChange() require.True(ts.T(), ts.linkedUser(id).IsBanned()) } -func (ts *SCIMUsersTestSuite) TestReplaceLinksUnlinkedRow() { +func (ts *SCIMTestSuite) TestReplaceLinksUnlinkedRow() { row, err := models.CreateSCIMUser(ts.API.db, ts.A.ID, []byte(`{"userName":"Alice@Example.com"}`)) require.NoError(ts.T(), err) @@ -793,7 +793,7 @@ func (ts *SCIMUsersTestSuite) TestReplaceLinksUnlinkedRow() { require.False(ts.T(), user.IsBanned()) } -func (ts *SCIMUsersTestSuite) TestDeleteLogsOutWithoutBanning() { +func (ts *SCIMTestSuite) TestDeleteLogsOutWithoutBanning() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) refreshToken := ts.refreshToken(user) @@ -808,7 +808,7 @@ func (ts *SCIMUsersTestSuite) TestDeleteLogsOutWithoutBanning() { require.Equal(ts.T(), http.StatusBadRequest, ts.refresh(refreshToken)) } -func (ts *SCIMUsersTestSuite) TestCreateRefusesUserDeletedByProvider() { +func (ts *SCIMTestSuite) TestCreateRefusesUserDeletedByProvider() { for _, body := range []string{oktaUser, strings.Replace(oktaUser, `"userName": "Alice@Example.com"`, `"userName": "alice.new@example.com"`, 1)} { ts.SetupTest() id := ts.create(ts.TokenA, oktaUser) @@ -827,7 +827,7 @@ func (ts *SCIMUsersTestSuite) TestCreateRefusesUserDeletedByProvider() { } } -func (ts *SCIMUsersTestSuite) TestCreateAfterAdminDeletesProviderDeletedUser() { +func (ts *SCIMTestSuite) TestCreateAfterAdminDeletesProviderDeletedUser() { for _, soft := range []bool{false, true} { ts.SetupTest() id := ts.create(ts.TokenA, oktaUser) @@ -848,7 +848,7 @@ func (ts *SCIMUsersTestSuite) TestCreateAfterAdminDeletesProviderDeletedUser() { } } -func (ts *SCIMUsersTestSuite) TestCreateConcurrentSameEmailLinksToOneUser() { +func (ts *SCIMTestSuite) TestCreateConcurrentSameEmailLinksToOneUser() { body := func(userName, externalID string) string { return `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"` + userName + `","externalId":"` + externalID + `","emails":[{"primary":true,"value":"race@example.com"}]}` } @@ -888,12 +888,12 @@ func (ts *SCIMUsersTestSuite) TestCreateConcurrentSameEmailLinksToOneUser() { require.Equal(ts.T(), 1, count) } -func (ts *SCIMUsersTestSuite) rename(id, userName string) (int, string) { +func (ts *SCIMTestSuite) rename(id, userName string) (int, string) { w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, strings.Replace(oktaUser, `"userName": "Alice@Example.com"`, `"userName": "`+userName+`"`, 1)) return w.Code, w.Body.String() } -func (ts *SCIMUsersTestSuite) TestReplaceRenamesSSOIdentity() { +func (ts *SCIMTestSuite) TestReplaceRenamesSSOIdentity() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) @@ -907,7 +907,7 @@ func (ts *SCIMUsersTestSuite) TestReplaceRenamesSSOIdentity() { require.Len(ts.T(), ts.identities(user), 1) } -func (ts *SCIMUsersTestSuite) providerIDs(user *models.User) []string { +func (ts *SCIMTestSuite) providerIDs(user *models.User) []string { ids := []string{} for _, identity := range ts.identities(user) { ids = append(ids, identity.ProviderID) @@ -915,7 +915,7 @@ func (ts *SCIMUsersTestSuite) providerIDs(user *models.User) []string { return ids } -func (ts *SCIMUsersTestSuite) TestReplaceRenamesCaseOnly() { +func (ts *SCIMTestSuite) TestReplaceRenamesCaseOnly() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) _, err := ts.samlLogin(ts.A, "alice@example.com", "alice@example.com") @@ -933,7 +933,7 @@ func (ts *SCIMUsersTestSuite) TestReplaceRenamesCaseOnly() { require.Equal(ts.T(), user.ID, signedIn.ID) } -func (ts *SCIMUsersTestSuite) TestReplaceRenameRemovesOldNameIDIdentity() { +func (ts *SCIMTestSuite) TestReplaceRenameRemovesOldNameIDIdentity() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) _, err := ts.samlLogin(ts.A, "alice@example.com", "alice@example.com") @@ -950,7 +950,7 @@ func (ts *SCIMUsersTestSuite) TestReplaceRenameRemovesOldNameIDIdentity() { require.NotEqual(ts.T(), user.ID, signedIn.ID) } -func (ts *SCIMUsersTestSuite) TestReplaceRejectsRenameToTakenIdentity() { +func (ts *SCIMTestSuite) TestReplaceRejectsRenameToTakenIdentity() { id := ts.create(ts.TokenA, oktaUser) ts.ssoUser(ts.A, "bob@example.com", "bob@example.com") @@ -964,7 +964,7 @@ func (ts *SCIMUsersTestSuite) TestReplaceRejectsRenameToTakenIdentity() { require.NoError(ts.T(), err) } -func (ts *SCIMUsersTestSuite) unlink(user *models.User, identity *models.Identity) *httptest.ResponseRecorder { +func (ts *SCIMTestSuite) unlink(user *models.User, identity *models.Identity) *httptest.ResponseRecorder { session, err := models.NewSession(user.ID, nil) require.NoError(ts.T(), err) require.NoError(ts.T(), ts.API.db.Create(session)) @@ -977,7 +977,7 @@ func (ts *SCIMUsersTestSuite) unlink(user *models.User, identity *models.Identit return w } -func (ts *SCIMUsersTestSuite) TestUnlinkRefusedForSCIMManagedIdentity() { +func (ts *SCIMTestSuite) TestUnlinkRefusedForSCIMManagedIdentity() { ts.API.config.Security.ManualLinkingEnabled = true defer func() { ts.API.config.Security.ManualLinkingEnabled = false }() id := ts.create(ts.TokenA, oktaUser) @@ -1005,7 +1005,7 @@ func (ts *SCIMUsersTestSuite) TestUnlinkRefusedForSCIMManagedIdentity() { require.Len(ts.T(), ts.identities(user), 1) } -func (ts *SCIMUsersTestSuite) TestUnlinkAllowedWhileSCIMFlagOff() { +func (ts *SCIMTestSuite) TestUnlinkAllowedWhileSCIMFlagOff() { ts.API.config.Security.ManualLinkingEnabled = true defer func() { ts.API.config.Security.ManualLinkingEnabled = false }() user := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) @@ -1022,7 +1022,7 @@ func (ts *SCIMUsersTestSuite) TestUnlinkAllowedWhileSCIMFlagOff() { require.Len(ts.T(), ts.identities(user), 1) } -func (ts *SCIMUsersTestSuite) TestRenameSkippedWhenIdentityMissing() { +func (ts *SCIMTestSuite) TestRenameSkippedWhenIdentityMissing() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) sso, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", "sso:"+ts.A.ID.String()) diff --git a/internal/api/scim_okta_spec_test.go b/internal/api/scim_okta_spec_test.go index c36f30870c..7ab60f2a80 100644 --- a/internal/api/scim_okta_spec_test.go +++ b/internal/api/scim_okta_spec_test.go @@ -14,7 +14,7 @@ import ( "github.com/supabase/auth/internal/models" ) -func (ts *SCIMUsersTestSuite) okta(method, path, body string, headers ...string) (int, map[string]any) { +func (ts *SCIMTestSuite) okta(method, path, body string, headers ...string) (int, map[string]any) { contentType := "application/scim+json; charset=utf-8" if method == http.MethodPost { contentType = "application/json" @@ -37,7 +37,7 @@ type replayRequest struct { Body json.RawMessage `json:"body"` } -func (ts *SCIMUsersTestSuite) replay(file, created string, ids []string, onRequest func(step string, request replayRequest, got map[string]any, id string), afterStep func(step, id string)) int { +func (ts *SCIMTestSuite) replay(file, created string, ids []string, onRequest func(step string, request replayRequest, got map[string]any, id string), afterStep func(step, id string)) int { raw, err := fs.ReadFile(os.DirFS("testdata/scim"), file) require.NoError(ts.T(), err) var steps []struct { @@ -69,7 +69,7 @@ func oktaFilter(userName string) string { return "/Users?" + url.Values{"filter": {`userName eq "` + userName + `"`}}.Encode() } -func (ts *SCIMUsersTestSuite) TestOktaSpec() { +func (ts *SCIMTestSuite) TestOktaSpec() { const ( userName = "okta.spec.user@example.com" givenName = "Okta" @@ -158,7 +158,7 @@ func (ts *SCIMUsersTestSuite) TestOktaSpec() { requireError(got, "404") } -func (ts *SCIMUsersTestSuite) TestOktaUserLifecycleReplay() { +func (ts *SCIMTestSuite) TestOktaUserLifecycleReplay() { const password = "okta-generated-password" type state struct { diff --git a/internal/api/scim_provider_delete_test.go b/internal/api/scim_provider_delete_test.go index c428b05d37..316e044693 100644 --- a/internal/api/scim_provider_delete_test.go +++ b/internal/api/scim_provider_delete_test.go @@ -16,7 +16,7 @@ func scimUser(name string) string { return userWith(name+"@example.com", name) } -func (ts *SCIMUsersTestSuite) deleteProvider(p *models.SSOProvider) { +func (ts *SCIMTestSuite) deleteProvider(p *models.SSOProvider) { token := adminJWT(ts.T(), ts.API.config.JWT.Secret) r := httptest.NewRequest(http.MethodDelete, "/admin/sso/providers/"+p.ID.String(), nil) r.Header.Set("Authorization", "Bearer "+token) @@ -25,23 +25,23 @@ func (ts *SCIMUsersTestSuite) deleteProvider(p *models.SSOProvider) { require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) } -func (ts *SCIMUsersTestSuite) reloadUser(id uuid.UUID) *models.User { +func (ts *SCIMTestSuite) reloadUser(id uuid.UUID) *models.User { user, err := models.FindUserByID(ts.API.db, id) require.NoError(ts.T(), err) return user } -func (ts *SCIMUsersTestSuite) countRows(model any, where string, args ...any) int { +func (ts *SCIMTestSuite) countRows(model any, where string, args ...any) int { count, err := ts.API.db.Q().Where(where, args...).Count(model) require.NoError(ts.T(), err) return count } -func (ts *SCIMUsersTestSuite) auditActions(action models.AuditAction) []models.AuditLogEntry { +func (ts *SCIMTestSuite) auditActions(action models.AuditAction) []models.AuditLogEntry { return queryAuditEntries(ts.T(), ts.API.db, "payload->>'action' = ?", string(action)) } -func (ts *SCIMUsersTestSuite) TestProviderDeleteBansDeprovisionedUsers() { +func (ts *SCIMTestSuite) TestProviderDeleteBansDeprovisionedUsers() { active := ts.linkedUser(ts.create(ts.TokenA, scimUser("active"))) deactivatedID := ts.create(ts.TokenA, scimUser("deactivated")) @@ -83,7 +83,7 @@ func (ts *SCIMUsersTestSuite) TestProviderDeleteBansDeprovisionedUsers() { require.Equal(ts.T(), http.StatusOK, w.Code) } -func (ts *SCIMUsersTestSuite) TestProviderDeleteKeepsLongerBan() { +func (ts *SCIMTestSuite) TestProviderDeleteKeepsLongerBan() { id := ts.create(ts.TokenA, scimUser("banned")) user := ts.linkedUser(id) ts.setActive(id, false) @@ -96,7 +96,7 @@ func (ts *SCIMUsersTestSuite) TestProviderDeleteKeepsLongerBan() { require.Empty(ts.T(), ts.auditActions(models.SCIMUsersBannedAction)) } -func (ts *SCIMUsersTestSuite) TestProviderDeleteClosesOAuthBypass() { +func (ts *SCIMTestSuite) TestProviderDeleteClosesOAuthBypass() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) identity, err := models.NewIdentity(user, "google", map[string]any{"sub": "google-sub", "email": "alice@example.com"}) @@ -122,7 +122,7 @@ func (ts *SCIMUsersTestSuite) TestProviderDeleteClosesOAuthBypass() { require.Zero(ts.T(), ts.sessions(user)) } -func (ts *SCIMUsersTestSuite) TestProviderDeleteAudit() { +func (ts *SCIMTestSuite) TestProviderDeleteAudit() { ts.setActive(ts.create(ts.TokenA, scimUser("audited")), false) tokens, err := models.FindActiveSCIMTokensBySSOProvider(ts.API.db, ts.A.ID) require.NoError(ts.T(), err) @@ -149,7 +149,7 @@ func (ts *SCIMUsersTestSuite) TestProviderDeleteAudit() { require.Equal(ts.T(), ts.A.ID.String(), traits["sso_provider_id"]) } -func (ts *SCIMUsersTestSuite) TestProviderDeleteWritesNoGroupEvents() { +func (ts *SCIMTestSuite) TestProviderDeleteWritesNoGroupEvents() { alice := ts.create(ts.TokenA, scimUser("alice")) ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) before := ts.countRows(&models.AuditLogEntry{}, "payload->>'action' LIKE 'scim_group_%'") @@ -160,7 +160,7 @@ func (ts *SCIMUsersTestSuite) TestProviderDeleteWritesNoGroupEvents() { require.Zero(ts.T(), ts.countRows(&models.SCIMGroup{}, "sso_provider_id = ?", ts.A.ID)) } -func (ts *SCIMUsersTestSuite) TestProviderDeleteAuditWithExpiredTokens() { +func (ts *SCIMTestSuite) TestProviderDeleteAuditWithExpiredTokens() { ts.setActive(ts.create(ts.TokenA, scimUser("expired")), false) require.NoError(ts.T(), ts.API.db.RawQuery( "UPDATE "+(&models.SCIMToken{}).TableName()+" SET created_at = now() - interval '2 hours', expires_at = now() - interval '1 hour' WHERE sso_provider_id = ?", ts.A.ID, @@ -175,7 +175,7 @@ func (ts *SCIMUsersTestSuite) TestProviderDeleteAuditWithExpiredTokens() { require.Len(ts.T(), ts.auditActions(models.SCIMUsersBannedAction), 1) } -func (ts *SCIMUsersTestSuite) TestProviderDeleteWithoutSCIMEnabled() { +func (ts *SCIMTestSuite) TestProviderDeleteWithoutSCIMEnabled() { provider := createSSOProvider(ts.T(), ts.API.db) token, _, err := models.CreateSCIMToken(ts.API.db, provider, nil) require.NoError(ts.T(), err) @@ -188,7 +188,7 @@ func (ts *SCIMUsersTestSuite) TestProviderDeleteWithoutSCIMEnabled() { require.Equal(ts.T(), token.Prefix, revoked[0].Payload["traits"].(map[string]any)["token_prefix"]) } -func (ts *SCIMUsersTestSuite) TestProviderDeleteAfterSCIMDisabled() { +func (ts *SCIMTestSuite) TestProviderDeleteAfterSCIMDisabled() { id := ts.create(ts.TokenA, scimUser("disabled")) user := ts.linkedUser(id) ts.setActive(id, false) @@ -202,7 +202,7 @@ func (ts *SCIMUsersTestSuite) TestProviderDeleteAfterSCIMDisabled() { require.True(ts.T(), ts.reloadUser(user.ID).IsBanned()) } -func (ts *SCIMUsersTestSuite) TestProviderDeleteStillBansWhileSCIMFlagOff() { +func (ts *SCIMTestSuite) TestProviderDeleteStillBansWhileSCIMFlagOff() { id := ts.create(ts.TokenA, scimUser("flagoff")) user := ts.linkedUser(id) ts.setActive(id, false) diff --git a/internal/api/scim_test.go b/internal/api/scim_test.go index 6ee9215c25..7c29346259 100644 --- a/internal/api/scim_test.go +++ b/internal/api/scim_test.go @@ -463,7 +463,7 @@ func TestSCIMServer(t *testing.T) { }) } -type SCIMUsersTestSuite struct { +type SCIMTestSuite struct { suite.Suite API *API TokenA string @@ -472,30 +472,30 @@ type SCIMUsersTestSuite struct { B *models.SSOProvider } -func TestSCIMUsers(t *testing.T) { +func TestSCIMSuite(t *testing.T) { api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { config.RateLimitScim = 1_000_000 }) defer api.db.Close() - suite.Run(t, &SCIMUsersTestSuite{API: api}) + suite.Run(t, &SCIMTestSuite{API: api}) } -func (ts *SCIMUsersTestSuite) SetupTest() { +func (ts *SCIMTestSuite) SetupTest() { require.NoError(ts.T(), models.TruncateAll(ts.API.db)) ts.A, ts.TokenA = ts.provider() ts.B, ts.TokenB = ts.provider() } -func (ts *SCIMUsersTestSuite) provider() (*models.SSOProvider, string) { +func (ts *SCIMTestSuite) provider() (*models.SSOProvider, string) { return createSSOProviderWithSCIMToken(ts.T(), ts.API.db) } -func (ts *SCIMUsersTestSuite) do(token, method, path, body string) (*httptest.ResponseRecorder, map[string]any) { +func (ts *SCIMTestSuite) do(token, method, path, body string) (*httptest.ResponseRecorder, map[string]any) { return ts.doAs(protocol.MediaType, token, method, path, body) } -func (ts *SCIMUsersTestSuite) doAs(contentType, token, method, path, body string, headers ...string) (*httptest.ResponseRecorder, map[string]any) { +func (ts *SCIMTestSuite) doAs(contentType, token, method, path, body string, headers ...string) (*httptest.ResponseRecorder, map[string]any) { w := ts.serve(contentType, token, method, path, body, headers...) var decoded map[string]any @@ -505,7 +505,7 @@ func (ts *SCIMUsersTestSuite) doAs(contentType, token, method, path, body string return w, decoded } -func (ts *SCIMUsersTestSuite) serve(contentType, token, method, path, body string, headers ...string) *httptest.ResponseRecorder { +func (ts *SCIMTestSuite) serve(contentType, token, method, path, body string, headers ...string) *httptest.ResponseRecorder { r := httptest.NewRequest(method, "/scim/v2"+path, strings.NewReader(body)) r.Header.Set("Authorization", "Bearer "+token) r.Header.Set("Content-Type", contentType) diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index 95d6e54d0e..ccce1aafc6 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -38,13 +38,13 @@ const oktaUser = `{ "active": true }` -func (ts *SCIMUsersTestSuite) create(token, body string) string { +func (ts *SCIMTestSuite) create(token, body string) string { w, created := ts.do(token, http.MethodPost, "/Users", body) require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) return created["id"].(string) } -func (ts *SCIMUsersTestSuite) list(token, filter string) map[string]any { +func (ts *SCIMTestSuite) list(token, filter string) map[string]any { w, body := ts.do(token, http.MethodGet, "/Users?"+url.Values{"filter": {filter}, "startIndex": {"1"}, "count": {"100"}}.Encode(), "") require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) return body @@ -54,7 +54,7 @@ func userWith(userName, externalID string) string { return `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"` + userName + `","externalId":"` + externalID + `","emails":[{"primary":true,"value":"` + userName + `"}]}` } -func (ts *SCIMUsersTestSuite) repository() (context.Context, server.Repository[*core.User]) { +func (ts *SCIMTestSuite) repository() (context.Context, server.Repository[*core.User]) { ctx, err := newSCIMTokenValidator(ts.API.db)(context.Background(), ts.TokenA) require.NoError(ts.T(), err) ctx = scimRequestKey.WithValue(ctx, httptest.NewRequest(http.MethodPost, "/scim/v2/Users", nil)) @@ -65,7 +65,7 @@ func emails(value string) []core.Email { return []core.Email{{Value: value, Primary: new(true)}} } -func (ts *SCIMUsersTestSuite) TestOktaLifecycle() { +func (ts *SCIMTestSuite) TestOktaLifecycle() { require.EqualValues(ts.T(), 0, ts.list(ts.TokenA, `userName eq "alice@example.com"`)["totalResults"]) w, created := ts.do(ts.TokenA, http.MethodPost, "/Users", oktaUser) @@ -134,7 +134,7 @@ func (ts *SCIMUsersTestSuite) TestOktaLifecycle() { require.NotNil(ts.T(), stored.DeletedAt) } -func (ts *SCIMUsersTestSuite) TestOktaContentTypesAndReactivate() { +func (ts *SCIMTestSuite) TestOktaContentTypesAndReactivate() { for i, contentType := range []string{"application/scim+json; charset=utf-8", "application/json", "application/json; charset=utf-8"} { name := string(rune('a'+i)) + "@example.com" w, created := ts.doAs(contentType, ts.TokenA, http.MethodPost, "/Users", userWith(name, name)) @@ -152,7 +152,7 @@ func (ts *SCIMUsersTestSuite) TestOktaContentTypesAndReactivate() { } } -func (ts *SCIMUsersTestSuite) TestUnsupportedEndpointsReturnNotImplemented() { +func (ts *SCIMTestSuite) TestUnsupportedEndpointsReturnNotImplemented() { for _, tc := range []struct{ method, path string }{ {http.MethodGet, "/Me"}, {http.MethodPost, "/Bulk"}, @@ -167,7 +167,7 @@ func (ts *SCIMUsersTestSuite) TestUnsupportedEndpointsReturnNotImplemented() { } } -func (ts *SCIMUsersTestSuite) TestUniquenessWithinProvider() { +func (ts *SCIMTestSuite) TestUniquenessWithinProvider() { ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) for _, body := range []string{userWith("ALICE@example.com", "a-2"), userWith("bob@example.com", "a-1")} { @@ -179,7 +179,7 @@ func (ts *SCIMUsersTestSuite) TestUniquenessWithinProvider() { ts.create(ts.TokenB, userWith("alice@example.com", "a-1")) } -func (ts *SCIMUsersTestSuite) TestUniqueIndexIsTheBackstop() { +func (ts *SCIMTestSuite) TestUniqueIndexIsTheBackstop() { ctx, users := ts.repository() _, err := users.Create(ctx, &core.User{UserName: "alice@example.com", Emails: emails("alice@example.com")}) @@ -191,7 +191,7 @@ func (ts *SCIMUsersTestSuite) TestUniqueIndexIsTheBackstop() { require.Equal(ts.T(), http.StatusConflict, scimErr.StatusCode()) } -func (ts *SCIMUsersTestSuite) TestReplaceRejectsStaleVersion() { +func (ts *SCIMTestSuite) TestReplaceRejectsStaleVersion() { ctx, users := ts.repository() created, err := users.Create(ctx, &core.User{UserName: "alice@example.com", Emails: emails("alice@example.com")}) @@ -230,7 +230,7 @@ func (ts *SCIMUsersTestSuite) TestReplaceRejectsStaleVersion() { require.Contains(ts.T(), string(stored.Resource), "winner") } -func (ts *SCIMUsersTestSuite) TestDeleteRejectsStaleVersion() { +func (ts *SCIMTestSuite) TestDeleteRejectsStaleVersion() { ctx, users := ts.repository() created, err := users.Create(ctx, &core.User{UserName: "alice@example.com", Emails: emails("alice@example.com")}) @@ -264,7 +264,7 @@ func (ts *SCIMUsersTestSuite) TestDeleteRejectsStaleVersion() { require.NotNil(ts.T(), stored.DeletedAt) } -func (ts *SCIMUsersTestSuite) whileLocked(lock, finish func(tx *storage.Connection) error, method, path, body string) (int, error) { +func (ts *SCIMTestSuite) whileLocked(lock, finish func(tx *storage.Connection) error, method, path, body string) (int, error) { locked, release := make(chan struct{}), make(chan struct{}) held := make(chan error, 1) go func() { @@ -293,11 +293,11 @@ func (ts *SCIMUsersTestSuite) whileLocked(lock, finish func(tx *storage.Connecti return <-code, <-held } -func (ts *SCIMUsersTestSuite) scimAuditEntries() []models.AuditLogEntry { +func (ts *SCIMTestSuite) scimAuditEntries() []models.AuditLogEntry { return queryAuditEntries(ts.T(), ts.API.db, "payload->>'log_type' = ?", "scim") } -func (ts *SCIMUsersTestSuite) TestAuditLog() { +func (ts *SCIMTestSuite) TestAuditLog() { id := ts.create(ts.TokenA, oktaUser) w, _ := ts.do(ts.TokenA, http.MethodPost, "/Users", oktaUser) @@ -340,7 +340,7 @@ func (ts *SCIMUsersTestSuite) TestAuditLog() { }, actions) } -func (ts *SCIMUsersTestSuite) TestRolesRoundTrip() { +func (ts *SCIMTestSuite) TestRolesRoundTrip() { body := `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"alice@example.com","emails":[{"primary":true,"value":"alice@example.com"}],"roles":[{"value":"admin","primary":true},{"value":"billing"}]}` id := ts.create(ts.TokenA, body) @@ -353,7 +353,7 @@ func (ts *SCIMUsersTestSuite) TestRolesRoundTrip() { require.Equal(ts.T(), []string{"admin", "billing"}, roles) } -func (ts *SCIMUsersTestSuite) TestConcurrentCreateWithinProvider() { +func (ts *SCIMTestSuite) TestConcurrentCreateWithinProvider() { const attempts = 8 codes := make(chan int, attempts) start := make(chan struct{}) @@ -381,7 +381,7 @@ func (ts *SCIMUsersTestSuite) TestConcurrentCreateWithinProvider() { require.EqualValues(ts.T(), 1, ts.list(ts.TokenA, `userName eq "race@example.com"`)["totalResults"]) } -func (ts *SCIMUsersTestSuite) TestSort() { +func (ts *SCIMTestSuite) TestSort() { ids := map[string]string{} for _, name := range []string{"carol@example.com", "Alice@example.com", "bob@example.com"} { ids[name] = ts.create(ts.TokenA, userWith(name, name)) @@ -426,7 +426,7 @@ func (ts *SCIMUsersTestSuite) TestSort() { } } -func (ts *SCIMUsersTestSuite) TestAttributeProjection() { +func (ts *SCIMTestSuite) TestAttributeProjection() { id := ts.create(ts.TokenA, oktaUser) for _, path := range []string{"/Users/" + id, "/Users"} { @@ -460,7 +460,7 @@ func (ts *SCIMUsersTestSuite) TestAttributeProjection() { } } -func (ts *SCIMUsersTestSuite) TestWriteResponseProjection() { +func (ts *SCIMTestSuite) TestWriteResponseProjection() { w, created := ts.do(ts.TokenA, http.MethodPost, "/Users?attributes=userName", userWith("alice@example.com", "a-1")) require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) id := created["id"].(string) @@ -486,7 +486,7 @@ func (ts *SCIMUsersTestSuite) TestWriteResponseProjection() { require.Equal(ts.T(), []string{id}, memberValues(got)) } -func (ts *SCIMUsersTestSuite) TestETagAndIfMatch() { +func (ts *SCIMTestSuite) TestETagAndIfMatch() { w, created := ts.do(ts.TokenA, http.MethodPost, "/Users", oktaUser) require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) id := created["id"].(string) @@ -522,7 +522,7 @@ func (ts *SCIMUsersTestSuite) TestETagAndIfMatch() { require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) } -func (ts *SCIMUsersTestSuite) TestPatchAttributesOutsideTheMinimalSchema() { +func (ts *SCIMTestSuite) TestPatchAttributesOutsideTheMinimalSchema() { id := ts.create(ts.TokenA, oktaUser) patch := `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ @@ -542,7 +542,7 @@ func (ts *SCIMUsersTestSuite) TestPatchAttributesOutsideTheMinimalSchema() { require.ElementsMatch(ts.T(), []any{string(core.SchemaUser), string(core.SchemaEnterpriseUser)}, read["schemas"]) } -func (ts *SCIMUsersTestSuite) TestActiveDefaultsToTrue() { +func (ts *SCIMTestSuite) TestActiveDefaultsToTrue() { w, created := ts.do(ts.TokenA, http.MethodPost, "/Users", userWith("alice@example.com", "a-1")) require.Equal(ts.T(), http.StatusCreated, w.Code) require.Equal(ts.T(), true, created["active"]) @@ -552,7 +552,7 @@ func (ts *SCIMUsersTestSuite) TestActiveDefaultsToTrue() { require.Equal(ts.T(), true, replaced["active"]) } -func (ts *SCIMUsersTestSuite) TestPagination() { +func (ts *SCIMTestSuite) TestPagination() { for _, name := range []string{"a", "b", "c"} { ts.create(ts.TokenA, userWith(name+"@example.com", name)) } @@ -575,7 +575,7 @@ func (ts *SCIMUsersTestSuite) TestPagination() { require.Empty(ts.T(), page["Resources"]) } -func (ts *SCIMUsersTestSuite) TestPageSizeCap() { +func (ts *SCIMTestSuite) TestPageSizeCap() { for i := range 101 { _, err := models.CreateSCIMUser(ts.API.db, ts.A.ID, []byte(`{"userName":"user`+strconv.Itoa(i)+`@example.com"}`)) require.NoError(ts.T(), err) @@ -592,7 +592,7 @@ func (ts *SCIMUsersTestSuite) TestPageSizeCap() { } } -func (ts *SCIMUsersTestSuite) TestSortTieBreaksOnID() { +func (ts *SCIMTestSuite) TestSortTieBreaksOnID() { ids := []string{ ts.create(ts.TokenA, userWith("alice@example.com", "a-1")), ts.create(ts.TokenA, userWith("bob@example.com", "b-1")), @@ -619,7 +619,7 @@ func (ts *SCIMUsersTestSuite) TestSortTieBreaksOnID() { } } -func (ts *SCIMUsersTestSuite) TestUnsupportedFilters() { +func (ts *SCIMTestSuite) TestUnsupportedFilters() { for _, filter := range []string{ `userName co "alice"`, `userName ne "alice"`, @@ -636,14 +636,14 @@ func (ts *SCIMUsersTestSuite) TestUnsupportedFilters() { } } -func (ts *SCIMUsersTestSuite) TestUnknownID() { +func (ts *SCIMTestSuite) TestUnknownID() { for _, id := range []string{"not-a-uuid", "00000000-0000-0000-0000-000000000000"} { w, _ := ts.do(ts.TokenA, http.MethodGet, "/Users/"+id, "") require.Equal(ts.T(), http.StatusNotFound, w.Code, id) } } -func (ts *SCIMUsersTestSuite) TestRequiresSSOProviderOnContext() { +func (ts *SCIMTestSuite) TestRequiresSSOProviderOnContext() { users := &scimUserRepository{api: ts.API} _, _, err := users.List(context.Background(), &protocol.SearchRequest{Count: 10}) @@ -685,7 +685,7 @@ func TestSCIMUserFields(t *testing.T) { }) } -func (ts *SCIMUsersTestSuite) TestUsersGroupsAttribute() { +func (ts *SCIMTestSuite) TestUsersGroupsAttribute() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) ops := ts.createGroup(ts.TokenA, groupWith("Ops", "g-2", alice)) From cb776b107462b3cee58ccfa85dcacdecef06ebb3 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 01:06:39 -0600 Subject: [PATCH 61/88] chore(scim): drop duplicate scim-go import aliases in SCIM tests --- internal/api/scim_test.go | 61 ++++++++++++++++----------------- internal/api/scim_users_test.go | 9 +++-- 2 files changed, 34 insertions(+), 36 deletions(-) diff --git a/internal/api/scim_test.go b/internal/api/scim_test.go index 7c29346259..fccafb4a2e 100644 --- a/internal/api/scim_test.go +++ b/internal/api/scim_test.go @@ -18,9 +18,8 @@ import ( logrustest "github.com/sirupsen/logrus/hooks/test" "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" - scimCore "github.com/supabase-community/scim-go/pkg/core" + "github.com/supabase-community/scim-go/pkg/core" "github.com/supabase-community/scim-go/pkg/protocol" - scimProtocol "github.com/supabase-community/scim-go/pkg/protocol" "github.com/supabase-community/scim-go/pkg/server" "github.com/supabase/auth/internal/conf" "github.com/supabase/auth/internal/models" @@ -66,7 +65,7 @@ func TestSCIM(t *testing.T) { api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusNotFound, w.Code) - require.NotContains(t, w.Body.String(), scimProtocol.SchemaError) + require.NotContains(t, w.Body.String(), protocol.SchemaError) }) }) @@ -92,8 +91,8 @@ func TestSCIM(t *testing.T) { api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusOK, w.Code) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) - require.Contains(t, w.Body.String(), scimCore.SchemaServiceProviderConfig) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) + require.Contains(t, w.Body.String(), core.SchemaServiceProviderConfig) }) for _, path := range []string{scimResourceTypesPath, scimSchemasPath} { @@ -105,8 +104,8 @@ func TestSCIM(t *testing.T) { api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusOK, w.Code) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) - require.Contains(t, w.Body.String(), scimProtocol.SchemaListResponse) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) + require.Contains(t, w.Body.String(), protocol.SchemaListResponse) }) t.Run(path+" rejects filter query parameter", func(t *testing.T) { @@ -118,8 +117,8 @@ func TestSCIM(t *testing.T) { api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusForbidden, w.Code) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) - require.Contains(t, w.Body.String(), scimProtocol.SchemaError) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) + require.Contains(t, w.Body.String(), protocol.SchemaError) }) } @@ -129,7 +128,7 @@ func TestSCIM(t *testing.T) { {http.MethodGet, scimResourceTypesPath}, {http.MethodGet, scimResourceTypesPath + "/User"}, {http.MethodGet, scimSchemasPath}, - {http.MethodGet, scimSchemasPath + "/" + string(scimCore.SchemaUser)}, + {http.MethodGet, scimSchemasPath + "/" + string(core.SchemaUser)}, {http.MethodGet, scimUsersPath}, {http.MethodPost, scimUsersPath}, {http.MethodGet, scimUsersPath + "/missing"}, @@ -144,7 +143,7 @@ func TestSCIM(t *testing.T) { r.Header.Set("Authorization", "Bearer "+token) api.handler.ServeHTTP(w, r) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type"), w.Body.String()) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type"), w.Body.String()) }) } }) @@ -180,7 +179,7 @@ func TestSCIM(t *testing.T) { api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusUnauthorized, w.Code) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) require.True(t, strings.HasPrefix(w.Header().Get("WWW-Authenticate"), "Bearer")) }) } @@ -195,7 +194,7 @@ func TestSCIM(t *testing.T) { api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusOK, w.Code) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) }) } }) @@ -221,8 +220,8 @@ func TestSCIM(t *testing.T) { api.handler.ServeHTTP(w, r) require.Equal(t, http.StatusNotFound, w.Code) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) - require.Contains(t, w.Body.String(), scimProtocol.SchemaError) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) + require.Contains(t, w.Body.String(), protocol.SchemaError) }) t.Run("Returns a SCIM 405 for an unsupported method", func(t *testing.T) { @@ -237,7 +236,7 @@ func TestSCIM(t *testing.T) { require.Equal(t, http.StatusMethodNotAllowed, w.Code) require.Equal(t, "GET, HEAD", w.Header().Get("Allow")) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) }) } } @@ -258,7 +257,7 @@ func TestSCIM(t *testing.T) { require.Equal(t, http.StatusMethodNotAllowed, w.Code) require.Equal(t, tc.allow, w.Header().Get("Allow")) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) }) } }) @@ -285,7 +284,7 @@ func newSCIMServerFor(externalURL string) *server.Server { func scimServe(t *testing.T, srv *server.Server, method, path, body string, headers ...string) *httptest.ResponseRecorder { r := httptest.NewRequest(method, path, strings.NewReader(body)) - r.Header.Set("Content-Type", scimProtocol.MediaType) + r.Header.Set("Content-Type", protocol.MediaType) r.Header.Set("Authorization", "Bearer "+scimValidToken) for i := 0; i+1 < len(headers); i += 2 { r.Header.Set(headers[i], headers[i+1]) @@ -316,7 +315,7 @@ func TestSCIMServer(t *testing.T) { w := scimServe(t, srv, http.MethodGet, scimBasePath+"/ServiceProviderConfig", "") require.Equal(t, http.StatusOK, w.Code) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) require.JSONEq(t, scimFixture(t, "service_provider_config.json"), w.Body.String()) }) @@ -324,7 +323,7 @@ func TestSCIMServer(t *testing.T) { w := scimServe(t, srv, http.MethodGet, scimBasePath+"/ResourceTypes", "") require.Equal(t, http.StatusOK, w.Code) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) body := scimDecode(t, w) require.EqualValues(t, 2, body["totalResults"]) resources := map[string]map[string]any{} @@ -333,12 +332,12 @@ func TestSCIMServer(t *testing.T) { } user := resources["User"] require.Equal(t, "/Users", user["endpoint"]) - require.Equal(t, string(scimCore.SchemaUser), user["schema"]) + require.Equal(t, string(core.SchemaUser), user["schema"]) extension := user["schemaExtensions"].([]any)[0].(map[string]any) - require.Equal(t, string(scimCore.SchemaEnterpriseUser), extension["schema"]) + require.Equal(t, string(core.SchemaEnterpriseUser), extension["schema"]) group := resources["Group"] require.Equal(t, "/Groups", group["endpoint"]) - require.Equal(t, string(scimCore.SchemaGroup), group["schema"]) + require.Equal(t, string(core.SchemaGroup), group["schema"]) require.Empty(t, group["schemaExtensions"]) }) @@ -355,18 +354,18 @@ func TestSCIMServer(t *testing.T) { w := scimServe(t, srv, http.MethodGet, scimBasePath+"/Schemas", "") require.Equal(t, http.StatusOK, w.Code) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) body := scimDecode(t, w) require.EqualValues(t, 3, body["totalResults"]) ids := []string{} for _, resource := range body["Resources"].([]any) { ids = append(ids, resource.(map[string]any)["id"].(string)) } - require.ElementsMatch(t, []string{string(scimCore.SchemaUser), string(scimCore.SchemaEnterpriseUser), string(scimCore.SchemaGroup)}, ids) + require.ElementsMatch(t, []string{string(core.SchemaUser), string(core.SchemaEnterpriseUser), string(core.SchemaGroup)}, ids) }) t.Run("Schemas/{id}", func(t *testing.T) { - for _, id := range []scimCore.SchemaURI{scimCore.SchemaUser, scimCore.SchemaEnterpriseUser, scimCore.SchemaGroup} { + for _, id := range []core.SchemaURI{core.SchemaUser, core.SchemaEnterpriseUser, core.SchemaGroup} { w := scimServe(t, srv, http.MethodGet, scimBasePath+"/Schemas/"+string(id), "") require.Equal(t, http.StatusOK, w.Code) @@ -379,7 +378,7 @@ func TestSCIMServer(t *testing.T) { }) t.Run("Group schema members reference only Users", func(t *testing.T) { - w := scimServe(t, srv, http.MethodGet, scimBasePath+"/Schemas/"+string(scimCore.SchemaGroup), "") + w := scimServe(t, srv, http.MethodGet, scimBasePath+"/Schemas/"+string(core.SchemaGroup), "") require.Equal(t, http.StatusOK, w.Code) sub := map[string]map[string]any{} @@ -396,16 +395,16 @@ func TestSCIMServer(t *testing.T) { }) t.Run("Schemas/{id} location uses the external URL prefix", func(t *testing.T) { - w := scimServe(t, newSCIMServerFor("https://project.supabase.co/auth/v1"), http.MethodGet, scimBasePath+"/Schemas/"+string(scimCore.SchemaUser), "") + w := scimServe(t, newSCIMServerFor("https://project.supabase.co/auth/v1"), http.MethodGet, scimBasePath+"/Schemas/"+string(core.SchemaUser), "") require.Equal(t, http.StatusOK, w.Code) - location := "https://project.supabase.co/auth/v1" + scimBasePath + "/Schemas/" + string(scimCore.SchemaUser) + location := "https://project.supabase.co/auth/v1" + scimBasePath + "/Schemas/" + string(core.SchemaUser) require.Equal(t, location, scimDecode(t, w)["meta"].(map[string]any)["location"]) require.Equal(t, location, w.Header().Get("Content-Location")) }) t.Run("Schemas/User advertises the full RFC 7643 User attributes", func(t *testing.T) { - w := scimServe(t, srv, http.MethodGet, scimBasePath+"/Schemas/"+string(scimCore.SchemaUser), "") + w := scimServe(t, srv, http.MethodGet, scimBasePath+"/Schemas/"+string(core.SchemaUser), "") require.Equal(t, http.StatusOK, w.Code) names := []string{} @@ -456,7 +455,7 @@ func TestSCIMServer(t *testing.T) { w := scimServe(t, srv, http.MethodGet, scimBasePath+"/Users", "", "Authorization", tc.authorization) require.Equal(t, tc.status, w.Code) - require.Equal(t, scimProtocol.MediaType, w.Header().Get("Content-Type")) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) require.Equal(t, tc.challenge, w.Header().Get("WWW-Authenticate")) }) } diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index ccce1aafc6..0eb9197f62 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -17,7 +17,6 @@ import ( "github.com/gofrs/uuid" "github.com/stretchr/testify/require" "github.com/supabase-community/scim-go/pkg/core" - scimCore "github.com/supabase-community/scim-go/pkg/core" "github.com/supabase-community/scim-go/pkg/protocol" "github.com/supabase-community/scim-go/pkg/scimerrors" "github.com/supabase-community/scim-go/pkg/server" @@ -658,20 +657,20 @@ func TestSCIMUserFields(t *testing.T) { }) t.Run("prefers the primary email", func(t *testing.T) { - require.Equal(t, "home@example.com", scimPrimaryEmail([]scimCore.Email{ + require.Equal(t, "home@example.com", scimPrimaryEmail([]core.Email{ {Value: "work@example.com"}, {Value: "home@example.com", Primary: new(true)}, })) }) t.Run("falls back to the first email", func(t *testing.T) { - require.Equal(t, "work@example.com", scimPrimaryEmail([]scimCore.Email{{Value: "work@example.com"}, {Value: "home@example.com"}})) + require.Equal(t, "work@example.com", scimPrimaryEmail([]core.Email{{Value: "work@example.com"}, {Value: "home@example.com"}})) }) t.Run("drops id, meta and password from the resource", func(t *testing.T) { - user := &scimCore.User{UserName: "alice", Password: "secret"} + user := &core.User{UserName: "alice", Password: "secret"} user.ID = "abc" - user.Meta = scimCore.Meta{Version: `W/"1"`} + user.Meta = core.Meta{Version: `W/"1"`} encoded, err := scimUserResource(user) require.NoError(t, err) From bedcb82f9cd85ffdab334ea71a21ee1b6172bb65 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 01:11:23 -0600 Subject: [PATCH 62/88] chore(scim): share admin and SCIM request helpers in SCIM tests --- internal/api/scim_admin_test.go | 12 +--------- internal/api/scim_link_test.go | 29 +++++++---------------- internal/api/scim_provider_delete_test.go | 6 +---- internal/api/scim_ratelimit_test.go | 19 ++++++--------- internal/api/scim_test.go | 14 +++++++++++ internal/api/scim_users_test.go | 7 +----- 6 files changed, 33 insertions(+), 54 deletions(-) diff --git a/internal/api/scim_admin_test.go b/internal/api/scim_admin_test.go index d6ca29e36f..056cd41a56 100644 --- a/internal/api/scim_admin_test.go +++ b/internal/api/scim_admin_test.go @@ -1,7 +1,6 @@ package api import ( - "bytes" "context" "encoding/json" "maps" @@ -55,16 +54,7 @@ func (ts *SCIMTokensTestSuite) tokensPath(provider *models.SSOProvider) string { } func (ts *SCIMTokensTestSuite) request(method, path string, body any) *httptest.ResponseRecorder { - var buf bytes.Buffer - if body != nil { - require.NoError(ts.T(), json.NewEncoder(&buf).Encode(body)) - } - r := httptest.NewRequest(method, path, &buf) - r.Header.Set("Authorization", "Bearer "+ts.AdminJWT) - r.Header.Set("Content-Type", "application/json") - w := httptest.NewRecorder() - ts.API.handler.ServeHTTP(w, r) - return w + return serveAdmin(ts.T(), ts.API, method, path, body) } func (ts *SCIMTokensTestSuite) create(provider *models.SSOProvider, body any) AdminSCIMTokenCreateResponse { diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index 21866e40e9..50745cd25a 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -104,14 +104,8 @@ func (ts *SCIMTestSuite) passkeyRegistrationOptions(user *models.User) int { ChallengeExpiryDuration: 5 * time.Minute, } - session, err := models.NewSession(user.ID, nil) - require.NoError(ts.T(), err) - require.NoError(ts.T(), ts.API.db.Create(session)) - token, _, err := ts.API.generateAccessToken(httptest.NewRequest(http.MethodPost, "/passkeys", nil), ts.API.db, user, &session.ID, models.PasswordGrant) - require.NoError(ts.T(), err) - r := httptest.NewRequest(http.MethodPost, "/passkeys/registration/options", nil) - r.Header.Set("Authorization", "Bearer "+token) + r.Header.Set("Authorization", "Bearer "+ts.accessToken(user)) w := httptest.NewRecorder() ts.API.handler.ServeHTTP(w, r) return w.Code @@ -834,11 +828,7 @@ func (ts *SCIMTestSuite) TestCreateAfterAdminDeletesProviderDeletedUser() { user := ts.linkedUser(id) w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") require.Equal(ts.T(), http.StatusNoContent, w.Code) - r := httptest.NewRequest(http.MethodDelete, "/admin/users/"+user.ID.String(), strings.NewReader(`{"should_soft_delete":`+strconv.FormatBool(soft)+`}`)) - r.Header.Set("Authorization", "Bearer "+adminJWT(ts.T(), ts.API.config.JWT.Secret)) - r.Header.Set("Content-Type", "application/json") - w = httptest.NewRecorder() - ts.API.handler.ServeHTTP(w, r) + w = serveAdmin(ts.T(), ts.API, http.MethodDelete, "/admin/users/"+user.ID.String(), map[string]any{"should_soft_delete": soft}) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) created := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) @@ -862,12 +852,7 @@ func (ts *SCIMTestSuite) TestCreateConcurrentSameEmailLinksToOneUser() { go func(i int) { defer wg.Done() <-start - r := httptest.NewRequest(http.MethodPost, "/scim/v2/Users", strings.NewReader(bodies[i])) - r.Header.Set("Authorization", "Bearer "+ts.TokenA) - r.Header.Set("Content-Type", protocol.MediaType) - w := httptest.NewRecorder() - ts.API.handler.ServeHTTP(w, r) - codes[i] = w.Code + codes[i] = ts.serve(protocol.MediaType, ts.TokenA, http.MethodPost, "/Users", bodies[i]).Code }(i) } close(start) @@ -964,14 +949,18 @@ func (ts *SCIMTestSuite) TestReplaceRejectsRenameToTakenIdentity() { require.NoError(ts.T(), err) } -func (ts *SCIMTestSuite) unlink(user *models.User, identity *models.Identity) *httptest.ResponseRecorder { +func (ts *SCIMTestSuite) accessToken(user *models.User) string { session, err := models.NewSession(user.ID, nil) require.NoError(ts.T(), err) require.NoError(ts.T(), ts.API.db.Create(session)) token, _, err := ts.API.generateAccessToken(httptest.NewRequest(http.MethodPost, "/token", nil), ts.API.db, user, &session.ID, models.PasswordGrant) require.NoError(ts.T(), err) + return token +} + +func (ts *SCIMTestSuite) unlink(user *models.User, identity *models.Identity) *httptest.ResponseRecorder { r := httptest.NewRequest(http.MethodDelete, "/user/identities/"+identity.ID.String(), nil) - r.Header.Set("Authorization", "Bearer "+token) + r.Header.Set("Authorization", "Bearer "+ts.accessToken(user)) w := httptest.NewRecorder() ts.API.handler.ServeHTTP(w, r) return w diff --git a/internal/api/scim_provider_delete_test.go b/internal/api/scim_provider_delete_test.go index 316e044693..b3950d687b 100644 --- a/internal/api/scim_provider_delete_test.go +++ b/internal/api/scim_provider_delete_test.go @@ -17,11 +17,7 @@ func scimUser(name string) string { } func (ts *SCIMTestSuite) deleteProvider(p *models.SSOProvider) { - token := adminJWT(ts.T(), ts.API.config.JWT.Secret) - r := httptest.NewRequest(http.MethodDelete, "/admin/sso/providers/"+p.ID.String(), nil) - r.Header.Set("Authorization", "Bearer "+token) - w := httptest.NewRecorder() - ts.API.handler.ServeHTTP(w, r) + w := serveAdmin(ts.T(), ts.API, http.MethodDelete, "/admin/sso/providers/"+p.ID.String(), nil) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) } diff --git a/internal/api/scim_ratelimit_test.go b/internal/api/scim_ratelimit_test.go index 5108888236..d9ed777f92 100644 --- a/internal/api/scim_ratelimit_test.go +++ b/internal/api/scim_ratelimit_test.go @@ -25,8 +25,8 @@ func TestSCIMRateLimit(t *testing.T) { } tokenA, tokenB := token(), token() - get := func(token, ip string) *httptest.ResponseRecorder { - r := httptest.NewRequest(http.MethodGet, "/scim/v2/Users", nil) + send := func(method, path, token, ip string) *httptest.ResponseRecorder { + r := httptest.NewRequest(method, path, nil) if token != "" { r.Header.Set("Authorization", "Bearer "+token) } @@ -35,6 +35,9 @@ func TestSCIMRateLimit(t *testing.T) { api.handler.ServeHTTP(w, r) return w } + get := func(token, ip string) *httptest.ResponseRecorder { + return send(http.MethodGet, "/scim/v2/Users", token, ip) + } limited := `{"schemas":["urn:ietf:params:scim:api:messages:2.0:Error"],"detail":"Request rate limit reached","status":"429"}` const ip = "192.0.2.1" @@ -74,19 +77,11 @@ func TestSCIMRateLimit(t *testing.T) { {http.MethodDelete, "/scim/v2/Users", http.StatusUnauthorized}, } { ip := fmt.Sprintf("203.0.113.%d", i+1) - send := func() *httptest.ResponseRecorder { - r := httptest.NewRequest(tc.method, tc.path, nil) - r.Header.Set("Authorization", "Bearer scim_invalid") - r.Header.Set(api.config.RateLimitHeader, ip) - w := httptest.NewRecorder() - api.handler.ServeHTTP(w, r) - return w - } for range 30 { - w := send() + w := send(tc.method, tc.path, "scim_invalid", ip) require.Equal(t, tc.status, w.Code, tc.method+" "+tc.path) } - w := send() + w := send(tc.method, tc.path, "scim_invalid", ip) require.Equal(t, http.StatusTooManyRequests, w.Code, tc.method+" "+tc.path) require.JSONEq(t, limited, w.Body.String()) } diff --git a/internal/api/scim_test.go b/internal/api/scim_test.go index fccafb4a2e..b9a3ca3108 100644 --- a/internal/api/scim_test.go +++ b/internal/api/scim_test.go @@ -1,6 +1,7 @@ package api import ( + "bytes" "context" "encoding/json" "io/fs" @@ -555,6 +556,19 @@ func adminJWT(t require.TestingT, secret string) string { return token } +func serveAdmin(t require.TestingT, api *API, method, path string, body any) *httptest.ResponseRecorder { + var buf bytes.Buffer + if body != nil { + require.NoError(t, json.NewEncoder(&buf).Encode(body)) + } + r := httptest.NewRequest(method, path, &buf) + r.Header.Set("Authorization", "Bearer "+adminJWT(t, api.config.JWT.Secret)) + r.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + api.handler.ServeHTTP(w, r) + return w +} + func queryAuditEntries(t require.TestingT, db *storage.Connection, where string, args ...any) []models.AuditLogEntry { entries := []models.AuditLogEntry{} require.NoError(t, db.Q().Where(where, args...).Order("created_at asc").All(&entries)) diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index 0eb9197f62..6807b34dd6 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -360,12 +360,7 @@ func (ts *SCIMTestSuite) TestConcurrentCreateWithinProvider() { for range attempts { wg.Go(func() { <-start - r := httptest.NewRequest(http.MethodPost, "/scim/v2/Users", strings.NewReader(userWith("race@example.com", ""))) - r.Header.Set("Authorization", "Bearer "+ts.TokenA) - r.Header.Set("Content-Type", protocol.MediaType) - w := httptest.NewRecorder() - ts.API.handler.ServeHTTP(w, r) - codes <- w.Code + codes <- ts.serve(protocol.MediaType, ts.TokenA, http.MethodPost, "/Users", userWith("race@example.com", "")).Code }) } close(start) From 0395c31038b04b3cbce704f3e1d656cdbab2c8da Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 01:13:48 -0600 Subject: [PATCH 63/88] chore(scim): build PatchOp and Okta user fixtures with helpers in SCIM tests --- internal/api/scim_groups_test.go | 46 +++++++----------------- internal/api/scim_isolation_test.go | 10 +++--- internal/api/scim_link_test.go | 38 ++++++++------------ internal/api/scim_users_test.go | 55 +++++++++++++++++------------ 4 files changed, 64 insertions(+), 85 deletions(-) diff --git a/internal/api/scim_groups_test.go b/internal/api/scim_groups_test.go index ffae9c9e39..1596eb3c9b 100644 --- a/internal/api/scim_groups_test.go +++ b/internal/api/scim_groups_test.go @@ -81,17 +81,12 @@ func (ts *SCIMTestSuite) TestGroupsLifecycle() { require.Equal(ts.T(), []string{bob}, memberValues(replaced)) require.Equal(ts.T(), meta["created"], replaced["meta"].(map[string]any)["created"]) - w, patched := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ - {"op":"add","path":"members","value":[{"value":"`+alice+`"}]}, - {"op":"replace","path":"displayName","value":"Platform"} - ]}`) + w, patched := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"add","path":"members","value":[{"value":"`+alice+`"}]}`, `{"op":"replace","path":"displayName","value":"Platform"}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.ElementsMatch(ts.T(), []string{alice, bob}, memberValues(patched)) require.Equal(ts.T(), "Platform", patched["displayName"]) - w, patched = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ - {"op":"remove","path":"members[value eq \"`+bob+`\"]"} - ]}`) + w, patched = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"remove","path":"members[value eq \"`+bob+`\"]"}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.Equal(ts.T(), []string{alice}, memberValues(patched)) @@ -164,9 +159,7 @@ func (ts *SCIMTestSuite) TestPatchReplaceMembers() { id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, bob)) before := len(ts.scimAuditEntries()) - w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ - {"op":"replace","path":"members","value":[{"value":"`+bob+`"},{"value":"`+carol+`"}]} - ]}`) + w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"replace","path":"members","value":[{"value":"`+bob+`"},{"value":"`+carol+`"}]}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.ElementsMatch(ts.T(), []string{bob, carol}, memberValues(got)) @@ -194,7 +187,7 @@ func (ts *SCIMTestSuite) TestGroupMemberEventsCarryUserID() { before := len(ts.scimAuditEntries()) id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, bob)) - w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"remove","path":"members[value eq \"`+alice+`\"]"}]}`) + w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"remove","path":"members[value eq \"`+alice+`\"]"}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Users/"+bob, "") require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) @@ -242,7 +235,7 @@ func (ts *SCIMTestSuite) TestExcludedMembersKeepsWrites() { require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.NotContains(ts.T(), got, "groups") - w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id+"?excludedAttributes=members", `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"add","path":"members","value":[{"value":"`+carol+`"}]}]}`) + w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id+"?excludedAttributes=members", patchOp(`{"op":"add","path":"members","value":[{"value":"`+carol+`"}]}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) w, _ = ts.do(ts.TokenA, http.MethodPut, "/Groups/"+id+"?excludedAttributes=members", groupWith("Platform", "", alice, bob, carol)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) @@ -259,7 +252,7 @@ func (ts *SCIMTestSuite) TestPatchRemoveAbsentMember() { id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice)) before := len(ts.auditActions(models.SCIMGroupMemberRemovedAction)) - w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"remove","path":"members[value eq \"`+bob+`\"]"}]}`) + w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"remove","path":"members[value eq \"`+bob+`\"]"}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.Equal(ts.T(), []string{alice}, memberValues(got)) require.Len(ts.T(), ts.auditActions(models.SCIMGroupMemberRemovedAction), before) @@ -274,7 +267,7 @@ func (ts *SCIMTestSuite) TestPatchRejectsRemoveWithValue() { "/Groups/" + id: `{"op":"Remove","path":"members","value":[{"$ref":null,"value":"` + bob + `"}]}`, "/Users/" + alice: `{"op":"remove","path":"emails","value":[{"value":"alice@example.com"}]}`, } { - w, got := ts.do(ts.TokenA, http.MethodPatch, path, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[`+body+`]}`) + w, got := ts.do(ts.TokenA, http.MethodPatch, path, patchOp(body)) require.Equal(ts.T(), http.StatusBadRequest, w.Code, path+" "+w.Body.String()) require.Equal(ts.T(), "invalidSyntax", got["scimType"], path) } @@ -286,9 +279,7 @@ func (ts *SCIMTestSuite) TestPatchRejectsRemoveWithValue() { require.Equal(ts.T(), http.StatusOK, w.Code) require.NotEmpty(ts.T(), got["emails"]) - w, got = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ - {"op":"remove","path":"members[value eq \"`+bob+`\"]","value":null} - ]}`) + w, got = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"remove","path":"members[value eq \"`+bob+`\"]","value":null}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.Equal(ts.T(), []string{alice}, memberValues(got)) } @@ -309,7 +300,7 @@ func (ts *SCIMTestSuite) TestGroupsETagAndIfMatch() { id := created["id"].(string) stale := w.Header().Get("ETag") - patch := `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","path":"displayName","value":"Platform"}]}` + patch := patchOp(`{"op":"replace","path":"displayName","value":"Platform"}`) w, _ = ts.doAs(protocol.MediaType, ts.TokenA, http.MethodPatch, "/Groups/"+id, patch, "If-Match", stale) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) current := w.Header().Get("ETag") @@ -433,14 +424,10 @@ func (ts *SCIMTestSuite) TestGroupsKeepDeactivatedMembers() { id := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) before := len(ts.scimAuditEntries()) - w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+alice, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ - {"op":"replace","path":"active","value":false} - ]}`) + w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+alice, patchOp(`{"op":"replace","path":"active","value":false}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ - {"op":"add","path":"members","value":[{"value":"`+bob+`"}]} - ]}`) + w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"add","path":"members","value":[{"value":"`+bob+`"}]}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.ElementsMatch(ts.T(), []string{alice, bob}, memberValues(got)) @@ -506,17 +493,10 @@ func (ts *SCIMTestSuite) TestGroupsAuditLog() { w, _ = ts.do(ts.TokenA, http.MethodPost, "/Groups", groupWith("Invalid", "g-2", uuid.Must(uuid.NewV4()).String())) require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) - w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ - {"op":"add","path":"members","value":[{"value":"`+bob+`"}]}, - {"op":"remove","path":"members[value eq \"`+alice+`\"]"}, - {"op":"replace","path":"displayName","value":"Platform"} - ]}`) + w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"add","path":"members","value":[{"value":"`+bob+`"}]}`, `{"op":"remove","path":"members[value eq \"`+alice+`\"]"}`, `{"op":"replace","path":"displayName","value":"Platform"}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ - {"op":"replace","path":"displayName","value":"Rejected"}, - {"op":"add","path":"members","value":[{"value":"`+uuid.Must(uuid.NewV4()).String()+`"}]} - ]}`) + w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"replace","path":"displayName","value":"Rejected"}`, `{"op":"add","path":"members","value":[{"value":"`+uuid.Must(uuid.NewV4()).String()+`"}]}`)) require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Groups/"+id, "") diff --git a/internal/api/scim_isolation_test.go b/internal/api/scim_isolation_test.go index 7a5fa3fdf7..d22a49a98b 100644 --- a/internal/api/scim_isolation_test.go +++ b/internal/api/scim_isolation_test.go @@ -23,7 +23,7 @@ func (ts *SCIMTestSuite) TestTenantIsolation() { for _, tc := range []struct{ method, body string }{ {http.MethodGet, ""}, {http.MethodPut, userWith("bob@example.com", "b-1")}, - {http.MethodPatch, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","value":{"active":false}}]}`}, + {http.MethodPatch, patchOp(`{"op":"replace","value":{"active":false}}`)}, {http.MethodDelete, ""}, } { w, _ := ts.do(ts.TokenA, tc.method, "/Users/"+idB, tc.body) @@ -105,13 +105,13 @@ func (ts *SCIMTestSuite) TestRevokedAndExpiredTokensRefusedEverywhere() { {http.MethodPost, "/Users", userWith("bob@example.com", "b-1")}, {http.MethodGet, "/Users/" + user, ""}, {http.MethodPut, "/Users/" + user, userWith("alice@example.com", "a-2")}, - {http.MethodPatch, "/Users/" + user, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","value":{"active":false}}]}`}, + {http.MethodPatch, "/Users/" + user, patchOp(`{"op":"replace","value":{"active":false}}`)}, {http.MethodDelete, "/Users/" + user, ""}, {http.MethodGet, "/Groups", ""}, {http.MethodPost, "/Groups", groupWith("Platform", "g-2")}, {http.MethodGet, "/Groups/" + group, ""}, {http.MethodPut, "/Groups/" + group, groupWith("Owned", "g-1")}, - {http.MethodPatch, "/Groups/" + group, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","path":"displayName","value":"Owned"}]}`}, + {http.MethodPatch, "/Groups/" + group, patchOp(`{"op":"replace","path":"displayName","value":"Owned"}`)}, {http.MethodDelete, "/Groups/" + group, ""}, } for name, token := range map[string]string{"revoked": ts.TokenA, "expired": expiredToken} { @@ -157,7 +157,7 @@ func (ts *SCIMTestSuite) TestGroupsTenantIsolation() { for _, tc := range []struct{ method, body string }{ {http.MethodGet, ""}, {http.MethodPut, groupWith("Engineering", "g-1")}, - {http.MethodPatch, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","path":"displayName","value":"Owned"}]}`}, + {http.MethodPatch, patchOp(`{"op":"replace","path":"displayName","value":"Owned"}`)}, {http.MethodDelete, ""}, } { w, _ := ts.do(ts.TokenB, tc.method, "/Groups/"+groupA, tc.body) @@ -179,7 +179,7 @@ func (ts *SCIMTestSuite) TestGroupsTenantIsolation() { for _, tc := range []struct{ method, body string }{ {http.MethodPut, groupWith("Engineering", "g-1", aliceA, bobB)}, - {http.MethodPatch, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"add","path":"members","value":[{"value":"` + bobB + `"}]}]}`}, + {http.MethodPatch, patchOp(`{"op":"add","path":"members","value":[{"value":"` + bobB + `"}]}`)}, } { w, body := ts.do(ts.TokenA, tc.method, "/Groups/"+groupA, tc.body) require.Equal(ts.T(), http.StatusBadRequest, w.Code, tc.method+" "+w.Body.String()) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index 50745cd25a..9d7b784f58 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -201,7 +201,7 @@ func (ts *SCIMTestSuite) TestOldEmailCannotSignInAfterEmailChange() { } func (ts *SCIMTestSuite) withEmail(email string) string { - return strings.Replace(oktaUser, `"value": "alice@example.com"`, `"value": "`+email+`"`, 1) + return oktaUserWith("value", email) } func (ts *SCIMTestSuite) TestReplaceChangesEmail() { @@ -241,7 +241,7 @@ func (ts *SCIMTestSuite) TestReplaceRenamesAndChangesEmail() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) - body := strings.Replace(ts.withEmail("alice.smith@example.com"), `"userName": "Alice@Example.com"`, `"userName": "alice.smith@example.com"`, 1) + body := withField(ts.withEmail("alice.smith@example.com"), "userName", "alice.smith@example.com") w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, body) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) @@ -276,7 +276,7 @@ func (ts *SCIMTestSuite) TestReplaceAllowsEmailTakenInAnotherProvider() { } func (ts *SCIMTestSuite) TestRemovingEmailsKeepsUserEmail() { - id := ts.create(ts.TokenA, strings.Replace(oktaUser, `"userName": "Alice@Example.com"`, `"userName": "alice.smith"`, 1)) + id := ts.create(ts.TokenA, oktaUserWith("userName", "alice.smith")) user := ts.linkedUser(id) for _, userName := range []string{"alice.smith", "asmith"} { @@ -298,7 +298,7 @@ func (ts *SCIMTestSuite) TestCreateInactiveLogsOutWithoutBanning() { existing := ts.ssoUser(ts.A, "Alice@Example.com", "alice@example.com") ts.session(existing) - user := ts.linkedUser(ts.create(ts.TokenA, strings.Replace(oktaUser, `"active": true`, `"active": false`, 1))) + user := ts.linkedUser(ts.create(ts.TokenA, oktaUserWith("active", false))) require.False(ts.T(), user.IsBanned()) require.Zero(ts.T(), ts.sessions(user)) @@ -307,7 +307,7 @@ func (ts *SCIMTestSuite) TestCreateInactiveLogsOutWithoutBanning() { func (ts *SCIMTestSuite) TestCreateRejectsSharedUser() { ts.create(ts.TokenA, oktaUser) - w, body := ts.do(ts.TokenA, http.MethodPost, "/Users", strings.Replace(oktaUser, `"userName": "Alice@Example.com"`, `"userName": "alice.smith"`, 1)) + w, body := ts.do(ts.TokenA, http.MethodPost, "/Users", oktaUserWith("userName", "alice.smith")) require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) require.Equal(ts.T(), "uniqueness", body["scimType"]) require.EqualValues(ts.T(), 0, ts.list(ts.TokenA, `userName eq "alice.smith"`)["totalResults"]) @@ -330,7 +330,7 @@ func (ts *SCIMTestSuite) TestCreateFallsBackToEmailUserName() { } func (ts *SCIMTestSuite) TestRejectsInvalidEmailsValue() { - invalid := strings.Replace(oktaUser, `"value": "alice@example.com"`, `"value": "not-an-email"`, 1) + invalid := oktaUserWith("value", "not-an-email") w, body := ts.do(ts.TokenA, http.MethodPost, "/Users", invalid) require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) require.Equal(ts.T(), "invalidValue", body["scimType"]) @@ -340,10 +340,7 @@ func (ts *SCIMTestSuite) TestRejectsInvalidEmailsValue() { require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) require.Equal(ts.T(), "invalidValue", body["scimType"]) - w, body = ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id, `{ - "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - "Operations": [{"op": "replace", "path": "emails", "value": [{"value": "not-an-email", "primary": true}]}] - }`) + w, body = ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id, patchOp(`{"op": "replace", "path": "emails", "value": [{"value": "not-an-email", "primary": true}]}`)) require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) require.Equal(ts.T(), "invalidValue", body["scimType"]) require.Equal(ts.T(), "alice@example.com", ts.linkedUser(id).GetEmail()) @@ -499,7 +496,7 @@ func (ts *SCIMTestSuite) TestSAMLLoginBlockedAfterDelete() { } func (ts *SCIMTestSuite) TestSAMLLoginBlockedWhenCreatedInactive() { - id := ts.create(ts.TokenA, strings.Replace(oktaUser, `"active": true`, `"active": false`, 1)) + id := ts.create(ts.TokenA, oktaUserWith("active", false)) ts.linkedUser(id) _, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") @@ -582,11 +579,7 @@ func (ts *SCIMTestSuite) TestSAMLLoginAllowsJITWhileSCIMFlagOff() { func (ts *SCIMTestSuite) TestSAMLLoginNotBlockedByOtherProvider() { id := ts.create(ts.TokenB, oktaUser) - w, _ := ts.do(ts.TokenB, http.MethodPatch, "/Users/"+id, `{ - "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - "Operations": [{"op": "replace", "value": {"active": false}}] - }`) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + ts.setActiveAs(ts.TokenB, id, false) existing := ts.ssoUser(ts.A, "Alice@Example.com", "alice@example.com") user, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") @@ -725,10 +718,7 @@ func (ts *SCIMTestSuite) setActive(id string, active bool) { } func (ts *SCIMTestSuite) setActiveAs(token, id string, active bool) { - w, _ := ts.do(token, http.MethodPatch, "/Users/"+id, `{ - "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - "Operations": [{"op": "replace", "value": {"active": `+strconv.FormatBool(active)+`}}] - }`) + w, _ := ts.do(token, http.MethodPatch, "/Users/"+id, patchOp(`{"op": "replace", "value": {"active": `+strconv.FormatBool(active)+`}}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) } @@ -759,7 +749,7 @@ func (ts *SCIMTestSuite) TestPutInactiveRevokesSessions() { user := ts.linkedUser(id) refreshToken := ts.refreshToken(user) - w, got := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, strings.Replace(oktaUser, `"active": true`, `"active": false`, 1)) + w, got := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, oktaUserWith("active", false)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.Equal(ts.T(), false, got["active"]) require.Zero(ts.T(), ts.sessions(user)) @@ -779,7 +769,7 @@ func (ts *SCIMTestSuite) TestReplaceLinksUnlinkedRow() { row, err := models.CreateSCIMUser(ts.API.db, ts.A.ID, []byte(`{"userName":"Alice@Example.com"}`)) require.NoError(ts.T(), err) - w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+row.ID.String(), strings.Replace(oktaUser, `"active": true`, `"active": false`, 1)) + w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+row.ID.String(), oktaUserWith("active", false)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) user := ts.linkedUser(row.ID.String()) @@ -803,7 +793,7 @@ func (ts *SCIMTestSuite) TestDeleteLogsOutWithoutBanning() { } func (ts *SCIMTestSuite) TestCreateRefusesUserDeletedByProvider() { - for _, body := range []string{oktaUser, strings.Replace(oktaUser, `"userName": "Alice@Example.com"`, `"userName": "alice.new@example.com"`, 1)} { + for _, body := range []string{oktaUser, oktaUserWith("userName", "alice.new@example.com")} { ts.SetupTest() id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) @@ -874,7 +864,7 @@ func (ts *SCIMTestSuite) TestCreateConcurrentSameEmailLinksToOneUser() { } func (ts *SCIMTestSuite) rename(id, userName string) (int, string) { - w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, strings.Replace(oktaUser, `"userName": "Alice@Example.com"`, `"userName": "`+userName+`"`, 1)) + w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, oktaUserWith("userName", userName)) return w.Code, w.Body.String() } diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index 6807b34dd6..2740b6f01b 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -7,6 +7,7 @@ import ( "net/http" "net/http/httptest" "net/url" + "regexp" "slices" "strconv" "strings" @@ -96,15 +97,12 @@ func (ts *SCIMTestSuite) TestOktaLifecycle() { require.Equal(ts.T(), http.StatusOK, w.Code) require.Equal(ts.T(), "Alice", got["name"].(map[string]any)["givenName"]) - w, replaced := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, strings.Replace(oktaUser, `"givenName": "Alice"`, `"givenName": "Alicia"`, 1)) + w, replaced := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, oktaUserWith("givenName", "Alicia")) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.Equal(ts.T(), "Alicia", replaced["name"].(map[string]any)["givenName"]) require.Equal(ts.T(), created["meta"].(map[string]any)["created"], replaced["meta"].(map[string]any)["created"]) - w, patched := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id, `{ - "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - "Operations": [{"op": "replace", "value": {"active": false}}] - }`) + w, patched := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id, patchOp(`{"op": "replace", "value": {"active": false}}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.Equal(ts.T(), false, patched["active"]) require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&stored)) @@ -116,7 +114,7 @@ func (ts *SCIMTestSuite) TestOktaLifecycle() { for _, tc := range []struct{ method, body string }{ {http.MethodGet, ""}, {http.MethodPut, oktaUser}, - {http.MethodPatch, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","value":{"active":true}}]}`}, + {http.MethodPatch, patchOp(`{"op":"replace","value":{"active":true}}`)}, {http.MethodDelete, ""}, } { w, _ = ts.do(ts.TokenA, tc.method, "/Users/"+id, tc.body) @@ -141,10 +139,7 @@ func (ts *SCIMTestSuite) TestOktaContentTypesAndReactivate() { id := created["id"].(string) for _, active := range []bool{false, true} { - w, patched := ts.doAs(contentType, ts.TokenA, http.MethodPatch, "/Users/"+id, `{ - "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - "Operations": [{"op": "replace", "value": {"active": `+strconv.FormatBool(active)+`}}] - }`) + w, patched := ts.doAs(contentType, ts.TokenA, http.MethodPatch, "/Users/"+id, patchOp(`{"op": "replace", "value": {"active": `+strconv.FormatBool(active)+`}}`)) require.Equal(ts.T(), http.StatusOK, w.Code, contentType+" "+w.Body.String()) require.Equal(ts.T(), active, patched["active"], contentType) } @@ -307,7 +302,7 @@ func (ts *SCIMTestSuite) TestAuditLog() { `{"op":"replace","path":"active","value":false}`, `{"op":"replace","path":"active","value":true}`, } { - w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[`+operation+`]}`) + w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id, patchOp(operation)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) } @@ -399,7 +394,7 @@ func (ts *SCIMTestSuite) TestSort() { require.Equal(ts.T(), []string{"Alice@example.com"}, sorted(url.Values{"sortBy": {"userName"}, "count": {"1"}})) require.Equal(ts.T(), []string{"bob@example.com"}, sorted(url.Values{"sortBy": {"userName"}, "startIndex": {"2"}, "count": {"1"}})) - _, body := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+ids["carol@example.com"], `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","path":"title","value":"Lead"}]}`) + _, body := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+ids["carol@example.com"], patchOp(`{"op":"replace","path":"title","value":"Lead"}`)) require.Equal(ts.T(), "Lead", body["title"]) require.Equal(ts.T(), "carol@example.com", sorted(url.Values{"sortBy": {"meta.lastModified"}, "sortOrder": {"descending"}})[0]) @@ -465,13 +460,13 @@ func (ts *SCIMTestSuite) TestWriteResponseProjection() { require.NotContains(ts.T(), got, "emails") require.Equal(ts.T(), "a-2", got["externalId"]) - w, got = ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id+"?excludedAttributes=emails", `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","path":"externalId","value":"a-3"}]}`) + w, got = ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id+"?excludedAttributes=emails", patchOp(`{"op":"replace","path":"externalId","value":"a-3"}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.NotContains(ts.T(), got, "emails") require.Equal(ts.T(), "a-3", got["externalId"]) group := ts.createGroup(ts.TokenA, groupWith("Engineering", "")) - w, got = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+group+"?attributes=displayName", `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"add","path":"members","value":[{"value":"`+id+`"}]}]}`) + w, got = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+group+"?attributes=displayName", patchOp(`{"op":"add","path":"members","value":[{"value":"`+id+`"}]}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.ElementsMatch(ts.T(), []string{"id", "schemas", "displayName"}, slices.Collect(maps.Keys(got))) @@ -492,7 +487,7 @@ func (ts *SCIMTestSuite) TestETagAndIfMatch() { require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.Equal(ts.T(), stale, w.Header().Get("ETag")) - patch := `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[{"op":"replace","path":"displayName","value":"Alice S."}]}` + patch := patchOp(`{"op":"replace","path":"displayName","value":"Alice S."}`) w, patched := ts.doAs(protocol.MediaType, ts.TokenA, http.MethodPatch, "/Users/"+id, patch, "If-Match", stale) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.Equal(ts.T(), "Alice S.", patched["displayName"]) @@ -519,11 +514,7 @@ func (ts *SCIMTestSuite) TestETagAndIfMatch() { func (ts *SCIMTestSuite) TestPatchAttributesOutsideTheMinimalSchema() { id := ts.create(ts.TokenA, oktaUser) - patch := `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ - {"op":"replace","path":"title","value":"Engineer"}, - {"op":"add","path":"phoneNumbers","value":[{"value":"555-0100","type":"work"}]}, - {"op":"replace","path":"urn:ietf:params:scim:schemas:extension:enterprise:2.0:User:department","value":"Auth"} - ]}` + patch := patchOp(`{"op":"replace","path":"title","value":"Engineer"}`, `{"op":"add","path":"phoneNumbers","value":[{"value":"555-0100","type":"work"}]}`, `{"op":"replace","path":"urn:ietf:params:scim:schemas:extension:enterprise:2.0:User:department","value":"Auth"}`) w, patched := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+id, patch) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.Equal(ts.T(), "Engineer", patched["title"]) @@ -724,9 +715,7 @@ func (ts *SCIMTestSuite) TestUsersGroupsAttribute() { require.NotContains(ts.T(), replaced, "groups") require.NotContains(ts.T(), storedResource(bob), "groups") - w, patched := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+alice, `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[ - {"op":"replace","path":"displayName","value":"Alice"} - ]}`) + w, patched := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+alice, patchOp(`{"op":"replace","path":"displayName","value":"Alice"}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.Len(ts.T(), groupsOf(patched), 2) require.NotContains(ts.T(), storedResource(alice), "groups") @@ -738,3 +727,23 @@ func (ts *SCIMTestSuite) TestUsersGroupsAttribute() { require.Len(ts.T(), groupsOf(got), 1) require.Equal(ts.T(), eng, groupsOf(got)[0]["value"]) } + +func oktaUserWith(field string, value any) string { + return withField(oktaUser, field, value) +} + +func withField(body, field string, value any) string { + encoded, err := json.Marshal(value) + if err != nil { + panic(err) + } + match := regexp.MustCompile(`"` + regexp.QuoteMeta(field) + `": ("[^"]*"|true|false)`).FindStringIndex(body) + if match == nil { + panic("no " + field + " in body") + } + return body[:match[0]] + `"` + field + `": ` + string(encoded) + body[match[1]:] +} + +func patchOp(ops ...string) string { + return `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[` + strings.Join(ops, ",") + `]}` +} From 51f408fe8f1263e7937a57ae4d7ab20a6f89a394 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 01:16:06 -0600 Subject: [PATCH 64/88] chore(scim): share audit window, row count and provider type helpers in SCIM tests --- internal/api/scim_groups_test.go | 132 ++++++++++++++-------------- internal/api/scim_isolation_test.go | 12 +-- internal/api/scim_link_test.go | 62 ++++++------- internal/api/scim_okta_spec_test.go | 57 ++++++------ internal/api/scim_users_test.go | 6 ++ 5 files changed, 128 insertions(+), 141 deletions(-) diff --git a/internal/api/scim_groups_test.go b/internal/api/scim_groups_test.go index 1596eb3c9b..87c12fb665 100644 --- a/internal/api/scim_groups_test.go +++ b/internal/api/scim_groups_test.go @@ -157,14 +157,14 @@ func (ts *SCIMTestSuite) TestPatchReplaceMembers() { bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) carol := ts.create(ts.TokenA, userWith("carol@example.com", "c-1")) id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, bob)) - before := len(ts.scimAuditEntries()) - - w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"replace","path":"members","value":[{"value":"`+bob+`"},{"value":"`+carol+`"}]}`)) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - require.ElementsMatch(ts.T(), []string{bob, carol}, memberValues(got)) + entries := ts.auditDuring(func() { + w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"replace","path":"members","value":[{"value":"`+bob+`"},{"value":"`+carol+`"}]}`)) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.ElementsMatch(ts.T(), []string{bob, carol}, memberValues(got)) + }) events := map[string][]string{} - for _, entry := range ts.scimAuditEntries()[before:] { + for _, entry := range entries { scimUserID, _ := entry.Payload["traits"].(map[string]any)["scim_user_id"].(string) events[entry.Payload["action"].(string)] = append(events[entry.Payload["action"].(string)], scimUserID) } @@ -184,19 +184,19 @@ func (ts *SCIMTestSuite) TestGroupMemberEventsCarryUserID() { require.NotNil(ts.T(), row.UserID) userIDs[id] = row.UserID.String() } - before := len(ts.scimAuditEntries()) - - id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, bob)) - w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"remove","path":"members[value eq \"`+alice+`\"]"}`)) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Users/"+bob, "") - require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + entries := ts.auditDuring(func() { + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, bob)) + w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"remove","path":"members[value eq \"`+alice+`\"]"}`)) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Users/"+bob, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + }) type event struct { action, scimUserID, userID string } events := []event{} - for _, entry := range ts.scimAuditEntries()[before:] { + for _, entry := range entries { traits := entry.Payload["traits"].(map[string]any) if _, ok := traits["scim_group_id"]; !ok { continue @@ -349,24 +349,22 @@ func (ts *SCIMTestSuite) TestGroupsRemoveDeletedMembers() { bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) eng := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice, bob)) ops := ts.createGroup(ts.TokenA, groupWith("Ops", "g-2", alice)) - before := len(ts.scimAuditEntries()) - - w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+alice, "") - require.Equal(ts.T(), http.StatusNoContent, w.Code) - - w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+eng, "") - require.Equal(ts.T(), http.StatusOK, w.Code) - require.Equal(ts.T(), []string{bob}, memberValues(got)) - w, got = ts.do(ts.TokenA, http.MethodGet, "/Groups/"+ops, "") - require.Equal(ts.T(), http.StatusOK, w.Code) - require.Empty(ts.T(), memberValues(got)) - - count, err := ts.API.db.Q().Where("scim_user_id = ?", alice).Count(&models.SCIMGroupMember{}) - require.NoError(ts.T(), err) - require.Zero(ts.T(), count) + entries := ts.auditDuring(func() { + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+alice, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + + w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+eng, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Equal(ts.T(), []string{bob}, memberValues(got)) + w, got = ts.do(ts.TokenA, http.MethodGet, "/Groups/"+ops, "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.Empty(ts.T(), memberValues(got)) + + require.Zero(ts.T(), ts.countRows(&models.SCIMGroupMember{}, "scim_user_id = ?", alice)) + }) removed := []string{} - for _, entry := range ts.scimAuditEntries()[before:] { + for _, entry := range entries { if entry.Payload["action"] != string(models.SCIMGroupMemberRemovedAction) { continue } @@ -422,21 +420,21 @@ func (ts *SCIMTestSuite) TestGroupsKeepDeactivatedMembers() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) id := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) - before := len(ts.scimAuditEntries()) - - w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+alice, patchOp(`{"op":"replace","path":"active","value":false}`)) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + entries := ts.auditDuring(func() { + w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+alice, patchOp(`{"op":"replace","path":"active","value":false}`)) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"add","path":"members","value":[{"value":"`+bob+`"}]}`)) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - require.ElementsMatch(ts.T(), []string{alice, bob}, memberValues(got)) + w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"add","path":"members","value":[{"value":"`+bob+`"}]}`)) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.ElementsMatch(ts.T(), []string{alice, bob}, memberValues(got)) - w, user := ts.do(ts.TokenA, http.MethodGet, "/Users/"+alice, "") - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - require.Equal(ts.T(), false, user["active"]) - require.Len(ts.T(), user["groups"], 1) + w, user := ts.do(ts.TokenA, http.MethodGet, "/Users/"+alice, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), false, user["active"]) + require.Len(ts.T(), user["groups"], 1) + }) - for _, entry := range ts.scimAuditEntries()[before:] { + for _, entry := range entries { require.NotEqual(ts.T(), string(models.SCIMGroupMemberRemovedAction), entry.Payload["action"]) } } @@ -484,23 +482,24 @@ func (ts *SCIMTestSuite) TestGroupsUnknownID() { func (ts *SCIMTestSuite) TestGroupsAuditLog() { alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) - before := len(ts.scimAuditEntries()) - - id := ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) + var id string + entries := ts.auditDuring(func() { + id = ts.createGroup(ts.TokenA, groupWith("Engineering", "g-1", alice)) - w, _ := ts.do(ts.TokenA, http.MethodPost, "/Groups", groupWith("Engineering", "g-1")) - require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) - w, _ = ts.do(ts.TokenA, http.MethodPost, "/Groups", groupWith("Invalid", "g-2", uuid.Must(uuid.NewV4()).String())) - require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + w, _ := ts.do(ts.TokenA, http.MethodPost, "/Groups", groupWith("Engineering", "g-1")) + require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) + w, _ = ts.do(ts.TokenA, http.MethodPost, "/Groups", groupWith("Invalid", "g-2", uuid.Must(uuid.NewV4()).String())) + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) - w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"add","path":"members","value":[{"value":"`+bob+`"}]}`, `{"op":"remove","path":"members[value eq \"`+alice+`\"]"}`, `{"op":"replace","path":"displayName","value":"Platform"}`)) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"add","path":"members","value":[{"value":"`+bob+`"}]}`, `{"op":"remove","path":"members[value eq \"`+alice+`\"]"}`, `{"op":"replace","path":"displayName","value":"Platform"}`)) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"replace","path":"displayName","value":"Rejected"}`, `{"op":"add","path":"members","value":[{"value":"`+uuid.Must(uuid.NewV4()).String()+`"}]}`)) - require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) + w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"replace","path":"displayName","value":"Rejected"}`, `{"op":"add","path":"members","value":[{"value":"`+uuid.Must(uuid.NewV4()).String()+`"}]}`)) + require.Equal(ts.T(), http.StatusBadRequest, w.Code, w.Body.String()) - w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Groups/"+id, "") - require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Groups/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + }) tokens, err := models.FindSCIMTokensBySSOProvider(ts.API.db, ts.A.ID) require.NoError(ts.T(), err) @@ -509,7 +508,7 @@ func (ts *SCIMTestSuite) TestGroupsAuditLog() { action, displayName, scimUserID string } events := []event{} - for _, entry := range ts.scimAuditEntries()[before:] { + for _, entry := range entries { require.Equal(ts.T(), uuid.Nil.String(), entry.Payload["actor_id"]) require.Equal(ts.T(), "scim:"+tokens[0].Prefix, entry.Payload["actor_username"]) traits := entry.Payload["traits"].(map[string]any) @@ -551,7 +550,6 @@ func (ts *SCIMTestSuite) TestGroupsPushReplay() { "deactivate bjensen": {"Group A", []string{bjensen}, false}, "reactivate bjensen and reassign app": {"Group A", []string{bjensen}, true}, } - before := len(ts.scimAuditEntries()) last := map[string]struct { body string version any @@ -566,21 +564,23 @@ func (ts *SCIMTestSuite) TestGroupsPushReplay() { version any }{string(request.Body), version} } - played := ts.replay("okta_group_push.json", rfcGroup, []string{rfcBjensen, bjensen, rfcJsmith, jsmith}, onRequest, func(step, group string) { - want, ok := expected[step] - require.True(ts.T(), ok, step) - ts.requireGroup(step, group, want.displayName, want.members) - w, user := ts.do(ts.TokenA, http.MethodGet, "/Users/"+bjensen, "") - require.Equal(ts.T(), http.StatusOK, w.Code, step) - require.Equal(ts.T(), want.bjensenActive, user["active"], step) + entries := ts.auditDuring(func() { + played := ts.replay("okta_group_push.json", rfcGroup, []string{rfcBjensen, bjensen, rfcJsmith, jsmith}, onRequest, func(step, group string) { + want, ok := expected[step] + require.True(ts.T(), ok, step) + ts.requireGroup(step, group, want.displayName, want.members) + w, user := ts.do(ts.TokenA, http.MethodGet, "/Users/"+bjensen, "") + require.Equal(ts.T(), http.StatusOK, w.Code, step) + require.Equal(ts.T(), want.bjensenActive, user["active"], step) + }) + require.Equal(ts.T(), len(expected), played) }) - require.Equal(ts.T(), len(expected), played) type event struct { action, subject string } events := []event{} - for _, entry := range ts.scimAuditEntries()[before:] { + for _, entry := range entries { traits := entry.Payload["traits"].(map[string]any) subject, _ := traits["scim_user_id"].(string) if name, ok := traits["display_name"].(string); ok { diff --git a/internal/api/scim_isolation_test.go b/internal/api/scim_isolation_test.go index d22a49a98b..daa3fe8922 100644 --- a/internal/api/scim_isolation_test.go +++ b/internal/api/scim_isolation_test.go @@ -129,19 +129,13 @@ func (ts *SCIMTestSuite) TestRevokedAndExpiredTokensRefusedEverywhere() { require.True(ts.T(), row.Active) require.Nil(ts.T(), row.DeletedAt) require.Contains(ts.T(), string(row.Resource), `"a-1"`) - users, err := ts.API.db.Q().Where("sso_provider_id = ?", ts.A.ID).Count(&models.SCIMUser{}) - require.NoError(ts.T(), err) - require.Equal(ts.T(), 1, users) + require.Equal(ts.T(), 1, ts.countRows(&models.SCIMUser{}, "sso_provider_id = ?", ts.A.ID)) var stored models.SCIMGroup require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", group).First(&stored)) require.Contains(ts.T(), string(stored.Resource), "Engineering") - groups, err := ts.API.db.Q().Where("sso_provider_id = ?", ts.A.ID).Count(&models.SCIMGroup{}) - require.NoError(ts.T(), err) - require.Equal(ts.T(), 1, groups) - members, err := ts.API.db.Q().Where("group_id = ?", group).Count(&models.SCIMGroupMember{}) - require.NoError(ts.T(), err) - require.Equal(ts.T(), 1, members) + require.Equal(ts.T(), 1, ts.countRows(&models.SCIMGroup{}, "sso_provider_id = ?", ts.A.ID)) + require.Equal(ts.T(), 1, ts.countRows(&models.SCIMGroupMember{}, "group_id = ?", group)) w, _ := ts.do(ts.TokenB, http.MethodGet, "/Users", "") require.Equal(ts.T(), http.StatusOK, w.Code) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index 9d7b784f58..bdffea5e3b 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -24,7 +24,7 @@ func (ts *SCIMTestSuite) ssoUser(provider *models.SSOProvider, sub, email string require.NoError(ts.T(), err) user.IsSSOUser = true require.NoError(ts.T(), ts.API.db.Create(user)) - identity, err := models.NewIdentity(user, "sso:"+provider.ID.String(), map[string]any{"sub": sub, "email": email}) + identity, err := models.NewIdentity(user, scimProviderType(provider.ID), map[string]any{"sub": sub, "email": email}) require.NoError(ts.T(), err) require.NoError(ts.T(), ts.API.db.Create(identity)) return user @@ -53,11 +53,11 @@ func (ts *SCIMTestSuite) TestCreateProvisionsSSOUser() { require.Equal(ts.T(), ts.API.config.JWT.Aud, user.Aud) require.NotNil(ts.T(), user.EmailConfirmedAt) require.False(ts.T(), user.IsBanned()) - require.Equal(ts.T(), []any{"sso:" + ts.A.ID.String()}, user.AppMetaData["providers"]) + require.Equal(ts.T(), []any{scimProviderType(ts.A.ID)}, user.AppMetaData["providers"]) identities := ts.identities(user) require.Len(ts.T(), identities, 1) - require.Equal(ts.T(), "sso:"+ts.A.ID.String(), identities[0].Provider) + require.Equal(ts.T(), scimProviderType(ts.A.ID), identities[0].Provider) require.Equal(ts.T(), "Alice@Example.com", identities[0].ProviderID) } @@ -139,7 +139,7 @@ func (ts *SCIMTestSuite) requireNonSSOUserNeverLinked(sub string) { password, err := models.NewUser("", "alice@example.com", "", ts.API.config.JWT.Aud, nil) require.NoError(ts.T(), err) require.NoError(ts.T(), ts.API.db.Create(password)) - identity, err := models.NewIdentity(password, "sso:"+ts.A.ID.String(), map[string]any{"sub": sub, "email": "alice@example.com", "email_verified": true}) + identity, err := models.NewIdentity(password, scimProviderType(ts.A.ID), map[string]any{"sub": sub, "email": "alice@example.com", "email_verified": true}) require.NoError(ts.T(), err) require.NoError(ts.T(), ts.API.db.Create(identity)) @@ -192,9 +192,7 @@ func (ts *SCIMTestSuite) TestOldEmailCannotSignInAfterEmailChange() { require.Nil(ts.T(), reloaded.RecoverySentAt, req.path) require.Empty(ts.T(), reloaded.RecoveryToken, req.path) require.Empty(ts.T(), reloaded.ConfirmationToken, req.path) - count, err := ts.API.db.Q().Where("user_id = ?", user.ID).Count(&models.OneTimeToken{}) - require.NoError(ts.T(), err) - require.Zero(ts.T(), count, req.path) + require.Zero(ts.T(), ts.countRows(&models.OneTimeToken{}, "user_id = ?", user.ID), req.path) require.Zero(ts.T(), ts.sessions(user), req.path) } } @@ -214,7 +212,7 @@ func (ts *SCIMTestSuite) TestReplaceChangesEmail() { reloaded := ts.linkedUser(id) require.Equal(ts.T(), "alice.smith@example.com", reloaded.GetEmail()) require.Equal(ts.T(), "Alice.Smith@example.com", reloaded.UserMetaData["email"]) - identity, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", "sso:"+ts.A.ID.String()) + identity, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", scimProviderType(ts.A.ID)) require.NoError(ts.T(), err) require.Equal(ts.T(), "Alice.Smith@example.com", identity.IdentityData["email"]) @@ -284,7 +282,7 @@ func (ts *SCIMTestSuite) TestRemovingEmailsKeepsUserEmail() { require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.Equal(ts.T(), "alice@example.com", ts.linkedUser(id).GetEmail(), userName) - identity, err := models.FindIdentityByIdAndProvider(ts.API.db, userName, "sso:"+ts.A.ID.String()) + identity, err := models.FindIdentityByIdAndProvider(ts.API.db, userName, scimProviderType(ts.A.ID)) require.NoError(ts.T(), err, userName) require.Equal(ts.T(), "alice@example.com", identity.IdentityData["email"], userName) @@ -358,9 +356,7 @@ func (ts *SCIMTestSuite) TestCreateLeavesNoUserOnConflict() { w, _ := ts.do(ts.TokenA, http.MethodPost, "/Users", userWith("bob@example.com", "a-1")) require.Equal(ts.T(), http.StatusConflict, w.Code) - count, err := ts.API.db.Q().Where("email = ?", "bob@example.com").Count(&models.User{}) - require.NoError(ts.T(), err) - require.Zero(ts.T(), count) + require.Zero(ts.T(), ts.users("bob@example.com")) } func (ts *SCIMTestSuite) session(user *models.User) { @@ -384,9 +380,7 @@ func (ts *SCIMTestSuite) refresh(token string) int { } func (ts *SCIMTestSuite) sessions(user *models.User) int { - count, err := ts.API.db.Q().Where("user_id = ?", user.ID).Count(&models.Session{}) - require.NoError(ts.T(), err) - return count + return ts.countRows(&models.Session{}, "user_id = ?", user.ID) } func (ts *SCIMTestSuite) samlLogin(ssoProvider *models.SSOProvider, sub, email string) (*models.User, error) { @@ -411,7 +405,7 @@ func (ts *SCIMTestSuite) samlLoginWith(ssoProvider *models.SSOProvider, userData var user *models.User err := ts.API.db.Transaction(func(tx *storage.Connection) error { var terr error - _, user, terr = ts.API.createAccountFromExternalIdentity(tx, r, userData, "sso:"+ssoProvider.ID.String(), false) + _, user, terr = ts.API.createAccountFromExternalIdentity(tx, r, userData, scimProviderType(ssoProvider.ID), false) return terr }) return user, err @@ -425,7 +419,7 @@ func (ts *SCIMTestSuite) TestSAMLLoginLocksVerifiedEmailWithoutMetadataEmail() { conn, err := ts.API.db.NewTransaction() require.NoError(ts.T(), err) tx := &storage.Connection{Connection: conn} - require.NoError(ts.T(), models.LockAccountLinking(tx, "sso:"+ts.A.ID.String(), "alice@example.com")) + require.NoError(ts.T(), models.LockAccountLinking(tx, scimProviderType(ts.A.ID), "alice@example.com")) done := make(chan error, 1) go func() { @@ -530,13 +524,11 @@ func (ts *SCIMTestSuite) TestSAMLLoginLinksDivergedNameIDToSCIMUser() { require.Equal(ts.T(), linked.ID, user.ID) } - count, err := ts.API.db.Q().Where("email = ?", "alice@example.com").Count(&models.User{}) - require.NoError(ts.T(), err) - require.Equal(ts.T(), 1, count) + require.Equal(ts.T(), 1, ts.users("alice@example.com")) providerIDs := []string{} for _, identity := range ts.identities(linked) { - require.Equal(ts.T(), "sso:"+ts.A.ID.String(), identity.Provider) + require.Equal(ts.T(), scimProviderType(ts.A.ID), identity.Provider) providerIDs = append(providerIDs, identity.ProviderID) } require.ElementsMatch(ts.T(), []string{"Alice@Example.com", "saml-name-id"}, providerIDs) @@ -553,9 +545,7 @@ func (ts *SCIMTestSuite) TestSAMLLoginBlockedForDivergedNameIDWhileInactive() { } func (ts *SCIMTestSuite) users(email string) int { - count, err := ts.API.db.Q().Where("email = ?", email).Count(&models.User{}) - require.NoError(ts.T(), err) - return count + return ts.countRows(&models.User{}, "email = ?", email) } func (ts *SCIMTestSuite) TestSAMLLoginAllowsJITWithoutSCIMToken() { @@ -858,9 +848,7 @@ func (ts *SCIMTestSuite) TestCreateConcurrentSameEmailLinksToOneUser() { } require.Equal(ts.T(), 1, created) - count, err := ts.API.db.Q().Where("email = ?", "race@example.com").Count(&models.User{}) - require.NoError(ts.T(), err) - require.Equal(ts.T(), 1, count) + require.Equal(ts.T(), 1, ts.users("race@example.com")) } func (ts *SCIMTestSuite) rename(id, userName string) (int, string) { @@ -875,7 +863,7 @@ func (ts *SCIMTestSuite) TestReplaceRenamesSSOIdentity() { code, body := ts.rename(id, "alice2@example.com") require.Equal(ts.T(), http.StatusOK, code, body) - identity, err := models.FindIdentityByIdAndProvider(ts.API.db, "alice2@example.com", "sso:"+ts.A.ID.String()) + identity, err := models.FindIdentityByIdAndProvider(ts.API.db, "alice2@example.com", scimProviderType(ts.A.ID)) require.NoError(ts.T(), err) require.Equal(ts.T(), user.ID, identity.UserID) require.Equal(ts.T(), "alice2@example.com", identity.IdentityData["sub"]) @@ -935,7 +923,7 @@ func (ts *SCIMTestSuite) TestReplaceRejectsRenameToTakenIdentity() { var row models.SCIMUser require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&row)) require.Equal(ts.T(), "alice@example.com", row.UserName) - _, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", "sso:"+ts.A.ID.String()) + _, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", scimProviderType(ts.A.ID)) require.NoError(ts.T(), err) } @@ -964,7 +952,7 @@ func (ts *SCIMTestSuite) TestUnlinkRefusedForSCIMManagedIdentity() { google, err := models.NewIdentity(user, "google", map[string]any{"sub": "google-1", "email": "alice@example.com"}) require.NoError(ts.T(), err) require.NoError(ts.T(), ts.API.db.Create(google)) - sso, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", "sso:"+ts.A.ID.String()) + sso, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", scimProviderType(ts.A.ID)) require.NoError(ts.T(), err) w := ts.unlink(user, sso) @@ -974,7 +962,7 @@ func (ts *SCIMTestSuite) TestUnlinkRefusedForSCIMManagedIdentity() { code, body := ts.rename(id, "alice2@example.com") require.Equal(ts.T(), http.StatusOK, code, body) - sso, err = models.FindIdentityByIdAndProvider(ts.API.db, "alice2@example.com", "sso:"+ts.A.ID.String()) + sso, err = models.FindIdentityByIdAndProvider(ts.API.db, "alice2@example.com", scimProviderType(ts.A.ID)) require.NoError(ts.T(), err) w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") @@ -991,7 +979,7 @@ func (ts *SCIMTestSuite) TestUnlinkAllowedWhileSCIMFlagOff() { google, err := models.NewIdentity(user, "google", map[string]any{"sub": "google-1", "email": "alice@example.com"}) require.NoError(ts.T(), err) require.NoError(ts.T(), ts.API.db.Create(google)) - sso, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", "sso:"+ts.A.ID.String()) + sso, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", scimProviderType(ts.A.ID)) require.NoError(ts.T(), err) ts.API.config.SSO.SCIM.Enabled = false defer func() { ts.API.config.SSO.SCIM.Enabled = true }() @@ -1004,21 +992,21 @@ func (ts *SCIMTestSuite) TestUnlinkAllowedWhileSCIMFlagOff() { func (ts *SCIMTestSuite) TestRenameSkippedWhenIdentityMissing() { id := ts.create(ts.TokenA, oktaUser) user := ts.linkedUser(id) - sso, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", "sso:"+ts.A.ID.String()) + sso, err := models.FindIdentityByIdAndProvider(ts.API.db, "Alice@Example.com", scimProviderType(ts.A.ID)) require.NoError(ts.T(), err) require.NoError(ts.T(), ts.API.db.Destroy(sso)) - before := len(ts.scimAuditEntries()) hook := logrustest.NewGlobal() defer hook.Reset() - code, body := ts.rename(id, "alice2@example.com") - require.Equal(ts.T(), http.StatusOK, code, body) + entries := ts.auditDuring(func() { + code, body := ts.rename(id, "alice2@example.com") + require.Equal(ts.T(), http.StatusOK, code, body) + }) var row models.SCIMUser require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&row)) require.Equal(ts.T(), "alice2@example.com", row.UserName) require.Empty(ts.T(), ts.identities(user)) - entries := ts.scimAuditEntries()[before:] require.Len(ts.T(), entries, 1) require.Equal(ts.T(), string(models.SCIMUserUpdatedAction), entries[0].Payload["action"]) warned := false diff --git a/internal/api/scim_okta_spec_test.go b/internal/api/scim_okta_spec_test.go index 7ab60f2a80..befa78e7c5 100644 --- a/internal/api/scim_okta_spec_test.go +++ b/internal/api/scim_okta_spec_test.go @@ -172,38 +172,37 @@ func (ts *SCIMTestSuite) TestOktaUserLifecycleReplay() { "unassign": {"Jensen", false, -1}, "reassign (PUT sent twice)": {"Jensen", true, 1}, } - before := len(ts.scimAuditEntries()) id, version := "", "" - played := ts.replay("okta_user_lifecycle.json", rfcBjensen, nil, func(step string, request replayRequest, got map[string]any, created string) { - if strings.Contains(request.Path, "filter=") { - want := expected[step] - require.EqualValues(ts.T(), want.found, got["totalResults"], step) - if want.found > 0 { - require.Equal(ts.T(), created, got["Resources"].([]any)[0].(map[string]any)["id"], step) + entries := ts.auditDuring(func() { + played := ts.replay("okta_user_lifecycle.json", rfcBjensen, nil, func(step string, request replayRequest, got map[string]any, created string) { + if strings.Contains(request.Path, "filter=") { + want := expected[step] + require.EqualValues(ts.T(), want.found, got["totalResults"], step) + if want.found > 0 { + require.Equal(ts.T(), created, got["Resources"].([]any)[0].(map[string]any)["id"], step) + } } - } - if request.Method == http.MethodPost { - require.Contains(ts.T(), string(request.Body), password) - require.NotContains(ts.T(), got, "password") - } - }, func(step, created string) { - want, ok := expected[step] - require.True(ts.T(), ok, step) - id = created - status, got := ts.okta(http.MethodGet, "/Users/"+id, "") - require.Equal(ts.T(), http.StatusOK, status, step) - require.Equal(ts.T(), want.familyName, got["name"].(map[string]any)["familyName"], step) - require.Equal(ts.T(), want.active, got["active"], step) - current := got["meta"].(map[string]any)["version"].(string) - require.NotEqual(ts.T(), version, current, step) - version = current - - count, err := ts.API.db.Q().Where("sso_provider_id = ?", ts.A.ID).Count(&models.SCIMUser{}) - require.NoError(ts.T(), err) - require.Equal(ts.T(), 1, count, step) + if request.Method == http.MethodPost { + require.Contains(ts.T(), string(request.Body), password) + require.NotContains(ts.T(), got, "password") + } + }, func(step, created string) { + want, ok := expected[step] + require.True(ts.T(), ok, step) + id = created + status, got := ts.okta(http.MethodGet, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusOK, status, step) + require.Equal(ts.T(), want.familyName, got["name"].(map[string]any)["familyName"], step) + require.Equal(ts.T(), want.active, got["active"], step) + current := got["meta"].(map[string]any)["version"].(string) + require.NotEqual(ts.T(), version, current, step) + version = current + + require.Equal(ts.T(), 1, ts.countRows(&models.SCIMUser{}, "sso_provider_id = ?", ts.A.ID), step) + }) + require.Equal(ts.T(), len(expected), played) }) - require.Equal(ts.T(), len(expected), played) var stored models.SCIMUser require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&stored)) @@ -213,7 +212,7 @@ func (ts *SCIMTestSuite) TestOktaUserLifecycleReplay() { require.False(ts.T(), user.HasPassword()) actions := []string{} - for _, entry := range ts.scimAuditEntries()[before:] { + for _, entry := range entries { payload, err := json.Marshal(entry.Payload) require.NoError(ts.T(), err) require.NotContains(ts.T(), string(payload), password) diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index 2740b6f01b..550252298c 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -291,6 +291,12 @@ func (ts *SCIMTestSuite) scimAuditEntries() []models.AuditLogEntry { return queryAuditEntries(ts.T(), ts.API.db, "payload->>'log_type' = ?", "scim") } +func (ts *SCIMTestSuite) auditDuring(fn func()) []models.AuditLogEntry { + before := len(ts.scimAuditEntries()) + fn() + return ts.scimAuditEntries()[before:] +} + func (ts *SCIMTestSuite) TestAuditLog() { id := ts.create(ts.TokenA, oktaUser) From 146161e13fbd57d1ba34458c6d1fd1a4025b45fd Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 01:16:47 -0600 Subject: [PATCH 65/88] chore(scim): share the SCIM models test bootstrap --- internal/models/scim_group_test.go | 8 +------- internal/models/scim_settings_test.go | 8 +------- internal/models/scim_test.go | 26 ++++++++++++++++++++++++++ internal/models/scim_token_test.go | 16 +--------------- 4 files changed, 29 insertions(+), 29 deletions(-) create mode 100644 internal/models/scim_test.go diff --git a/internal/models/scim_group_test.go b/internal/models/scim_group_test.go index b6e8feceb2..fa6705bf3e 100644 --- a/internal/models/scim_group_test.go +++ b/internal/models/scim_group_test.go @@ -8,9 +8,7 @@ import ( "github.com/gofrs/uuid" "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" - "github.com/supabase/auth/internal/conf/confload" "github.com/supabase/auth/internal/storage" - "github.com/supabase/auth/internal/storage/test" ) type SCIMGroupTestSuite struct { @@ -20,11 +18,7 @@ type SCIMGroupTestSuite struct { } func TestSCIMGroup(t *testing.T) { - globalConfig, err := confload.LoadGlobal(modelsTestConfig) - require.NoError(t, err) - conn, err := test.SetupDBConnection(globalConfig) - require.NoError(t, err) - ts := &SCIMGroupTestSuite{db: conn} + ts := &SCIMGroupTestSuite{db: setupSCIMTestDB(t)} defer ts.db.Close() suite.Run(t, ts) } diff --git a/internal/models/scim_settings_test.go b/internal/models/scim_settings_test.go index 6c37570286..6bd5300299 100644 --- a/internal/models/scim_settings_test.go +++ b/internal/models/scim_settings_test.go @@ -7,9 +7,7 @@ import ( "github.com/gofrs/uuid" "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" - "github.com/supabase/auth/internal/conf/confload" "github.com/supabase/auth/internal/storage" - "github.com/supabase/auth/internal/storage/test" ) type SCIMSettingsTestSuite struct { @@ -19,11 +17,7 @@ type SCIMSettingsTestSuite struct { } func TestSCIMSettings(t *testing.T) { - globalConfig, err := confload.LoadGlobal(modelsTestConfig) - require.NoError(t, err) - conn, err := test.SetupDBConnection(globalConfig) - require.NoError(t, err) - ts := &SCIMSettingsTestSuite{db: conn} + ts := &SCIMSettingsTestSuite{db: setupSCIMTestDB(t)} defer ts.db.Close() suite.Run(t, ts) } diff --git a/internal/models/scim_test.go b/internal/models/scim_test.go new file mode 100644 index 0000000000..68fb8f7d26 --- /dev/null +++ b/internal/models/scim_test.go @@ -0,0 +1,26 @@ +package models + +import ( + "testing" + + "github.com/stretchr/testify/require" + "github.com/supabase/auth/internal/conf/confload" + "github.com/supabase/auth/internal/storage" + "github.com/supabase/auth/internal/storage/test" +) + +func setupSCIMTestDB(t *testing.T) *storage.Connection { + globalConfig, err := confload.LoadGlobal(modelsTestConfig) + require.NoError(t, err) + conn, err := test.SetupDBConnection(globalConfig) + require.NoError(t, err) + return conn +} + +func createSCIMTestProvider(t require.TestingT, db *storage.Connection) *SSOProvider { + provider := &SSOProvider{} + require.NoError(t, db.Create(provider)) + _, err := EnableSCIM(db, provider.ID) + require.NoError(t, err) + return provider +} diff --git a/internal/models/scim_token_test.go b/internal/models/scim_token_test.go index 1dbae59042..d4cc6a40b0 100644 --- a/internal/models/scim_token_test.go +++ b/internal/models/scim_token_test.go @@ -7,9 +7,7 @@ import ( "github.com/gofrs/uuid" "github.com/stretchr/testify/require" "github.com/stretchr/testify/suite" - "github.com/supabase/auth/internal/conf/confload" "github.com/supabase/auth/internal/storage" - "github.com/supabase/auth/internal/storage/test" ) type SCIMTokenTestSuite struct { @@ -19,11 +17,7 @@ type SCIMTokenTestSuite struct { } func TestSCIMToken(t *testing.T) { - globalConfig, err := confload.LoadGlobal(modelsTestConfig) - require.NoError(t, err) - conn, err := test.SetupDBConnection(globalConfig) - require.NoError(t, err) - ts := &SCIMTokenTestSuite{db: conn} + ts := &SCIMTokenTestSuite{db: setupSCIMTestDB(t)} defer ts.db.Close() suite.Run(t, ts) } @@ -37,14 +31,6 @@ func (ts *SCIMTokenTestSuite) createProvider() *SSOProvider { return createSCIMTestProvider(ts.T(), ts.db) } -func createSCIMTestProvider(t require.TestingT, db *storage.Connection) *SSOProvider { - provider := &SSOProvider{} - require.NoError(t, db.Create(provider)) - _, err := EnableSCIM(db, provider.ID) - require.NoError(t, err) - return provider -} - func (ts *SCIMTokenTestSuite) createToken(expiresAt *time.Time) (*SCIMToken, string) { token, plaintext, err := CreateSCIMToken(ts.db, ts.provider, expiresAt) require.NoError(ts.T(), err) From 9104f65e2e7dc01ebce0ada995d6c4364c15fee4 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 01:17:48 -0600 Subject: [PATCH 66/88] chore(scim): split looped SCIM link tests instead of resetting the suite --- internal/api/scim_link_test.go | 66 +++++++++++++++++++--------------- 1 file changed, 38 insertions(+), 28 deletions(-) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index bdffea5e3b..d6ee345112 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -782,40 +782,50 @@ func (ts *SCIMTestSuite) TestDeleteLogsOutWithoutBanning() { require.Equal(ts.T(), http.StatusBadRequest, ts.refresh(refreshToken)) } -func (ts *SCIMTestSuite) TestCreateRefusesUserDeletedByProvider() { - for _, body := range []string{oktaUser, oktaUserWith("userName", "alice.new@example.com")} { - ts.SetupTest() - id := ts.create(ts.TokenA, oktaUser) - user := ts.linkedUser(id) - w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") - require.Equal(ts.T(), http.StatusNoContent, w.Code) +func (ts *SCIMTestSuite) TestCreateRefusesUserDeletedByProviderSameUserName() { + ts.requireCreateRefusedAfterProviderDelete(oktaUser) +} - w, got := ts.do(ts.TokenA, http.MethodPost, "/Users", body) +func (ts *SCIMTestSuite) TestCreateRefusesUserDeletedByProviderNewUserName() { + ts.requireCreateRefusedAfterProviderDelete(oktaUserWith("userName", "alice.new@example.com")) +} - require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) - require.Equal(ts.T(), "uniqueness", got["scimType"]) - require.False(ts.T(), ts.reloadUser(user.ID).IsBanned()) - require.Zero(ts.T(), ts.countRows(&models.SCIMUser{}, "user_id = ? AND deleted_at IS NULL", user.ID)) - require.Len(ts.T(), ts.identities(user), 1) - require.Equal(ts.T(), 1, ts.users("alice@example.com")) - } +func (ts *SCIMTestSuite) requireCreateRefusedAfterProviderDelete(body string) { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + + w, got := ts.do(ts.TokenA, http.MethodPost, "/Users", body) + + require.Equal(ts.T(), http.StatusConflict, w.Code, w.Body.String()) + require.Equal(ts.T(), "uniqueness", got["scimType"]) + require.False(ts.T(), ts.reloadUser(user.ID).IsBanned()) + require.Zero(ts.T(), ts.countRows(&models.SCIMUser{}, "user_id = ? AND deleted_at IS NULL", user.ID)) + require.Len(ts.T(), ts.identities(user), 1) + require.Equal(ts.T(), 1, ts.users("alice@example.com")) } -func (ts *SCIMTestSuite) TestCreateAfterAdminDeletesProviderDeletedUser() { - for _, soft := range []bool{false, true} { - ts.SetupTest() - id := ts.create(ts.TokenA, oktaUser) - user := ts.linkedUser(id) - w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") - require.Equal(ts.T(), http.StatusNoContent, w.Code) - w = serveAdmin(ts.T(), ts.API, http.MethodDelete, "/admin/users/"+user.ID.String(), map[string]any{"should_soft_delete": soft}) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) +func (ts *SCIMTestSuite) TestCreateAfterAdminHardDeletesProviderDeletedUser() { + ts.requireCreateAfterAdminDelete(false) +} - created := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) +func (ts *SCIMTestSuite) TestCreateAfterAdminSoftDeletesProviderDeletedUser() { + ts.requireCreateAfterAdminDelete(true) +} - require.NotEqual(ts.T(), user.ID, created.ID, soft) - require.True(ts.T(), created.IsSSOUser, soft) - } +func (ts *SCIMTestSuite) requireCreateAfterAdminDelete(soft bool) { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+id, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code) + w = serveAdmin(ts.T(), ts.API, http.MethodDelete, "/admin/users/"+user.ID.String(), map[string]any{"should_soft_delete": soft}) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + + created := ts.linkedUser(ts.create(ts.TokenA, oktaUser)) + + require.NotEqual(ts.T(), user.ID, created.ID, soft) + require.True(ts.T(), created.IsSSOUser, soft) } func (ts *SCIMTestSuite) TestCreateConcurrentSameEmailLinksToOneUser() { From 576a01c9cf5daf6cc0f7680ed5981f72f48229d1 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 01:21:15 -0600 Subject: [PATCH 67/88] chore(scim): rename scimCanLink to scimRequireSSOUser --- internal/api/scim_user_linking.go | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/internal/api/scim_user_linking.go b/internal/api/scim_user_linking.go index 45a290524f..87e5b22a1a 100644 --- a/internal/api/scim_user_linking.go +++ b/internal/api/scim_user_linking.go @@ -39,11 +39,11 @@ func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SC linked := decision.User switch decision.Decision { case models.AccountExists: - if err = scimCanLink(linked); err != nil { + if err = scimRequireSSOUser(linked); err != nil { return nil, false, err } case models.LinkAccount: - if err = scimCanLink(linked); err != nil { + if err = scimRequireSSOUser(linked); err != nil { return nil, false, err } if _, err = s.api.createNewIdentity(tx, linked, providerType, scimIdentityData(user)); err != nil { @@ -66,7 +66,7 @@ func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SC return linked, false, models.LinkSCIMUser(tx, row, linked.ID) } -func scimCanLink(linked *models.User) error { +func scimRequireSSOUser(linked *models.User) error { if !linked.IsSSOUser { return scimerrors.ErrUniqueness("user is not an SSO user") } From ea3d9f9ae8046ded1d3c6972147b0cf75fc2bd92 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 01:22:36 -0600 Subject: [PATCH 68/88] chore(scim): pass SCIM audit fields as a scimAuditEvent --- internal/api/scim.go | 15 ++++++++--- internal/api/scim_admin.go | 42 ++++++++++++++++++++++++++----- internal/api/scim_groups.go | 18 ++++++++++--- internal/api/scim_user_cleanup.go | 14 +++++++++-- internal/api/scim_users.go | 7 +++++- 5 files changed, 79 insertions(+), 17 deletions(-) diff --git a/internal/api/scim.go b/internal/api/scim.go index 0de124edc0..5515e911b6 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -36,6 +36,13 @@ type scimLister[Row, Resource any] struct { render func(*storage.Connection, uuid.UUID, []Row, protocol.Projection) ([]Resource, error) } +type scimAuditEvent struct { + actor *models.User + action models.AuditAction + providerID uuid.UUID + traits map[string]any +} + var ( errMissingSSOProvider = errors.New("scim: request has no SSO provider") @@ -119,10 +126,10 @@ func (a *API) withSCIMRequest(w http.ResponseWriter, req *http.Request) (context return scimRequestKey.WithValue(req.Context(), req), nil } -func (a *API) auditSCIM(tx *storage.Connection, r *http.Request, actor *models.User, action models.AuditAction, providerID uuid.UUID, traits map[string]any) error { - traits["sso_provider_id"] = providerID - traits["outcome"] = "success" - return models.NewAuditLogEntry(a.config.AuditLog, r, tx, actor, action, utilities.GetIPAddress(r), traits) +func (a *API) auditSCIM(tx *storage.Connection, r *http.Request, event scimAuditEvent) error { + event.traits["sso_provider_id"] = event.providerID + event.traits["outcome"] = "success" + return models.NewAuditLogEntry(a.config.AuditLog, r, tx, event.actor, event.action, utilities.GetIPAddress(r), event.traits) } func scimBaseURL(config *conf.GlobalConfiguration) string { diff --git a/internal/api/scim_admin.go b/internal/api/scim_admin.go index 8e2a6164df..7332e77ffd 100644 --- a/internal/api/scim_admin.go +++ b/internal/api/scim_admin.go @@ -53,7 +53,12 @@ func (a *API) adminSCIMEnable(w http.ResponseWriter, r *http.Request) error { if err != nil || !changed { return err } - return a.auditSCIM(tx, r, getAdminUser(ctx), models.SCIMEnabledAction, provider.ID, map[string]any{}) + return a.auditSCIM(tx, r, scimAuditEvent{ + actor: getAdminUser(ctx), + action: models.SCIMEnabledAction, + providerID: provider.ID, + traits: map[string]any{}, + }) }); err != nil { return apierrors.NewInternalServerError("Error enabling SCIM").WithInternalError(err) } @@ -106,7 +111,12 @@ func (a *API) adminSCIMTokensCreate(w http.ResponseWriter, r *http.Request) erro if token, plaintext, err = models.CreateSCIMToken(tx, provider, params.ExpiresAt); err != nil { return err } - return a.auditSCIM(tx, r, getAdminUser(ctx), models.SCIMTokenCreatedAction, provider.ID, map[string]any{scimTokenPrefixTrait: token.Prefix}) + return a.auditSCIM(tx, r, scimAuditEvent{ + actor: getAdminUser(ctx), + action: models.SCIMTokenCreatedAction, + providerID: provider.ID, + traits: map[string]any{scimTokenPrefixTrait: token.Prefix}, + }) }); err != nil { if errors.Is(err, models.SCIMTokenExpiryError{}) { return apierrors.NewBadRequestError(apierrors.ErrorCodeValidationFailed, "expires_at must be in the future") @@ -180,7 +190,12 @@ func (a *API) revokeSCIMToken(tx *storage.Connection, r *http.Request, providerI if err := token.Revoke(tx); err != nil { return nil, err } - return token, a.auditSCIM(tx, r, getAdminUser(r.Context()), models.SCIMTokenRevokedAction, providerID, map[string]any{scimTokenPrefixTrait: token.Prefix}) + return token, a.auditSCIM(tx, r, scimAuditEvent{ + actor: getAdminUser(r.Context()), + action: models.SCIMTokenRevokedAction, + providerID: providerID, + traits: map[string]any{scimTokenPrefixTrait: token.Prefix}, + }) } func (a *API) sendSCIMStatus(w http.ResponseWriter, db *storage.Connection, provider *models.SSOProvider) error { @@ -226,7 +241,12 @@ func (a *API) deprovisionSCIM(tx *storage.Connection, r *http.Request, provider if err != nil || banned == 0 { return err } - return a.auditSCIM(tx, r, actor, models.SCIMUsersBannedAction, provider.ID, map[string]any{"banned_user_count": banned}) + return a.auditSCIM(tx, r, scimAuditEvent{ + actor: actor, + action: models.SCIMUsersBannedAction, + providerID: provider.ID, + traits: map[string]any{"banned_user_count": banned}, + }) } func (a *API) revokeSCIMTokens(tx *storage.Connection, r *http.Request, actor *models.User, providerID uuid.UUID) ([]string, error) { @@ -243,7 +263,12 @@ func (a *API) revokeSCIMTokens(tx *storage.Connection, r *http.Request, actor *m if err := tokens[i].Revoke(tx); err != nil { return nil, err } - if err := a.auditSCIM(tx, r, actor, models.SCIMTokenRevokedAction, providerID, map[string]any{scimTokenPrefixTrait: tokens[i].Prefix}); err != nil { + if err := a.auditSCIM(tx, r, scimAuditEvent{ + actor: actor, + action: models.SCIMTokenRevokedAction, + providerID: providerID, + traits: map[string]any{scimTokenPrefixTrait: tokens[i].Prefix}, + }); err != nil { return nil, err } } @@ -251,5 +276,10 @@ func (a *API) revokeSCIMTokens(tx *storage.Connection, r *http.Request, actor *m } func (a *API) auditSCIMDisabled(tx *storage.Connection, r *http.Request, providerID uuid.UUID, prefixes []string) error { - return a.auditSCIM(tx, r, getAdminUser(r.Context()), models.SCIMDisabledAction, providerID, map[string]any{"token_prefixes": prefixes}) + return a.auditSCIM(tx, r, scimAuditEvent{ + actor: getAdminUser(r.Context()), + action: models.SCIMDisabledAction, + providerID: providerID, + traits: map[string]any{"token_prefixes": prefixes}, + }) } diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index a0fb69e56d..298824dcc3 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -213,9 +213,14 @@ func (s *scimGroupRepository) audit(tx *storage.Connection, r *http.Request, act if err := json.Unmarshal(row.Resource, &resource); err != nil { return err } - return s.api.auditSCIM(tx, r, scimActor(r), action, row.SSOProviderID, map[string]any{ - "scim_group_id": row.ID, - "display_name": resource.DisplayName, + return s.api.auditSCIM(tx, r, scimAuditEvent{ + actor: scimActor(r), + action: action, + providerID: row.SSOProviderID, + traits: map[string]any{ + "scim_group_id": row.ID, + "display_name": resource.DisplayName, + }, }) } @@ -236,7 +241,12 @@ func (s *scimGroupRepository) auditMembers(tx *storage.Connection, r *http.Reque if linked, ok := links[id]; ok { userID = &linked } - if err := s.api.auditSCIM(tx, r, scimActor(r), change.action, row.SSOProviderID, scimMemberTraits(row.ID, id, userID)); err != nil { + if err := s.api.auditSCIM(tx, r, scimAuditEvent{ + actor: scimActor(r), + action: change.action, + providerID: row.SSOProviderID, + traits: scimMemberTraits(row.ID, id, userID), + }); err != nil { return err } } diff --git a/internal/api/scim_user_cleanup.go b/internal/api/scim_user_cleanup.go index bff85bbe9b..ad41b61c26 100644 --- a/internal/api/scim_user_cleanup.go +++ b/internal/api/scim_user_cleanup.go @@ -17,7 +17,12 @@ func (a *API) deleteSCIMUsers(tx *storage.Connection, r *http.Request, actor *mo if err := a.removeSCIMUserFromGroups(tx, r, actor, &rows[i]); err != nil { return err } - if err := a.auditSCIM(tx, r, actor, models.SCIMUserDeletedAction, rows[i].SSOProviderID, scimUserTraits(&rows[i])); err != nil { + if err := a.auditSCIM(tx, r, scimAuditEvent{ + actor: actor, + action: models.SCIMUserDeletedAction, + providerID: rows[i].SSOProviderID, + traits: scimUserTraits(&rows[i]), + }); err != nil { return err } } @@ -30,7 +35,12 @@ func (a *API) removeSCIMUserFromGroups(tx *storage.Connection, r *http.Request, return err } for _, groupID := range groupIDs { - if err := a.auditSCIM(tx, r, actor, models.SCIMGroupMemberRemovedAction, row.SSOProviderID, scimMemberTraits(groupID, row.ID, row.UserID)); err != nil { + if err := a.auditSCIM(tx, r, scimAuditEvent{ + actor: actor, + action: models.SCIMGroupMemberRemovedAction, + providerID: row.SSOProviderID, + traits: scimMemberTraits(groupID, row.ID, row.UserID), + }); err != nil { return err } } diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index fef0e2d127..ef86b295c8 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -342,7 +342,12 @@ func (s *scimUserRepository) changeEmail(tx *storage.Connection, change scimUser } func (s *scimUserRepository) audit(tx *storage.Connection, r *http.Request, action models.AuditAction, row *models.SCIMUser) error { - return s.api.auditSCIM(tx, r, scimActor(r), action, row.SSOProviderID, scimUserTraits(row)) + return s.api.auditSCIM(tx, r, scimAuditEvent{ + actor: scimActor(r), + action: action, + providerID: row.SSOProviderID, + traits: scimUserTraits(row), + }) } func scimUserResource(user *core.User) ([]byte, error) { From 77b7296462a575526625ef2a63ecf8d0c1a7fe21 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 01:22:56 -0600 Subject: [PATCH 69/88] chore(scim): name the SCIM user and group write callback types --- internal/api/scim_groups.go | 4 +++- internal/api/scim_users.go | 4 +++- 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index 298824dcc3..78ee10c408 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -17,6 +17,8 @@ type scimGroupRepository struct { api *API } +type scimGroupWrite func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error) + type scimGroupChange struct { r *http.Request action models.AuditAction @@ -99,7 +101,7 @@ func (s *scimGroupRepository) Delete(ctx context.Context, id, version string) er })) } -func (s *scimGroupRepository) save(ctx context.Context, action models.AuditAction, group *core.Group, write func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error)) (*core.Group, error) { +func (s *scimGroupRepository) save(ctx context.Context, action models.AuditAction, group *core.Group, write scimGroupWrite) (*core.Group, error) { members, err := scimMemberIDs(group.Members) if err != nil { return nil, err diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index ef86b295c8..d07c63961c 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -20,6 +20,8 @@ type scimUserRepository struct { api *API } +type scimUserWrite func(tx *storage.Connection, change scimUserChange) (*models.SCIMUser, *models.User, error) + type scimUserChange struct { r *http.Request target models.SCIMTarget @@ -120,7 +122,7 @@ func (s *scimUserRepository) Delete(ctx context.Context, id, version string) err })) } -func (s *scimUserRepository) save(db *storage.Connection, change scimUserChange, write func(*storage.Connection, scimUserChange) (*models.SCIMUser, *models.User, error)) (*core.User, error) { +func (s *scimUserRepository) save(db *storage.Connection, change scimUserChange, write scimUserWrite) (*core.User, error) { var ( row *models.SCIMUser created *models.User From 976c8459e5db572bb01f38b22c3c2661a4149721 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 08:24:05 -0600 Subject: [PATCH 70/88] fix(scim): send Retry-After on SCIM rate limit responses --- internal/api/scim_ratelimit.go | 15 +++++++++++---- internal/api/scim_ratelimit_test.go | 4 ++++ 2 files changed, 15 insertions(+), 4 deletions(-) diff --git a/internal/api/scim_ratelimit.go b/internal/api/scim_ratelimit.go index 98ed74b487..2c3ba5442e 100644 --- a/internal/api/scim_ratelimit.go +++ b/internal/api/scim_ratelimit.go @@ -2,8 +2,10 @@ package api import ( "context" + "math" "net/http" "path" + "strconv" "github.com/didip/tollbooth/v5" "github.com/didip/tollbooth/v5/limiter" @@ -16,7 +18,7 @@ func (a *API) limitSCIMByIP(lmt *limiter.Limiter) func(http.Handler) http.Handle return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if a.scimSkipsTokenValidator(r) && a.performRateLimiting(lmt, r) != nil { - handler(scimTooManyRequests)(w, r) + handler(scimTooManyRequests(lmt))(w, r) return } next.ServeHTTP(w, r) @@ -48,7 +50,7 @@ func (a *API) limitSCIMByProvider(lmt *limiter.Limiter) func(http.Handler) http. return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if providerID, err := scimProviderID(r.Context()); err == nil && tollbooth.LimitByKeys(lmt, []string{providerID.String()}) != nil { - handler(scimTooManyRequests)(w, r) + handler(scimTooManyRequests(lmt))(w, r) return } next.ServeHTTP(w, r) @@ -56,6 +58,11 @@ func (a *API) limitSCIMByProvider(lmt *limiter.Limiter) func(http.Handler) http. } } -func scimTooManyRequests(w http.ResponseWriter, r *http.Request) error { - return protocol.SendError(w, errSCIMTooManyRequests()) +func scimTooManyRequests(lmt *limiter.Limiter) apiHandler { + return func(w http.ResponseWriter, r *http.Request) error { + if perSecond := lmt.GetMax(); perSecond > 0 { + w.Header().Set("Retry-After", strconv.Itoa(int(math.Ceil(1/perSecond)))) + } + return protocol.SendError(w, errSCIMTooManyRequests()) + } } diff --git a/internal/api/scim_ratelimit_test.go b/internal/api/scim_ratelimit_test.go index d9ed777f92..351ac9ecfc 100644 --- a/internal/api/scim_ratelimit_test.go +++ b/internal/api/scim_ratelimit_test.go @@ -49,6 +49,7 @@ func TestSCIMRateLimit(t *testing.T) { w := get(tokenA, ip) require.Equal(t, http.StatusTooManyRequests, w.Code) require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) + require.Equal(t, "300", w.Header().Get("Retry-After")) require.JSONEq(t, limited, w.Body.String()) w = get(tokenB, ip) require.Equal(t, http.StatusOK, w.Code, w.Body.String()) @@ -84,6 +85,9 @@ func TestSCIMRateLimit(t *testing.T) { w := send(tc.method, tc.path, "scim_invalid", ip) require.Equal(t, http.StatusTooManyRequests, w.Code, tc.method+" "+tc.path) require.JSONEq(t, limited, w.Body.String()) + if tc.status == http.StatusOK { + require.Equal(t, "300", w.Header().Get("Retry-After"), tc.method+" "+tc.path) + } } } From ced63bf790e65794b953a5cc626d350e60da9983 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 08:43:29 -0600 Subject: [PATCH 71/88] chore(scim): add hack/scim-demo.sh to exercise SCIM against a local server --- hack/scim-demo.sh | 97 +++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 97 insertions(+) create mode 100755 hack/scim-demo.sh diff --git a/hack/scim-demo.sh b/hack/scim-demo.sh new file mode 100755 index 0000000000..1d796f1ef7 --- /dev/null +++ b/hack/scim-demo.sh @@ -0,0 +1,97 @@ +#!/usr/bin/env bash + +set -euo pipefail + +ROOT="$(cd "$(dirname "$0")/.." && pwd)" +BASE_URL="${GOTRUE_URL:-http://localhost:9999}" +SECRET="${GOTRUE_JWT_SECRET:-$(sed -n 's/^GOTRUE_JWT_SECRET=//p' "$ROOT/.env" 2>/dev/null | tr -d '"')}" +RUN="$(date +%s)" +TMP="$(mktemp -d)" +trap 'rm -rf "$TMP"' EXIT + +[ -n "$SECRET" ] || { echo "set GOTRUE_JWT_SECRET or add it to $ROOT/.env" >&2; exit 1; } + +b64url() { + openssl base64 -e -A | tr '+/' '-_' | tr -d '=' +} + +admin_jwt() { + local now header payload sig + now="$(date +%s)" + header="$(printf '%s' '{"alg":"HS256","typ":"JWT"}' | b64url)" + payload="$(printf '{"role":"service_role","iat":%s,"exp":%s}' "$now" "$((now + 600))" | b64url)" + sig="$(printf '%s.%s' "$header" "$payload" | openssl dgst -sha256 -hmac "$SECRET" -binary | b64url)" + printf '%s.%s.%s' "$header" "$payload" "$sig" +} + +ADMIN_TOKEN="$(admin_jwt)" + +call() { + local expected="$1" method="$2" url="$3" token="$4" body="${5:-}" status + local args=(-s -X "$method" -o "$TMP/body" -w '%{http_code}' -H "Authorization: Bearer $token") + [ -n "$body" ] && args+=(-H 'Content-Type: application/scim+json' --data "$body") + printf '\n\033[1m%s %s\033[0m\n' "$method" "${url#"$BASE_URL"}" + status="$(curl "${args[@]}" "$url")" + [ -s "$TMP/body" ] && jq -C . < "$TMP/body" + if [ "$status" != "$expected" ]; then + printf '\033[1;31mHTTP %s, expected %s\033[0m\n' "$status" "$expected" + exit 1 + fi + printf '\033[1;32mHTTP %s\033[0m\n' "$status" +} + +field() { + jq -r "$1" < "$TMP/body" +} + +openssl req -x509 -newkey rsa:2048 -nodes -keyout "$TMP/key.pem" -subj "/CN=scim-demo.example" -days 1 -outform DER -out "$TMP/cert.der" 2>/dev/null +CERT="$(openssl base64 -e -A < "$TMP/cert.der")" +METADATA="$CERT" + +call 201 POST "$BASE_URL/admin/sso/providers" "$ADMIN_TOKEN" "$(jq -n --arg xml "$METADATA" '{type: "saml", metadata_xml: $xml}')" +PROVIDER="$(field .id)" +ADMIN="$BASE_URL/admin/sso/providers/$PROVIDER" + +call 200 POST "$ADMIN/scim" "$ADMIN_TOKEN" +call 201 POST "$ADMIN/scim/tokens" "$ADMIN_TOKEN" '{}' +SCIM_TOKEN="$(field .token)" +SCIM="$BASE_URL/scim/v2" + +call 200 GET "$SCIM/ServiceProviderConfig" "$SCIM_TOKEN" + +USERNAME="bjensen+$RUN@example.com" +call 201 POST "$SCIM/Users" "$SCIM_TOKEN" "{ + \"schemas\": [\"urn:ietf:params:scim:schemas:core:2.0:User\"], + \"userName\": \"$USERNAME\", + \"name\": {\"givenName\": \"Barbara\", \"familyName\": \"Jensen\"}, + \"emails\": [{\"value\": \"$USERNAME\", \"primary\": true}], + \"active\": true +}" +USER="$(field .id)" + +call 200 GET "$SCIM/Users?filter=$(jq -rn --arg f "userName eq \"$USERNAME\"" '$f | @uri')" "$SCIM_TOKEN" + +call 200 PATCH "$SCIM/Users/$USER" "$SCIM_TOKEN" '{ + "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + "Operations": [{"op": "replace", "path": "name.familyName", "value": "Jensen-Smith"}] +}' + +call 201 POST "$SCIM/Groups" "$SCIM_TOKEN" "{ + \"schemas\": [\"urn:ietf:params:scim:schemas:core:2.0:Group\"], + \"displayName\": \"Tour Guides $RUN\", + \"members\": [{\"value\": \"$USER\"}] +}" +GROUP="$(field .id)" + +call 200 PATCH "$SCIM/Users/$USER" "$SCIM_TOKEN" '{ + "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + "Operations": [{"op": "replace", "path": "active", "value": false}] +}' + +call 204 DELETE "$SCIM/Groups/$GROUP" "$SCIM_TOKEN" +call 204 DELETE "$SCIM/Users/$USER" "$SCIM_TOKEN" +call 404 GET "$SCIM/Users/$USER" "$SCIM_TOKEN" + +call 200 DELETE "$ADMIN/scim" "$ADMIN_TOKEN" +call 401 GET "$SCIM/Users" "$SCIM_TOKEN" +call 200 DELETE "$ADMIN" "$ADMIN_TOKEN" From aab7e635682af4e59b78e01ddf564bf071e1df97 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 08:51:56 -0600 Subject: [PATCH 72/88] chore(scim): cover discovery, paging, groups and errors in hack/scim-demo.sh --- hack/scim-demo.sh | 98 ++++++++++++++++++++++++++++++++--------------- 1 file changed, 68 insertions(+), 30 deletions(-) diff --git a/hack/scim-demo.sh b/hack/scim-demo.sh index 1d796f1ef7..6df4d1edcb 100755 --- a/hack/scim-demo.sh +++ b/hack/scim-demo.sh @@ -27,12 +27,12 @@ admin_jwt() { ADMIN_TOKEN="$(admin_jwt)" call() { - local expected="$1" method="$2" url="$3" token="$4" body="${5:-}" status + local expected="$1" method="$2" url="$3" token="$4" body="${5:-}" show="${6:-.}" status local args=(-s -X "$method" -o "$TMP/body" -w '%{http_code}' -H "Authorization: Bearer $token") [ -n "$body" ] && args+=(-H 'Content-Type: application/scim+json' --data "$body") printf '\n\033[1m%s %s\033[0m\n' "$method" "${url#"$BASE_URL"}" status="$(curl "${args[@]}" "$url")" - [ -s "$TMP/body" ] && jq -C . < "$TMP/body" + [ -s "$TMP/body" ] && jq -C "$show" < "$TMP/body" if [ "$status" != "$expected" ]; then printf '\033[1;31mHTTP %s, expected %s\033[0m\n' "$status" "$expected" exit 1 @@ -44,10 +44,33 @@ field() { jq -r "$1" < "$TMP/body" } +uri() { + jq -rn --arg v "$1" '$v | @uri' +} + +section() { + printf '\n\033[1;34m== %s ==\033[0m\n' "$1" +} + +user() { + jq -n --arg u "$1" --arg g "$2" --arg f "$3" --argjson active "${4:-true}" '{ + schemas: ["urn:ietf:params:scim:schemas:core:2.0:User"], + userName: $u, + name: {givenName: $g, familyName: $f}, + emails: [{value: $u, primary: true}], + active: $active + }' +} + +patch() { + jq -n --argjson ops "[$1]" '{schemas: ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], Operations: $ops}' +} + openssl req -x509 -newkey rsa:2048 -nodes -keyout "$TMP/key.pem" -subj "/CN=scim-demo.example" -days 1 -outform DER -out "$TMP/cert.der" 2>/dev/null CERT="$(openssl base64 -e -A < "$TMP/cert.der")" METADATA="$CERT" +section "Admin: provider, SCIM and tokens" call 201 POST "$BASE_URL/admin/sso/providers" "$ADMIN_TOKEN" "$(jq -n --arg xml "$METADATA" '{type: "saml", metadata_xml: $xml}')" PROVIDER="$(field .id)" ADMIN="$BASE_URL/admin/sso/providers/$PROVIDER" @@ -55,43 +78,58 @@ ADMIN="$BASE_URL/admin/sso/providers/$PROVIDER" call 200 POST "$ADMIN/scim" "$ADMIN_TOKEN" call 201 POST "$ADMIN/scim/tokens" "$ADMIN_TOKEN" '{}' SCIM_TOKEN="$(field .token)" +call 201 POST "$ADMIN/scim/tokens" "$ADMIN_TOKEN" '{}' +SPARE_TOKEN="$(field .token)" +SPARE_PREFIX="$(field .prefix)" +call 200 GET "$ADMIN/scim/tokens" "$ADMIN_TOKEN" +call 200 DELETE "$ADMIN/scim/tokens/$SPARE_PREFIX" "$ADMIN_TOKEN" +call 200 GET "$ADMIN/scim" "$ADMIN_TOKEN" SCIM="$BASE_URL/scim/v2" +call 401 GET "$SCIM/Users" "$SPARE_TOKEN" +section "Discovery" call 200 GET "$SCIM/ServiceProviderConfig" "$SCIM_TOKEN" +call 200 GET "$SCIM/ResourceTypes" "$SCIM_TOKEN" "" '[.Resources[] | {name, endpoint, schema}]' +call 200 GET "$SCIM/Schemas" "$SCIM_TOKEN" "" '[.Resources[].id]' -USERNAME="bjensen+$RUN@example.com" -call 201 POST "$SCIM/Users" "$SCIM_TOKEN" "{ - \"schemas\": [\"urn:ietf:params:scim:schemas:core:2.0:User\"], - \"userName\": \"$USERNAME\", - \"name\": {\"givenName\": \"Barbara\", \"familyName\": \"Jensen\"}, - \"emails\": [{\"value\": \"$USERNAME\", \"primary\": true}], - \"active\": true -}" +section "Users" +BJENSEN="bjensen+$RUN@example.com" +JSMITH="jsmith+$RUN@example.com" +call 201 POST "$SCIM/Users" "$SCIM_TOKEN" "$(user "$BJENSEN" Barbara Jensen)" USER="$(field .id)" - -call 200 GET "$SCIM/Users?filter=$(jq -rn --arg f "userName eq \"$USERNAME\"" '$f | @uri')" "$SCIM_TOKEN" - -call 200 PATCH "$SCIM/Users/$USER" "$SCIM_TOKEN" '{ - "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - "Operations": [{"op": "replace", "path": "name.familyName", "value": "Jensen-Smith"}] -}' - -call 201 POST "$SCIM/Groups" "$SCIM_TOKEN" "{ - \"schemas\": [\"urn:ietf:params:scim:schemas:core:2.0:Group\"], - \"displayName\": \"Tour Guides $RUN\", - \"members\": [{\"value\": \"$USER\"}] -}" +call 201 POST "$SCIM/Users" "$SCIM_TOKEN" "$(user "$JSMITH" John Smith)" +OTHER="$(field .id)" +call 200 GET "$SCIM/Users?filter=$(uri "userName eq \"$BJENSEN\"")" "$SCIM_TOKEN" +call 200 GET "$SCIM/Users?sortBy=userName&sortOrder=descending" "$SCIM_TOKEN" +call 200 GET "$SCIM/Users?sortBy=userName&startIndex=2&count=1" "$SCIM_TOKEN" +call 200 PATCH "$SCIM/Users/$USER" "$SCIM_TOKEN" "$(patch '{"op": "replace", "path": "name.familyName", "value": "Jensen-Smith"}')" +call 200 PUT "$SCIM/Users/$USER" "$SCIM_TOKEN" "$(user "$BJENSEN" Babs Jensen)" +call 200 PATCH "$SCIM/Users/$USER" "$SCIM_TOKEN" "$(patch '{"op": "replace", "path": "active", "value": false}')" +call 200 PATCH "$SCIM/Users/$USER" "$SCIM_TOKEN" "$(patch '{"op": "replace", "path": "active", "value": true}')" + +section "Groups" +call 201 POST "$SCIM/Groups" "$SCIM_TOKEN" "$(jq -n --arg n "Tour Guides $RUN" --arg m "$USER" '{ + schemas: ["urn:ietf:params:scim:schemas:core:2.0:Group"], + displayName: $n, + members: [{value: $m}] +}')" GROUP="$(field .id)" - -call 200 PATCH "$SCIM/Users/$USER" "$SCIM_TOKEN" '{ - "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - "Operations": [{"op": "replace", "path": "active", "value": false}] -}' - +call 200 PATCH "$SCIM/Groups/$GROUP" "$SCIM_TOKEN" "$(patch "{\"op\": \"add\", \"path\": \"members\", \"value\": [{\"value\": \"$OTHER\"}]}")" +call 200 PATCH "$SCIM/Groups/$GROUP" "$SCIM_TOKEN" "$(patch "{\"op\": \"remove\", \"path\": \"members[value eq \\\"$USER\\\"]\"}")" +call 200 GET "$SCIM/Groups/$GROUP" "$SCIM_TOKEN" + +section "Errors" +call 409 POST "$SCIM/Users" "$SCIM_TOKEN" "$(user "$BJENSEN" Barbara Jensen)" +call 400 GET "$SCIM/Users?filter=$(uri 'userName sw "bjensen"')" "$SCIM_TOKEN" +call 400 GET "$SCIM/Users?sortBy=title" "$SCIM_TOKEN" +call 404 GET "$SCIM/Users/00000000-0000-0000-0000-000000000000" "$SCIM_TOKEN" +call 401 GET "$SCIM/Users" "not-a-token" + +section "Cleanup" call 204 DELETE "$SCIM/Groups/$GROUP" "$SCIM_TOKEN" call 204 DELETE "$SCIM/Users/$USER" "$SCIM_TOKEN" +call 204 DELETE "$SCIM/Users/$OTHER" "$SCIM_TOKEN" call 404 GET "$SCIM/Users/$USER" "$SCIM_TOKEN" - call 200 DELETE "$ADMIN/scim" "$ADMIN_TOKEN" call 401 GET "$SCIM/Users" "$SCIM_TOKEN" call 200 DELETE "$ADMIN" "$ADMIN_TOKEN" From 6eb37b429470882cd4db4d00854c4b9058892a7b Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 09:31:42 -0600 Subject: [PATCH 73/88] fix(scim): write SCIM group member audit entries in one insert --- internal/api/scim.go | 14 +++- internal/api/scim_groups.go | 41 +++++++----- internal/api/scim_user_cleanup.go | 9 ++- internal/models/audit_log_entry.go | 66 +++++++++++++++---- internal/models/audit_log_entry_test.go | 85 +++++++++++++++++++++++++ 5 files changed, 179 insertions(+), 36 deletions(-) create mode 100644 internal/models/audit_log_entry_test.go diff --git a/internal/api/scim.go b/internal/api/scim.go index 5515e911b6..c56612d3dc 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -127,9 +127,17 @@ func (a *API) withSCIMRequest(w http.ResponseWriter, req *http.Request) (context } func (a *API) auditSCIM(tx *storage.Connection, r *http.Request, event scimAuditEvent) error { - event.traits["sso_provider_id"] = event.providerID - event.traits["outcome"] = "success" - return models.NewAuditLogEntry(a.config.AuditLog, r, tx, event.actor, event.action, utilities.GetIPAddress(r), event.traits) + return a.auditSCIMEvents(tx, r, []scimAuditEvent{event}) +} + +func (a *API) auditSCIMEvents(tx *storage.Connection, r *http.Request, events []scimAuditEvent) error { + entries := make([]models.AuditEvent, len(events)) + for i, event := range events { + event.traits["sso_provider_id"] = event.providerID + event.traits["outcome"] = "success" + entries[i] = models.AuditEvent{Actor: event.actor, Action: event.action, Traits: event.traits} + } + return models.NewAuditLogEntries(a.config.AuditLog, r, tx, utilities.GetIPAddress(r), entries) } func scimBaseURL(config *conf.GlobalConfiguration) string { diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index 78ee10c408..fb8c64bf11 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -94,10 +94,15 @@ func (s *scimGroupRepository) Delete(ctx context.Context, id, version string) er if row, err = models.DeleteSCIMGroup(tx, target); err != nil { return err } - if err := s.auditMembers(tx, r, row, nil, removed); err != nil { + events, err := s.memberEvents(tx, r, row, nil, removed) + if err != nil { + return err + } + deleted, err := s.groupEvent(r, models.SCIMGroupDeletedAction, row) + if err != nil { return err } - return s.audit(tx, r, models.SCIMGroupDeletedAction, row) + return s.api.auditSCIMEvents(tx, r, append(events, deleted)) })) } @@ -141,9 +146,12 @@ func (s *scimGroupRepository) applyMembers(tx *storage.Connection, change scimGr if err != nil { return nil, err } + events := []scimAuditEvent{} switch { case changed: - err = s.audit(tx, change.r, change.action, row) + var event scimAuditEvent + event, err = s.groupEvent(change.r, change.action, row) + events = append(events, event) case len(added) > 0 || len(removed) > 0: row, err = models.TouchSCIMGroup(tx, row) default: @@ -152,7 +160,11 @@ func (s *scimGroupRepository) applyMembers(tx *storage.Connection, change scimGr if err != nil { return nil, err } - return row, s.auditMembers(tx, change.r, row, added, removed) + members, err := s.memberEvents(tx, change.r, row, added, removed) + if err != nil { + return nil, err + } + return row, s.api.auditSCIMEvents(tx, change.r, append(events, members...)) } func (s *scimGroupRepository) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMGroup, projection protocol.Projection) ([]*core.Group, error) { @@ -208,14 +220,14 @@ func (s *scimGroupRepository) renderOne(tx *storage.Connection, providerID uuid. return groups[0], nil } -func (s *scimGroupRepository) audit(tx *storage.Connection, r *http.Request, action models.AuditAction, row *models.SCIMGroup) error { +func (s *scimGroupRepository) groupEvent(r *http.Request, action models.AuditAction, row *models.SCIMGroup) (scimAuditEvent, error) { var resource struct { DisplayName string `json:"displayName"` } if err := json.Unmarshal(row.Resource, &resource); err != nil { - return err + return scimAuditEvent{}, err } - return s.api.auditSCIM(tx, r, scimAuditEvent{ + return scimAuditEvent{ actor: scimActor(r), action: action, providerID: row.SSOProviderID, @@ -223,14 +235,15 @@ func (s *scimGroupRepository) audit(tx *storage.Connection, r *http.Request, act "scim_group_id": row.ID, "display_name": resource.DisplayName, }, - }) + }, nil } -func (s *scimGroupRepository) auditMembers(tx *storage.Connection, r *http.Request, row *models.SCIMGroup, added, removed []uuid.UUID) error { +func (s *scimGroupRepository) memberEvents(tx *storage.Connection, r *http.Request, row *models.SCIMGroup, added, removed []uuid.UUID) ([]scimAuditEvent, error) { links, err := models.FindSCIMUserLinks(tx, slices.Concat(added, removed)) if err != nil { - return err + return nil, err } + events := make([]scimAuditEvent, 0, len(added)+len(removed)) for _, change := range []struct { action models.AuditAction ids []uuid.UUID @@ -243,17 +256,15 @@ func (s *scimGroupRepository) auditMembers(tx *storage.Connection, r *http.Reque if linked, ok := links[id]; ok { userID = &linked } - if err := s.api.auditSCIM(tx, r, scimAuditEvent{ + events = append(events, scimAuditEvent{ actor: scimActor(r), action: change.action, providerID: row.SSOProviderID, traits: scimMemberTraits(row.ID, id, userID), - }); err != nil { - return err - } + }) } } - return nil + return events, nil } func scimMemberIDs(members []core.Member) ([]uuid.UUID, error) { diff --git a/internal/api/scim_user_cleanup.go b/internal/api/scim_user_cleanup.go index ad41b61c26..dee8109c91 100644 --- a/internal/api/scim_user_cleanup.go +++ b/internal/api/scim_user_cleanup.go @@ -34,15 +34,14 @@ func (a *API) removeSCIMUserFromGroups(tx *storage.Connection, r *http.Request, if err != nil { return err } - for _, groupID := range groupIDs { - if err := a.auditSCIM(tx, r, scimAuditEvent{ + events := make([]scimAuditEvent, len(groupIDs)) + for i, groupID := range groupIDs { + events[i] = scimAuditEvent{ actor: actor, action: models.SCIMGroupMemberRemovedAction, providerID: row.SSOProviderID, traits: scimMemberTraits(groupID, row.ID, row.UserID), - }); err != nil { - return err } } - return nil + return a.auditSCIMEvents(tx, r, events) } diff --git a/internal/models/audit_log_entry.go b/internal/models/audit_log_entry.go index 76b4291183..45ae89c845 100644 --- a/internal/models/audit_log_entry.go +++ b/internal/models/audit_log_entry.go @@ -2,6 +2,7 @@ package models import ( "bytes" + "encoding/json" "fmt" "net/http" "time" @@ -137,7 +138,55 @@ func (AuditLogEntry) TableName() string { } func NewAuditLogEntry(config conf.AuditLogConfiguration, r *http.Request, tx *storage.Connection, actor *User, action AuditAction, ipAddress string, traits map[string]interface{}) error { + l, _ := logAuditEvent(r, AuditEvent{Actor: actor, Action: action, Traits: traits}, ipAddress) + + if config.DisablePostgres { + return nil + } + + if err := tx.Create(&l); err != nil { + return errors.Wrap(err, "Database error creating audit log entry") + } + + return nil +} + +type AuditEvent struct { + Actor *User + Action AuditAction + Traits map[string]interface{} +} + +func NewAuditLogEntries(config conf.AuditLogConfiguration, r *http.Request, tx *storage.Connection, ipAddress string, events []AuditEvent) error { + ids := make([]uuid.UUID, len(events)) + payloads := make([]string, len(events)) + createdAt := make([]time.Time, len(events)) + for i, event := range events { + l, at := logAuditEvent(r, event, ipAddress) + payload, err := json.Marshal(l.Payload) + if err != nil { + return errors.Wrap(err, "Error encoding audit log entry") + } + ids[i], payloads[i], createdAt[i] = l.ID, string(payload), at + } + + if config.DisablePostgres || len(events) == 0 { + return nil + } + + if err := tx.RawQuery( + fmt.Sprintf("INSERT INTO %q (instance_id, id, payload, created_at, ip_address) SELECT ?, e.id, e.payload::json, e.created_at, ? FROM unnest(?::uuid[], ?::text[], ?::timestamptz[]) AS e(id, payload, created_at)", AuditLogEntry{}.TableName()), + uuid.Nil, ipAddress, ids, payloads, createdAt, + ).Exec(); err != nil { + return errors.Wrap(err, "Database error creating audit log entries") + } + + return nil +} + +func logAuditEvent(r *http.Request, event AuditEvent, ipAddress string) (AuditLogEntry, time.Time) { id := uuid.Must(uuid.NewV4()) + actor, action, traits := event.Actor, event.Action, event.Traits username := actor.GetEmail() @@ -184,7 +233,8 @@ func NewAuditLogEntry(config conf.AuditLogConfiguration, r *http.Request, tx *st maps.Copy(auditLogPayload, payload) auditLogPayload["audit_log_id"] = id auditLogPayload["ip_address"] = ipAddress - auditLogPayload["created_at"] = time.Now().UTC() + createdAt := time.Now().UTC() + auditLogPayload["created_at"] = createdAt if requestID := utilities.GetRequestID(r.Context()); requestID != "" { auditLogPayload["request_id"] = requestID @@ -196,21 +246,11 @@ func NewAuditLogEntry(config conf.AuditLogConfiguration, r *http.Request, tx *st "auth_audit_event": auditLogPayload, }).Info("audit_event") - if config.DisablePostgres { - return nil - } - - l := AuditLogEntry{ + return AuditLogEntry{ ID: id, Payload: JSONMap(payload), IPAddress: ipAddress, - } - - if err := tx.Create(&l); err != nil { - return errors.Wrap(err, "Database error creating audit log entry") - } - - return nil + }, createdAt } func FindAuditLogEntries(tx *storage.Connection, filterColumns []string, filterValue string, pageParams *Pagination) ([]*AuditLogEntry, error) { diff --git a/internal/models/audit_log_entry_test.go b/internal/models/audit_log_entry_test.go new file mode 100644 index 0000000000..332646d2b3 --- /dev/null +++ b/internal/models/audit_log_entry_test.go @@ -0,0 +1,85 @@ +package models + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/require" + "github.com/stretchr/testify/suite" + "github.com/supabase/auth/internal/conf" + "github.com/supabase/auth/internal/storage" +) + +type AuditLogEntryTestSuite struct { + suite.Suite + db *storage.Connection + r *http.Request + actor *User +} + +func TestAuditLogEntry(t *testing.T) { + ts := &AuditLogEntryTestSuite{db: setupSCIMTestDB(t)} + defer ts.db.Close() + suite.Run(t, ts) +} + +func (ts *AuditLogEntryTestSuite) SetupTest() { + require.NoError(ts.T(), TruncateAll(ts.db)) + ts.r = httptest.NewRequest(http.MethodGet, "/", nil) + ts.r.Header.Set("User-Agent", "audit-test") + actor, err := NewUser("", "bjensen@example.com", "", "authenticated", map[string]any{"full_name": "Barbara Jensen"}) + require.NoError(ts.T(), err) + ts.actor = actor +} + +func (ts *AuditLogEntryTestSuite) event(action AuditAction, n int) AuditEvent { + return AuditEvent{Actor: ts.actor, Action: action, Traits: map[string]any{"n": n}} +} + +func (ts *AuditLogEntryTestSuite) entries() []*AuditLogEntry { + entries, err := FindAuditLogEntries(ts.db, nil, "", nil) + require.NoError(ts.T(), err) + return entries +} + +func (ts *AuditLogEntryTestSuite) TestBatchMatchesSingleEntry() { + single := ts.event(SCIMGroupMemberAddedAction, 1) + require.NoError(ts.T(), NewAuditLogEntry(conf.AuditLogConfiguration{}, ts.r, ts.db, single.Actor, single.Action, "192.0.2.1", single.Traits)) + require.NoError(ts.T(), NewAuditLogEntries(conf.AuditLogConfiguration{}, ts.r, ts.db, "192.0.2.1", []AuditEvent{ts.event(SCIMGroupMemberAddedAction, 1)})) + + entries := ts.entries() + require.Len(ts.T(), entries, 2) + require.Equal(ts.T(), entries[1].Payload, entries[0].Payload) + require.Equal(ts.T(), entries[1].IPAddress, entries[0].IPAddress) + require.NotEqual(ts.T(), entries[1].ID, entries[0].ID) + require.False(ts.T(), entries[0].CreatedAt.IsZero()) +} + +func (ts *AuditLogEntryTestSuite) TestBatchKeepsEventOrder() { + events := make([]AuditEvent, 500) + for i := range events { + events[i] = ts.event(SCIMGroupMemberRemovedAction, i) + } + require.NoError(ts.T(), NewAuditLogEntries(conf.AuditLogConfiguration{}, ts.r, ts.db, "192.0.2.1", events)) + + entries := ts.entries() + require.Len(ts.T(), entries, len(events)) + for i, entry := range entries { + require.EqualValues(ts.T(), len(events)-1-i, entry.Payload["traits"].(map[string]any)["n"]) + if i > 0 { + require.True(ts.T(), entries[i-1].CreatedAt.After(entry.CreatedAt)) + } + } +} + +func (ts *AuditLogEntryTestSuite) TestBatchWithoutEvents() { + require.NoError(ts.T(), NewAuditLogEntries(conf.AuditLogConfiguration{}, ts.r, ts.db, "192.0.2.1", nil)) + require.Empty(ts.T(), ts.entries()) +} + +func (ts *AuditLogEntryTestSuite) TestBatchSkipsPostgresWhenDisabled() { + events := []AuditEvent{ts.event(SCIMGroupMemberAddedAction, 1)} + require.NoError(ts.T(), NewAuditLogEntries(conf.AuditLogConfiguration{DisablePostgres: true}, ts.r, ts.db, "192.0.2.1", events)) + require.Empty(ts.T(), ts.entries()) +} From f0847f1d569e1e21e4c05d2a5f2a00fb19f2bdd2 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 09:40:05 -0600 Subject: [PATCH 74/88] fix(scim): preallocate SCIM group members when rendering --- internal/api/scim_groups.go | 12 ++++++++++-- internal/api/scim_users.go | 5 +++-- 2 files changed, 13 insertions(+), 4 deletions(-) diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index fb8c64bf11..f654953d19 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -201,11 +201,19 @@ func (s *scimGroupRepository) members(tx *storage.Connection, providerID uuid.UU if err != nil { return nil, err } + counts := make(map[uuid.UUID]int, len(rows)) + for _, m := range memberships { + counts[m.GroupID]++ + } base := scimBaseURL(s.api.config) for _, m := range memberships { + if members[m.GroupID] == nil { + members[m.GroupID] = make([]core.Member, 0, counts[m.GroupID]) + } + id := m.SCIMUserID.String() members[m.GroupID] = append(members[m.GroupID], core.Member{ - Value: m.SCIMUserID.String(), - Ref: base + "/Users/" + m.SCIMUserID.String(), + Value: id, + Ref: base + "/Users/" + id, Type: scimResourceTypeUser, }) } diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index d07c63961c..33df11f857 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -179,9 +179,10 @@ func (s *scimUserRepository) groupMemberships(tx *storage.Connection, providerID } base := scimBaseURL(s.api.config) for _, m := range memberships { + id := m.GroupID.String() groups[m.SCIMUserID] = append(groups[m.SCIMUserID], core.GroupMembership{ - Value: m.GroupID.String(), - Ref: base + "/Groups/" + m.GroupID.String(), + Value: id, + Ref: base + "/Groups/" + id, Display: m.Display, Type: "direct", }) From 3518767ff9b5521b860564b4dcb58b5516743e6f Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 09:55:42 -0600 Subject: [PATCH 75/88] chore(scim): rename SCIM helpers to say which table and scope they act on --- internal/api/external.go | 10 +++++----- internal/api/scim_admin.go | 12 ++++++------ internal/api/scim_errors.go | 2 +- internal/api/scim_groups.go | 2 +- internal/api/scim_link_test.go | 2 +- internal/api/scim_user_linking.go | 2 +- internal/api/scim_users.go | 8 ++++---- internal/models/audit_log_entry.go | 6 +++--- internal/models/scim.go | 18 +++++++++--------- internal/models/scim_group.go | 26 +++++++++++++------------- internal/models/scim_group_test.go | 14 +++++++------- internal/models/scim_user.go | 22 +++++++++++----------- 12 files changed, 62 insertions(+), 62 deletions(-) diff --git a/internal/api/external.go b/internal/api/external.go index 0ce737a853..e5c41f55df 100644 --- a/internal/api/external.go +++ b/internal/api/external.go @@ -312,12 +312,12 @@ func (a *API) createAccountFromExternalIdentity(tx *storage.Connection, r *http. } ssoProviderID, isSSO, perr := models.SSOProviderID(providerType) - scimOn := isSSO && config.SSO.SCIM.Enabled - if scimOn && perr != nil { + isSCIMProvider := isSSO && config.SSO.SCIM.Enabled + if isSCIMProvider && perr != nil { return 0, nil, apierrors.NewInternalServerError("Invalid SSO provider id in provider type").WithInternalError(perr) } - if scimOn { + if isSCIMProvider { if terr := models.LockAccountLinkingEmails(tx, providerType, models.VerifiedEmails(config, userData.Emails)); terr != nil { return 0, nil, terr } @@ -419,8 +419,8 @@ func (a *API) createAccountFromExternalIdentity(tx *storage.Connection, r *http. return 0, nil, apierrors.NewForbiddenError(apierrors.ErrorCodeUserBanned, "User is banned") } - if scimOn { - deprovisioned, terr := models.IsSCIMDeprovisioned(tx, ssoProviderID, user.ID) + if isSCIMProvider { + deprovisioned, terr := models.IsSCIMUserDeprovisionedByProvider(tx, ssoProviderID, user.ID) if terr != nil { return 0, nil, terr } diff --git a/internal/api/scim_admin.go b/internal/api/scim_admin.go index 7332e77ffd..d3c0008d74 100644 --- a/internal/api/scim_admin.go +++ b/internal/api/scim_admin.go @@ -14,8 +14,8 @@ import ( ) const ( - scimProviderDeletedBan = 100 * 365 * 24 * time.Hour - scimTokenPrefixTrait = "token_prefix" + scimDeprovisionedBanDuration = 100 * 365 * 24 * time.Hour + scimTokenPrefixTrait = "token_prefix" ) type AdminSCIMTokenCreateParams struct { @@ -77,7 +77,7 @@ func (a *API) adminSCIMDisable(w http.ResponseWriter, r *http.Request) error { if err != nil { return err } - prefixes, err := a.revokeSCIMTokens(tx, r, actor, provider.ID) + prefixes, err := a.revokeActiveSCIMTokens(tx, r, actor, provider.ID) if err != nil || !changed { return err } @@ -228,7 +228,7 @@ func (a *API) deprovisionSCIM(tx *storage.Connection, r *http.Request, provider return err } actor := getAdminUser(r.Context()) - prefixes, err := a.revokeSCIMTokens(tx, r, actor, provider.ID) + prefixes, err := a.revokeActiveSCIMTokens(tx, r, actor, provider.ID) if err != nil { return err } @@ -237,7 +237,7 @@ func (a *API) deprovisionSCIM(tx *storage.Connection, r *http.Request, provider return err } } - banned, err := models.BanDeprovisionedSCIMUsers(tx, provider.ID, a.Now().Add(scimProviderDeletedBan)) + banned, err := models.BanDeprovisionedSCIMUsers(tx, provider.ID, a.Now().Add(scimDeprovisionedBanDuration)) if err != nil || banned == 0 { return err } @@ -249,7 +249,7 @@ func (a *API) deprovisionSCIM(tx *storage.Connection, r *http.Request, provider }) } -func (a *API) revokeSCIMTokens(tx *storage.Connection, r *http.Request, actor *models.User, providerID uuid.UUID) ([]string, error) { +func (a *API) revokeActiveSCIMTokens(tx *storage.Connection, r *http.Request, actor *models.User, providerID uuid.UUID) ([]string, error) { if err := models.LockSCIMTokens(tx, providerID); err != nil { return nil, err } diff --git a/internal/api/scim_errors.go b/internal/api/scim_errors.go index 2b9e7cdc67..cea0e53956 100644 --- a/internal/api/scim_errors.go +++ b/internal/api/scim_errors.go @@ -30,7 +30,7 @@ func scimError(err error) error { } func errSCIMNotFound() error { - return scimerrors.ErrNotFound("Resource not found") + return scimerrors.ErrNotFound("resource not found") } func errSCIMStale() error { diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index f654953d19..f0a048b747 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -197,7 +197,7 @@ func (s *scimGroupRepository) members(tx *storage.Connection, providerID uuid.UU for i, row := range rows { ids[i] = row.ID } - memberships, err := models.FindSCIMGroupMembers(tx, providerID, ids) + memberships, err := models.FindSCIMMembershipsByGroup(tx, providerID, ids) if err != nil { return nil, err } diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index d6ee345112..295851d7cd 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -1021,7 +1021,7 @@ func (ts *SCIMTestSuite) TestRenameSkippedWhenIdentityMissing() { require.Equal(ts.T(), string(models.SCIMUserUpdatedAction), entries[0].Payload["action"]) warned := false for _, entry := range hook.AllEntries() { - warned = warned || entry.Level == logrus.WarnLevel && strings.Contains(entry.Message, "SCIM identity") + warned = warned || entry.Level == logrus.WarnLevel && strings.Contains(entry.Message, "identity not found") } require.True(ts.T(), warned) diff --git a/internal/api/scim_user_linking.go b/internal/api/scim_user_linking.go index 87e5b22a1a..0ad33f808c 100644 --- a/internal/api/scim_user_linking.go +++ b/internal/api/scim_user_linking.go @@ -24,7 +24,7 @@ func (s *scimUserRepository) provisionAuthUser(tx *storage.Connection, row *mode created = linked } if !row.Active { - return created, models.LogoutSCIMUser(tx, linked.ID) + return created, models.LogoutUserForSCIM(tx, linked.ID) } return created, nil } diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index 33df11f857..a4837f677c 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -173,7 +173,7 @@ func (s *scimUserRepository) groupMemberships(tx *storage.Connection, providerID for i, row := range rows { ids[i] = row.ID } - memberships, err := models.FindSCIMGroupsForUsers(tx, providerID, ids) + memberships, err := models.FindSCIMMembershipsByUser(tx, providerID, ids) if err != nil { return nil, err } @@ -293,7 +293,7 @@ func (s *scimUserRepository) syncAuthUser(tx *storage.Connection, change scimUse return nil, err } if old.Active && !row.Active { - return nil, models.LogoutSCIMUser(tx, linked.ID) + return nil, models.LogoutUserForSCIM(tx, linked.ID) } return nil, nil } @@ -316,7 +316,7 @@ func (s *scimUserRepository) renameIdentity(tx *storage.Connection, change scimU Data: data, }) if errors.Is(err, models.SCIMIdentityNotFoundError{}) { - observability.GetLogEntry(change.r).Entry.WithField("user_id", userID).WithField("sso_provider_id", providerID).Warn("scim: SCIM identity not found, rename skipped") + observability.GetLogEntry(change.r).Entry.WithField("user_id", userID).WithField("sso_provider_id", providerID).Warn("scim: identity not found, rename skipped") return nil } return err @@ -375,7 +375,7 @@ func logoutSCIMLinkedUser(tx *storage.Connection, userID *uuid.UUID) error { if userID == nil { return nil } - return models.LogoutSCIMUser(tx, *userID) + return models.LogoutUserForSCIM(tx, *userID) } func scimUserEmail(user *core.User) string { diff --git a/internal/models/audit_log_entry.go b/internal/models/audit_log_entry.go index 45ae89c845..37e7159af4 100644 --- a/internal/models/audit_log_entry.go +++ b/internal/models/audit_log_entry.go @@ -138,7 +138,7 @@ func (AuditLogEntry) TableName() string { } func NewAuditLogEntry(config conf.AuditLogConfiguration, r *http.Request, tx *storage.Connection, actor *User, action AuditAction, ipAddress string, traits map[string]interface{}) error { - l, _ := logAuditEvent(r, AuditEvent{Actor: actor, Action: action, Traits: traits}, ipAddress) + l, _ := buildAuditLogEntry(r, AuditEvent{Actor: actor, Action: action, Traits: traits}, ipAddress) if config.DisablePostgres { return nil @@ -162,7 +162,7 @@ func NewAuditLogEntries(config conf.AuditLogConfiguration, r *http.Request, tx * payloads := make([]string, len(events)) createdAt := make([]time.Time, len(events)) for i, event := range events { - l, at := logAuditEvent(r, event, ipAddress) + l, at := buildAuditLogEntry(r, event, ipAddress) payload, err := json.Marshal(l.Payload) if err != nil { return errors.Wrap(err, "Error encoding audit log entry") @@ -184,7 +184,7 @@ func NewAuditLogEntries(config conf.AuditLogConfiguration, r *http.Request, tx * return nil } -func logAuditEvent(r *http.Request, event AuditEvent, ipAddress string) (AuditLogEntry, time.Time) { +func buildAuditLogEntry(r *http.Request, event AuditEvent, ipAddress string) (AuditLogEntry, time.Time) { id := uuid.Must(uuid.NewV4()) actor, action, traits := event.Actor, event.Action, event.Traits diff --git a/internal/models/scim.go b/internal/models/scim.go index 7d0dd19dc8..1767cab7d5 100644 --- a/internal/models/scim.go +++ b/internal/models/scim.go @@ -48,7 +48,7 @@ type SCIMTarget struct { } type scimTable struct { - name string + tableName string label string columns string nameColumn string @@ -62,17 +62,17 @@ func findSCIMPage[T any](tx *storage.Connection, table scimTable, providerID uui where, args := table.where(providerID, query.Filter) total, err := tx.Q().Where(where, args...).Count(new(T)) if err != nil { - return nil, 0, errors.Wrapf(err, "error counting %s", table.name) + return nil, 0, errors.Wrapf(err, "error counting %s", table.tableName) } rows := []T{} if query.Limit <= 0 || query.Offset >= total { return rows, total, nil } if err := tx.RawQuery( - fmt.Sprintf("SELECT %s FROM %q WHERE %s ORDER BY %s OFFSET ? LIMIT ?", table.columns, table.name, where, table.orderBy(query.Order)), + fmt.Sprintf("SELECT %s FROM %q WHERE %s ORDER BY %s OFFSET ? LIMIT ?", table.columns, table.tableName, where, table.orderBy(query.Order)), append(args, query.Offset, query.Limit)..., ).All(&rows); err != nil { - return nil, 0, errors.Wrapf(err, "error finding %s", table.name) + return nil, 0, errors.Wrapf(err, "error finding %s", table.tableName) } return rows, total, nil } @@ -80,7 +80,7 @@ func findSCIMPage[T any](tx *storage.Connection, table scimTable, providerID uui func createSCIMRow[T any](tx *storage.Connection, table scimTable, providerID uuid.UUID, resource []byte) (*T, error) { row := new(T) if err := tx.RawQuery( - fmt.Sprintf("INSERT INTO %q (id, sso_provider_id, resource) VALUES (?, ?, ?::jsonb) RETURNING "+table.columns, table.name), + fmt.Sprintf("INSERT INTO %q (id, sso_provider_id, resource) VALUES (?, ?, ?::jsonb) RETURNING "+table.columns, table.tableName), uuid.Must(uuid.NewV4()), providerID, string(resource), ).First(row); err != nil { return nil, table.wrapError(err, "creating") @@ -95,7 +95,7 @@ func findSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarg } row := new(T) if err := tx.RawQuery( - fmt.Sprintf("SELECT %s FROM %q WHERE %s%s", table.columns, table.name, table.targetClause(), lock), + fmt.Sprintf("SELECT %s FROM %q WHERE %s%s", table.columns, table.tableName, table.targetClause(), lock), target.ID, target.ProviderID, ).First(row); err != nil { return nil, table.wrapError(err, "finding") @@ -106,7 +106,7 @@ func findSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarg func findUnchangedSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarget, resource []byte) (*T, error) { row := new(T) if err := tx.RawQuery( - fmt.Sprintf("SELECT %s FROM %q WHERE %s AND resource = ?::jsonb AND "+scimVersionClause+" FOR UPDATE", table.columns, table.name, table.targetClause()), + fmt.Sprintf("SELECT %s FROM %q WHERE %s AND resource = ?::jsonb AND "+scimVersionClause+" FOR UPDATE", table.columns, table.tableName, table.targetClause()), target.ID, target.ProviderID, string(resource), target.UpdatedAt, target.UpdatedAt, ).First(row); err != nil { if errors.Is(err, sql.ErrNoRows) { @@ -120,7 +120,7 @@ func findUnchangedSCIMRow[T any](tx *storage.Connection, table scimTable, target func replaceSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarget, resource []byte) (*T, error) { row := new(T) if err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET resource = ?::jsonb, updated_at = now() WHERE %s AND "+scimVersionClause+" RETURNING %s", table.name, table.targetClause(), table.columns), + fmt.Sprintf("UPDATE %q SET resource = ?::jsonb, updated_at = now() WHERE %s AND "+scimVersionClause+" RETURNING %s", table.tableName, table.targetClause(), table.columns), string(resource), target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(row); err != nil { return nil, table.writeError(tx, target, err, "replacing") @@ -157,7 +157,7 @@ func (t scimTable) exists(tx *storage.Connection, target SCIMTarget) (bool, erro Exists bool `db:"exists"` }{} if err := tx.RawQuery( - fmt.Sprintf("SELECT EXISTS(SELECT 1 FROM %q WHERE %s) AS exists", t.name, t.targetClause()), + fmt.Sprintf("SELECT EXISTS(SELECT 1 FROM %q WHERE %s) AS exists", t.tableName, t.targetClause()), target.ID, target.ProviderID, ).First(&result); err != nil { return false, errors.Wrapf(err, "error finding %s", t.label) diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go index 979009947c..f21630066f 100644 --- a/internal/models/scim_group.go +++ b/internal/models/scim_group.go @@ -42,7 +42,7 @@ type SCIMGroupMembership struct { } var scimGroupsTable = scimTable{ - name: SCIMGroup{}.TableName(), + tableName: SCIMGroup{}.TableName(), label: "SCIM group", columns: scimGroupColumns, nameColumn: "display_name", @@ -78,7 +78,7 @@ func FindUnchangedSCIMGroup(tx *storage.Connection, target SCIMTarget, resource func TouchSCIMGroup(tx *storage.Connection, group *SCIMGroup) (*SCIMGroup, error) { touched := &SCIMGroup{} if err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET updated_at = now() WHERE id = ? RETURNING "+scimGroupColumns, scimGroupsTable.name), + fmt.Sprintf("UPDATE %q SET updated_at = now() WHERE id = ? RETURNING "+scimGroupColumns, scimGroupsTable.tableName), group.ID, ).First(touched); err != nil { return nil, errors.Wrap(err, "error updating SCIM group") @@ -89,7 +89,7 @@ func TouchSCIMGroup(tx *storage.Connection, group *SCIMGroup) (*SCIMGroup, error func DeleteSCIMGroup(tx *storage.Connection, target SCIMTarget) (*SCIMGroup, error) { group := &SCIMGroup{} err := tx.RawQuery( - fmt.Sprintf("DELETE FROM %q WHERE id = ? AND sso_provider_id = ? AND "+scimVersionClause+" RETURNING "+scimGroupColumns, scimGroupsTable.name), + fmt.Sprintf("DELETE FROM %q WHERE id = ? AND sso_provider_id = ? AND "+scimVersionClause+" RETURNING "+scimGroupColumns, scimGroupsTable.tableName), target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(group) if err != nil { @@ -98,13 +98,13 @@ func DeleteSCIMGroup(tx *storage.Connection, target SCIMTarget) (*SCIMGroup, err return group, nil } -func FindSCIMGroupMembers(tx *storage.Connection, providerID uuid.UUID, groupIDs []uuid.UUID) ([]SCIMGroupMembership, error) { +func FindSCIMMembershipsByGroup(tx *storage.Connection, providerID uuid.UUID, groupIDs []uuid.UUID) ([]SCIMGroupMembership, error) { members := []SCIMGroupMembership{} if len(groupIDs) == 0 { return members, nil } err := tx.RawQuery( - fmt.Sprintf("SELECT m.group_id, m.scim_user_id FROM %q m JOIN %q u ON u.id = m.scim_user_id WHERE m.group_id = ANY(?::uuid[]) AND u.sso_provider_id = ? AND u.deleted_at IS NULL ORDER BY m.group_id, m.created_at, m.scim_user_id", SCIMGroupMember{}.TableName(), scimUsersTable.name), + fmt.Sprintf("SELECT m.group_id, m.scim_user_id FROM %q m JOIN %q u ON u.id = m.scim_user_id WHERE m.group_id = ANY(?::uuid[]) AND u.sso_provider_id = ? AND u.deleted_at IS NULL ORDER BY m.group_id, m.created_at, m.scim_user_id", SCIMGroupMember{}.TableName(), scimUsersTable.tableName), groupIDs, providerID, ).All(&members) if err != nil { @@ -113,19 +113,19 @@ func FindSCIMGroupMembers(tx *storage.Connection, providerID uuid.UUID, groupIDs return members, nil } -func FindSCIMGroupsForUsers(tx *storage.Connection, providerID uuid.UUID, scimUserIDs []uuid.UUID) ([]SCIMGroupMembership, error) { - groups := []SCIMGroupMembership{} +func FindSCIMMembershipsByUser(tx *storage.Connection, providerID uuid.UUID, scimUserIDs []uuid.UUID) ([]SCIMGroupMembership, error) { + memberships := []SCIMGroupMembership{} if len(scimUserIDs) == 0 { - return groups, nil + return memberships, nil } err := tx.RawQuery( - fmt.Sprintf("SELECT m.group_id, m.scim_user_id, g.resource->>'displayName' AS display FROM %q m JOIN %q g ON g.id = m.group_id WHERE m.scim_user_id = ANY(?::uuid[]) AND g.sso_provider_id = ? ORDER BY m.scim_user_id, g.display_name COLLATE \"C\", g.id", SCIMGroupMember{}.TableName(), scimGroupsTable.name), + fmt.Sprintf("SELECT m.group_id, m.scim_user_id, g.resource->>'displayName' AS display FROM %q m JOIN %q g ON g.id = m.group_id WHERE m.scim_user_id = ANY(?::uuid[]) AND g.sso_provider_id = ? ORDER BY m.scim_user_id, g.display_name COLLATE \"C\", g.id", SCIMGroupMember{}.TableName(), scimGroupsTable.tableName), scimUserIDs, providerID, - ).All(&groups) + ).All(&memberships) if err != nil { return nil, errors.Wrap(err, "error finding SCIM groups for users") } - return groups, nil + return memberships, nil } func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserIDs []uuid.UUID) (added, removed []uuid.UUID, err error) { @@ -145,7 +145,7 @@ func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserI } func RemoveSCIMUserFromGroups(tx *storage.Connection, scimUserID uuid.UUID) ([]uuid.UUID, error) { - groups, members := scimGroupsTable.name, SCIMGroupMember{}.TableName() + groups, members := scimGroupsTable.tableName, SCIMGroupMember{}.TableName() if err := tx.RawQuery( fmt.Sprintf("SELECT id FROM %q WHERE id IN (SELECT group_id FROM %q WHERE scim_user_id = ?) ORDER BY id FOR UPDATE", groups, members), scimUserID, @@ -177,7 +177,7 @@ func RemoveSCIMUserFromGroups(tx *storage.Connection, scimUserID uuid.UUID) ([]u func lockLiveSCIMUserIDs(tx *storage.Connection, providerID uuid.UUID, ids []uuid.UUID) ([]uuid.UUID, error) { found := []uuid.UUID{} if err := tx.RawQuery( - fmt.Sprintf("SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL ORDER BY id FOR SHARE", scimUsersTable.name), + fmt.Sprintf("SELECT id FROM %q WHERE id = ANY(?::uuid[]) AND sso_provider_id = ? AND deleted_at IS NULL ORDER BY id FOR SHARE", scimUsersTable.tableName), ids, providerID, ).All(&found); err != nil { return nil, errors.Wrap(err, "error locking SCIM group members") diff --git a/internal/models/scim_group_test.go b/internal/models/scim_group_test.go index fa6705bf3e..fe2830bc9d 100644 --- a/internal/models/scim_group_test.go +++ b/internal/models/scim_group_test.go @@ -175,7 +175,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersDiffs() { require.Equal(ts.T(), []uuid.UUID{carol.ID}, added) require.Equal(ts.T(), []uuid.UUID{alice.ID}, removed) - members, err := FindSCIMGroupMembers(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) + members, err := FindSCIMMembershipsByGroup(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) require.NoError(ts.T(), err) require.Len(ts.T(), members, 2) require.ElementsMatch(ts.T(), []uuid.UUID{bob.ID, carol.ID}, []uuid.UUID{members[0].SCIMUserID, members[1].SCIMUserID}) @@ -194,7 +194,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersRejectsOtherProviderUsers() { _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, outsider.ID}) require.Equal(ts.T(), SCIMGroupMemberNotFoundError{IDs: []uuid.UUID{outsider.ID}}, err) - members, err := FindSCIMGroupMembers(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) + members, err := FindSCIMMembershipsByGroup(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) require.NoError(ts.T(), err) require.Empty(ts.T(), members) } @@ -274,7 +274,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersDoesNotLockExistingMembers() { require.Equal(ts.T(), []uuid.UUID{group.ID}, removed) require.NoError(ts.T(), deleting.TX.Commit()) - members, err := FindSCIMGroupMembers(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) + members, err := FindSCIMMembershipsByGroup(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) require.NoError(ts.T(), err) require.Len(ts.T(), members, 1) require.Equal(ts.T(), bob.ID, members[0].SCIMUserID) @@ -289,7 +289,7 @@ func (ts *SCIMGroupTestSuite) TestFindMembersHidesDeletedUsers() { _, err = DeleteSCIMUser(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) require.NoError(ts.T(), err) - members, err := FindSCIMGroupMembers(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) + members, err := FindSCIMMembershipsByGroup(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) require.NoError(ts.T(), err) require.Empty(ts.T(), members) } @@ -312,7 +312,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersValidatesOnlyAddedMembers() { require.Equal(ts.T(), SCIMGroupMemberNotFoundError{IDs: []uuid.UUID{uuid.Nil}}, err) } -func (ts *SCIMGroupTestSuite) TestFindGroupsForUsers() { +func (ts *SCIMGroupTestSuite) TestFindMembershipsByUser() { engineering := ts.createGroup(ts.provider.ID, "Engineering") admins := ts.createGroup(ts.provider.ID, "Admins") alice := ts.createUser(ts.provider.ID, "alice") @@ -322,7 +322,7 @@ func (ts *SCIMGroupTestSuite) TestFindGroupsForUsers() { _, _, err = ReplaceSCIMGroupMembers(ts.db, admins, []uuid.UUID{alice.ID}) require.NoError(ts.T(), err) - groups, err := FindSCIMGroupsForUsers(ts.db, ts.provider.ID, []uuid.UUID{alice.ID, bob.ID}) + groups, err := FindSCIMMembershipsByUser(ts.db, ts.provider.ID, []uuid.UUID{alice.ID, bob.ID}) require.NoError(ts.T(), err) require.Len(ts.T(), groups, 3) @@ -333,7 +333,7 @@ func (ts *SCIMGroupTestSuite) TestFindGroupsForUsers() { require.Equal(ts.T(), []string{"Admins", "Engineering"}, byUser[alice.ID]) require.Equal(ts.T(), []string{"Engineering"}, byUser[bob.ID]) - groups, err = FindSCIMGroupsForUsers(ts.db, ts.createProvider().ID, []uuid.UUID{alice.ID}) + groups, err = FindSCIMMembershipsByUser(ts.db, ts.createProvider().ID, []uuid.UUID{alice.ID}) require.NoError(ts.T(), err) require.Empty(ts.T(), groups) } diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index 15e9f4cd08..b53aeab768 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -45,7 +45,7 @@ type SCIMIdentityEmailChange struct { } var scimUsersTable = scimTable{ - name: SCIMUser{}.TableName(), + tableName: SCIMUser{}.TableName(), label: "SCIM user", columns: scimUserColumns, nameColumn: "user_name", @@ -82,7 +82,7 @@ func ReplaceSCIMUser(tx *storage.Connection, target SCIMTarget, resource []byte) func DeleteSCIMUser(tx *storage.Connection, target SCIMTarget) (*SCIMUser, error) { user := &SCIMUser{} err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE id = ? AND sso_provider_id = ? AND deleted_at IS NULL AND "+scimVersionClause+" RETURNING "+scimUserColumns, scimUsersTable.name), + fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE id = ? AND sso_provider_id = ? AND deleted_at IS NULL AND "+scimVersionClause+" RETURNING "+scimUserColumns, scimUsersTable.tableName), target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(user) if err != nil { @@ -101,7 +101,7 @@ func FindSCIMUserLinks(tx *storage.Connection, ids []uuid.UUID) (map[uuid.UUID]u UserID uuid.UUID `db:"user_id"` }{} if err := tx.RawQuery( - fmt.Sprintf("SELECT id, user_id FROM %q WHERE id = ANY(?::uuid[]) AND user_id IS NOT NULL", scimUsersTable.name), + fmt.Sprintf("SELECT id, user_id FROM %q WHERE id = ANY(?::uuid[]) AND user_id IS NOT NULL", scimUsersTable.tableName), ids, ).All(&rows); err != nil { return nil, errors.Wrap(err, "error finding SCIM user links") @@ -122,7 +122,7 @@ func LockUserForSCIM(tx *storage.Connection, userID uuid.UUID) error { return nil } -func LogoutSCIMUser(tx *storage.Connection, userID uuid.UUID) error { +func LogoutUserForSCIM(tx *storage.Connection, userID uuid.UUID) error { if err := LockUserForSCIM(tx, userID); err != nil { return err } @@ -132,7 +132,7 @@ func LogoutSCIMUser(tx *storage.Connection, userID uuid.UUID) error { func SoftDeleteSCIMUsersByUserID(tx *storage.Connection, userID uuid.UUID) ([]SCIMUser, error) { rows := []SCIMUser{} if err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE user_id = ? AND deleted_at IS NULL RETURNING "+scimUserColumns, scimUsersTable.name), + fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE user_id = ? AND deleted_at IS NULL RETURNING "+scimUserColumns, scimUsersTable.tableName), userID, ).All(&rows); err != nil { return nil, errors.Wrap(err, "error deleting SCIM users by user id") @@ -141,7 +141,7 @@ func SoftDeleteSCIMUsersByUserID(tx *storage.Connection, userID uuid.UUID) ([]SC } func BanDeprovisionedSCIMUsers(tx *storage.Connection, providerID uuid.UUID, until time.Time) (int, error) { - users, scimUsers := User{}.TableName(), scimUsersTable.name + users, scimUsers := User{}.TableName(), scimUsersTable.tableName count, err := tx.RawQuery( fmt.Sprintf( "UPDATE %[1]q u SET banned_until = ?, updated_at = now() "+ @@ -171,7 +171,7 @@ func LinkSCIMUser(tx *storage.Connection, user *SCIMUser, userID uuid.UUID) erro fmt.Sprintf( "SELECT EXISTS(SELECT 1 FROM %[1]q WHERE sso_provider_id = ? AND user_id = ? AND deleted_at IS NULL) AS live, "+ "EXISTS(SELECT 1 FROM %[1]q WHERE sso_provider_id = ? AND user_id = ? AND deleted_at IS NOT NULL) AS deleted", - scimUsersTable.name, + scimUsersTable.tableName, ), user.SSOProviderID, userID, user.SSOProviderID, userID, ).First(&existing); err != nil { @@ -185,7 +185,7 @@ func LinkSCIMUser(tx *storage.Connection, user *SCIMUser, userID uuid.UUID) erro } if err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET user_id = ? WHERE id = ?", scimUsersTable.name), + fmt.Sprintf("UPDATE %q SET user_id = ? WHERE id = ?", scimUsersTable.tableName), userID, user.ID, ).Exec(); err != nil { return errors.Wrap(err, "error linking SCIM user") @@ -202,7 +202,7 @@ func IsSCIMManaged(tx *storage.Connection, providerID, userID uuid.UUID) (bool, return managed, nil } -func IsSCIMDeprovisioned(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { +func IsSCIMUserDeprovisionedByProvider(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { result := struct { AnyRow bool `db:"any_row"` Live bool `db:"live"` @@ -211,7 +211,7 @@ func IsSCIMDeprovisioned(tx *storage.Connection, providerID, userID uuid.UUID) ( fmt.Sprintf( "SELECT EXISTS(SELECT 1 FROM %[1]q WHERE sso_provider_id = ? AND user_id = ?) AS any_row, "+ "EXISTS(SELECT 1 FROM %[1]q WHERE sso_provider_id = ? AND user_id = ? AND deleted_at IS NULL AND active) AS live", - scimUsersTable.name, + scimUsersTable.tableName, ), providerID, userID, providerID, userID, ).First(&result); err != nil { @@ -235,7 +235,7 @@ func IsSCIMUserDeprovisionedForUpdate(tx *storage.Connection, userID uuid.UUID) DeletedAt *time.Time `db:"deleted_at"` }{} if err := tx.RawQuery( - fmt.Sprintf("SELECT active, deleted_at FROM %q WHERE user_id = ?", scimUsersTable.name), + fmt.Sprintf("SELECT active, deleted_at FROM %q WHERE user_id = ?", scimUsersTable.tableName), userID, ).All(&rows); err != nil { return false, errors.Wrap(err, "error finding SCIM users") From 70c15b640d8d8ff1ba173a3d150b741fb67ff133 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 09:58:25 -0600 Subject: [PATCH 76/88] chore(scim): share SCIM replace, delete, lock and audit helpers --- internal/api/scim_admin.go | 30 ++++++++++++------------------ internal/api/scim_errors.go | 2 +- internal/api/scim_groups.go | 7 +------ internal/api/scim_user_linking.go | 28 +++++++++++++++------------- internal/api/scim_users.go | 5 +---- internal/models/errors.go | 14 ++++++++++++++ internal/models/linking.go | 7 +++++-- internal/models/scim.go | 13 +++++++++++-- internal/models/scim_group.go | 6 +++--- internal/models/scim_token.go | 2 +- internal/models/scim_user.go | 6 +++--- 11 files changed, 67 insertions(+), 53 deletions(-) diff --git a/internal/api/scim_admin.go b/internal/api/scim_admin.go index d3c0008d74..3328c643c5 100644 --- a/internal/api/scim_admin.go +++ b/internal/api/scim_admin.go @@ -111,12 +111,7 @@ func (a *API) adminSCIMTokensCreate(w http.ResponseWriter, r *http.Request) erro if token, plaintext, err = models.CreateSCIMToken(tx, provider, params.ExpiresAt); err != nil { return err } - return a.auditSCIM(tx, r, scimAuditEvent{ - actor: getAdminUser(ctx), - action: models.SCIMTokenCreatedAction, - providerID: provider.ID, - traits: map[string]any{scimTokenPrefixTrait: token.Prefix}, - }) + return a.auditSCIM(tx, r, scimTokenAudit(getAdminUser(ctx), models.SCIMTokenCreatedAction, token)) }); err != nil { if errors.Is(err, models.SCIMTokenExpiryError{}) { return apierrors.NewBadRequestError(apierrors.ErrorCodeValidationFailed, "expires_at must be in the future") @@ -190,12 +185,7 @@ func (a *API) revokeSCIMToken(tx *storage.Connection, r *http.Request, providerI if err := token.Revoke(tx); err != nil { return nil, err } - return token, a.auditSCIM(tx, r, scimAuditEvent{ - actor: getAdminUser(r.Context()), - action: models.SCIMTokenRevokedAction, - providerID: providerID, - traits: map[string]any{scimTokenPrefixTrait: token.Prefix}, - }) + return token, a.auditSCIM(tx, r, scimTokenAudit(getAdminUser(r.Context()), models.SCIMTokenRevokedAction, token)) } func (a *API) sendSCIMStatus(w http.ResponseWriter, db *storage.Connection, provider *models.SSOProvider) error { @@ -263,18 +253,22 @@ func (a *API) revokeActiveSCIMTokens(tx *storage.Connection, r *http.Request, ac if err := tokens[i].Revoke(tx); err != nil { return nil, err } - if err := a.auditSCIM(tx, r, scimAuditEvent{ - actor: actor, - action: models.SCIMTokenRevokedAction, - providerID: providerID, - traits: map[string]any{scimTokenPrefixTrait: tokens[i].Prefix}, - }); err != nil { + if err := a.auditSCIM(tx, r, scimTokenAudit(actor, models.SCIMTokenRevokedAction, &tokens[i])); err != nil { return nil, err } } return prefixes, nil } +func scimTokenAudit(actor *models.User, action models.AuditAction, token *models.SCIMToken) scimAuditEvent { + return scimAuditEvent{ + actor: actor, + action: action, + providerID: token.SSOProviderID, + traits: map[string]any{scimTokenPrefixTrait: token.Prefix}, + } +} + func (a *API) auditSCIMDisabled(tx *storage.Connection, r *http.Request, providerID uuid.UUID, prefixes []string) error { return a.auditSCIM(tx, r, scimAuditEvent{ actor: getAdminUser(r.Context()), diff --git a/internal/api/scim_errors.go b/internal/api/scim_errors.go index cea0e53956..99f846ae7d 100644 --- a/internal/api/scim_errors.go +++ b/internal/api/scim_errors.go @@ -13,7 +13,7 @@ func scimError(err error) error { switch { case models.IsNotFoundError(err): return errSCIMNotFound() - case errors.Is(err, models.SCIMUserStaleError{}), errors.Is(err, models.SCIMGroupStaleError{}): + case models.IsStaleError(err): return errSCIMStale() case errors.Is(err, models.SCIMGroupConflictError{}): return scimerrors.ErrUniqueness(`"externalId" must be unique`) diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index f0a048b747..2c6a03468e 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -64,12 +64,7 @@ func (s *scimGroupRepository) Replace(ctx context.Context, group *core.Group) (* return nil, err } return s.save(ctx, models.SCIMGroupUpdatedAction, group, func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error) { - unchanged, err := models.FindUnchangedSCIMGroup(tx, target, resource) - if err != nil || unchanged != nil { - return unchanged, false, err - } - row, err := models.ReplaceSCIMGroup(tx, target, resource) - return row, true, err + return models.ReplaceSCIMGroupIfChanged(tx, target, resource) }) } diff --git a/internal/api/scim_user_linking.go b/internal/api/scim_user_linking.go index 0ad33f808c..a242f0d3ac 100644 --- a/internal/api/scim_user_linking.go +++ b/internal/api/scim_user_linking.go @@ -36,34 +36,36 @@ func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SC return nil, false, err } - linked := decision.User + linked, isNew := decision.User, false switch decision.Decision { - case models.AccountExists: + case models.AccountExists, models.LinkAccount: if err = scimRequireSSOUser(linked); err != nil { return nil, false, err } - case models.LinkAccount: - if err = scimRequireSSOUser(linked); err != nil { - return nil, false, err - } - if _, err = s.api.createNewIdentity(tx, linked, providerType, scimIdentityData(user)); err != nil { - return nil, false, err - } - if err = linked.UpdateAppMetaDataProviders(tx); err != nil { - return nil, false, err + if decision.Decision == models.LinkAccount { + if err = s.linkIdentity(tx, linked, providerType, user); err != nil { + return nil, false, err + } } case models.CreateAccount: if linked, err = s.createAuthUser(tx, providerType, decision, user); err != nil { return nil, false, err } - return linked, true, models.LinkSCIMUser(tx, row, linked.ID) + isNew = true case models.MultipleAccounts: return nil, false, scimerrors.ErrUniqueness("multiple users share this email in the SSO provider") default: return nil, false, apierrors.NewInternalServerError("Unknown automatic linking decision: %v", decision.Decision) } - return linked, false, models.LinkSCIMUser(tx, row, linked.ID) + return linked, isNew, models.LinkSCIMUser(tx, row, linked.ID) +} + +func (s *scimUserRepository) linkIdentity(tx *storage.Connection, linked *models.User, providerType string, user *core.User) error { + if _, err := s.api.createNewIdentity(tx, linked, providerType, scimIdentityData(user)); err != nil { + return err + } + return linked.UpdateAppMetaDataProviders(tx) } func scimRequireSSOUser(linked *models.User) error { diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index a4837f677c..db34ff9af5 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -265,10 +265,7 @@ func (s *scimUserRepository) lockForReplace(tx *storage.Connection, target model func (s *scimUserRepository) replaceRow(tx *storage.Connection, target models.SCIMTarget, old *models.SCIMUser, resource []byte) (*models.SCIMUser, bool, error) { if old.UserID != nil { - unchanged, err := models.FindUnchangedSCIMUser(tx, target, resource) - if err != nil || unchanged != nil { - return unchanged, false, err - } + return models.ReplaceSCIMUserIfChanged(tx, target, resource) } row, err := models.ReplaceSCIMUser(tx, target, resource) return row, err == nil, err diff --git a/internal/models/errors.go b/internal/models/errors.go index 67716c51b7..a7238498d4 100644 --- a/internal/models/errors.go +++ b/internal/models/errors.go @@ -9,6 +9,8 @@ import ( // sentinel error for all not found errors. var errNotFound = errors.New("not found") +var errStale = errors.New("stale") + // sentinel error for unique constraint violations. var errUniqueConstraintViolated = errors.New("unique constraint violated") @@ -17,6 +19,10 @@ func IsNotFoundError(err error) bool { return errors.Is(err, errNotFound) } +func IsStaleError(err error) bool { + return errors.Is(err, errStale) +} + type SessionNotFoundError struct{} func (e SessionNotFoundError) Error() string { @@ -248,6 +254,10 @@ func (e SCIMUserStaleError) Error() string { return "SCIM user has changed since it was read" } +func (e SCIMUserStaleError) Is(target error) bool { + return target == errStale +} + type SCIMUserConflictError struct{} func (e SCIMUserConflictError) Error() string { @@ -282,6 +292,10 @@ func (e SCIMGroupStaleError) Error() string { return "SCIM group has changed since it was read" } +func (e SCIMGroupStaleError) Is(target error) bool { + return target == errStale +} + type SCIMGroupConflictError struct{} func (e SCIMGroupConflictError) Error() string { diff --git a/internal/models/linking.go b/internal/models/linking.go index e9a61a5d70..990520baf6 100644 --- a/internal/models/linking.go +++ b/internal/models/linking.go @@ -238,9 +238,12 @@ func LockAccountLinkingEmails(tx *storage.Connection, providerType string, email return nil } +func advisoryXactLock(tx *storage.Connection, key string) error { + return tx.RawQuery("SELECT pg_advisory_xact_lock(hashtextextended(?, 0))", key).Exec() +} + func LockAccountLinking(tx *storage.Connection, providerType, email string) error { - key := providerType + "|" + strings.ToLower(email) - if err := tx.RawQuery("SELECT pg_advisory_xact_lock(hashtextextended(?, 0))", key).Exec(); err != nil { + if err := advisoryXactLock(tx, providerType+"|"+strings.ToLower(email)); err != nil { return errors.Wrap(err, "error locking account linking") } return nil diff --git a/internal/models/scim.go b/internal/models/scim.go index 1767cab7d5..27e18ebeda 100644 --- a/internal/models/scim.go +++ b/internal/models/scim.go @@ -62,7 +62,7 @@ func findSCIMPage[T any](tx *storage.Connection, table scimTable, providerID uui where, args := table.where(providerID, query.Filter) total, err := tx.Q().Where(where, args...).Count(new(T)) if err != nil { - return nil, 0, errors.Wrapf(err, "error counting %s", table.tableName) + return nil, 0, errors.Wrapf(err, "error counting %ss", table.label) } rows := []T{} if query.Limit <= 0 || query.Offset >= total { @@ -72,7 +72,7 @@ func findSCIMPage[T any](tx *storage.Connection, table scimTable, providerID uui fmt.Sprintf("SELECT %s FROM %q WHERE %s ORDER BY %s OFFSET ? LIMIT ?", table.columns, table.tableName, where, table.orderBy(query.Order)), append(args, query.Offset, query.Limit)..., ).All(&rows); err != nil { - return nil, 0, errors.Wrapf(err, "error finding %s", table.tableName) + return nil, 0, errors.Wrapf(err, "error finding %ss", table.label) } return rows, total, nil } @@ -117,6 +117,15 @@ func findUnchangedSCIMRow[T any](tx *storage.Connection, table scimTable, target return row, nil } +func replaceSCIMRowIfChanged[T any](tx *storage.Connection, table scimTable, target SCIMTarget, resource []byte) (*T, bool, error) { + unchanged, err := findUnchangedSCIMRow[T](tx, table, target, resource) + if err != nil || unchanged != nil { + return unchanged, false, err + } + row, err := replaceSCIMRow[T](tx, table, target, resource) + return row, err == nil, err +} + func replaceSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarget, resource []byte) (*T, error) { row := new(T) if err := tx.RawQuery( diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go index f21630066f..03fbc5c10f 100644 --- a/internal/models/scim_group.go +++ b/internal/models/scim_group.go @@ -71,8 +71,8 @@ func ReplaceSCIMGroup(tx *storage.Connection, target SCIMTarget, resource []byte return replaceSCIMRow[SCIMGroup](tx, scimGroupsTable, target, resource) } -func FindUnchangedSCIMGroup(tx *storage.Connection, target SCIMTarget, resource []byte) (*SCIMGroup, error) { - return findUnchangedSCIMRow[SCIMGroup](tx, scimGroupsTable, target, resource) +func ReplaceSCIMGroupIfChanged(tx *storage.Connection, target SCIMTarget, resource []byte) (*SCIMGroup, bool, error) { + return replaceSCIMRowIfChanged[SCIMGroup](tx, scimGroupsTable, target, resource) } func TouchSCIMGroup(tx *storage.Connection, group *SCIMGroup) (*SCIMGroup, error) { @@ -89,7 +89,7 @@ func TouchSCIMGroup(tx *storage.Connection, group *SCIMGroup) (*SCIMGroup, error func DeleteSCIMGroup(tx *storage.Connection, target SCIMTarget) (*SCIMGroup, error) { group := &SCIMGroup{} err := tx.RawQuery( - fmt.Sprintf("DELETE FROM %q WHERE id = ? AND sso_provider_id = ? AND "+scimVersionClause+" RETURNING "+scimGroupColumns, scimGroupsTable.tableName), + fmt.Sprintf("DELETE FROM %q WHERE %s AND "+scimVersionClause+" RETURNING %s", scimGroupsTable.tableName, scimGroupsTable.targetClause(), scimGroupsTable.columns), target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(group) if err != nil { diff --git a/internal/models/scim_token.go b/internal/models/scim_token.go index 78bafdc7ab..5de0f2cf6e 100644 --- a/internal/models/scim_token.go +++ b/internal/models/scim_token.go @@ -94,7 +94,7 @@ func FindActiveSCIMTokensBySSOProvider(tx *storage.Connection, providerID uuid.U } func LockSCIMTokens(tx *storage.Connection, providerID uuid.UUID) error { - if err := tx.RawQuery("SELECT pg_advisory_xact_lock(hashtextextended(?, 0))", "scim_tokens|"+providerID.String()).Exec(); err != nil { + if err := advisoryXactLock(tx, "scim_tokens|"+providerID.String()); err != nil { return errors.Wrap(err, "error locking SCIM tokens") } return nil diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index b53aeab768..3fd7a13495 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -67,8 +67,8 @@ func FindSCIMUserForUpdate(tx *storage.Connection, providerID, id uuid.UUID) (*S return findSCIMRow[SCIMUser](tx, scimUsersTable, SCIMTarget{ProviderID: providerID, ID: id}, true) } -func FindUnchangedSCIMUser(tx *storage.Connection, target SCIMTarget, resource []byte) (*SCIMUser, error) { - return findUnchangedSCIMRow[SCIMUser](tx, scimUsersTable, target, resource) +func ReplaceSCIMUserIfChanged(tx *storage.Connection, target SCIMTarget, resource []byte) (*SCIMUser, bool, error) { + return replaceSCIMRowIfChanged[SCIMUser](tx, scimUsersTable, target, resource) } func FindSCIMUsers(tx *storage.Connection, providerID uuid.UUID, query SCIMQuery) ([]SCIMUser, int, error) { @@ -82,7 +82,7 @@ func ReplaceSCIMUser(tx *storage.Connection, target SCIMTarget, resource []byte) func DeleteSCIMUser(tx *storage.Connection, target SCIMTarget) (*SCIMUser, error) { user := &SCIMUser{} err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE id = ? AND sso_provider_id = ? AND deleted_at IS NULL AND "+scimVersionClause+" RETURNING "+scimUserColumns, scimUsersTable.tableName), + fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE %s AND "+scimVersionClause+" RETURNING %s", scimUsersTable.tableName, scimUsersTable.targetClause(), scimUsersTable.columns), target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(user) if err != nil { From 7e421afd06545e5723f5dbf799883831110a78d6 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 09:59:40 -0600 Subject: [PATCH 77/88] fix(scim): render SCIM write responses inside the write transaction --- internal/api/scim_groups.go | 16 ++++++++-------- internal/api/scim_users.go | 12 ++++++++---- 2 files changed, 16 insertions(+), 12 deletions(-) diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index 2c6a03468e..ff07582204 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -118,22 +118,22 @@ func (s *scimGroupRepository) save(ctx context.Context, action models.AuditActio } change := scimGroupChange{r: r, action: action, members: members} db := s.api.db.WithContext(ctx) - var row *models.SCIMGroup + var saved *core.Group err = db.Transaction(func(tx *storage.Connection) error { - var ( - changed bool - terr error - ) - if row, changed, terr = write(tx, resource); terr != nil { + row, changed, terr := write(tx, resource) + if terr != nil { return terr } - row, terr = s.applyMembers(tx, change, row, changed) + if row, terr = s.applyMembers(tx, change, row, changed); terr != nil { + return terr + } + saved, terr = s.renderOne(tx, row.SSOProviderID, row, protocol.Projection{}) return terr }) if err != nil { return nil, scimError(err) } - return s.renderOne(db, row.SSOProviderID, row, protocol.Projection{}) + return saved, nil } func (s *scimGroupRepository) applyMembers(tx *storage.Connection, change scimGroupChange, row *models.SCIMGroup, changed bool) (*models.SCIMGroup, error) { diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index db34ff9af5..3c3bc181b3 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -124,19 +124,23 @@ func (s *scimUserRepository) Delete(ctx context.Context, id, version string) err func (s *scimUserRepository) save(db *storage.Connection, change scimUserChange, write scimUserWrite) (*core.User, error) { var ( - row *models.SCIMUser + saved *core.User created *models.User ) err := db.Transaction(func(tx *storage.Connection) error { - var terr error - row, created, terr = write(tx, change) + row, user, terr := write(tx, change) + if terr != nil { + return terr + } + created = user + saved, terr = s.renderOne(tx, change.target.ProviderID, row, protocol.Projection{}) return terr }) if err != nil { return nil, scimError(err) } s.runAfterUserCreatedHook(change.r, db, created) - return s.renderOne(db, change.target.ProviderID, row, protocol.Projection{}) + return saved, nil } func (s *scimUserRepository) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMUser, projection protocol.Projection) ([]*core.User, error) { From f85f80838cef726be83e23880b1d525a38d056a6 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 10:01:25 -0600 Subject: [PATCH 78/88] fix(scim): refuse a SCIM user write when its link changed after the read --- internal/api/scim_link_test.go | 31 +++++++++++++++++++++++++++++++ internal/api/scim_users.go | 14 ++++++++++++-- 2 files changed, 43 insertions(+), 2 deletions(-) diff --git a/internal/api/scim_link_test.go b/internal/api/scim_link_test.go index 295851d7cd..1a150e006d 100644 --- a/internal/api/scim_link_test.go +++ b/internal/api/scim_link_test.go @@ -8,6 +8,7 @@ import ( "sync" "time" + "github.com/gofrs/uuid" "github.com/sirupsen/logrus" logrustest "github.com/sirupsen/logrus/hooks/test" "github.com/stretchr/testify/require" @@ -1030,3 +1031,33 @@ func (ts *SCIMTestSuite) TestRenameSkippedWhenIdentityMissing() { require.NotEqual(ts.T(), user.ID, signedIn.ID) require.Equal(ts.T(), 2, ts.users("alice@example.com")) } + +func (ts *SCIMTestSuite) relinkBehindRead(id string) (*scimUserRepository, models.SCIMTarget, *models.SCIMUser) { + existing, err := models.FindSCIMUser(ts.API.db, ts.A.ID, uuid.FromStringOrNil(id)) + require.NoError(ts.T(), err) + other := ts.ssoUser(ts.A, "relinked@example.com", "relinked@example.com") + require.NoError(ts.T(), ts.API.db.RawQuery("UPDATE "+existing.TableName()+" SET user_id = ? WHERE id = ?", other.ID, existing.ID).Exec()) + target := models.SCIMTarget{ProviderID: ts.A.ID, ID: existing.ID, UpdatedAt: &existing.UpdatedAt} + return &scimUserRepository{api: ts.API}, target, existing +} + +func (ts *SCIMTestSuite) TestReplaceRefusesUserRelinkedAfterRead() { + repo, target, existing := ts.relinkBehindRead(ts.create(ts.TokenA, scimUser("alice"))) + err := ts.API.db.Transaction(func(tx *storage.Connection) error { + _, err := repo.lockForReplace(tx, target, "alice@example.com", existing) + return err + }) + require.ErrorIs(ts.T(), err, models.SCIMUserStaleError{}) +} + +func (ts *SCIMTestSuite) TestDeleteRefusesUserRelinkedAfterRead() { + id := ts.create(ts.TokenA, scimUser("alice")) + repo, target, existing := ts.relinkBehindRead(id) + r := httptest.NewRequest(http.MethodDelete, "/scim/v2/Users/"+id, nil) + err := ts.API.db.Transaction(func(tx *storage.Connection) error { + return repo.delete(tx, r, target, existing) + }) + require.ErrorIs(ts.T(), err, models.SCIMUserStaleError{}) + _, err = models.FindSCIMUser(ts.API.db, ts.A.ID, existing.ID) + require.NoError(ts.T(), err) +} diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index 3c3bc181b3..46b91cd751 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -241,6 +241,9 @@ func (s *scimUserRepository) delete(tx *storage.Connection, r *http.Request, tar if err != nil { return err } + if !sameSCIMLink(row.UserID, existing.UserID) { + return models.SCIMUserStaleError{} + } if err := logoutSCIMLinkedUser(tx, row.UserID); err != nil { return err } @@ -261,12 +264,19 @@ func (s *scimUserRepository) lockForReplace(tx *storage.Connection, target model if err != nil { return nil, err } - if err := lockSCIMLinkedUser(tx, old.UserID); err != nil { - return nil, err + if !sameSCIMLink(old.UserID, existing.UserID) { + return nil, models.SCIMUserStaleError{} } return old, nil } +func sameSCIMLink(a, b *uuid.UUID) bool { + if a == nil || b == nil { + return a == b + } + return *a == *b +} + func (s *scimUserRepository) replaceRow(tx *storage.Connection, target models.SCIMTarget, old *models.SCIMUser, resource []byte) (*models.SCIMUser, bool, error) { if old.UserID != nil { return models.ReplaceSCIMUserIfChanged(tx, target, resource) From 4064f4a863b21ba7d6954e4a50b6b3d7dfa42e89 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 10:02:19 -0600 Subject: [PATCH 79/88] fix(scim): read only member ids and build the audit actor once per group write --- internal/api/scim_groups.go | 3 ++- internal/models/scim_group.go | 10 +++------- 2 files changed, 5 insertions(+), 8 deletions(-) diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index ff07582204..20011c9632 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -246,6 +246,7 @@ func (s *scimGroupRepository) memberEvents(tx *storage.Connection, r *http.Reque if err != nil { return nil, err } + actor := scimActor(r) events := make([]scimAuditEvent, 0, len(added)+len(removed)) for _, change := range []struct { action models.AuditAction @@ -260,7 +261,7 @@ func (s *scimGroupRepository) memberEvents(tx *storage.Connection, r *http.Reque userID = &linked } events = append(events, scimAuditEvent{ - actor: scimActor(r), + actor: actor, action: change.action, providerID: row.SSOProviderID, traits: scimMemberTraits(row.ID, id, userID), diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go index 03fbc5c10f..7275c8894b 100644 --- a/internal/models/scim_group.go +++ b/internal/models/scim_group.go @@ -189,17 +189,13 @@ func lockLiveSCIMUserIDs(tx *storage.Connection, providerID uuid.UUID, ids []uui } func findSCIMGroupMemberIDs(tx *storage.Connection, groupID uuid.UUID) ([]uuid.UUID, error) { - members := []SCIMGroupMember{} + ids := []uuid.UUID{} if err := tx.RawQuery( - fmt.Sprintf("SELECT group_id, scim_user_id, created_at FROM %q WHERE group_id = ?", SCIMGroupMember{}.TableName()), + fmt.Sprintf("SELECT scim_user_id FROM %q WHERE group_id = ?", SCIMGroupMember{}.TableName()), groupID, - ).All(&members); err != nil { + ).All(&ids); err != nil { return nil, errors.Wrap(err, "error finding SCIM group members") } - ids := make([]uuid.UUID, len(members)) - for i := range members { - ids[i] = members[i].SCIMUserID - } return ids, nil } From 19c7dc8b3838a45ffec8f642539bc90173c14c9c Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 10:07:10 -0600 Subject: [PATCH 80/88] fix(scim): write SCIM user delete and token revoke audit entries in one insert --- internal/api/scim_admin.go | 7 +++---- internal/api/scim_user_cleanup.go | 34 +++++++++++++++---------------- internal/api/scim_users.go | 5 +++-- 3 files changed, 23 insertions(+), 23 deletions(-) diff --git a/internal/api/scim_admin.go b/internal/api/scim_admin.go index 3328c643c5..8886278d4d 100644 --- a/internal/api/scim_admin.go +++ b/internal/api/scim_admin.go @@ -248,16 +248,15 @@ func (a *API) revokeActiveSCIMTokens(tx *storage.Connection, r *http.Request, ac return nil, err } prefixes := make([]string, len(tokens)) + events := make([]scimAuditEvent, len(tokens)) for i := range tokens { prefixes[i] = tokens[i].Prefix if err := tokens[i].Revoke(tx); err != nil { return nil, err } - if err := a.auditSCIM(tx, r, scimTokenAudit(actor, models.SCIMTokenRevokedAction, &tokens[i])); err != nil { - return nil, err - } + events[i] = scimTokenAudit(actor, models.SCIMTokenRevokedAction, &tokens[i]) } - return prefixes, nil + return prefixes, a.auditSCIMEvents(tx, r, events) } func scimTokenAudit(actor *models.User, action models.AuditAction, token *models.SCIMToken) scimAuditEvent { diff --git a/internal/api/scim_user_cleanup.go b/internal/api/scim_user_cleanup.go index dee8109c91..b222f6de6f 100644 --- a/internal/api/scim_user_cleanup.go +++ b/internal/api/scim_user_cleanup.go @@ -13,35 +13,35 @@ func (a *API) deleteSCIMUsers(tx *storage.Connection, r *http.Request, actor *mo if err != nil { return err } + events := []scimAuditEvent{} for i := range rows { - if err := a.removeSCIMUserFromGroups(tx, r, actor, &rows[i]); err != nil { - return err - } - if err := a.auditSCIM(tx, r, scimAuditEvent{ - actor: actor, - action: models.SCIMUserDeletedAction, - providerID: rows[i].SSOProviderID, - traits: scimUserTraits(&rows[i]), - }); err != nil { + removed, err := scimUserRemovalEvents(tx, actor, &rows[i]) + if err != nil { return err } + events = append(events, removed...) } - return nil + return a.auditSCIMEvents(tx, r, events) } -func (a *API) removeSCIMUserFromGroups(tx *storage.Connection, r *http.Request, actor *models.User, row *models.SCIMUser) error { +func scimUserRemovalEvents(tx *storage.Connection, actor *models.User, row *models.SCIMUser) ([]scimAuditEvent, error) { groupIDs, err := models.RemoveSCIMUserFromGroups(tx, row.ID) if err != nil { - return err + return nil, err } - events := make([]scimAuditEvent, len(groupIDs)) - for i, groupID := range groupIDs { - events[i] = scimAuditEvent{ + events := make([]scimAuditEvent, 0, len(groupIDs)+1) + for _, groupID := range groupIDs { + events = append(events, scimAuditEvent{ actor: actor, action: models.SCIMGroupMemberRemovedAction, providerID: row.SSOProviderID, traits: scimMemberTraits(groupID, row.ID, row.UserID), - } + }) } - return a.auditSCIMEvents(tx, r, events) + return append(events, scimAuditEvent{ + actor: actor, + action: models.SCIMUserDeletedAction, + providerID: row.SSOProviderID, + traits: scimUserTraits(row), + }), nil } diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index 46b91cd751..dcc14ee14d 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -247,10 +247,11 @@ func (s *scimUserRepository) delete(tx *storage.Connection, r *http.Request, tar if err := logoutSCIMLinkedUser(tx, row.UserID); err != nil { return err } - if err := s.api.removeSCIMUserFromGroups(tx, r, scimActor(r), row); err != nil { + events, err := scimUserRemovalEvents(tx, scimActor(r), row) + if err != nil { return err } - return s.audit(tx, r, models.SCIMUserDeletedAction, row) + return s.api.auditSCIMEvents(tx, r, events) } func (s *scimUserRepository) lockForReplace(tx *storage.Connection, target models.SCIMTarget, email string, existing *models.SCIMUser) (*models.SCIMUser, error) { From 5f79c9338e6d507f1bde5762f18b12feac834926 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 10:08:16 -0600 Subject: [PATCH 81/88] fix(scim): skip the SCIM list count query when the page is not full --- internal/api/scim_users_test.go | 5 +++++ internal/models/scim.go | 23 +++++++++++++---------- 2 files changed, 18 insertions(+), 10 deletions(-) diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index 550252298c..f63b35a5a1 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -564,6 +564,11 @@ func (ts *SCIMTestSuite) TestPagination() { require.Equal(ts.T(), http.StatusOK, w.Code) require.EqualValues(ts.T(), 3, page["totalResults"]) require.Empty(ts.T(), page["Resources"]) + + w, page = ts.do(ts.TokenA, http.MethodGet, "/Users?startIndex=2&count=5", "") + require.Equal(ts.T(), http.StatusOK, w.Code) + require.EqualValues(ts.T(), 3, page["totalResults"]) + require.Len(ts.T(), page["Resources"], 2) } func (ts *SCIMTestSuite) TestPageSizeCap() { diff --git a/internal/models/scim.go b/internal/models/scim.go index 27e18ebeda..68f841eed0 100644 --- a/internal/models/scim.go +++ b/internal/models/scim.go @@ -3,6 +3,7 @@ package models import ( "database/sql" "fmt" + "slices" "strings" "time" @@ -60,20 +61,22 @@ type scimTable struct { func findSCIMPage[T any](tx *storage.Connection, table scimTable, providerID uuid.UUID, query SCIMQuery) ([]T, int, error) { where, args := table.where(providerID, query.Filter) + rows := []T{} + if query.Limit > 0 { + if err := tx.RawQuery( + fmt.Sprintf("SELECT %s FROM %q WHERE %s ORDER BY %s OFFSET ? LIMIT ?", table.columns, table.tableName, where, table.orderBy(query.Order)), + append(slices.Clone(args), query.Offset, query.Limit)..., + ).All(&rows); err != nil { + return nil, 0, errors.Wrapf(err, "error finding %ss", table.label) + } + if len(rows) < query.Limit && (query.Offset == 0 || len(rows) > 0) { + return rows, query.Offset + len(rows), nil + } + } total, err := tx.Q().Where(where, args...).Count(new(T)) if err != nil { return nil, 0, errors.Wrapf(err, "error counting %ss", table.label) } - rows := []T{} - if query.Limit <= 0 || query.Offset >= total { - return rows, total, nil - } - if err := tx.RawQuery( - fmt.Sprintf("SELECT %s FROM %q WHERE %s ORDER BY %s OFFSET ? LIMIT ?", table.columns, table.tableName, where, table.orderBy(query.Order)), - append(args, query.Offset, query.Limit)..., - ).All(&rows); err != nil { - return nil, 0, errors.Wrapf(err, "error finding %ss", table.label) - } return rows, total, nil } From e02297e369839766be52f133ec20b8ecbc46debb Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 10:09:31 -0600 Subject: [PATCH 82/88] fix(scim): skip group and link checks for a SCIM user created in the same request --- internal/api/scim_user_linking.go | 5 ++++- internal/api/scim_users.go | 14 +++++++++++++- internal/api/scim_users_test.go | 10 ++++++++++ internal/models/scim_user.go | 3 +++ 4 files changed, 30 insertions(+), 2 deletions(-) diff --git a/internal/api/scim_user_linking.go b/internal/api/scim_user_linking.go index a242f0d3ac..5053b349cb 100644 --- a/internal/api/scim_user_linking.go +++ b/internal/api/scim_user_linking.go @@ -58,7 +58,10 @@ func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SC return nil, false, apierrors.NewInternalServerError("Unknown automatic linking decision: %v", decision.Decision) } - return linked, isNew, models.LinkSCIMUser(tx, row, linked.ID) + if isNew { + return linked, true, models.LinkNewSCIMUser(tx, row, linked.ID) + } + return linked, false, models.LinkSCIMUser(tx, row, linked.ID) } func (s *scimUserRepository) linkIdentity(tx *storage.Connection, linked *models.User, providerType string, user *core.User) error { diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index dcc14ee14d..037569041f 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "net/http" + "net/url" "strings" "github.com/badoux/checkmail" @@ -133,7 +134,7 @@ func (s *scimUserRepository) save(db *storage.Connection, change scimUserChange, return terr } created = user - saved, terr = s.renderOne(tx, change.target.ProviderID, row, protocol.Projection{}) + saved, terr = s.renderOne(tx, change.target.ProviderID, row, scimSavedUserProjection(change)) return terr }) if err != nil { @@ -194,6 +195,17 @@ func (s *scimUserRepository) groupMemberships(tx *storage.Connection, providerID return groups, nil } +func scimSavedUserProjection(change scimUserChange) protocol.Projection { + if change.target.ID != uuid.Nil { + return protocol.Projection{} + } + projection, err := protocol.ParseProjection(url.Values{"excludedAttributes": {"groups"}}, scimUserSchemas) + if err != nil { + return protocol.Projection{} + } + return projection +} + func (s *scimUserRepository) renderOne(tx *storage.Connection, providerID uuid.UUID, row *models.SCIMUser, projection protocol.Projection) (*core.User, error) { users, err := s.render(tx, providerID, []models.SCIMUser{*row}, projection) if err != nil { diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index f63b35a5a1..cfca9ef7ac 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -758,3 +758,13 @@ func withField(body, field string, value any) string { func patchOp(ops ...string) string { return `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[` + strings.Join(ops, ",") + `]}` } + +func TestSCIMSavedUserProjectionSkipsGroupsOnCreate(t *testing.T) { + created := scimUserChange{target: models.SCIMTarget{ProviderID: uuid.Must(uuid.NewV4())}} + require.False(t, scimSavedUserProjection(created).Returns("groups")) + require.True(t, scimSavedUserProjection(created).Returns("userName")) + + replaced := created + replaced.target.ID = uuid.Must(uuid.NewV4()) + require.True(t, scimSavedUserProjection(replaced).Returns("groups")) +} diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index 3fd7a13495..c8695fe768 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -183,7 +183,10 @@ func LinkSCIMUser(tx *storage.Connection, user *SCIMUser, userID uuid.UUID) erro if existing.Deleted { return SCIMUserDeletedError{} } + return LinkNewSCIMUser(tx, user, userID) +} +func LinkNewSCIMUser(tx *storage.Connection, user *SCIMUser, userID uuid.UUID) error { if err := tx.RawQuery( fmt.Sprintf("UPDATE %q SET user_id = ? WHERE id = ?", scimUsersTable.tableName), userID, user.ID, From 5066ed4bde27e8497bbf4d8ebe141a3db94e0d1c Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 10:15:55 -0600 Subject: [PATCH 83/88] fix(scim): render SCIM group writes from the member diff instead of rereading members --- internal/api/admin_test.go | 2 +- internal/api/scim_groups.go | 70 ++++++++++++++++++------------ internal/api/scim_groups_test.go | 21 ++++++++- internal/models/scim_group.go | 23 +++++++--- internal/models/scim_group_test.go | 63 ++++++++++++++++++++------- 5 files changed, 128 insertions(+), 51 deletions(-) diff --git a/internal/api/admin_test.go b/internal/api/admin_test.go index ad19d8f595..969eee7ce8 100644 --- a/internal/api/admin_test.go +++ b/internal/api/admin_test.go @@ -909,7 +909,7 @@ func (ts *AdminTestSuite) TestAdminUserDeleteSoftDeletesSCIMUser() { scimUser, u := ts.createLinkedSCIMUser(c.wantEmail) group, err := models.CreateSCIMGroup(ts.API.db, scimUser.SSOProviderID, []byte(`{"displayName":"Engineering"}`)) require.NoError(ts.T(), err) - _, _, err = models.ReplaceSCIMGroupMembers(ts.API.db, group, []uuid.UUID{scimUser.ID}) + _, _, _, err = models.ReplaceSCIMGroupMembers(ts.API.db, group, []uuid.UUID{scimUser.ID}) require.NoError(ts.T(), err) var buffer bytes.Buffer diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index 20011c9632..359331fee7 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -82,7 +82,7 @@ func (s *scimGroupRepository) Delete(ctx context.Context, id, version string) er if err != nil { return err } - _, removed, err := models.ReplaceSCIMGroupMembers(tx, row, nil) + removed, err := models.ClearSCIMGroupMembers(tx, row.ID) if err != nil { return err } @@ -124,10 +124,11 @@ func (s *scimGroupRepository) save(ctx context.Context, action models.AuditActio if terr != nil { return terr } - if row, terr = s.applyMembers(tx, change, row, changed); terr != nil { + row, members, terr := s.applyMembers(tx, change, row, changed) + if terr != nil { return terr } - saved, terr = s.renderOne(tx, row.SSOProviderID, row, protocol.Projection{}) + saved, terr = s.compose(*row, scimMembers(scimBaseURL(s.api.config), members)) return terr }) if err != nil { @@ -136,10 +137,10 @@ func (s *scimGroupRepository) save(ctx context.Context, action models.AuditActio return saved, nil } -func (s *scimGroupRepository) applyMembers(tx *storage.Connection, change scimGroupChange, row *models.SCIMGroup, changed bool) (*models.SCIMGroup, error) { - added, removed, err := models.ReplaceSCIMGroupMembers(tx, row, change.members) +func (s *scimGroupRepository) applyMembers(tx *storage.Connection, change scimGroupChange, row *models.SCIMGroup, changed bool) (*models.SCIMGroup, []uuid.UUID, error) { + added, removed, members, err := models.ReplaceSCIMGroupMembers(tx, row, change.members) if err != nil { - return nil, err + return nil, nil, err } events := []scimAuditEvent{} switch { @@ -150,16 +151,16 @@ func (s *scimGroupRepository) applyMembers(tx *storage.Connection, change scimGr case len(added) > 0 || len(removed) > 0: row, err = models.TouchSCIMGroup(tx, row) default: - return row, nil + return row, members, nil } if err != nil { - return nil, err + return nil, nil, err } - members, err := s.memberEvents(tx, change.r, row, added, removed) + memberEvents, err := s.memberEvents(tx, change.r, row, added, removed) if err != nil { - return nil, err + return nil, nil, err } - return row, s.api.auditSCIMEvents(tx, change.r, append(events, members...)) + return row, members, s.api.auditSCIMEvents(tx, change.r, append(events, memberEvents...)) } func (s *scimGroupRepository) render(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMGroup, projection protocol.Projection) ([]*core.Group, error) { @@ -167,22 +168,29 @@ func (s *scimGroupRepository) render(tx *storage.Connection, providerID uuid.UUI if err != nil { return nil, err } - base := scimBaseURL(s.api.config) groups := make([]*core.Group, 0, len(rows)) for _, row := range rows { - group := &core.Group{} - if err := json.Unmarshal(row.Resource, group); err != nil { + group, err := s.compose(row, members[row.ID]) + if err != nil { return nil, err } - group.ID = row.ID.String() - group.Schemas = []core.SchemaURI{core.SchemaGroup} - group.Meta = scimMeta(scimResourceTypeGroup, base+"/Groups/"+group.ID, row.CreatedAt, row.UpdatedAt) - group.Members = members[row.ID] groups = append(groups, group) } return groups, nil } +func (s *scimGroupRepository) compose(row models.SCIMGroup, members []core.Member) (*core.Group, error) { + group := &core.Group{} + if err := json.Unmarshal(row.Resource, group); err != nil { + return nil, err + } + group.ID = row.ID.String() + group.Schemas = []core.SchemaURI{core.SchemaGroup} + group.Meta = scimMeta(scimResourceTypeGroup, scimBaseURL(s.api.config)+"/Groups/"+group.ID, row.CreatedAt, row.UpdatedAt) + group.Members = members + return group, nil +} + func (s *scimGroupRepository) members(tx *storage.Connection, providerID uuid.UUID, rows []models.SCIMGroup, projection protocol.Projection) (map[uuid.UUID][]core.Member, error) { members := map[uuid.UUID][]core.Member{} if !projection.Returns("members") { @@ -200,21 +208,29 @@ func (s *scimGroupRepository) members(tx *storage.Connection, providerID uuid.UU for _, m := range memberships { counts[m.GroupID]++ } - base := scimBaseURL(s.api.config) + byGroup := make(map[uuid.UUID][]uuid.UUID, len(counts)) for _, m := range memberships { - if members[m.GroupID] == nil { - members[m.GroupID] = make([]core.Member, 0, counts[m.GroupID]) + if byGroup[m.GroupID] == nil { + byGroup[m.GroupID] = make([]uuid.UUID, 0, counts[m.GroupID]) } - id := m.SCIMUserID.String() - members[m.GroupID] = append(members[m.GroupID], core.Member{ - Value: id, - Ref: base + "/Users/" + id, - Type: scimResourceTypeUser, - }) + byGroup[m.GroupID] = append(byGroup[m.GroupID], m.SCIMUserID) + } + base := scimBaseURL(s.api.config) + for groupID, scimUserIDs := range byGroup { + members[groupID] = scimMembers(base, scimUserIDs) } return members, nil } +func scimMembers(base string, scimUserIDs []uuid.UUID) []core.Member { + members := make([]core.Member, len(scimUserIDs)) + for i, scimUserID := range scimUserIDs { + id := scimUserID.String() + members[i] = core.Member{Value: id, Ref: base + "/Users/" + id, Type: scimResourceTypeUser} + } + return members +} + func (s *scimGroupRepository) renderOne(tx *storage.Connection, providerID uuid.UUID, row *models.SCIMGroup, projection protocol.Projection) (*core.Group, error) { groups, err := s.render(tx, providerID, []models.SCIMGroup{*row}, projection) if err != nil { diff --git a/internal/api/scim_groups_test.go b/internal/api/scim_groups_test.go index 87c12fb665..d6c138bfba 100644 --- a/internal/api/scim_groups_test.go +++ b/internal/api/scim_groups_test.go @@ -403,7 +403,7 @@ func (ts *SCIMTestSuite) TestUserDeleteWaitsForGroupWrite() { return err }, func(tx *storage.Connection) error { - _, _, err := models.ReplaceSCIMGroupMembers(tx, group, []uuid.UUID{uuid.FromStringOrNil(bob)}) + _, _, _, err := models.ReplaceSCIMGroupMembers(tx, group, []uuid.UUID{uuid.FromStringOrNil(bob)}) return err }, http.MethodDelete, "/Users/"+alice, "", @@ -630,3 +630,22 @@ func (ts *SCIMTestSuite) requireGroup(step, group, displayName string, members [ require.Equal(ts.T(), displayName, got["displayName"], step) require.ElementsMatch(ts.T(), members, memberValues(got), step) } + +func (ts *SCIMTestSuite) TestWriteResponseMembersMatchGet() { + ids := []string{} + for _, name := range []string{"a", "b", "c", "d"} { + ids = append(ids, ts.create(ts.TokenA, scimUser(name))) + } + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", ids[2], ids[0])) + requireMatchesGet := func(method, body string) { + w, written := ts.do(ts.TokenA, method, "/Groups/"+id, body) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + require.Equal(ts.T(), got["members"], written["members"]) + } + requireMatchesGet(http.MethodPatch, patchOp(`{"op":"add","path":"members","value":[{"value":"`+ids[3]+`"},{"value":"`+ids[1]+`"}]}`)) + requireMatchesGet(http.MethodPatch, patchOp(`{"op":"remove","path":"members[value eq \"`+ids[2]+`\"]"}`)) + requireMatchesGet(http.MethodPut, groupWith("Engineering", "", ids[1], ids[2], ids[3])) + requireMatchesGet(http.MethodPatch, patchOp(`{"op":"replace","path":"displayName","value":"Platform"}`)) +} diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go index 7275c8894b..a70e67e711 100644 --- a/internal/models/scim_group.go +++ b/internal/models/scim_group.go @@ -128,20 +128,31 @@ func FindSCIMMembershipsByUser(tx *storage.Connection, providerID uuid.UUID, sci return memberships, nil } -func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserIDs []uuid.UUID) (added, removed []uuid.UUID, err error) { +func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserIDs []uuid.UUID) (added, removed, members []uuid.UUID, err error) { current, err := findSCIMGroupMemberIDs(tx, group.ID) if err != nil { - return nil, nil, err + return nil, nil, nil, err } removed = differenceUUIDs(current, scimUserIDs) if err := removeSCIMGroupMembers(tx, group.ID, removed); err != nil { - return nil, nil, err + return nil, nil, nil, err } added, err = addSCIMGroupMembers(tx, group, differenceUUIDs(scimUserIDs, current)) if err != nil { - return nil, nil, err + return nil, nil, nil, err } - return added, removed, nil + return added, removed, append(differenceUUIDs(current, removed), added...), nil +} + +func ClearSCIMGroupMembers(tx *storage.Connection, groupID uuid.UUID) ([]uuid.UUID, error) { + removed := []uuid.UUID{} + if err := tx.RawQuery( + fmt.Sprintf("DELETE FROM %q WHERE group_id = ? RETURNING scim_user_id", SCIMGroupMember{}.TableName()), + groupID, + ).All(&removed); err != nil { + return nil, errors.Wrap(err, "error removing SCIM group members") + } + return removed, nil } func RemoveSCIMUserFromGroups(tx *storage.Connection, scimUserID uuid.UUID) ([]uuid.UUID, error) { @@ -191,7 +202,7 @@ func lockLiveSCIMUserIDs(tx *storage.Connection, providerID uuid.UUID, ids []uui func findSCIMGroupMemberIDs(tx *storage.Connection, groupID uuid.UUID) ([]uuid.UUID, error) { ids := []uuid.UUID{} if err := tx.RawQuery( - fmt.Sprintf("SELECT scim_user_id FROM %q WHERE group_id = ?", SCIMGroupMember{}.TableName()), + fmt.Sprintf("SELECT scim_user_id FROM %q WHERE group_id = ? ORDER BY created_at, scim_user_id", SCIMGroupMember{}.TableName()), groupID, ).All(&ids); err != nil { return nil, errors.Wrap(err, "error finding SCIM group members") diff --git a/internal/models/scim_group_test.go b/internal/models/scim_group_test.go index fe2830bc9d..ccf9eea47c 100644 --- a/internal/models/scim_group_test.go +++ b/internal/models/scim_group_test.go @@ -1,6 +1,7 @@ package models import ( + "bytes" "fmt" "testing" "time" @@ -125,7 +126,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceChecksVersion() { func (ts *SCIMGroupTestSuite) TestDeleteRemovesMembers() { group := ts.createGroup(ts.provider.ID, "Engineering") user := ts.createUser(ts.provider.ID, "alice") - _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{user.ID}) + _, _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{user.ID}) require.NoError(ts.T(), err) _, err = DeleteSCIMGroup(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: group.ID}) @@ -165,12 +166,12 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersDiffs() { bob := ts.createUser(ts.provider.ID, "bob") carol := ts.createUser(ts.provider.ID, "carol") - added, removed, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, bob.ID, alice.ID}) + added, removed, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, bob.ID, alice.ID}) require.NoError(ts.T(), err) require.ElementsMatch(ts.T(), []uuid.UUID{alice.ID, bob.ID}, added) require.Empty(ts.T(), removed) - added, removed, err = ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{bob.ID, carol.ID}) + added, removed, _, err = ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{bob.ID, carol.ID}) require.NoError(ts.T(), err) require.Equal(ts.T(), []uuid.UUID{carol.ID}, added) require.Equal(ts.T(), []uuid.UUID{alice.ID}, removed) @@ -180,7 +181,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersDiffs() { require.Len(ts.T(), members, 2) require.ElementsMatch(ts.T(), []uuid.UUID{bob.ID, carol.ID}, []uuid.UUID{members[0].SCIMUserID, members[1].SCIMUserID}) - added, removed, err = ReplaceSCIMGroupMembers(ts.db, group, nil) + added, removed, _, err = ReplaceSCIMGroupMembers(ts.db, group, nil) require.NoError(ts.T(), err) require.Empty(ts.T(), added) require.ElementsMatch(ts.T(), []uuid.UUID{bob.ID, carol.ID}, removed) @@ -191,7 +192,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersRejectsOtherProviderUsers() { alice := ts.createUser(ts.provider.ID, "alice") outsider := ts.createUser(ts.createProvider().ID, "mallory") - _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, outsider.ID}) + _, _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, outsider.ID}) require.Equal(ts.T(), SCIMGroupMemberNotFoundError{IDs: []uuid.UUID{outsider.ID}}, err) members, err := FindSCIMMembershipsByGroup(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) @@ -205,10 +206,10 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersRejectsDeletedUsers() { _, err := DeleteSCIMUser(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) require.NoError(ts.T(), err) - _, _, err = ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) + _, _, _, err = ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) require.Equal(ts.T(), SCIMGroupMemberNotFoundError{IDs: []uuid.UUID{alice.ID}}, err) - _, _, err = ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{uuid.Must(uuid.NewV4())}) + _, _, _, err = ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{uuid.Must(uuid.NewV4())}) require.ErrorAs(ts.T(), err, &SCIMGroupMemberNotFoundError{}) } @@ -226,7 +227,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersWaitsForConcurrentUserDelete() { result := make(chan error, 1) go func() { result <- ts.db.Transaction(func(tx *storage.Connection) error { - _, _, err := ReplaceSCIMGroupMembers(tx, group, []uuid.UUID{alice.ID}) + _, _, _, err := ReplaceSCIMGroupMembers(tx, group, []uuid.UUID{alice.ID}) return err }) }() @@ -243,7 +244,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersDoesNotLockExistingMembers() { group := ts.createGroup(ts.provider.ID, "Engineering") alice := ts.createUser(ts.provider.ID, "alice") bob := ts.createUser(ts.provider.ID, "bob") - _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) + _, _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) require.NoError(ts.T(), err) deleting := ts.beginTx() @@ -258,7 +259,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersDoesNotLockExistingMembers() { if err != nil { return err } - _, _, err = ReplaceSCIMGroupMembers(tx, locked, []uuid.UUID{alice.ID, bob.ID}) + _, _, _, err = ReplaceSCIMGroupMembers(tx, locked, []uuid.UUID{alice.ID, bob.ID}) return err }) }() @@ -283,7 +284,7 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersDoesNotLockExistingMembers() { func (ts *SCIMGroupTestSuite) TestFindMembersHidesDeletedUsers() { group := ts.createGroup(ts.provider.ID, "Engineering") alice := ts.createUser(ts.provider.ID, "alice") - _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) + _, _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) require.NoError(ts.T(), err) _, err = DeleteSCIMUser(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) @@ -298,17 +299,17 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersValidatesOnlyAddedMembers() { group := ts.createGroup(ts.provider.ID, "Engineering") alice := ts.createUser(ts.provider.ID, "alice") bob := ts.createUser(ts.provider.ID, "bob") - _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) + _, _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}) require.NoError(ts.T(), err) _, err = DeleteSCIMUser(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) require.NoError(ts.T(), err) - added, removed, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, bob.ID}) + added, removed, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, bob.ID}) require.NoError(ts.T(), err) require.Equal(ts.T(), []uuid.UUID{bob.ID}, added) require.Empty(ts.T(), removed) - _, _, err = ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, bob.ID, uuid.Nil}) + _, _, _, err = ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, bob.ID, uuid.Nil}) require.Equal(ts.T(), SCIMGroupMemberNotFoundError{IDs: []uuid.UUID{uuid.Nil}}, err) } @@ -317,9 +318,9 @@ func (ts *SCIMGroupTestSuite) TestFindMembershipsByUser() { admins := ts.createGroup(ts.provider.ID, "Admins") alice := ts.createUser(ts.provider.ID, "alice") bob := ts.createUser(ts.provider.ID, "bob") - _, _, err := ReplaceSCIMGroupMembers(ts.db, engineering, []uuid.UUID{alice.ID, bob.ID}) + _, _, _, err := ReplaceSCIMGroupMembers(ts.db, engineering, []uuid.UUID{alice.ID, bob.ID}) require.NoError(ts.T(), err) - _, _, err = ReplaceSCIMGroupMembers(ts.db, admins, []uuid.UUID{alice.ID}) + _, _, _, err = ReplaceSCIMGroupMembers(ts.db, admins, []uuid.UUID{alice.ID}) require.NoError(ts.T(), err) groups, err := FindSCIMMembershipsByUser(ts.db, ts.provider.ID, []uuid.UUID{alice.ID, bob.ID}) @@ -351,3 +352,33 @@ func (ts *SCIMGroupTestSuite) lockWaiters() int { require.NoError(ts.T(), ts.db.RawQuery("SELECT count(*) AS count FROM pg_stat_activity WHERE datname = current_database() AND wait_event_type = 'Lock'").First(&row)) return row.Count } + +func (ts *SCIMGroupTestSuite) TestReplaceMembersReturnsMembersInRenderOrder() { + group := ts.createGroup(ts.provider.ID, "Engineering") + alice := ts.createUser(ts.provider.ID, "alice") + bob := ts.createUser(ts.provider.ID, "bob") + carol := ts.createUser(ts.provider.ID, "carol") + _, _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, bob.ID}) + require.NoError(ts.T(), err) + first, second := sortedUUIDs(alice.ID, bob.ID) + require.NoError(ts.T(), ts.db.RawQuery("UPDATE "+SCIMGroupMember{}.TableName()+" SET created_at = created_at - interval '1 day' WHERE scim_user_id = ?", second).Exec()) + + _, _, members, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, bob.ID, carol.ID}) + require.NoError(ts.T(), err) + require.Equal(ts.T(), []uuid.UUID{second, first, carol.ID}, members) + + rendered, err := FindSCIMMembershipsByGroup(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) + require.NoError(ts.T(), err) + ids := make([]uuid.UUID, len(rendered)) + for i, m := range rendered { + ids[i] = m.SCIMUserID + } + require.Equal(ts.T(), ids, members) +} + +func sortedUUIDs(a, b uuid.UUID) (uuid.UUID, uuid.UUID) { + if bytes.Compare(a.Bytes(), b.Bytes()) > 0 { + return b, a + } + return a, b +} From 6f6d867210023aad64600b43681b2e610fb729f4 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 10:35:37 -0600 Subject: [PATCH 84/88] chore(scim): fix golangci-lint findings in SCIM code and tests --- internal/api/admin_test.go | 6 +-- internal/api/external_test.go | 2 +- internal/api/scim.go | 3 +- internal/api/scim_admin_test.go | 2 +- internal/api/scim_groups.go | 60 ++++++++++++++----------- internal/api/scim_ratelimit_test.go | 4 +- internal/api/scim_test.go | 2 +- internal/api/scim_user_linking.go | 40 +++++++++-------- internal/api/scim_users.go | 18 +++++--- internal/conf/confload/confload_test.go | 10 ++--- internal/models/audit_log_entry.go | 5 ++- internal/models/audit_log_entry_test.go | 10 ++--- internal/models/scim_group_test.go | 2 +- internal/models/scim_settings_test.go | 2 +- internal/models/scim_token_test.go | 2 +- 15 files changed, 93 insertions(+), 75 deletions(-) diff --git a/internal/api/admin_test.go b/internal/api/admin_test.go index 969eee7ce8..a8103e8d4f 100644 --- a/internal/api/admin_test.go +++ b/internal/api/admin_test.go @@ -889,17 +889,17 @@ func (ts *AdminTestSuite) findSCIMUserByID(id uuid.UUID) *models.SCIMUser { func (ts *AdminTestSuite) TestAdminUserDeleteSoftDeletesSCIMUser() { cases := []struct { desc string - body map[string]interface{} + body map[string]any wantEmail string }{ { desc: "hard delete", - body: map[string]interface{}{"should_soft_delete": false}, + body: map[string]any{"should_soft_delete": false}, wantEmail: "scim-hard-delete@example.com", }, { desc: "soft delete", - body: map[string]interface{}{"should_soft_delete": true}, + body: map[string]any{"should_soft_delete": true}, wantEmail: "scim-soft-delete@example.com", }, } diff --git a/internal/api/external_test.go b/internal/api/external_test.go index f714398927..72d9ce85e4 100644 --- a/internal/api/external_test.go +++ b/internal/api/external_test.go @@ -151,7 +151,7 @@ func (ts *ExternalTestSuite) TestSSOConcurrentCreateSameEmailLinksToOneUser() { count, err := ts.API.db.Q().Where("email = ?", "sso-race@example.com").Count(&models.User{}) require.NoError(ts.T(), err) - require.EqualValues(ts.T(), 1, count) + require.Equal(ts.T(), 1, count) } func (ts *ExternalTestSuite) createUser(providerId string, email string, name string, avatar string, confirmationToken string) (*models.User, error) { diff --git a/internal/api/scim.go b/internal/api/scim.go index c56612d3dc..4f6a30a58c 100644 --- a/internal/api/scim.go +++ b/internal/api/scim.go @@ -20,7 +20,6 @@ import ( "github.com/supabase/auth/internal/models" "github.com/supabase/auth/internal/observability" "github.com/supabase/auth/internal/storage" - "github.com/supabase/auth/internal/utilities" ) const ( @@ -137,7 +136,7 @@ func (a *API) auditSCIMEvents(tx *storage.Connection, r *http.Request, events [] event.traits["outcome"] = "success" entries[i] = models.AuditEvent{Actor: event.actor, Action: event.action, Traits: event.traits} } - return models.NewAuditLogEntries(a.config.AuditLog, r, tx, utilities.GetIPAddress(r), entries) + return models.NewAuditLogEntries(a.config.AuditLog, r, tx, entries) } func scimBaseURL(config *conf.GlobalConfiguration) string { diff --git a/internal/api/scim_admin_test.go b/internal/api/scim_admin_test.go index 056cd41a56..8ddfcfcbb4 100644 --- a/internal/api/scim_admin_test.go +++ b/internal/api/scim_admin_test.go @@ -31,7 +31,7 @@ type SCIMTokensTestSuite struct { func TestSCIMTokens(t *testing.T) { api, config := setupSCIMAPI(t, nil) - defer api.db.Close() + defer func() { require.NoError(t, api.db.Close()) }() suite.Run(t, &SCIMTokensTestSuite{API: api, Config: config}) } diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index 359331fee7..ee5e69546e 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -78,29 +78,33 @@ func (s *scimGroupRepository) Delete(ctx context.Context, id, version string) er return err } return scimError(s.api.db.WithContext(ctx).Transaction(func(tx *storage.Connection) error { - row, err := models.FindSCIMGroupForUpdate(tx, target.ProviderID, target.ID) - if err != nil { - return err - } - removed, err := models.ClearSCIMGroupMembers(tx, row.ID) - if err != nil { - return err - } - if row, err = models.DeleteSCIMGroup(tx, target); err != nil { - return err - } - events, err := s.memberEvents(tx, r, row, nil, removed) - if err != nil { - return err - } - deleted, err := s.groupEvent(r, models.SCIMGroupDeletedAction, row) - if err != nil { - return err - } - return s.api.auditSCIMEvents(tx, r, append(events, deleted)) + return s.delete(tx, r, target) })) } +func (s *scimGroupRepository) delete(tx *storage.Connection, r *http.Request, target models.SCIMTarget) error { + row, err := models.FindSCIMGroupForUpdate(tx, target.ProviderID, target.ID) + if err != nil { + return err + } + removed, err := models.ClearSCIMGroupMembers(tx, row.ID) + if err != nil { + return err + } + if row, err = models.DeleteSCIMGroup(tx, target); err != nil { + return err + } + events, err := s.memberEvents(tx, r, row, scimMemberDiff{removed: removed}) + if err != nil { + return err + } + deleted, err := s.groupEvent(r, models.SCIMGroupDeletedAction, row) + if err != nil { + return err + } + return s.api.auditSCIMEvents(tx, r, append(events, deleted)) +} + func (s *scimGroupRepository) save(ctx context.Context, action models.AuditAction, group *core.Group, write scimGroupWrite) (*core.Group, error) { members, err := scimMemberIDs(group.Members) if err != nil { @@ -156,7 +160,7 @@ func (s *scimGroupRepository) applyMembers(tx *storage.Connection, change scimGr if err != nil { return nil, nil, err } - memberEvents, err := s.memberEvents(tx, change.r, row, added, removed) + memberEvents, err := s.memberEvents(tx, change.r, row, scimMemberDiff{added: added, removed: removed}) if err != nil { return nil, nil, err } @@ -257,19 +261,23 @@ func (s *scimGroupRepository) groupEvent(r *http.Request, action models.AuditAct }, nil } -func (s *scimGroupRepository) memberEvents(tx *storage.Connection, r *http.Request, row *models.SCIMGroup, added, removed []uuid.UUID) ([]scimAuditEvent, error) { - links, err := models.FindSCIMUserLinks(tx, slices.Concat(added, removed)) +type scimMemberDiff struct { + added, removed []uuid.UUID +} + +func (s *scimGroupRepository) memberEvents(tx *storage.Connection, r *http.Request, row *models.SCIMGroup, diff scimMemberDiff) ([]scimAuditEvent, error) { + links, err := models.FindSCIMUserLinks(tx, slices.Concat(diff.added, diff.removed)) if err != nil { return nil, err } actor := scimActor(r) - events := make([]scimAuditEvent, 0, len(added)+len(removed)) + events := make([]scimAuditEvent, 0, len(diff.added)+len(diff.removed)) for _, change := range []struct { action models.AuditAction ids []uuid.UUID }{ - {action: models.SCIMGroupMemberAddedAction, ids: added}, - {action: models.SCIMGroupMemberRemovedAction, ids: removed}, + {action: models.SCIMGroupMemberAddedAction, ids: diff.added}, + {action: models.SCIMGroupMemberRemovedAction, ids: diff.removed}, } { for _, id := range change.ids { var userID *uuid.UUID diff --git a/internal/api/scim_ratelimit_test.go b/internal/api/scim_ratelimit_test.go index 351ac9ecfc..14ffff350e 100644 --- a/internal/api/scim_ratelimit_test.go +++ b/internal/api/scim_ratelimit_test.go @@ -16,7 +16,7 @@ func TestSCIMRateLimit(t *testing.T) { api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { config.RateLimitScim = 1 }) - defer api.db.Close() + defer func() { require.NoError(t, api.db.Close()) }() require.NoError(t, models.TruncateAll(api.db)) token := func() string { @@ -95,7 +95,7 @@ func TestSCIMRateLimitBoundsEveryUnauthenticatedRequest(t *testing.T) { api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { config.RateLimitScim = 1 }) - defer api.db.Close() + defer func() { require.NoError(t, api.db.Close()) }() require.NoError(t, models.TruncateAll(api.db)) type request struct{ path, authorization string } diff --git a/internal/api/scim_test.go b/internal/api/scim_test.go index b9a3ca3108..9e5b248d66 100644 --- a/internal/api/scim_test.go +++ b/internal/api/scim_test.go @@ -476,7 +476,7 @@ func TestSCIMSuite(t *testing.T) { api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { config.RateLimitScim = 1_000_000 }) - defer api.db.Close() + defer func() { require.NoError(t, api.db.Close()) }() suite.Run(t, &SCIMTestSuite{API: api}) } diff --git a/internal/api/scim_user_linking.go b/internal/api/scim_user_linking.go index 5053b349cb..54d9ef3e0f 100644 --- a/internal/api/scim_user_linking.go +++ b/internal/api/scim_user_linking.go @@ -36,32 +36,36 @@ func (s *scimUserRepository) linkAuthUser(tx *storage.Connection, row *models.SC return nil, false, err } - linked, isNew := decision.User, false + if decision.Decision == models.CreateAccount { + linked, err := s.createAuthUser(tx, providerType, decision, user) + if err != nil { + return nil, false, err + } + return linked, true, models.LinkNewSCIMUser(tx, row, linked.ID) + } + linked, err := s.existingAuthUser(tx, providerType, decision, user) + if err != nil { + return nil, false, err + } + return linked, false, models.LinkSCIMUser(tx, row, linked.ID) +} + +func (s *scimUserRepository) existingAuthUser(tx *storage.Connection, providerType string, decision models.AccountLinkingResult, user *core.User) (*models.User, error) { switch decision.Decision { case models.AccountExists, models.LinkAccount: - if err = scimRequireSSOUser(linked); err != nil { - return nil, false, err + if err := scimRequireSSOUser(decision.User); err != nil { + return nil, err } if decision.Decision == models.LinkAccount { - if err = s.linkIdentity(tx, linked, providerType, user); err != nil { - return nil, false, err + if err := s.linkIdentity(tx, decision.User, providerType, user); err != nil { + return nil, err } } - case models.CreateAccount: - if linked, err = s.createAuthUser(tx, providerType, decision, user); err != nil { - return nil, false, err - } - isNew = true + return decision.User, nil case models.MultipleAccounts: - return nil, false, scimerrors.ErrUniqueness("multiple users share this email in the SSO provider") - default: - return nil, false, apierrors.NewInternalServerError("Unknown automatic linking decision: %v", decision.Decision) + return nil, scimerrors.ErrUniqueness("multiple users share this email in the SSO provider") } - - if isNew { - return linked, true, models.LinkNewSCIMUser(tx, row, linked.ID) - } - return linked, false, models.LinkSCIMUser(tx, row, linked.ID) + return nil, apierrors.NewInternalServerError("Unknown automatic linking decision: %v", decision.Decision) } func (s *scimUserRepository) linkIdentity(tx *storage.Connection, linked *models.User, providerType string, user *core.User) error { diff --git a/internal/api/scim_users.go b/internal/api/scim_users.go index 037569041f..f201536d09 100644 --- a/internal/api/scim_users.go +++ b/internal/api/scim_users.go @@ -17,6 +17,12 @@ import ( "github.com/supabase/auth/internal/storage" ) +const ( + scimClaimSub = "sub" + scimClaimEmail = "email" + scimClaimEmailVerified = "email_verified" +) + type scimUserRepository struct { api *API } @@ -328,9 +334,9 @@ func (s *scimUserRepository) renameIdentity(tx *storage.Connection, change scimU if err != nil || from == user.UserName { return err } - data := map[string]any{"sub": user.UserName} + data := map[string]any{scimClaimSub: user.UserName} if email := scimUserEmail(user); email != "" { - data["email"] = email + data[scimClaimEmail] = email } err = models.RenameSCIMIdentity(tx, models.SCIMIdentityRename{ UserID: userID, @@ -365,7 +371,7 @@ func (s *scimUserRepository) changeEmail(tx *storage.Connection, change scimUser if err := linked.ClearAllPendingTokens(tx); err != nil { return err } - return linked.UpdateUserMetaData(tx, map[string]any{"email": email}) + return linked.UpdateUserMetaData(tx, map[string]any{scimClaimEmail: email}) } func (s *scimUserRepository) audit(tx *storage.Connection, r *http.Request, action models.AuditAction, row *models.SCIMUser) error { @@ -437,9 +443,9 @@ func scimPrimaryEmail(emails []core.Email) string { func scimIdentityData(user *core.User) map[string]any { return map[string]any{ - "sub": user.UserName, - "email": scimUserEmail(user), - "email_verified": true, + scimClaimSub: user.UserName, + scimClaimEmail: scimUserEmail(user), + scimClaimEmailVerified: true, } } diff --git a/internal/conf/confload/confload_test.go b/internal/conf/confload/confload_test.go index 443802c5a8..2224d3b510 100644 --- a/internal/conf/confload/confload_test.go +++ b/internal/conf/confload/confload_test.go @@ -281,25 +281,25 @@ func TestSCIMEnabled(t *testing.T) { cfg, err := LoadGlobalFromEnv() require.NoError(t, err) require.NotNil(t, cfg) - assert.Equal(t, false, cfg.SSO.SCIM.Enabled) + assert.False(t, cfg.SSO.SCIM.Enabled) } { baseEnv() - os.Setenv("GOTRUE_SSO_SCIM_ENABLED", "true") + t.Setenv("GOTRUE_SSO_SCIM_ENABLED", "true") cfg, err := LoadGlobalFromEnv() require.NoError(t, err) require.NotNil(t, cfg) - assert.Equal(t, true, cfg.SSO.SCIM.Enabled) + assert.True(t, cfg.SSO.SCIM.Enabled) } { baseEnv() - os.Setenv("GOTRUE_EXPERIMENTAL_SCIM_ENABLED", "true") + t.Setenv("GOTRUE_EXPERIMENTAL_SCIM_ENABLED", "true") cfg, err := LoadGlobalFromEnv() require.NoError(t, err) require.NotNil(t, cfg) - assert.Equal(t, false, cfg.SSO.SCIM.Enabled) + assert.False(t, cfg.SSO.SCIM.Enabled) } } diff --git a/internal/models/audit_log_entry.go b/internal/models/audit_log_entry.go index 37e7159af4..45998ad5af 100644 --- a/internal/models/audit_log_entry.go +++ b/internal/models/audit_log_entry.go @@ -154,10 +154,11 @@ func NewAuditLogEntry(config conf.AuditLogConfiguration, r *http.Request, tx *st type AuditEvent struct { Actor *User Action AuditAction - Traits map[string]interface{} + Traits map[string]any } -func NewAuditLogEntries(config conf.AuditLogConfiguration, r *http.Request, tx *storage.Connection, ipAddress string, events []AuditEvent) error { +func NewAuditLogEntries(config conf.AuditLogConfiguration, r *http.Request, tx *storage.Connection, events []AuditEvent) error { + ipAddress := utilities.GetIPAddress(r) ids := make([]uuid.UUID, len(events)) payloads := make([]string, len(events)) createdAt := make([]time.Time, len(events)) diff --git a/internal/models/audit_log_entry_test.go b/internal/models/audit_log_entry_test.go index 332646d2b3..1aab8c45fe 100644 --- a/internal/models/audit_log_entry_test.go +++ b/internal/models/audit_log_entry_test.go @@ -20,7 +20,7 @@ type AuditLogEntryTestSuite struct { func TestAuditLogEntry(t *testing.T) { ts := &AuditLogEntryTestSuite{db: setupSCIMTestDB(t)} - defer ts.db.Close() + defer func() { require.NoError(t, ts.db.Close()) }() suite.Run(t, ts) } @@ -46,7 +46,7 @@ func (ts *AuditLogEntryTestSuite) entries() []*AuditLogEntry { func (ts *AuditLogEntryTestSuite) TestBatchMatchesSingleEntry() { single := ts.event(SCIMGroupMemberAddedAction, 1) require.NoError(ts.T(), NewAuditLogEntry(conf.AuditLogConfiguration{}, ts.r, ts.db, single.Actor, single.Action, "192.0.2.1", single.Traits)) - require.NoError(ts.T(), NewAuditLogEntries(conf.AuditLogConfiguration{}, ts.r, ts.db, "192.0.2.1", []AuditEvent{ts.event(SCIMGroupMemberAddedAction, 1)})) + require.NoError(ts.T(), NewAuditLogEntries(conf.AuditLogConfiguration{}, ts.r, ts.db, []AuditEvent{ts.event(SCIMGroupMemberAddedAction, 1)})) entries := ts.entries() require.Len(ts.T(), entries, 2) @@ -61,7 +61,7 @@ func (ts *AuditLogEntryTestSuite) TestBatchKeepsEventOrder() { for i := range events { events[i] = ts.event(SCIMGroupMemberRemovedAction, i) } - require.NoError(ts.T(), NewAuditLogEntries(conf.AuditLogConfiguration{}, ts.r, ts.db, "192.0.2.1", events)) + require.NoError(ts.T(), NewAuditLogEntries(conf.AuditLogConfiguration{}, ts.r, ts.db, events)) entries := ts.entries() require.Len(ts.T(), entries, len(events)) @@ -74,12 +74,12 @@ func (ts *AuditLogEntryTestSuite) TestBatchKeepsEventOrder() { } func (ts *AuditLogEntryTestSuite) TestBatchWithoutEvents() { - require.NoError(ts.T(), NewAuditLogEntries(conf.AuditLogConfiguration{}, ts.r, ts.db, "192.0.2.1", nil)) + require.NoError(ts.T(), NewAuditLogEntries(conf.AuditLogConfiguration{}, ts.r, ts.db, nil)) require.Empty(ts.T(), ts.entries()) } func (ts *AuditLogEntryTestSuite) TestBatchSkipsPostgresWhenDisabled() { events := []AuditEvent{ts.event(SCIMGroupMemberAddedAction, 1)} - require.NoError(ts.T(), NewAuditLogEntries(conf.AuditLogConfiguration{DisablePostgres: true}, ts.r, ts.db, "192.0.2.1", events)) + require.NoError(ts.T(), NewAuditLogEntries(conf.AuditLogConfiguration{DisablePostgres: true}, ts.r, ts.db, events)) require.Empty(ts.T(), ts.entries()) } diff --git a/internal/models/scim_group_test.go b/internal/models/scim_group_test.go index ccf9eea47c..86b30d152b 100644 --- a/internal/models/scim_group_test.go +++ b/internal/models/scim_group_test.go @@ -20,7 +20,7 @@ type SCIMGroupTestSuite struct { func TestSCIMGroup(t *testing.T) { ts := &SCIMGroupTestSuite{db: setupSCIMTestDB(t)} - defer ts.db.Close() + defer func() { require.NoError(t, ts.db.Close()) }() suite.Run(t, ts) } diff --git a/internal/models/scim_settings_test.go b/internal/models/scim_settings_test.go index 6bd5300299..26735ea246 100644 --- a/internal/models/scim_settings_test.go +++ b/internal/models/scim_settings_test.go @@ -18,7 +18,7 @@ type SCIMSettingsTestSuite struct { func TestSCIMSettings(t *testing.T) { ts := &SCIMSettingsTestSuite{db: setupSCIMTestDB(t)} - defer ts.db.Close() + defer func() { require.NoError(t, ts.db.Close()) }() suite.Run(t, ts) } diff --git a/internal/models/scim_token_test.go b/internal/models/scim_token_test.go index d4cc6a40b0..8fb47709aa 100644 --- a/internal/models/scim_token_test.go +++ b/internal/models/scim_token_test.go @@ -18,7 +18,7 @@ type SCIMTokenTestSuite struct { func TestSCIMToken(t *testing.T) { ts := &SCIMTokenTestSuite{db: setupSCIMTestDB(t)} - defer ts.db.Close() + defer func() { require.NoError(t, ts.db.Close()) }() suite.Run(t, ts) } From 374e36cdaffa448cbb85525913b6cbf33319b0aa Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 10:40:37 -0600 Subject: [PATCH 85/88] chore(scim): order SCIM helpers after exported functions --- internal/api/admin_test.go | 44 +++---- internal/api/scim_admin_test.go | 158 ++++++++++++------------ internal/api/scim_test.go | 68 +++++----- internal/api/scim_users_test.go | 36 +++--- internal/models/audit_log_entry.go | 68 +++++----- internal/models/audit_log_entry_test.go | 20 +-- internal/models/linking.go | 8 +- internal/models/scim_group_test.go | 60 ++++----- internal/models/scim_token_test.go | 32 ++--- 9 files changed, 247 insertions(+), 247 deletions(-) diff --git a/internal/api/admin_test.go b/internal/api/admin_test.go index a8103e8d4f..1d282451d7 100644 --- a/internal/api/admin_test.go +++ b/internal/api/admin_test.go @@ -864,28 +864,6 @@ func (ts *AdminTestSuite) TestAdminUserDelete() { } } -func (ts *AdminTestSuite) createLinkedSCIMUser(email string) (*models.SCIMUser, *models.User) { - provider := &models.SSOProvider{} - require.NoError(ts.T(), ts.API.db.Create(provider)) - - u, err := models.NewUser("", email, "", ts.Config.JWT.Aud, nil) - require.NoError(ts.T(), err) - u.IsSSOUser = true - require.NoError(ts.T(), ts.API.db.Create(u)) - - scimUser, err := models.CreateSCIMUser(ts.API.db, provider.ID, []byte(`{"userName":"`+email+`"}`)) - require.NoError(ts.T(), err) - require.NoError(ts.T(), models.LinkSCIMUser(ts.API.db, scimUser, u.ID)) - - return scimUser, u -} - -func (ts *AdminTestSuite) findSCIMUserByID(id uuid.UUID) *models.SCIMUser { - var row models.SCIMUser - require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&row)) - return &row -} - func (ts *AdminTestSuite) TestAdminUserDeleteSoftDeletesSCIMUser() { cases := []struct { desc string @@ -1268,3 +1246,25 @@ func (ts *AdminTestSuite) TestAdminUserCreateValidationErrors() { } } + +func (ts *AdminTestSuite) createLinkedSCIMUser(email string) (*models.SCIMUser, *models.User) { + provider := &models.SSOProvider{} + require.NoError(ts.T(), ts.API.db.Create(provider)) + + u, err := models.NewUser("", email, "", ts.Config.JWT.Aud, nil) + require.NoError(ts.T(), err) + u.IsSSOUser = true + require.NoError(ts.T(), ts.API.db.Create(u)) + + scimUser, err := models.CreateSCIMUser(ts.API.db, provider.ID, []byte(`{"userName":"`+email+`"}`)) + require.NoError(ts.T(), err) + require.NoError(ts.T(), models.LinkSCIMUser(ts.API.db, scimUser, u.ID)) + + return scimUser, u +} + +func (ts *AdminTestSuite) findSCIMUserByID(id uuid.UUID) *models.SCIMUser { + var row models.SCIMUser + require.NoError(ts.T(), ts.API.db.Q().Where("id = ?", id).First(&row)) + return &row +} diff --git a/internal/api/scim_admin_test.go b/internal/api/scim_admin_test.go index 8ddfcfcbb4..144e4a2e50 100644 --- a/internal/api/scim_admin_test.go +++ b/internal/api/scim_admin_test.go @@ -45,35 +45,6 @@ func (ts *SCIMTokensTestSuite) SetupTest() { ts.Provider = ts.createProvider() } -func (ts *SCIMTokensTestSuite) createProvider() *models.SSOProvider { - return createSCIMEnabledProvider(ts.T(), ts.API.db) -} - -func (ts *SCIMTokensTestSuite) tokensPath(provider *models.SSOProvider) string { - return "/admin/sso/providers/" + provider.ID.String() + "/scim/tokens" -} - -func (ts *SCIMTokensTestSuite) request(method, path string, body any) *httptest.ResponseRecorder { - return serveAdmin(ts.T(), ts.API, method, path, body) -} - -func (ts *SCIMTokensTestSuite) create(provider *models.SSOProvider, body any) AdminSCIMTokenCreateResponse { - w := ts.request(http.MethodPost, ts.tokensPath(provider), body) - require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) - - var response AdminSCIMTokenCreateResponse - require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &response)) - return response -} - -func (ts *SCIMTokensTestSuite) scimRequest(token string) *httptest.ResponseRecorder { - r := httptest.NewRequest(http.MethodGet, "/scim/v2/Users", nil) - r.Header.Set("Authorization", "Bearer "+token) - w := httptest.NewRecorder() - ts.API.handler.ServeHTTP(w, r) - return w -} - func (ts *SCIMTokensTestSuite) TestCreate() { w := ts.request(http.MethodPost, ts.tokensPath(ts.Provider), map[string]any{}) require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) @@ -280,19 +251,6 @@ func (ts *SCIMTokensTestSuite) TestDisabled() { } } -func (ts *SCIMTokensTestSuite) scimPath(provider *models.SSOProvider) string { - return "/admin/sso/providers/" + provider.ID.String() + "/scim" -} - -func (ts *SCIMTokensTestSuite) status(method string, provider *models.SSOProvider) AdminSCIMStatusResponse { - w := ts.request(method, ts.scimPath(provider), nil) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - - var response AdminSCIMStatusResponse - require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &response)) - return response -} - func (ts *SCIMTokensTestSuite) TestStatus() { status := ts.status(http.MethodGet, createSSOProvider(ts.T(), ts.API.db)) require.False(ts.T(), status.Enabled) @@ -464,15 +422,6 @@ func (ts *SCIMTokensTestSuite) TestConcurrentEnableAndDisable() { require.Equal(ts.T(), []string{string(models.SCIMEnabledAction), string(models.SCIMDisabledAction)}, ts.scimActions(provider)) } -func (ts *SCIMTokensTestSuite) scimActions(provider *models.SSOProvider) []string { - entries := queryAuditEntries(ts.T(), ts.API.db, "payload->>'log_type' = ? AND payload->'traits'->>'sso_provider_id' = ?", "scim", provider.ID.String()) - actions := []string{} - for _, entry := range entries { - actions = append(actions, entry.Payload["action"].(string)) - } - return actions -} - func (ts *SCIMTokensTestSuite) TestStatusForUnknownProvider() { for _, method := range []string{http.MethodGet, http.MethodPost, http.MethodDelete} { w := ts.request(method, "/admin/sso/providers/"+uuid.Must(uuid.NewV4()).String()+"/scim", nil) @@ -482,10 +431,6 @@ func (ts *SCIMTokensTestSuite) TestStatusForUnknownProvider() { require.Empty(ts.T(), ts.tokenEvents()) } -func (ts *SCIMTokensTestSuite) setProviderDisabled(disabled bool) { - require.NoError(ts.T(), ts.API.db.RawQuery("UPDATE "+ts.Provider.TableName()+" SET disabled = ? WHERE id = ?", disabled, ts.Provider.ID).Exec()) -} - func (ts *SCIMTokensTestSuite) TestStatusForDisabledProvider() { first := ts.create(ts.Provider, map[string]any{}) ts.setProviderDisabled(true) @@ -522,30 +467,6 @@ func (ts *SCIMTokensTestSuite) TestStatusForDisabledProvider() { type scimTokenEvent struct{ action, prefix string } -func (ts *SCIMTokensTestSuite) tokenEvents() []scimTokenEvent { - entries := queryAuditEntries(ts.T(), ts.API.db, "payload->>'log_type' = ?", "scim") - - events := []scimTokenEvent{} - for _, entry := range entries { - require.Equal(ts.T(), "supabase_admin", entry.Payload["actor_username"]) - traits := entry.Payload["traits"].(map[string]any) - require.Equal(ts.T(), ts.Provider.ID.String(), traits["sso_provider_id"]) - require.Equal(ts.T(), "success", traits["outcome"]) - prefix, _ := traits["token_prefix"].(string) - if prefixes, ok := traits["token_prefixes"].([]any); ok && len(prefixes) > 0 { - require.Len(ts.T(), prefixes, 1) - prefix = prefixes[0].(string) - } - events = append(events, scimTokenEvent{entry.Payload["action"].(string), prefix}) - } - return events -} - -func (ts *SCIMTokensTestSuite) revoke(prefix string) { - w := ts.request(http.MethodDelete, ts.tokensPath(ts.Provider)+"/"+prefix, nil) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) -} - func (ts *SCIMTokensTestSuite) TestAuditLog() { first := ts.create(ts.Provider, map[string]any{}) second := ts.create(ts.Provider, map[string]any{}) @@ -592,3 +513,82 @@ func (ts *SCIMTokensTestSuite) TestEnableSSODisabledProvider() { require.NoError(ts.T(), ts.API.db.RawQuery("UPDATE "+provider.TableName()+" SET disabled = false WHERE id = ?", provider.ID).Exec()) require.True(ts.T(), ts.status(http.MethodGet, provider).Enabled) } + +func (ts *SCIMTokensTestSuite) createProvider() *models.SSOProvider { + return createSCIMEnabledProvider(ts.T(), ts.API.db) +} + +func (ts *SCIMTokensTestSuite) tokensPath(provider *models.SSOProvider) string { + return "/admin/sso/providers/" + provider.ID.String() + "/scim/tokens" +} + +func (ts *SCIMTokensTestSuite) request(method, path string, body any) *httptest.ResponseRecorder { + return serveAdmin(ts.T(), ts.API, method, path, body) +} + +func (ts *SCIMTokensTestSuite) create(provider *models.SSOProvider, body any) AdminSCIMTokenCreateResponse { + w := ts.request(http.MethodPost, ts.tokensPath(provider), body) + require.Equal(ts.T(), http.StatusCreated, w.Code, w.Body.String()) + + var response AdminSCIMTokenCreateResponse + require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &response)) + return response +} + +func (ts *SCIMTokensTestSuite) scimRequest(token string) *httptest.ResponseRecorder { + r := httptest.NewRequest(http.MethodGet, "/scim/v2/Users", nil) + r.Header.Set("Authorization", "Bearer "+token) + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + return w +} + +func (ts *SCIMTokensTestSuite) scimPath(provider *models.SSOProvider) string { + return "/admin/sso/providers/" + provider.ID.String() + "/scim" +} + +func (ts *SCIMTokensTestSuite) status(method string, provider *models.SSOProvider) AdminSCIMStatusResponse { + w := ts.request(method, ts.scimPath(provider), nil) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + + var response AdminSCIMStatusResponse + require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &response)) + return response +} + +func (ts *SCIMTokensTestSuite) scimActions(provider *models.SSOProvider) []string { + entries := queryAuditEntries(ts.T(), ts.API.db, "payload->>'log_type' = ? AND payload->'traits'->>'sso_provider_id' = ?", "scim", provider.ID.String()) + actions := []string{} + for _, entry := range entries { + actions = append(actions, entry.Payload["action"].(string)) + } + return actions +} + +func (ts *SCIMTokensTestSuite) setProviderDisabled(disabled bool) { + require.NoError(ts.T(), ts.API.db.RawQuery("UPDATE "+ts.Provider.TableName()+" SET disabled = ? WHERE id = ?", disabled, ts.Provider.ID).Exec()) +} + +func (ts *SCIMTokensTestSuite) tokenEvents() []scimTokenEvent { + entries := queryAuditEntries(ts.T(), ts.API.db, "payload->>'log_type' = ?", "scim") + + events := []scimTokenEvent{} + for _, entry := range entries { + require.Equal(ts.T(), "supabase_admin", entry.Payload["actor_username"]) + traits := entry.Payload["traits"].(map[string]any) + require.Equal(ts.T(), ts.Provider.ID.String(), traits["sso_provider_id"]) + require.Equal(ts.T(), "success", traits["outcome"]) + prefix, _ := traits["token_prefix"].(string) + if prefixes, ok := traits["token_prefixes"].([]any); ok && len(prefixes) > 0 { + require.Len(ts.T(), prefixes, 1) + prefix = prefixes[0].(string) + } + events = append(events, scimTokenEvent{entry.Payload["action"].(string), prefix}) + } + return events +} + +func (ts *SCIMTokensTestSuite) revoke(prefix string) { + w := ts.request(http.MethodDelete, ts.tokensPath(ts.Provider)+"/"+prefix, nil) + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) +} diff --git a/internal/api/scim_test.go b/internal/api/scim_test.go index 9e5b248d66..23ee07529a 100644 --- a/internal/api/scim_test.go +++ b/internal/api/scim_test.go @@ -267,40 +267,6 @@ func TestSCIM(t *testing.T) { const scimValidToken = "scim_valid" -func scimFixture(t *testing.T, file string) string { - data, err := fs.ReadFile(os.DirFS("testdata/scim"), file) - require.NoError(t, err) - return string(data) -} - -func newSCIMServerFor(externalURL string) *server.Server { - validate := func(ctx context.Context, candidate string) (context.Context, error) { - if candidate != scimValidToken { - return ctx, server.ErrInvalidToken - } - return ctx, nil - } - return (&API{config: &conf.GlobalConfiguration{API: conf.APIConfiguration{ExternalURL: externalURL}}}).newSCIMServer(validate, nil) -} - -func scimServe(t *testing.T, srv *server.Server, method, path, body string, headers ...string) *httptest.ResponseRecorder { - r := httptest.NewRequest(method, path, strings.NewReader(body)) - r.Header.Set("Content-Type", protocol.MediaType) - r.Header.Set("Authorization", "Bearer "+scimValidToken) - for i := 0; i+1 < len(headers); i += 2 { - r.Header.Set(headers[i], headers[i+1]) - } - w := httptest.NewRecorder() - srv.ServeHTTP(w, r) - return w -} - -func scimDecode(t *testing.T, w *httptest.ResponseRecorder) map[string]any { - var body map[string]any - require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body)) - return body -} - func TestSCIMServer(t *testing.T) { srv := newSCIMServerFor("http://localhost:9999") require.NotNil(t, srv) @@ -574,3 +540,37 @@ func queryAuditEntries(t require.TestingT, db *storage.Connection, where string, require.NoError(t, db.Q().Where(where, args...).Order("created_at asc").All(&entries)) return entries } + +func scimFixture(t *testing.T, file string) string { + data, err := fs.ReadFile(os.DirFS("testdata/scim"), file) + require.NoError(t, err) + return string(data) +} + +func newSCIMServerFor(externalURL string) *server.Server { + validate := func(ctx context.Context, candidate string) (context.Context, error) { + if candidate != scimValidToken { + return ctx, server.ErrInvalidToken + } + return ctx, nil + } + return (&API{config: &conf.GlobalConfiguration{API: conf.APIConfiguration{ExternalURL: externalURL}}}).newSCIMServer(validate, nil) +} + +func scimServe(t *testing.T, srv *server.Server, method, path, body string, headers ...string) *httptest.ResponseRecorder { + r := httptest.NewRequest(method, path, strings.NewReader(body)) + r.Header.Set("Content-Type", protocol.MediaType) + r.Header.Set("Authorization", "Bearer "+scimValidToken) + for i := 0; i+1 < len(headers); i += 2 { + r.Header.Set(headers[i], headers[i+1]) + } + w := httptest.NewRecorder() + srv.ServeHTTP(w, r) + return w +} + +func scimDecode(t *testing.T, w *httptest.ResponseRecorder) map[string]any { + var body map[string]any + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body)) + return body +} diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go index cfca9ef7ac..2e4a7d2b3d 100644 --- a/internal/api/scim_users_test.go +++ b/internal/api/scim_users_test.go @@ -50,10 +50,6 @@ func (ts *SCIMTestSuite) list(token, filter string) map[string]any { return body } -func userWith(userName, externalID string) string { - return `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"` + userName + `","externalId":"` + externalID + `","emails":[{"primary":true,"value":"` + userName + `"}]}` -} - func (ts *SCIMTestSuite) repository() (context.Context, server.Repository[*core.User]) { ctx, err := newSCIMTokenValidator(ts.API.db)(context.Background(), ts.TokenA) require.NoError(ts.T(), err) @@ -61,10 +57,6 @@ func (ts *SCIMTestSuite) repository() (context.Context, server.Repository[*core. return ctx, &scimUserRepository{api: ts.API} } -func emails(value string) []core.Email { - return []core.Email{{Value: value, Primary: new(true)}} -} - func (ts *SCIMTestSuite) TestOktaLifecycle() { require.EqualValues(ts.T(), 0, ts.list(ts.TokenA, `userName eq "alice@example.com"`)["totalResults"]) @@ -739,6 +731,24 @@ func (ts *SCIMTestSuite) TestUsersGroupsAttribute() { require.Equal(ts.T(), eng, groupsOf(got)[0]["value"]) } +func TestSCIMSavedUserProjectionSkipsGroupsOnCreate(t *testing.T) { + created := scimUserChange{target: models.SCIMTarget{ProviderID: uuid.Must(uuid.NewV4())}} + require.False(t, scimSavedUserProjection(created).Returns("groups")) + require.True(t, scimSavedUserProjection(created).Returns("userName")) + + replaced := created + replaced.target.ID = uuid.Must(uuid.NewV4()) + require.True(t, scimSavedUserProjection(replaced).Returns("groups")) +} + +func userWith(userName, externalID string) string { + return `{"schemas":["urn:ietf:params:scim:schemas:core:2.0:User"],"userName":"` + userName + `","externalId":"` + externalID + `","emails":[{"primary":true,"value":"` + userName + `"}]}` +} + +func emails(value string) []core.Email { + return []core.Email{{Value: value, Primary: new(true)}} +} + func oktaUserWith(field string, value any) string { return withField(oktaUser, field, value) } @@ -758,13 +768,3 @@ func withField(body, field string, value any) string { func patchOp(ops ...string) string { return `{"schemas":["urn:ietf:params:scim:api:messages:2.0:PatchOp"],"Operations":[` + strings.Join(ops, ",") + `]}` } - -func TestSCIMSavedUserProjectionSkipsGroupsOnCreate(t *testing.T) { - created := scimUserChange{target: models.SCIMTarget{ProviderID: uuid.Must(uuid.NewV4())}} - require.False(t, scimSavedUserProjection(created).Returns("groups")) - require.True(t, scimSavedUserProjection(created).Returns("userName")) - - replaced := created - replaced.target.ID = uuid.Must(uuid.NewV4()) - require.True(t, scimSavedUserProjection(replaced).Returns("groups")) -} diff --git a/internal/models/audit_log_entry.go b/internal/models/audit_log_entry.go index 45998ad5af..071383c951 100644 --- a/internal/models/audit_log_entry.go +++ b/internal/models/audit_log_entry.go @@ -185,6 +185,40 @@ func NewAuditLogEntries(config conf.AuditLogConfiguration, r *http.Request, tx * return nil } +func FindAuditLogEntries(tx *storage.Connection, filterColumns []string, filterValue string, pageParams *Pagination) ([]*AuditLogEntry, error) { + q := tx.Q().Order("created_at desc").Where("instance_id = ?", uuid.Nil) + + if len(filterColumns) > 0 && filterValue != "" { + lf := "%" + filterValue + "%" + + builder := bytes.NewBufferString("(") + values := make([]any, len(filterColumns)) + + for idx, col := range filterColumns { + fmt.Fprintf(builder, "payload->>'%s' ILIKE ?", col) + values[idx] = lf + + if idx+1 < len(filterColumns) { + builder.WriteString(" OR ") + } + } + builder.WriteString(")") + + q = q.Where(builder.String(), values...) + } + + logs := []*AuditLogEntry{} + var err error + if pageParams != nil { + err = q.Paginate(int(pageParams.Page), int(pageParams.PerPage)).All(&logs) // #nosec G115 + pageParams.Count = uint64(q.Paginator.TotalEntriesSize) // #nosec G115 + } else { + err = q.All(&logs) + } + + return logs, err +} + func buildAuditLogEntry(r *http.Request, event AuditEvent, ipAddress string) (AuditLogEntry, time.Time) { id := uuid.Must(uuid.NewV4()) actor, action, traits := event.Actor, event.Action, event.Traits @@ -253,37 +287,3 @@ func buildAuditLogEntry(r *http.Request, event AuditEvent, ipAddress string) (Au IPAddress: ipAddress, }, createdAt } - -func FindAuditLogEntries(tx *storage.Connection, filterColumns []string, filterValue string, pageParams *Pagination) ([]*AuditLogEntry, error) { - q := tx.Q().Order("created_at desc").Where("instance_id = ?", uuid.Nil) - - if len(filterColumns) > 0 && filterValue != "" { - lf := "%" + filterValue + "%" - - builder := bytes.NewBufferString("(") - values := make([]interface{}, len(filterColumns)) - - for idx, col := range filterColumns { - fmt.Fprintf(builder, "payload->>'%s' ILIKE ?", col) - values[idx] = lf - - if idx+1 < len(filterColumns) { - builder.WriteString(" OR ") - } - } - builder.WriteString(")") - - q = q.Where(builder.String(), values...) - } - - logs := []*AuditLogEntry{} - var err error - if pageParams != nil { - err = q.Paginate(int(pageParams.Page), int(pageParams.PerPage)).All(&logs) // #nosec G115 - pageParams.Count = uint64(q.Paginator.TotalEntriesSize) // #nosec G115 - } else { - err = q.All(&logs) - } - - return logs, err -} diff --git a/internal/models/audit_log_entry_test.go b/internal/models/audit_log_entry_test.go index 1aab8c45fe..9640e68d91 100644 --- a/internal/models/audit_log_entry_test.go +++ b/internal/models/audit_log_entry_test.go @@ -33,16 +33,6 @@ func (ts *AuditLogEntryTestSuite) SetupTest() { ts.actor = actor } -func (ts *AuditLogEntryTestSuite) event(action AuditAction, n int) AuditEvent { - return AuditEvent{Actor: ts.actor, Action: action, Traits: map[string]any{"n": n}} -} - -func (ts *AuditLogEntryTestSuite) entries() []*AuditLogEntry { - entries, err := FindAuditLogEntries(ts.db, nil, "", nil) - require.NoError(ts.T(), err) - return entries -} - func (ts *AuditLogEntryTestSuite) TestBatchMatchesSingleEntry() { single := ts.event(SCIMGroupMemberAddedAction, 1) require.NoError(ts.T(), NewAuditLogEntry(conf.AuditLogConfiguration{}, ts.r, ts.db, single.Actor, single.Action, "192.0.2.1", single.Traits)) @@ -83,3 +73,13 @@ func (ts *AuditLogEntryTestSuite) TestBatchSkipsPostgresWhenDisabled() { require.NoError(ts.T(), NewAuditLogEntries(conf.AuditLogConfiguration{DisablePostgres: true}, ts.r, ts.db, events)) require.Empty(ts.T(), ts.entries()) } + +func (ts *AuditLogEntryTestSuite) event(action AuditAction, n int) AuditEvent { + return AuditEvent{Actor: ts.actor, Action: action, Traits: map[string]any{"n": n}} +} + +func (ts *AuditLogEntryTestSuite) entries() []*AuditLogEntry { + entries, err := FindAuditLogEntries(ts.db, nil, "", nil) + require.NoError(ts.T(), err) + return entries +} diff --git a/internal/models/linking.go b/internal/models/linking.go index 990520baf6..a7aa23ec5a 100644 --- a/internal/models/linking.go +++ b/internal/models/linking.go @@ -238,13 +238,13 @@ func LockAccountLinkingEmails(tx *storage.Connection, providerType string, email return nil } -func advisoryXactLock(tx *storage.Connection, key string) error { - return tx.RawQuery("SELECT pg_advisory_xact_lock(hashtextextended(?, 0))", key).Exec() -} - func LockAccountLinking(tx *storage.Connection, providerType, email string) error { if err := advisoryXactLock(tx, providerType+"|"+strings.ToLower(email)); err != nil { return errors.Wrap(err, "error locking account linking") } return nil } + +func advisoryXactLock(tx *storage.Connection, key string) error { + return tx.RawQuery("SELECT pg_advisory_xact_lock(hashtextextended(?, 0))", key).Exec() +} diff --git a/internal/models/scim_group_test.go b/internal/models/scim_group_test.go index 86b30d152b..1de33aba09 100644 --- a/internal/models/scim_group_test.go +++ b/internal/models/scim_group_test.go @@ -29,22 +29,6 @@ func (ts *SCIMGroupTestSuite) SetupTest() { ts.provider = ts.createProvider() } -func (ts *SCIMGroupTestSuite) createProvider() *SSOProvider { - return createSCIMTestProvider(ts.T(), ts.db) -} - -func (ts *SCIMGroupTestSuite) createGroup(providerID uuid.UUID, displayName string) *SCIMGroup { - group, err := CreateSCIMGroup(ts.db, providerID, []byte(fmt.Sprintf(`{"displayName":%q}`, displayName))) - require.NoError(ts.T(), err) - return group -} - -func (ts *SCIMGroupTestSuite) createUser(providerID uuid.UUID, userName string) *SCIMUser { - user, err := CreateSCIMUser(ts.db, providerID, []byte(fmt.Sprintf(`{"userName":%q}`, userName))) - require.NoError(ts.T(), err) - return user -} - func (ts *SCIMGroupTestSuite) TestCreate() { group, err := CreateSCIMGroup(ts.db, ts.provider.ID, []byte(`{"displayName":"Engineering","externalId":"ext-1"}`)) require.NoError(ts.T(), err) @@ -339,20 +323,6 @@ func (ts *SCIMGroupTestSuite) TestFindMembershipsByUser() { require.Empty(ts.T(), groups) } -func (ts *SCIMGroupTestSuite) beginTx() *storage.Connection { - tx, err := ts.db.NewTransaction() - require.NoError(ts.T(), err) - return &storage.Connection{Connection: tx} -} - -func (ts *SCIMGroupTestSuite) lockWaiters() int { - row := struct { - Count int `db:"count"` - }{} - require.NoError(ts.T(), ts.db.RawQuery("SELECT count(*) AS count FROM pg_stat_activity WHERE datname = current_database() AND wait_event_type = 'Lock'").First(&row)) - return row.Count -} - func (ts *SCIMGroupTestSuite) TestReplaceMembersReturnsMembersInRenderOrder() { group := ts.createGroup(ts.provider.ID, "Engineering") alice := ts.createUser(ts.provider.ID, "alice") @@ -382,3 +352,33 @@ func sortedUUIDs(a, b uuid.UUID) (uuid.UUID, uuid.UUID) { } return a, b } + +func (ts *SCIMGroupTestSuite) createProvider() *SSOProvider { + return createSCIMTestProvider(ts.T(), ts.db) +} + +func (ts *SCIMGroupTestSuite) createGroup(providerID uuid.UUID, displayName string) *SCIMGroup { + group, err := CreateSCIMGroup(ts.db, providerID, []byte(fmt.Sprintf(`{"displayName":%q}`, displayName))) + require.NoError(ts.T(), err) + return group +} + +func (ts *SCIMGroupTestSuite) createUser(providerID uuid.UUID, userName string) *SCIMUser { + user, err := CreateSCIMUser(ts.db, providerID, []byte(fmt.Sprintf(`{"userName":%q}`, userName))) + require.NoError(ts.T(), err) + return user +} + +func (ts *SCIMGroupTestSuite) beginTx() *storage.Connection { + tx, err := ts.db.NewTransaction() + require.NoError(ts.T(), err) + return &storage.Connection{Connection: tx} +} + +func (ts *SCIMGroupTestSuite) lockWaiters() int { + row := struct { + Count int `db:"count"` + }{} + require.NoError(ts.T(), ts.db.RawQuery("SELECT count(*) AS count FROM pg_stat_activity WHERE datname = current_database() AND wait_event_type = 'Lock'").First(&row)) + return row.Count +} diff --git a/internal/models/scim_token_test.go b/internal/models/scim_token_test.go index 8fb47709aa..0ca617b262 100644 --- a/internal/models/scim_token_test.go +++ b/internal/models/scim_token_test.go @@ -27,16 +27,6 @@ func (ts *SCIMTokenTestSuite) SetupTest() { ts.provider = ts.createProvider() } -func (ts *SCIMTokenTestSuite) createProvider() *SSOProvider { - return createSCIMTestProvider(ts.T(), ts.db) -} - -func (ts *SCIMTokenTestSuite) createToken(expiresAt *time.Time) (*SCIMToken, string) { - token, plaintext, err := CreateSCIMToken(ts.db, ts.provider, expiresAt) - require.NoError(ts.T(), err) - return token, plaintext -} - func (ts *SCIMTokenTestSuite) TestCreate() { token, plaintext := ts.createToken(nil) @@ -227,12 +217,6 @@ func (ts *SCIMTokenTestSuite) TestAuthenticateRejects() { } } -func (ts *SCIMTokenTestSuite) expire(token *SCIMToken) { - require.NoError(ts.T(), ts.db.RawQuery( - "UPDATE scim_tokens SET created_at = now() - interval '2 hours', expires_at = now() - interval '1 hour' WHERE id = ?", token.ID, - ).Exec()) -} - func (ts *SCIMTokenTestSuite) TestFindActiveBySSOProvider() { active, _ := ts.createToken(nil) revoked, _ := ts.createToken(nil) @@ -251,3 +235,19 @@ func (ts *SCIMTokenTestSuite) TestFindActiveBySSOProvider() { require.NoError(ts.T(), err) require.Empty(ts.T(), tokens) } + +func (ts *SCIMTokenTestSuite) createProvider() *SSOProvider { + return createSCIMTestProvider(ts.T(), ts.db) +} + +func (ts *SCIMTokenTestSuite) createToken(expiresAt *time.Time) (*SCIMToken, string) { + token, plaintext, err := CreateSCIMToken(ts.db, ts.provider, expiresAt) + require.NoError(ts.T(), err) + return token, plaintext +} + +func (ts *SCIMTokenTestSuite) expire(token *SCIMToken) { + require.NoError(ts.T(), ts.db.RawQuery( + "UPDATE scim_tokens SET created_at = now() - interval '2 hours', expires_at = now() - interval '1 hour' WHERE id = ?", token.ID, + ).Exec()) +} From 62dc31f4ed8d9c92b16ac295eccf8113d3da0809 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 10:46:23 -0600 Subject: [PATCH 86/88] chore(scim): bump scim-go to v0.9.0 --- go.mod | 2 +- go.sum | 4 ++-- hack/scim-demo.sh | 4 ++-- internal/api/scim_groups_test.go | 30 +++++++++++++++++------------- 4 files changed, 22 insertions(+), 18 deletions(-) diff --git a/go.mod b/go.mod index b3e282de34..4b6c4cf765 100644 --- a/go.mod +++ b/go.mod @@ -163,7 +163,7 @@ require ( github.com/spf13/cobra v1.8.1 github.com/standard-webhooks/standard-webhooks/libraries v0.0.0-20240303152453-e0e82adf1721 github.com/stretchr/testify v1.12.1 - github.com/supabase-community/scim-go v0.8.4 + github.com/supabase-community/scim-go v0.9.0 github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 github.com/xeipuuv/gojsonschema v1.2.0 go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.64.0 diff --git a/go.sum b/go.sum index 359c873f60..5fd14da8ca 100644 --- a/go.sum +++ b/go.sum @@ -498,8 +498,8 @@ github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= -github.com/supabase-community/scim-go v0.8.4 h1:9FBo8r9c2u3rL//ZXa8YbTQod5MS4/2l6UjQjoBj9JI= -github.com/supabase-community/scim-go v0.8.4/go.mod h1:oEMij9JuKtAl0wl0jeyIHDKqHvPpUmTXegKeCBKyxXw= +github.com/supabase-community/scim-go v0.9.0 h1:7flgOmRbi67NYEYBILxwio1XzqMFL0pCwRKffMHBqBg= +github.com/supabase-community/scim-go v0.9.0/go.mod h1:oEMij9JuKtAl0wl0jeyIHDKqHvPpUmTXegKeCBKyxXw= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869 h1:VDuRtwen5Z7QQ5ctuHUse4wAv/JozkKZkdic5vUV4Lg= github.com/supabase/hibp v0.0.0-20231124125943-d225752ae869/go.mod h1:eHX5nlSMSnyPjUrbYzeqrA8snCe2SKyfizKjU3dkfOw= github.com/supranational/blst v0.3.16-0.20250831170142-f48500c1fdbe h1:nbdqkIGOGfUAD54q1s2YBcBz/WcsxCO9HUQ4aGV5hUw= diff --git a/hack/scim-demo.sh b/hack/scim-demo.sh index 6df4d1edcb..d8f1226003 100755 --- a/hack/scim-demo.sh +++ b/hack/scim-demo.sh @@ -114,8 +114,8 @@ call 201 POST "$SCIM/Groups" "$SCIM_TOKEN" "$(jq -n --arg n "Tour Guides $RUN" - members: [{value: $m}] }')" GROUP="$(field .id)" -call 200 PATCH "$SCIM/Groups/$GROUP" "$SCIM_TOKEN" "$(patch "{\"op\": \"add\", \"path\": \"members\", \"value\": [{\"value\": \"$OTHER\"}]}")" -call 200 PATCH "$SCIM/Groups/$GROUP" "$SCIM_TOKEN" "$(patch "{\"op\": \"remove\", \"path\": \"members[value eq \\\"$USER\\\"]\"}")" +call 204 PATCH "$SCIM/Groups/$GROUP" "$SCIM_TOKEN" "$(patch "{\"op\": \"add\", \"path\": \"members\", \"value\": [{\"value\": \"$OTHER\"}]}")" +call 204 PATCH "$SCIM/Groups/$GROUP" "$SCIM_TOKEN" "$(patch "{\"op\": \"remove\", \"path\": \"members[value eq \\\"$USER\\\"]\"}")" call 200 GET "$SCIM/Groups/$GROUP" "$SCIM_TOKEN" section "Errors" diff --git a/internal/api/scim_groups_test.go b/internal/api/scim_groups_test.go index d6c138bfba..67325c936e 100644 --- a/internal/api/scim_groups_test.go +++ b/internal/api/scim_groups_test.go @@ -86,8 +86,7 @@ func (ts *SCIMTestSuite) TestGroupsLifecycle() { require.ElementsMatch(ts.T(), []string{alice, bob}, memberValues(patched)) require.Equal(ts.T(), "Platform", patched["displayName"]) - w, patched = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"remove","path":"members[value eq \"`+bob+`\"]"}`)) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + patched = ts.patchMembers(id, `{"op":"remove","path":"members[value eq \"`+bob+`\"]"}`) require.Equal(ts.T(), []string{alice}, memberValues(patched)) w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") @@ -186,9 +185,8 @@ func (ts *SCIMTestSuite) TestGroupMemberEventsCarryUserID() { } entries := ts.auditDuring(func() { id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, bob)) - w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"remove","path":"members[value eq \"`+alice+`\"]"}`)) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - w, _ = ts.do(ts.TokenA, http.MethodDelete, "/Users/"+bob, "") + ts.patchMembers(id, `{"op":"remove","path":"members[value eq \"`+alice+`\"]"}`) + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+bob, "") require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) }) @@ -252,8 +250,7 @@ func (ts *SCIMTestSuite) TestPatchRemoveAbsentMember() { id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice)) before := len(ts.auditActions(models.SCIMGroupMemberRemovedAction)) - w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"remove","path":"members[value eq \"`+bob+`\"]"}`)) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + got := ts.patchMembers(id, `{"op":"remove","path":"members[value eq \"`+bob+`\"]"}`) require.Equal(ts.T(), []string{alice}, memberValues(got)) require.Len(ts.T(), ts.auditActions(models.SCIMGroupMemberRemovedAction), before) } @@ -279,8 +276,7 @@ func (ts *SCIMTestSuite) TestPatchRejectsRemoveWithValue() { require.Equal(ts.T(), http.StatusOK, w.Code) require.NotEmpty(ts.T(), got["emails"]) - w, got = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"remove","path":"members[value eq \"`+bob+`\"]","value":null}`)) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + got = ts.patchMembers(id, `{"op":"remove","path":"members[value eq \"`+bob+`\"]","value":null}`) require.Equal(ts.T(), []string{alice}, memberValues(got)) } @@ -424,8 +420,7 @@ func (ts *SCIMTestSuite) TestGroupsKeepDeactivatedMembers() { w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Users/"+alice, patchOp(`{"op":"replace","path":"active","value":false}`)) require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) - w, got := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(`{"op":"add","path":"members","value":[{"value":"`+bob+`"}]}`)) - require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + got := ts.patchMembers(id, `{"op":"add","path":"members","value":[{"value":"`+bob+`"}]}`) require.ElementsMatch(ts.T(), []string{alice, bob}, memberValues(got)) w, user := ts.do(ts.TokenA, http.MethodGet, "/Users/"+alice, "") @@ -644,8 +639,17 @@ func (ts *SCIMTestSuite) TestWriteResponseMembersMatchGet() { require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) require.Equal(ts.T(), got["members"], written["members"]) } - requireMatchesGet(http.MethodPatch, patchOp(`{"op":"add","path":"members","value":[{"value":"`+ids[3]+`"},{"value":"`+ids[1]+`"}]}`)) - requireMatchesGet(http.MethodPatch, patchOp(`{"op":"remove","path":"members[value eq \"`+ids[2]+`\"]"}`)) + requireMatchesGet(http.MethodPatch, patchOp(`{"op":"add","path":"members","value":[{"value":"`+ids[3]+`"},{"value":"`+ids[1]+`"}]}`, `{"op":"replace","path":"displayName","value":"Platform"}`)) + requireMatchesGet(http.MethodPatch, patchOp(`{"op":"remove","path":"members[value eq \"`+ids[2]+`\"]"}`, `{"op":"replace","path":"displayName","value":"Engineering"}`)) requireMatchesGet(http.MethodPut, groupWith("Engineering", "", ids[1], ids[2], ids[3])) requireMatchesGet(http.MethodPatch, patchOp(`{"op":"replace","path":"displayName","value":"Platform"}`)) } + +func (ts *SCIMTestSuite) patchMembers(id string, ops ...string) map[string]any { + w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(ops...)) + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + require.Empty(ts.T(), w.Body.String()) + w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + require.Equal(ts.T(), http.StatusOK, w.Code, w.Body.String()) + return got +} From b5ead48a2da52c75564eb1a38b720b091582537c Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 11:11:57 -0600 Subject: [PATCH 87/88] fix(scim): apply member-only group PATCH as a delta instead of rereading the group --- internal/api/scim_groups.go | 79 +++++++++++++++++++++++++ internal/api/scim_groups_test.go | 95 ++++++++++++++++++++++++++++++ internal/models/scim.go | 11 ++++ internal/models/scim_group.go | 42 +++++++++---- internal/models/scim_group_test.go | 37 ++++++++++++ 5 files changed, 252 insertions(+), 12 deletions(-) diff --git a/internal/api/scim_groups.go b/internal/api/scim_groups.go index ee5e69546e..c69d0f8923 100644 --- a/internal/api/scim_groups.go +++ b/internal/api/scim_groups.go @@ -5,14 +5,20 @@ import ( "encoding/json" "net/http" "slices" + "strconv" + "strings" "github.com/gofrs/uuid" "github.com/supabase-community/scim-go/pkg/core" "github.com/supabase-community/scim-go/pkg/protocol" + "github.com/supabase-community/scim-go/pkg/scimerrors" + "github.com/supabase-community/scim-go/pkg/server" "github.com/supabase/auth/internal/models" "github.com/supabase/auth/internal/storage" ) +var _ server.AttributePatcher[*core.Group] = (*scimGroupRepository)(nil) + type scimGroupRepository struct { api *API } @@ -68,6 +74,37 @@ func (s *scimGroupRepository) Replace(ctx context.Context, group *core.Group) (* }) } +func (s *scimGroupRepository) PatchAttribute(ctx context.Context, id, version string, delta server.AttributeDelta) (core.Meta, bool, error) { + if !strings.EqualFold(delta.Attribute, "members") { + return core.Meta{}, false, scimerrors.ErrInvalidPath(strconv.Quote(delta.Attribute) + " is unknown") + } + target, err := scimTarget(ctx, id, version) + if err != nil { + return core.Meta{}, false, err + } + add, err := scimAddedMemberIDs(delta.Added) + if err != nil { + return core.Meta{}, false, err + } + r, err := scimRequest(ctx) + if err != nil { + return core.Meta{}, false, err + } + var ( + row *models.SCIMGroup + changed bool + ) + err = s.api.db.WithContext(ctx).Transaction(func(tx *storage.Connection) error { + var terr error + row, changed, terr = s.patchMembers(tx, r, target, scimMemberDiff{added: add, removed: scimRemovedMemberIDs(delta.Removed)}) + return terr + }) + if err != nil { + return core.Meta{}, false, scimError(err) + } + return scimMeta(scimResourceTypeGroup, scimBaseURL(s.api.config)+"/Groups/"+row.ID.String(), row.CreatedAt, row.UpdatedAt), changed, nil +} + func (s *scimGroupRepository) Delete(ctx context.Context, id, version string) error { target, err := scimTarget(ctx, id, version) if err != nil { @@ -105,6 +142,25 @@ func (s *scimGroupRepository) delete(tx *storage.Connection, r *http.Request, ta return s.api.auditSCIMEvents(tx, r, append(events, deleted)) } +func (s *scimGroupRepository) patchMembers(tx *storage.Connection, r *http.Request, target models.SCIMTarget, delta scimMemberDiff) (*models.SCIMGroup, bool, error) { + row, err := models.LockSCIMGroup(tx, target) + if err != nil { + return nil, false, err + } + added, removed, err := models.PatchSCIMGroupMembers(tx, row, delta.added, delta.removed) + if err != nil || len(added)+len(removed) == 0 { + return row, false, err + } + if row, err = models.TouchSCIMGroup(tx, row); err != nil { + return nil, false, err + } + events, err := s.memberEvents(tx, r, row, scimMemberDiff{added: added, removed: removed}) + if err != nil { + return nil, false, err + } + return row, true, s.api.auditSCIMEvents(tx, r, events) +} + func (s *scimGroupRepository) save(ctx context.Context, action models.AuditAction, group *core.Group, write scimGroupWrite) (*core.Group, error) { members, err := scimMemberIDs(group.Members) if err != nil { @@ -306,3 +362,26 @@ func scimMemberIDs(members []core.Member) ([]uuid.UUID, error) { } return ids, nil } + +func scimAddedMemberIDs(added []core.Object) ([]uuid.UUID, error) { + ids := make([]uuid.UUID, 0, len(added)) + for _, object := range added { + value, _ := object.Get("value").(string) + id, err := uuid.FromString(value) + if err != nil { + return nil, errSCIMMemberNotFound() + } + ids = append(ids, id) + } + return ids, nil +} + +func scimRemovedMemberIDs(removed []string) []uuid.UUID { + ids := make([]uuid.UUID, 0, len(removed)) + for _, value := range removed { + if id, err := uuid.FromString(value); err == nil { + ids = append(ids, id) + } + } + return ids +} diff --git a/internal/api/scim_groups_test.go b/internal/api/scim_groups_test.go index 67325c936e..85c4c924dc 100644 --- a/internal/api/scim_groups_test.go +++ b/internal/api/scim_groups_test.go @@ -645,6 +645,101 @@ func (ts *SCIMTestSuite) TestWriteResponseMembersMatchGet() { requireMatchesGet(http.MethodPatch, patchOp(`{"op":"replace","path":"displayName","value":"Platform"}`)) } +func (ts *SCIMTestSuite) TestPatchMembersDeltaMatchesFullPatch() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + carol := ts.create(ts.TokenA, userWith("carol@example.com", "c-1")) + deleted := ts.create(ts.TokenA, userWith("dave@example.com", "d-1")) + w, _ := ts.do(ts.TokenA, http.MethodDelete, "/Users/"+deleted, "") + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + other := ts.create(ts.TokenB, userWith("erin@example.com", "e-1")) + add := func(ids ...string) string { + values := make([]string, len(ids)) + for i, id := range ids { + values[i] = `{"value":"` + id + `"}` + } + return `{"op":"add","path":"members","value":[` + strings.Join(values, ",") + `]}` + } + remove := func(id string) string { + return `{"op":"remove","path":"members[value eq \"` + id + `\"]"}` + } + + type outcome struct { + code int + scimType any + members []string + events []string + versioned bool + } + apply := func(query string, ops []string) outcome { + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, bob)) + w, _ := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + before := w.Header().Get("ETag") + var result outcome + entries := ts.auditDuring(func() { + var got map[string]any + w, got = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id+query, patchOp(ops...)) + result.code, result.scimType = w.Code, got["scimType"] + }) + for _, entry := range entries { + traits := entry.Payload["traits"].(map[string]any) + result.events = append(result.events, entry.Payload["action"].(string)+" "+traits["scim_user_id"].(string)) + } + w, got := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + result.members = memberValues(got) + result.versioned = w.Header().Get("ETag") != before + return result + } + + for name, ops := range map[string][]string{ + "add new": {add(carol)}, + "add present": {add(alice)}, + "add duplicate": {add(carol, carol)}, + "remove present": {remove(alice)}, + "remove absent": {remove(carol)}, + "add and remove": {add(carol), remove(bob)}, + "uppercase": {add(strings.ToUpper(carol)), remove(strings.ToUpper(alice))}, + "remove non uuid": {remove("nope")}, + "add non uuid": {add("nope")}, + "add deleted user": {add(deleted)}, + "add other provider": {add(other)}, + } { + delta, full := apply("", ops), apply("?attributes=members", ops) + if full.code == http.StatusOK { + require.Equal(ts.T(), http.StatusNoContent, delta.code, name) + delta.code = full.code + } + require.Equal(ts.T(), full.code, delta.code, name) + require.Equal(ts.T(), full.scimType, delta.scimType, name) + require.Equal(ts.T(), full.members, delta.members, name) + require.ElementsMatch(ts.T(), full.events, delta.events, name) + require.Equal(ts.T(), full.versioned, delta.versioned, name) + } +} + +func (ts *SCIMTestSuite) TestPatchMembersDeltaChecksIfMatch() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "")) + w, _ := ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + current := w.Header().Get("ETag") + body := patchOp(`{"op":"add","path":"members","value":[{"value":"` + alice + `"}]}`) + + w, _ = ts.doAs(protocol.MediaType, ts.TokenA, http.MethodPatch, "/Groups/"+id, body, "If-Match", current) + require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) + next := w.Header().Get("ETag") + require.NotEqual(ts.T(), current, next) + w, _ = ts.do(ts.TokenA, http.MethodGet, "/Groups/"+id, "") + require.Equal(ts.T(), next, w.Header().Get("ETag")) + + w, _ = ts.doAs(protocol.MediaType, ts.TokenA, http.MethodPatch, "/Groups/"+id, body, "If-Match", current) + require.Equal(ts.T(), http.StatusPreconditionFailed, w.Code, w.Body.String()) + + w, _ = ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+uuid.Must(uuid.NewV4()).String(), body) + require.Equal(ts.T(), http.StatusNotFound, w.Code, w.Body.String()) + w, _ = ts.do(ts.TokenB, http.MethodPatch, "/Groups/"+id, body) + require.Equal(ts.T(), http.StatusNotFound, w.Code, w.Body.String()) +} + func (ts *SCIMTestSuite) patchMembers(id string, ops ...string) map[string]any { w, _ := ts.do(ts.TokenA, http.MethodPatch, "/Groups/"+id, patchOp(ops...)) require.Equal(ts.T(), http.StatusNoContent, w.Code, w.Body.String()) diff --git a/internal/models/scim.go b/internal/models/scim.go index 68f841eed0..0b7b5b17af 100644 --- a/internal/models/scim.go +++ b/internal/models/scim.go @@ -106,6 +106,17 @@ func findSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarg return row, nil } +func lockSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarget) (*T, error) { + row := new(T) + if err := tx.RawQuery( + fmt.Sprintf("SELECT %s FROM %q WHERE %s AND "+scimVersionClause+" FOR UPDATE", table.columns, table.tableName, table.targetClause()), + target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, + ).First(row); err != nil { + return nil, table.writeError(tx, target, err, "locking") + } + return row, nil +} + func findUnchangedSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarget, resource []byte) (*T, error) { row := new(T) if err := tx.RawQuery( diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go index a70e67e711..c702dd2e8c 100644 --- a/internal/models/scim_group.go +++ b/internal/models/scim_group.go @@ -1,7 +1,9 @@ package models import ( + "bytes" "fmt" + "slices" "time" "github.com/gofrs/uuid" @@ -67,6 +69,10 @@ func FindSCIMGroups(tx *storage.Connection, providerID uuid.UUID, query SCIMQuer return findSCIMPage[SCIMGroup](tx, scimGroupsTable, providerID, query) } +func LockSCIMGroup(tx *storage.Connection, target SCIMTarget) (*SCIMGroup, error) { + return lockSCIMRow[SCIMGroup](tx, scimGroupsTable, target) +} + func ReplaceSCIMGroup(tx *storage.Connection, target SCIMTarget, resource []byte) (*SCIMGroup, error) { return replaceSCIMRow[SCIMGroup](tx, scimGroupsTable, target, resource) } @@ -133,8 +139,7 @@ func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserI if err != nil { return nil, nil, nil, err } - removed = differenceUUIDs(current, scimUserIDs) - if err := removeSCIMGroupMembers(tx, group.ID, removed); err != nil { + if removed, err = removeSCIMGroupMembers(tx, group.ID, differenceUUIDs(current, scimUserIDs)); err != nil { return nil, nil, nil, err } added, err = addSCIMGroupMembers(tx, group, differenceUUIDs(scimUserIDs, current)) @@ -144,6 +149,16 @@ func ReplaceSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserI return added, removed, append(differenceUUIDs(current, removed), added...), nil } +func PatchSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, add, remove []uuid.UUID) (added, removed []uuid.UUID, err error) { + if removed, err = removeSCIMGroupMembers(tx, group.ID, remove); err != nil { + return nil, nil, err + } + if added, err = addSCIMGroupMembers(tx, group, add); err != nil { + return nil, nil, err + } + return added, removed, nil +} + func ClearSCIMGroupMembers(tx *storage.Connection, groupID uuid.UUID) ([]uuid.UUID, error) { removed := []uuid.UUID{} if err := tx.RawQuery( @@ -210,32 +225,35 @@ func findSCIMGroupMemberIDs(tx *storage.Connection, groupID uuid.UUID) ([]uuid.U return ids, nil } -func removeSCIMGroupMembers(tx *storage.Connection, groupID uuid.UUID, scimUserIDs []uuid.UUID) error { +func removeSCIMGroupMembers(tx *storage.Connection, groupID uuid.UUID, scimUserIDs []uuid.UUID) ([]uuid.UUID, error) { + removed := []uuid.UUID{} if len(scimUserIDs) == 0 { - return nil + return removed, nil } if err := tx.RawQuery( - fmt.Sprintf("DELETE FROM %q WHERE group_id = ? AND scim_user_id = ANY(?::uuid[])", SCIMGroupMember{}.TableName()), + fmt.Sprintf("DELETE FROM %q WHERE group_id = ? AND scim_user_id = ANY(?::uuid[]) RETURNING scim_user_id", SCIMGroupMember{}.TableName()), groupID, scimUserIDs, - ).Exec(); err != nil { - return errors.Wrap(err, "error removing SCIM group members") + ).All(&removed); err != nil { + return nil, errors.Wrap(err, "error removing SCIM group members") } - return nil + return removed, nil } func addSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserIDs []uuid.UUID) ([]uuid.UUID, error) { + added := []uuid.UUID{} if len(scimUserIDs) == 0 { - return []uuid.UUID{}, nil + return added, nil } locked, err := lockLiveSCIMUserIDs(tx, group.SSOProviderID, scimUserIDs) if err != nil { return nil, err } if err := tx.RawQuery( - fmt.Sprintf("INSERT INTO %q (group_id, scim_user_id) SELECT ?, unnest(?::uuid[])", SCIMGroupMember{}.TableName()), + fmt.Sprintf("INSERT INTO %q (group_id, scim_user_id) SELECT ?, unnest(?::uuid[]) ON CONFLICT DO NOTHING RETURNING scim_user_id", SCIMGroupMember{}.TableName()), group.ID, locked, - ).Exec(); err != nil { + ).All(&added); err != nil { return nil, errors.Wrap(err, "error adding SCIM group members") } - return locked, nil + slices.SortFunc(added, func(a, b uuid.UUID) int { return bytes.Compare(a[:], b[:]) }) + return added, nil } diff --git a/internal/models/scim_group_test.go b/internal/models/scim_group_test.go index 1de33aba09..01e3e750a9 100644 --- a/internal/models/scim_group_test.go +++ b/internal/models/scim_group_test.go @@ -297,6 +297,43 @@ func (ts *SCIMGroupTestSuite) TestReplaceMembersValidatesOnlyAddedMembers() { require.Equal(ts.T(), SCIMGroupMemberNotFoundError{IDs: []uuid.UUID{uuid.Nil}}, err) } +func (ts *SCIMGroupTestSuite) TestLockChecksVersion() { + group := ts.createGroup(ts.provider.ID, "Engineering") + + locked, err := LockSCIMGroup(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: group.ID, UpdatedAt: &group.UpdatedAt}) + require.NoError(ts.T(), err) + require.Equal(ts.T(), group.ID, locked.ID) + + stale := group.UpdatedAt.Add(-time.Second) + _, err = LockSCIMGroup(ts.db, SCIMTarget{ProviderID: ts.provider.ID, ID: group.ID, UpdatedAt: &stale}) + require.ErrorIs(ts.T(), err, SCIMGroupStaleError{}) + + _, err = LockSCIMGroup(ts.db, SCIMTarget{ProviderID: ts.createProvider().ID, ID: group.ID}) + require.ErrorIs(ts.T(), err, SCIMGroupNotFoundError{}) +} + +func (ts *SCIMGroupTestSuite) TestPatchMembersReturnsOnlyChanges() { + group := ts.createGroup(ts.provider.ID, "Engineering") + alice := ts.createUser(ts.provider.ID, "alice") + bob := ts.createUser(ts.provider.ID, "bob") + carol := ts.createUser(ts.provider.ID, "carol") + _, _, _, err := ReplaceSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, bob.ID}) + require.NoError(ts.T(), err) + + added, removed, err := PatchSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID, carol.ID, carol.ID}, []uuid.UUID{bob.ID, uuid.Must(uuid.NewV4())}) + require.NoError(ts.T(), err) + require.Equal(ts.T(), []uuid.UUID{carol.ID}, added) + require.Equal(ts.T(), []uuid.UUID{bob.ID}, removed) + + added, removed, err = PatchSCIMGroupMembers(ts.db, group, []uuid.UUID{alice.ID}, []uuid.UUID{bob.ID}) + require.NoError(ts.T(), err) + require.Empty(ts.T(), added) + require.Empty(ts.T(), removed) + + _, _, err = PatchSCIMGroupMembers(ts.db, group, []uuid.UUID{uuid.Nil}, nil) + require.Equal(ts.T(), SCIMGroupMemberNotFoundError{IDs: []uuid.UUID{uuid.Nil}}, err) +} + func (ts *SCIMGroupTestSuite) TestFindMembershipsByUser() { engineering := ts.createGroup(ts.provider.ID, "Engineering") admins := ts.createGroup(ts.provider.ID, "Admins") From 20197f531faa4f2032c963987ec8f018651557a0 Mon Sep 17 00:00:00 2001 From: mo khan Date: Thu, 1 Oct 2026 13:19:15 -0600 Subject: [PATCH 88/88] fix(scim): take SCIM versions from clock_timestamp so each write gets a new version --- internal/models/scim.go | 2 +- internal/models/scim_group.go | 4 ++-- internal/models/scim_group_test.go | 13 +++++++++++++ internal/models/scim_user.go | 4 ++-- 4 files changed, 18 insertions(+), 5 deletions(-) diff --git a/internal/models/scim.go b/internal/models/scim.go index 0b7b5b17af..4125fa41b2 100644 --- a/internal/models/scim.go +++ b/internal/models/scim.go @@ -143,7 +143,7 @@ func replaceSCIMRowIfChanged[T any](tx *storage.Connection, table scimTable, tar func replaceSCIMRow[T any](tx *storage.Connection, table scimTable, target SCIMTarget, resource []byte) (*T, error) { row := new(T) if err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET resource = ?::jsonb, updated_at = now() WHERE %s AND "+scimVersionClause+" RETURNING %s", table.tableName, table.targetClause(), table.columns), + fmt.Sprintf("UPDATE %q SET resource = ?::jsonb, updated_at = clock_timestamp() WHERE %s AND "+scimVersionClause+" RETURNING %s", table.tableName, table.targetClause(), table.columns), string(resource), target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(row); err != nil { return nil, table.writeError(tx, target, err, "replacing") diff --git a/internal/models/scim_group.go b/internal/models/scim_group.go index c702dd2e8c..092fd153d0 100644 --- a/internal/models/scim_group.go +++ b/internal/models/scim_group.go @@ -84,7 +84,7 @@ func ReplaceSCIMGroupIfChanged(tx *storage.Connection, target SCIMTarget, resour func TouchSCIMGroup(tx *storage.Connection, group *SCIMGroup) (*SCIMGroup, error) { touched := &SCIMGroup{} if err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET updated_at = now() WHERE id = ? RETURNING "+scimGroupColumns, scimGroupsTable.tableName), + fmt.Sprintf("UPDATE %q SET updated_at = clock_timestamp() WHERE id = ? RETURNING "+scimGroupColumns, scimGroupsTable.tableName), group.ID, ).First(touched); err != nil { return nil, errors.Wrap(err, "error updating SCIM group") @@ -191,7 +191,7 @@ func RemoveSCIMUserFromGroups(tx *storage.Connection, scimUserID uuid.UUID) ([]u } if len(groupIDs) > 0 { if err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET updated_at = now() WHERE id = ANY(?::uuid[])", groups), + fmt.Sprintf("UPDATE %q SET updated_at = clock_timestamp() WHERE id = ANY(?::uuid[])", groups), groupIDs, ).Exec(); err != nil { return nil, errors.Wrap(err, "error updating SCIM groups") diff --git a/internal/models/scim_group_test.go b/internal/models/scim_group_test.go index 01e3e750a9..acc1c79acd 100644 --- a/internal/models/scim_group_test.go +++ b/internal/models/scim_group_test.go @@ -107,6 +107,19 @@ func (ts *SCIMGroupTestSuite) TestReplaceChecksVersion() { require.ErrorIs(ts.T(), err, SCIMGroupNotFoundError{}) } +func (ts *SCIMGroupTestSuite) TestEachWriteInATransactionGetsANewVersion() { + group := ts.createGroup(ts.provider.ID, "Engineering") + + require.NoError(ts.T(), ts.db.Transaction(func(tx *storage.Connection) error { + replaced, err := ReplaceSCIMGroup(tx, SCIMTarget{ProviderID: ts.provider.ID, ID: group.ID}, []byte(`{"displayName":"Platform"}`)) + require.NoError(ts.T(), err) + touched, err := TouchSCIMGroup(tx, replaced) + require.NoError(ts.T(), err) + require.True(ts.T(), touched.UpdatedAt.After(replaced.UpdatedAt)) + return nil + })) +} + func (ts *SCIMGroupTestSuite) TestDeleteRemovesMembers() { group := ts.createGroup(ts.provider.ID, "Engineering") user := ts.createUser(ts.provider.ID, "alice") diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go index c8695fe768..96d4cfaab1 100644 --- a/internal/models/scim_user.go +++ b/internal/models/scim_user.go @@ -82,7 +82,7 @@ func ReplaceSCIMUser(tx *storage.Connection, target SCIMTarget, resource []byte) func DeleteSCIMUser(tx *storage.Connection, target SCIMTarget) (*SCIMUser, error) { user := &SCIMUser{} err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE %s AND "+scimVersionClause+" RETURNING %s", scimUsersTable.tableName, scimUsersTable.targetClause(), scimUsersTable.columns), + fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = clock_timestamp() WHERE %s AND "+scimVersionClause+" RETURNING %s", scimUsersTable.tableName, scimUsersTable.targetClause(), scimUsersTable.columns), target.ID, target.ProviderID, target.UpdatedAt, target.UpdatedAt, ).First(user) if err != nil { @@ -132,7 +132,7 @@ func LogoutUserForSCIM(tx *storage.Connection, userID uuid.UUID) error { func SoftDeleteSCIMUsersByUserID(tx *storage.Connection, userID uuid.UUID) ([]SCIMUser, error) { rows := []SCIMUser{} if err := tx.RawQuery( - fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = now() WHERE user_id = ? AND deleted_at IS NULL RETURNING "+scimUserColumns, scimUsersTable.tableName), + fmt.Sprintf("UPDATE %q SET deleted_at = now(), updated_at = clock_timestamp() WHERE user_id = ? AND deleted_at IS NULL RETURNING "+scimUserColumns, scimUsersTable.tableName), userID, ).All(&rows); err != nil { return nil, errors.Wrap(err, "error deleting SCIM users by user id")