diff --git a/README.md b/README.md index 7b103730b0..d1bf238a5c 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`, 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..4b6c4cf765 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.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 ef418ed404..5fd14da8ca 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.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= @@ -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/hack/scim-demo.sh b/hack/scim-demo.sh new file mode 100755 index 0000000000..d8f1226003 --- /dev/null +++ b/hack/scim-demo.sh @@ -0,0 +1,135 @@ +#!/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:-}" 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 "$show" < "$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" +} + +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" + +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]' + +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 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 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" +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" 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..1d282451d7 100644 --- a/internal/api/admin_test.go +++ b/internal/api/admin_test.go @@ -864,6 +864,68 @@ func (ts *AdminTestSuite) TestAdminUserDelete() { } } +func (ts *AdminTestSuite) TestAdminUserDeleteSoftDeletesSCIMUser() { + cases := []struct { + desc string + body map[string]any + wantEmail string + }{ + { + desc: "hard delete", + body: map[string]any{"should_soft_delete": false}, + wantEmail: "scim-hard-delete@example.com", + }, + { + desc: "soft delete", + body: map[string]any{"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"}) @@ -1184,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/api.go b/internal/api/api.go index 96f6f6d015..2ed37c94cb 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,10 @@ func NewAPIWithVersion(globalConfig *conf.GlobalConfiguration, db *storage.Conne api.oauthServer = oauthserver.NewServer(globalConfig, db, api.tokenService) } - api.scim = scim.NewServer(globalConfig) + api.scim = api.newSCIMServer( + api.limitSCIMInvalidToken(newSCIMTokenValidator(db), api.limiterOpts.SCIMIP), + api.limitSCIMByProvider(api.limiterOpts.SCIM), + ) if api.config.Password.HIBP.Enabled { httpClient := &http.Client{ @@ -404,6 +407,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 +478,11 @@ 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.limitSCIMByIP(api.limiterOpts.SCIMIP)) + r.Handle("/*", 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..92dd1c4df5 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" @@ -29,6 +30,8 @@ var ( 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/external.go b/internal/api/external.go index f32420ee75..e5c41f55df 100644 --- a/internal/api/external.go +++ b/internal/api/external.go @@ -311,6 +311,18 @@ func (a *API) createAccountFromExternalIdentity(tx *storage.Connection, r *http. identityData = structs.Map(userData.Metadata) } + ssoProviderID, isSSO, perr := models.SSOProviderID(providerType) + 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 isSCIMProvider { + if terr := models.LockAccountLinkingEmails(tx, providerType, models.VerifiedEmails(config, userData.Emails)); 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 @@ -407,6 +419,16 @@ func (a *API) createAccountFromExternalIdentity(tx *storage.Connection, r *http. return 0, nil, apierrors.NewForbiddenError(apierrors.ErrorCodeUserBanned, "User is banned") } + if isSCIMProvider { + deprovisioned, terr := models.IsSCIMUserDeprovisionedByProvider(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..72d9ce85e4 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.Equal(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..68a37ac1b9 100644 --- a/internal/api/identity.go +++ b/internal/api/identity.go @@ -51,10 +51,21 @@ 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 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{ "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..114b6aff6a 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) Handle(pattern string, h http.Handler) { + r.chi.Handle(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..4f6a30a58c --- /dev/null +++ b/internal/api/scim.go @@ -0,0 +1,273 @@ +package api + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "strconv" + "strings" + "time" + + "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" +) + +const ( + scimBasePath = "/scim/v2" + scimResourceTypeUser = "User" + 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) +} + +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") + + scimUserSchemas = core.Schemas{ + core.NewSchema(core.SchemaUser).With(core.UserAttributes()...), + core.NewSchema(core.SchemaEnterpriseUser).With(core.EnterpriseUserAttributes()...), + } + 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 { + 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(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(&scimUserRepository{api: a})), + server.WithResource(server.NewResource[*core.Group](scimResourceTypeGroup, "/Groups", core.SchemaGroup, scimGroupSchemas.Base().Attributes...). + WithRepository(&scimGroupRepository{api: a})), + server.WithAuthentication(core.NewOAuthBearerToken().AsPrimary(), authenticate), + ) +} + +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 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 + } + return scimTokenKey.WithValue(ctx, token), nil + } +} + +func (a *API) withSCIMRequest(w http.ResponseWriter, req *http.Request) (context.Context, error) { + return scimRequestKey.WithValue(req.Context(), req), nil +} + +func (a *API) auditSCIM(tx *storage.Connection, r *http.Request, event scimAuditEvent) error { + 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, entries) +} + +func scimBaseURL(config *conf.GlobalConfiguration) string { + return strings.TrimRight(config.API.ExternalURL, "/") + scimBasePath +} + +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 != "" { + 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 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) (models.SCIMTarget, error) { + providerID, err := scimProviderID(ctx) + if err != nil { + return models.SCIMTarget{}, err + } + resourceID, err := uuid.FromString(id) + if err != nil { + return models.SCIMTarget{}, errSCIMNotFound() + } + updatedAt, err := scimParseVersion(version) + if err != nil { + return models.SCIMTarget{}, err + } + return models.SCIMTarget{ProviderID: providerID, ID: resourceID, UpdatedAt: updatedAt}, nil +} + +func scimProviderID(ctx context.Context) (uuid.UUID, error) { + token := scimTokenKey.Value(ctx) + if token == nil || token.SSOProviderID == uuid.Nil { + return uuid.Nil, errMissingSSOProvider + } + 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 { + 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 scimActor(r *http.Request) *models.User { + prefix := "" + if token := scimTokenKey.Value(r.Context()); token != nil { + prefix = token.Prefix + } + return &models.User{Email: storage.NullString("scim:" + prefix)} +} + +func scimMemberTraits(groupID, scimUserID uuid.UUID, userID *uuid.UUID) map[string]any { + traits := map[string]any{ + "scim_group_id": groupID, + "scim_user_id": scimUserID, + } + if userID != nil { + traits["user_id"] = *userID + } + return 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_found.json b/internal/api/scim/testdata/not_found.json deleted file mode 100644 index 4d241ba672..0000000000 --- a/internal/api/scim/testdata/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" -} 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..8886278d4d --- /dev/null +++ b/internal/api/scim_admin.go @@ -0,0 +1,278 @@ +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 ( + scimDeprovisionedBanDuration = 100 * 365 * 24 * time.Hour + scimTokenPrefixTrait = "token_prefix" +) + +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 { + 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, scimAuditEvent{ + actor: getAdminUser(ctx), + action: models.SCIMEnabledAction, + providerID: provider.ID, + traits: 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 { + 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.revokeActiveSCIMTokens(tx, r, actor, provider.ID) + if err != nil || !changed { + return err + } + return a.auditSCIMDisabled(tx, r, provider.ID, 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 { + ctx := r.Context() + db := a.db.WithContext(ctx) + provider := getSSOProvider(ctx) + + params, err := a.scimTokenCreateParams(r) + if err != nil { + return err + } + + 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, 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") + } + 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 { + var err error + 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") + } + return apierrors.NewInternalServerError("Error revoking SCIM token").WithInternalError(err) + } + + 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, scimTokenAudit(getAdminUser(r.Context()), models.SCIMTokenRevokedAction, token)) +} + +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.revokeActiveSCIMTokens(tx, r, actor, provider.ID) + if err != nil { + return err + } + if enabled && a.config.SSO.SCIM.Enabled { + if err := a.auditSCIMDisabled(tx, r, provider.ID, prefixes); err != nil { + return err + } + } + banned, err := models.BanDeprovisionedSCIMUsers(tx, provider.ID, a.Now().Add(scimDeprovisionedBanDuration)) + if err != nil || banned == 0 { + return err + } + 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) 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 + } + tokens, err := models.FindActiveSCIMTokensBySSOProvider(tx, providerID) + if err != nil { + 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 + } + events[i] = scimTokenAudit(actor, models.SCIMTokenRevokedAction, &tokens[i]) + } + return prefixes, a.auditSCIMEvents(tx, r, events) +} + +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()), + action: models.SCIMDisabledAction, + providerID: providerID, + traits: map[string]any{"token_prefixes": prefixes}, + }) +} diff --git a/internal/api/scim_admin_test.go b/internal/api/scim_admin_test.go new file mode 100644 index 0000000000..144e4a2e50 --- /dev/null +++ b/internal/api/scim_admin_test.go @@ -0,0 +1,594 @@ +package api + +import ( + "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 func() { require.NoError(t, 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 + + ts.AdminJWT = adminJWT(ts.T(), ts.Config.JWT.Secret) + + ts.Provider = ts.createProvider() +} + +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, 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) + require.Nil(ts.T(), scimTokenKey.Value(ctx)) + + 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) 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) TestDisableRevokesActiveTokens() { + active := ts.create(ts.Provider, map[string]any{}) + expiring := ts.create(ts.Provider, map[string]any{"expires_at": time.Now().Add(time.Hour)}) + 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) + + 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(), 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() { + token := ts.create(ts.Provider, map[string]any{}) + + status := ts.status(http.MethodDelete, ts.Provider) + require.False(ts.T(), status.Enabled) + 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) + 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) TestReenableKeepsOldTokensRevoked() { + 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.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() { + 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.Empty(ts.T(), status.Tokens) +} + +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) 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) 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) 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) +} + +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_errors.go b/internal/api/scim_errors.go new file mode 100644 index 0000000000..99f846ae7d --- /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 models.IsStaleError(err): + 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_filter.go b/internal/api/scim_filter.go new file mode 100644 index 0000000000..ef1c2e0f31 --- /dev/null +++ b/internal/api/scim_filter.go @@ -0,0 +1,54 @@ +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, isString := value.(string) + isEquals := op == filter.OpEquals + isTopLevel := attribute.Parent == nil + if !isEquals || !isTopLevel || !isString { + 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..c69d0f8923 --- /dev/null +++ b/internal/api/scim_groups.go @@ -0,0 +1,387 @@ +package api + +import ( + "context" + "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 +} + +type scimGroupWrite func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error) + +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, + 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) { + target, err := scimTarget(ctx, id, "") + if err != nil { + return nil, err + } + db := s.api.db.WithContext(ctx) + row, err := models.FindSCIMGroup(db, target.ProviderID, target.ID) + if err != nil { + return nil, scimError(err) + } + return s.renderOne(db, target.ProviderID, row, protocol.ProjectionFrom(ctx)) +} + +func (s *scimGroupRepository) Create(ctx context.Context, group *core.Group) (*core.Group, error) { + providerID, err := scimProviderID(ctx) + if err != nil { + return nil, err + } + 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 + }) +} + +func (s *scimGroupRepository) Replace(ctx context.Context, group *core.Group) (*core.Group, error) { + target, err := scimTarget(ctx, group.ID, group.Meta.Version) + if err != nil { + return nil, err + } + return s.save(ctx, models.SCIMGroupUpdatedAction, group, func(tx *storage.Connection, resource []byte) (*models.SCIMGroup, bool, error) { + return models.ReplaceSCIMGroupIfChanged(tx, target, resource) + }) +} + +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 { + return err + } + r, err := scimRequest(ctx) + if err != nil { + return err + } + return scimError(s.api.db.WithContext(ctx).Transaction(func(tx *storage.Connection) error { + 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) 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 { + return nil, err + } + attributes := *group + attributes.Members = nil + resource, err := scimEncode(&attributes, "id", "meta", "members") + if err != nil { + return nil, err + } + r, err := scimRequest(ctx) + if err != nil { + return nil, err + } + change := scimGroupChange{r: r, action: action, members: members} + db := s.api.db.WithContext(ctx) + var saved *core.Group + err = db.Transaction(func(tx *storage.Connection) error { + row, changed, terr := write(tx, resource) + if terr != nil { + return terr + } + row, members, terr := s.applyMembers(tx, change, row, changed) + if terr != nil { + return terr + } + saved, terr = s.compose(*row, scimMembers(scimBaseURL(s.api.config), members)) + return terr + }) + if err != nil { + return nil, scimError(err) + } + return saved, nil +} + +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, nil, err + } + events := []scimAuditEvent{} + switch { + case changed: + 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: + return row, members, nil + } + if err != nil { + return nil, nil, err + } + memberEvents, err := s.memberEvents(tx, change.r, row, scimMemberDiff{added: added, removed: removed}) + if err != nil { + return nil, nil, err + } + 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) { + members, err := s.members(tx, providerID, rows, projection) + if err != nil { + return nil, err + } + groups := make([]*core.Group, 0, len(rows)) + for _, row := range rows { + group, err := s.compose(row, members[row.ID]) + if err != nil { + return nil, err + } + 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") { + return members, nil + } + ids := make([]uuid.UUID, len(rows)) + for i, row := range rows { + ids[i] = row.ID + } + memberships, err := models.FindSCIMMembershipsByGroup(tx, providerID, ids) + if err != nil { + return nil, err + } + counts := make(map[uuid.UUID]int, len(rows)) + for _, m := range memberships { + counts[m.GroupID]++ + } + byGroup := make(map[uuid.UUID][]uuid.UUID, len(counts)) + for _, m := range memberships { + if byGroup[m.GroupID] == nil { + byGroup[m.GroupID] = make([]uuid.UUID, 0, counts[m.GroupID]) + } + 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 { + return nil, err + } + return groups[0], nil +} + +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 scimAuditEvent{}, err + } + return scimAuditEvent{ + actor: scimActor(r), + action: action, + providerID: row.SSOProviderID, + traits: map[string]any{ + "scim_group_id": row.ID, + "display_name": resource.DisplayName, + }, + }, nil +} + +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(diff.added)+len(diff.removed)) + for _, change := range []struct { + action models.AuditAction + ids []uuid.UUID + }{ + {action: models.SCIMGroupMemberAddedAction, ids: diff.added}, + {action: models.SCIMGroupMemberRemovedAction, ids: diff.removed}, + } { + for _, id := range change.ids { + var userID *uuid.UUID + if linked, ok := links[id]; ok { + userID = &linked + } + events = append(events, scimAuditEvent{ + actor: actor, + action: change.action, + providerID: row.SSOProviderID, + traits: scimMemberTraits(row.ID, id, userID), + }) + } + } + return events, nil +} + +func scimMemberIDs(members []core.Member) ([]uuid.UUID, error) { + ids := make([]uuid.UUID, 0, len(members)) + for _, member := range members { + id, err := uuid.FromString(member.Value) + if err != nil { + return nil, errSCIMMemberNotFound() + } + ids = append(ids, id) + } + 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 new file mode 100644 index 0000000000..85c4c924dc --- /dev/null +++ b/internal/api/scim_groups_test.go @@ -0,0 +1,750 @@ +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 *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 *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 +} + +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 *SCIMTestSuite) 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, 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"]) + + 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, "") + 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 *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") + + ts.createGroup(ts.TokenA, groupWith("Empty", "e-2")) + require.EqualValues(ts.T(), 2, ts.listGroups(ts.TokenA, `displayName eq "Empty"`)["totalResults"]) +} + +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, "") + 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 *SCIMTestSuite) 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 *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")) + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, bob)) + 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 entries { + 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 *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{} + 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() + } + entries := ts.auditDuring(func() { + id := ts.createGroup(ts.TokenA, groupWith("Engineering", "", alice, 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()) + }) + + type event struct { + action, scimUserID, userID string + } + events := []event{} + for _, entry := range entries { + 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 *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")) + 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", 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()) + + 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 *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)) + before := len(ts.auditActions(models.SCIMGroupMemberRemovedAction)) + + 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) +} + +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)) + + 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, patchOp(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"]) + + got = ts.patchMembers(id, `{"op":"remove","path":"members[value eq \"`+bob+`\"]","value":null}`) + require.Equal(ts.T(), []string{alice}, memberValues(got)) +} + +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")) + 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 *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) + stale := w.Header().Get("ETag") + + 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") + 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 *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"), + "/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 *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)) + ops := ts.createGroup(ts.TokenA, groupWith("Ops", "g-2", alice)) + 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 entries { + 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 *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, "") + 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 *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)) + 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 *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)) + 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()) + + 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, "") + 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 entries { + require.NotEqual(ts.T(), string(models.SCIMGroupMemberRemovedAction), entry.Payload["action"]) + } +} + +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")) + + 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 *SCIMTestSuite) 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 *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 *SCIMTestSuite) TestGroupsAuditLog() { + alice := ts.create(ts.TokenA, userWith("alice@example.com", "a-1")) + bob := ts.create(ts.TokenA, userWith("bob@example.com", "b-1")) + 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.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.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 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) + 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 *SCIMTestSuite) 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}, + } + 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} + } + 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) + }) + + type event struct { + action, subject string + } + events := []event{} + 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 { + 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 *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")) + + 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 *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) + 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]+`"}]}`, `{"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) 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()) + 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 +} diff --git a/internal/api/scim_isolation_test.go b/internal/api/scim_isolation_test.go new file mode 100644 index 0000000000..daa3fe8922 --- /dev/null +++ b/internal/api/scim_isolation_test.go @@ -0,0 +1,193 @@ +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 *SCIMTestSuite) 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, patchOp(`{"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 *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]) + 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 *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, "") + 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 *SCIMTestSuite) 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, 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, patchOp(`{"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"`) + 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") + 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) +} + +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)) + + 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, patchOp(`{"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, 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()) + 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..1a150e006d --- /dev/null +++ b/internal/api/scim_link_test.go @@ -0,0 +1,1063 @@ +package api + +import ( + "net/http" + "net/http/httptest" + "strconv" + "strings" + "sync" + "time" + + "github.com/gofrs/uuid" + "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/conf" + "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/storage" +) + +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 + require.NoError(ts.T(), ts.API.db.Create(user)) + 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 +} + +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) + user, err := models.FindUserByID(ts.API.db, *row.UserID) + require.NoError(ts.T(), err) + return user +} + +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 *SCIMTestSuite) 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{scimProviderType(ts.A.ID)}, user.AppMetaData["providers"]) + + identities := ts.identities(user) + require.Len(ts.T(), identities, 1) + require.Equal(ts.T(), scimProviderType(ts.A.ID), identities[0].Provider) + require.Equal(ts.T(), "Alice@Example.com", identities[0].ProviderID) +} + +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)) + 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.Empty(ts.T(), ts.identities(password)) + require.Len(ts.T(), ts.identities(other), 1) +} + +func (ts *SCIMTestSuite) 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 *SCIMTestSuite) 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 *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 + ts.API.config.WebAuthn = conf.WebAuthnConfiguration{ + RPID: "localhost", + RPDisplayName: "Test App", + RPOrigins: []string{"http://localhost:3000"}, + ChallengeExpiryDuration: 5 * time.Minute, + } + + r := httptest.NewRequest(http.MethodPost, "/passkeys/registration/options", nil) + r.Header.Set("Authorization", "Bearer "+ts.accessToken(user)) + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + return w.Code +} + +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)) + + 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.Empty(ts.T(), ts.identities(reloaded)) +} + +func (ts *SCIMTestSuite) TestNonSSOUserWithSSOIdentityEmailIsNeverLinked() { + ts.requireNonSSOUserNeverLinked("saml-name-id") +} + +func (ts *SCIMTestSuite) TestNonSSOUserWithSSOIdentitySubjectIsNeverLinked() { + ts.requireNonSSOUserNeverLinked("Alice@Example.com") +} + +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, 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)) + + 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 *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)) + 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.Empty(ts.T(), ts.identities(password)) +} + +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()) + user := ts.linkedUser(id) + 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) + require.Zero(ts.T(), ts.countRows(&models.OneTimeToken{}, "user_id = ?", user.ID), req.path) + require.Zero(ts.T(), ts.sessions(user), req.path) + } + } +} + +func (ts *SCIMTestSuite) withEmail(email string) string { + return oktaUserWith("value", email) +} + +func (ts *SCIMTestSuite) 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", scimProviderType(ts.A.ID)) + 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 *SCIMTestSuite) 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 *SCIMTestSuite) TestReplaceRenamesAndChangesEmail() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + + 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()) + + 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 *SCIMTestSuite) 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 *SCIMTestSuite) 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 *SCIMTestSuite) TestRemovingEmailsKeepsUserEmail() { + id := ts.create(ts.TokenA, oktaUserWith("userName", "alice.smith")) + 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, scimProviderType(ts.A.ID)) + 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 *SCIMTestSuite) TestCreateInactiveLogsOutWithoutBanning() { + existing := ts.ssoUser(ts.A, "Alice@Example.com", "alice@example.com") + ts.session(existing) + + user := ts.linkedUser(ts.create(ts.TokenA, oktaUserWith("active", false))) + + require.False(ts.T(), user.IsBanned()) + require.Zero(ts.T(), ts.sessions(user)) +} + +func (ts *SCIMTestSuite) TestCreateRejectsSharedUser() { + ts.create(ts.TokenA, oktaUser) + + 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"]) +} + +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 *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()) + + 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 *SCIMTestSuite) TestRejectsInvalidEmailsValue() { + 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"]) + + 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, 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()) +} + +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 *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")) + require.Equal(ts.T(), http.StatusConflict, w.Code) + require.Zero(ts.T(), ts.users("bob@example.com")) +} + +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 *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 *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() + ts.API.handler.ServeHTTP(w, r) + return w.Code +} + +func (ts *SCIMTestSuite) sessions(user *models.User) int { + return ts.countRows(&models.Session{}, "user_id = ?", user.ID) +} + +func (ts *SCIMTestSuite) 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, + }}, + } + return ts.samlLoginWith(ssoProvider, userData) +} + +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 + err := ts.API.db.Transaction(func(tx *storage.Connection) error { + var terr error + _, user, terr = ts.API.createAccountFromExternalIdentity(tx, r, userData, scimProviderType(ssoProvider.ID), false) + return terr + }) + return user, err +} + +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}}, + } + conn, err := ts.API.db.NewTransaction() + require.NoError(ts.T(), err) + tx := &storage.Connection{Connection: conn} + require.NoError(ts.T(), models.LockAccountLinking(tx, scimProviderType(ts.A.ID), "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 *SCIMTestSuite) 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 *SCIMTestSuite) 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 *SCIMTestSuite) 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 *SCIMTestSuite) 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) + + 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 *SCIMTestSuite) TestSAMLLoginBlockedWhenCreatedInactive() { + id := ts.create(ts.TokenA, oktaUserWith("active", false)) + ts.linkedUser(id) + + _, err := ts.samlLogin(ts.A, "Alice@Example.com", "alice@example.com") + require.Error(ts.T(), err) +} + +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) + + 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 *SCIMTestSuite) 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 *SCIMTestSuite) 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) + } + + require.Equal(ts.T(), 1, ts.users("alice@example.com")) + + providerIDs := []string{} + for _, identity := range ts.identities(linked) { + 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) +} + +func (ts *SCIMTestSuite) 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 *SCIMTestSuite) users(email string) int { + return ts.countRows(&models.User{}, "email = ?", email) +} + +func (ts *SCIMTestSuite) 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 *SCIMTestSuite) 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 *SCIMTestSuite) TestSAMLLoginNotBlockedByOtherProvider() { + id := ts.create(ts.TokenB, oktaUser) + 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") + + require.NoError(ts.T(), err) + require.Equal(ts.T(), existing.ID, user.ID) +} + +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 *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 *SCIMTestSuite) 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)) + + 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 *SCIMTestSuite) 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 *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"}) + 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 *SCIMTestSuite) 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 *SCIMTestSuite) 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 *SCIMTestSuite) TestSessionAllowedForSSOUserWithoutSCIMRow() { + user := ts.ssoUser(ts.A, "saml-sub", "carol@example.com") + require.NoError(ts.T(), ts.issueSession(ts.API.db, user)) +} + +func (ts *SCIMTestSuite) setActive(id string, active bool) { + ts.setActiveAs(ts.TokenA, id, active) +} + +func (ts *SCIMTestSuite) setActiveAs(token, id string, active bool) { + 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()) +} + +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)) + + 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 *SCIMTestSuite) 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, 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)) + require.Equal(ts.T(), http.StatusBadRequest, ts.refresh(refreshToken)) + require.Len(ts.T(), ts.auditActions(models.SCIMUserDeactivatedAction), 1) +} + +func (ts *SCIMTestSuite) 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 *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(), oktaUserWith("active", false)) + 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 *SCIMTestSuite) 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 *SCIMTestSuite) TestCreateRefusesUserDeletedByProviderSameUserName() { + ts.requireCreateRefusedAfterProviderDelete(oktaUser) +} + +func (ts *SCIMTestSuite) TestCreateRefusesUserDeletedByProviderNewUserName() { + ts.requireCreateRefusedAfterProviderDelete(oktaUserWith("userName", "alice.new@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) TestCreateAfterAdminHardDeletesProviderDeletedUser() { + ts.requireCreateAfterAdminDelete(false) +} + +func (ts *SCIMTestSuite) TestCreateAfterAdminSoftDeletesProviderDeletedUser() { + ts.requireCreateAfterAdminDelete(true) +} + +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() { + 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 + codes[i] = ts.serve(protocol.MediaType, ts.TokenA, http.MethodPost, "/Users", bodies[i]).Code + }(i) + } + close(start) + wg.Wait() + + 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) + + require.Equal(ts.T(), 1, ts.users("race@example.com")) +} + +func (ts *SCIMTestSuite) rename(id, userName string) (int, string) { + w, _ := ts.do(ts.TokenA, http.MethodPut, "/Users/"+id, oktaUserWith("userName", userName)) + return w.Code, w.Body.String() +} + +func (ts *SCIMTestSuite) 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", 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"]) + require.Len(ts.T(), ts.identities(user), 1) +} + +func (ts *SCIMTestSuite) providerIDs(user *models.User) []string { + ids := []string{} + for _, identity := range ts.identities(user) { + ids = append(ids, identity.ProviderID) + } + return ids +} + +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") + 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 *SCIMTestSuite) 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 *SCIMTestSuite) 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", scimProviderType(ts.A.ID)) + require.NoError(ts.T(), err) +} + +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 "+ts.accessToken(user)) + w := httptest.NewRecorder() + ts.API.handler.ServeHTTP(w, r) + return w +} + +func (ts *SCIMTestSuite) 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", scimProviderType(ts.A.ID)) + 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", scimProviderType(ts.A.ID)) + 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 *SCIMTestSuite) 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", 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 }() + + 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 *SCIMTestSuite) TestRenameSkippedWhenIdentityMissing() { + id := ts.create(ts.TokenA, oktaUser) + user := ts.linkedUser(id) + 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)) + hook := logrustest.NewGlobal() + defer hook.Reset() + + 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)) + 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, "identity not found") + } + 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")) +} + +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_okta_spec_test.go b/internal/api/scim_okta_spec_test.go new file mode 100644 index 0000000000..befa78e7c5 --- /dev/null +++ b/internal/api/scim_okta_spec_test.go @@ -0,0 +1,227 @@ +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 *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" + } + 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 *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 { + 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 *SCIMTestSuite) 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 *SCIMTestSuite) 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}, + } + + id, version := "", "" + 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 + + require.Equal(ts.T(), 1, ts.countRows(&models.SCIMUser{}, "sso_provider_id = ?", ts.A.ID), 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 entries { + 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..b3950d687b --- /dev/null +++ b/internal/api/scim_provider_delete_test.go @@ -0,0 +1,215 @@ +package api + +import ( + "net/http" + "net/http/httptest" + "time" + + "github.com/gofrs/uuid" + "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 *SCIMTestSuite) deleteProvider(p *models.SSOProvider) { + 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()) +} + +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 *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 *SCIMTestSuite) auditActions(action models.AuditAction) []models.AuditLogEntry { + return queryAuditEntries(ts.T(), ts.API.db, "payload->>'action' = ?", string(action)) +} + +func (ts *SCIMTestSuite) 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) + 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)) + + 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 *SCIMTestSuite) 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 *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"}) + 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 *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) + 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 *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_%'") + + 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 *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, + ).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 *SCIMTestSuite) 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 *SCIMTestSuite) 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 *SCIMTestSuite) 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_ratelimit.go b/internal/api/scim_ratelimit.go new file mode 100644 index 0000000000..2c3ba5442e --- /dev/null +++ b/internal/api/scim_ratelimit.go @@ -0,0 +1,68 @@ +package api + +import ( + "context" + "math" + "net/http" + "path" + "strconv" + + "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(lmt))(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(lmt))(w, r) + return + } + next.ServeHTTP(w, r) + }) + } +} + +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 new file mode 100644 index 0000000000..14ffff350e --- /dev/null +++ b/internal/api/scim_ratelimit_test.go @@ -0,0 +1,125 @@ +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 func() { require.NoError(t, api.db.Close()) }() + require.NoError(t, models.TruncateAll(api.db)) + + token := func() string { + _, token := createSSOProviderWithSCIMToken(t, api.db) + return token + } + tokenA, tokenB := token(), token() + + send := func(method, path, token, ip string) *httptest.ResponseRecorder { + r := httptest.NewRequest(method, path, 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 + } + 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" + + 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.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()) + + 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) + for range 30 { + w := send(tc.method, tc.path, "scim_invalid", ip) + require.Equal(t, tc.status, w.Code, tc.method+" "+tc.path) + } + 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) + } + } +} + +func TestSCIMRateLimitBoundsEveryUnauthenticatedRequest(t *testing.T) { + api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { + config.RateLimitScim = 1 + }) + defer func() { require.NoError(t, 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 a6a966823d..23ee07529a 100644 --- a/internal/api/scim_test.go +++ b/internal/api/scim_test.go @@ -1,15 +1,30 @@ package api import ( + "bytes" + "context" + "encoding/json" + "io/fs" "net/http" "net/http/httptest" "net/url" + "os" + "strings" "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" - scimCore "github.com/supabase/auth/internal/api/scim/core" - scimProtocol "github.com/supabase/auth/internal/api/scim/protocol" + "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/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 +32,7 @@ const ( scimServiceProviderConfigPath = "/scim/v2/ServiceProviderConfig" scimResourceTypesPath = "/scim/v2/ResourceTypes" scimSchemasPath = "/scim/v2/Schemas" + scimUsersPath = "/scim/v2/Users" ) var scimPaths = []string{ @@ -30,7 +46,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) @@ -50,29 +66,34 @@ 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) }) }) 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) - 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} { @@ -80,11 +101,12 @@ 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) - 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) { @@ -92,23 +114,115 @@ 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) - 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("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(core.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, protocol.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, protocol.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, protocol.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) - 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) { @@ -118,13 +232,345 @@ 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) - require.Equal(t, []string{http.MethodGet}, w.Header().Values("Allow")) + require.Equal(t, "GET, HEAD", w.Header().Get("Allow")) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) }) } } + + for _, tc := range []struct { + method, path string + allow string + }{ + {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) + w := httptest.NewRecorder() + + r.Header.Set("Authorization", "Bearer "+token) + api.handler.ServeHTTP(w, r) + + require.Equal(t, http.StatusMethodNotAllowed, w.Code) + require.Equal(t, tc.allow, w.Header().Get("Allow")) + require.Equal(t, protocol.MediaType, w.Header().Get("Content-Type")) + }) + } }) }) } + +const scimValidToken = "scim_valid" + +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, protocol.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, protocol.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(core.SchemaUser), user["schema"]) + extension := user["schemaExtensions"].([]any)[0].(map[string]any) + require.Equal(t, string(core.SchemaEnterpriseUser), extension["schema"]) + group := resources["Group"] + require.Equal(t, "/Groups", group["endpoint"]) + require.Equal(t, string(core.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, 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(core.SchemaUser), string(core.SchemaEnterpriseUser), string(core.SchemaGroup)}, ids) + }) + + t.Run("Schemas/{id}", func(t *testing.T) { + 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) + 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(core.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(core.SchemaUser), "") + + require.Equal(t, http.StatusOK, w.Code) + 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(core.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, protocol.MediaType, w.Header().Get("Content-Type")) + require.Equal(t, tc.challenge, w.Header().Get("WWW-Authenticate")) + }) + } + }) +} + +type SCIMTestSuite struct { + suite.Suite + API *API + TokenA string + TokenB string + A *models.SSOProvider + B *models.SSOProvider +} + +func TestSCIMSuite(t *testing.T) { + api, _ := setupSCIMAPI(t, func(config *conf.GlobalConfiguration) { + config.RateLimitScim = 1_000_000 + }) + defer func() { require.NoError(t, api.db.Close()) }() + + suite.Run(t, &SCIMTestSuite{API: api}) +} + +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 *SCIMTestSuite) provider() (*models.SSOProvider, string) { + return createSSOProviderWithSCIMToken(ts.T(), ts.API.db) +} + +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 *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 + if w.Body.Len() > 0 { + require.NoError(ts.T(), json.Unmarshal(w.Body.Bytes(), &decoded), w.Body.String()) + } + return w, decoded +} + +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) + 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 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)) + 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_user_cleanup.go b/internal/api/scim_user_cleanup.go new file mode 100644 index 0000000000..b222f6de6f --- /dev/null +++ b/internal/api/scim_user_cleanup.go @@ -0,0 +1,47 @@ +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 + } + events := []scimAuditEvent{} + for i := range rows { + removed, err := scimUserRemovalEvents(tx, actor, &rows[i]) + if err != nil { + return err + } + events = append(events, removed...) + } + return a.auditSCIMEvents(tx, r, events) +} + +func scimUserRemovalEvents(tx *storage.Connection, actor *models.User, row *models.SCIMUser) ([]scimAuditEvent, error) { + groupIDs, err := models.RemoveSCIMUserFromGroups(tx, row.ID) + if err != nil { + return nil, err + } + 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 append(events, scimAuditEvent{ + actor: actor, + action: models.SCIMUserDeletedAction, + providerID: row.SSOProviderID, + traits: scimUserTraits(row), + }), nil +} diff --git a/internal/api/scim_user_linking.go b/internal/api/scim_user_linking.go new file mode 100644 index 0000000000..54d9ef3e0f --- /dev/null +++ b/internal/api/scim_user_linking.go @@ -0,0 +1,151 @@ +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.LogoutUserForSCIM(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 + } + + 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(decision.User); err != nil { + return nil, err + } + if decision.Decision == models.LinkAccount { + if err := s.linkIdentity(tx, decision.User, providerType, user); err != nil { + return nil, err + } + } + return decision.User, nil + case models.MultipleAccounts: + return nil, scimerrors.ErrUniqueness("multiple users share this email in the SSO provider") + } + 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 { + if _, err := s.api.createNewIdentity(tx, linked, providerType, scimIdentityData(user)); err != nil { + return err + } + return linked.UpdateAppMetaDataProviders(tx) +} + +func scimRequireSSOUser(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 new file mode 100644 index 0000000000..f201536d09 --- /dev/null +++ b/internal/api/scim_users.go @@ -0,0 +1,482 @@ +package api + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/url" + "strings" + + "github.com/badoux/checkmail" + "github.com/gofrs/uuid" + "github.com/supabase-community/scim-go/pkg/core" + "github.com/supabase-community/scim-go/pkg/protocol" + "github.com/supabase/auth/internal/models" + "github.com/supabase/auth/internal/observability" + "github.com/supabase/auth/internal/storage" +) + +const ( + scimClaimSub = "sub" + scimClaimEmail = "email" + scimClaimEmailVerified = "email_verified" +) + +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 + 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, + 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) { + target, err := scimTarget(ctx, id, "") + if err != nil { + return nil, err + } + db := s.api.db.WithContext(ctx) + row, err := models.FindSCIMUser(db, target.ProviderID, target.ID) + if err != nil { + return nil, scimError(err) + } + return s.renderOne(db, target.ProviderID, row, protocol.ProjectionFrom(ctx)) +} + +func (s *scimUserRepository) 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 + } + r, err := scimRequest(ctx) + if err != nil { + return nil, err + } + db := s.api.db.WithContext(ctx) + 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} + return s.save(db, change, s.create) +} + +func (s *scimUserRepository) Replace(ctx context.Context, user *core.User) (*core.User, error) { + target, err := scimTarget(ctx, user.ID, user.Meta.Version) + if err != nil { + return nil, err + } + resource, err := scimUserResource(user) + if err != nil { + return nil, err + } + r, err := scimRequest(ctx) + if err != nil { + return nil, err + } + db := s.api.db.WithContext(ctx) + existing, err := models.FindSCIMUser(db, target.ProviderID, target.ID) + if err != nil { + return nil, scimError(err) + } + if existing.UserID == nil { + if err := s.beforeProvision(r, db, target.ProviderID, user); err != nil { + return nil, err + } + } + + change := scimUserChange{r: r, target: target, resource: resource, user: user} + return s.save(db, change, func(tx *storage.Connection, change scimUserChange) (*models.SCIMUser, *models.User, error) { + return s.replace(tx, change, existing) + }) +} + +func (s *scimUserRepository) Delete(ctx context.Context, id, version string) error { + target, 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, target.ProviderID, target.ID) + if err != nil { + return scimError(err) + } + return scimError(db.Transaction(func(tx *storage.Connection) error { + return s.delete(tx, r, target, existing) + })) +} + +func (s *scimUserRepository) save(db *storage.Connection, change scimUserChange, write scimUserWrite) (*core.User, error) { + var ( + saved *core.User + created *models.User + ) + err := db.Transaction(func(tx *storage.Connection) error { + row, user, terr := write(tx, change) + if terr != nil { + return terr + } + created = user + saved, terr = s.renderOne(tx, change.target.ProviderID, row, scimSavedUserProjection(change)) + return terr + }) + if err != nil { + return nil, scimError(err) + } + s.runAfterUserCreatedHook(change.r, db, created) + return saved, nil +} + +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 { + return nil, err + } + base := scimBaseURL(s.api.config) + 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 *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.FindSCIMMembershipsByUser(tx, providerID, ids) + if err != nil { + return nil, err + } + base := scimBaseURL(s.api.config) + for _, m := range memberships { + id := m.GroupID.String() + groups[m.SCIMUserID] = append(groups[m.SCIMUserID], core.GroupMembership{ + Value: id, + Ref: base + "/Groups/" + id, + Display: m.Display, + Type: "direct", + }) + } + 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 { + return nil, err + } + return users[0], nil +} + +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 !sameSCIMLink(row.UserID, existing.UserID) { + return models.SCIMUserStaleError{} + } + if err := logoutSCIMLinkedUser(tx, row.UserID); err != nil { + return err + } + events, err := scimUserRemovalEvents(tx, scimActor(r), row) + if err != nil { + return err + } + 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) { + 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 !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) + } + 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(change.user) == "" { + return nil, errSCIMEmailRequired() + } + return s.provisionAuthUser(tx, row, change.user) + } + + linked, err := models.FindUserByID(tx, *old.UserID) + if err != nil { + return nil, err + } + 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.LogoutUserForSCIM(tx, linked.ID) + } + return nil, 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) + if err != nil || from == user.UserName { + return err + } + data := map[string]any{scimClaimSub: user.UserName} + if email := scimUserEmail(user); email != "" { + data[scimClaimEmail] = 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{}) { + 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 +} + +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 + } + if err := linked.ClearAllPendingTokens(tx); err != nil { + return err + } + 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 { + 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) { + 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.LogoutUserForSCIM(tx, *userID) +} + +func scimUserEmail(user *core.User) string { + if email := scimPrimaryEmail(user.Emails); email != "" { + return email + } + if isEmailAddress(user.UserName) { + return user.UserName + } + return "" +} + +func scimValidatePrimaryEmail(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 { + return email.Value + } + } + if len(emails) > 0 { + return emails[0].Value + } + return "" +} + +func scimIdentityData(user *core.User) map[string]any { + return map[string]any{ + scimClaimSub: user.UserName, + scimClaimEmail: scimUserEmail(user), + scimClaimEmailVerified: true, + } +} + +func scimUserName(resource []byte) (string, error) { + var r struct { + UserName string `json:"userName"` + } + if err := json.Unmarshal(resource, &r); err != nil { + return "", err + } + return r.UserName, nil +} + +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 +} diff --git a/internal/api/scim_users_test.go b/internal/api/scim_users_test.go new file mode 100644 index 0000000000..2e4a7d2b3d --- /dev/null +++ b/internal/api/scim_users_test.go @@ -0,0 +1,770 @@ +package api + +import ( + "context" + "encoding/json" + "maps" + "net/http" + "net/http/httptest" + "net/url" + "regexp" + "slices" + "strconv" + "strings" + "sync" + "testing" + "time" + + "github.com/gofrs/uuid" + "github.com/stretchr/testify/require" + "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" +) + +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 +}` + +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 *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 +} + +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)) + return ctx, &scimUserRepository{api: ts.API} +} + +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) + 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, 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, 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)) + 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, patchOp(`{"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) + } + + 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) +} + +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)) + 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, 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) + } + } +} + +func (ts *SCIMTestSuite) 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 *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")} { + 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 *SCIMTestSuite) 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 *SCIMTestSuite) 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 *SCIMTestSuite) 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 *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() { + 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() { + code <- ts.serve(protocol.MediaType, ts.TokenA, method, path, body).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 *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) + + 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, patchOp(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 *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) + + 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 *SCIMTestSuite) TestConcurrentCreateWithinProvider() { + const attempts = 8 + codes := make(chan int, attempts) + start := make(chan struct{}) + var wg sync.WaitGroup + for range attempts { + wg.Go(func() { + <-start + codes <- ts.serve(protocol.MediaType, ts.TokenA, http.MethodPost, "/Users", userWith("race@example.com", "")).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 *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)) + } + 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"], 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]) + + 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 *SCIMTestSuite) 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 *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) + 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", 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", 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))) + + 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 *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) + 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 := 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"]) + 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 *SCIMTestSuite) TestPatchAttributesOutsideTheMinimalSchema() { + id := ts.create(ts.TokenA, oktaUser) + + 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"]) + + 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 *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"]) + + 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 *SCIMTestSuite) 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"]) + + 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() { + 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 *SCIMTestSuite) 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 *SCIMTestSuite) 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 *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 *SCIMTestSuite) TestRequiresSSOProviderOnContext() { + users := &scimUserRepository{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 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([]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([]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 := &core.User{UserName: "alice", Password: "secret"} + user.ID = "abc" + user.Meta = core.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"]) + }) +} + +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)) + 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, 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") + + 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 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) +} + +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, ",") + `]}` +} 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/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..2224d3b510 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,16 +281,25 @@ func TestExperimentalScimEndpoints(t *testing.T) { cfg, err := LoadGlobalFromEnv() require.NoError(t, err) require.NotNil(t, cfg) - assert.Equal(t, false, cfg.Experimental.ScimEnabled) + assert.False(t, cfg.SSO.SCIM.Enabled) } { baseEnv() - os.Setenv("GOTRUE_EXPERIMENTAL_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.Experimental.ScimEnabled) + assert.True(t, cfg.SSO.SCIM.Enabled) + } + + { + baseEnv() + t.Setenv("GOTRUE_EXPERIMENTAL_SCIM_ENABLED", "true") + cfg, err := LoadGlobalFromEnv() + require.NoError(t, err) + require.NotNil(t, cfg) + assert.False(t, cfg.SSO.SCIM.Enabled) } } diff --git a/internal/models/audit_log_entry.go b/internal/models/audit_log_entry.go index f8b493a783..071383c951 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" @@ -50,12 +51,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 +105,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. @@ -106,7 +138,90 @@ 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, _ := buildAuditLogEntry(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]any +} + +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)) + for i, event := range events { + l, at := buildAuditLogEntry(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 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 username := actor.GetEmail() @@ -153,7 +268,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 @@ -165,53 +281,9 @@ 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 -} - -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 + }, createdAt } diff --git a/internal/models/audit_log_entry_test.go b/internal/models/audit_log_entry_test.go new file mode 100644 index 0000000000..9640e68d91 --- /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 func() { require.NoError(t, 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) 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, []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, 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, 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, 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/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..a7238498d4 100644 --- a/internal/models/errors.go +++ b/internal/models/errors.go @@ -1,10 +1,16 @@ package models -import "errors" +import ( + "errors" + + "github.com/gofrs/uuid" +) // 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") @@ -13,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 { @@ -211,3 +221,101 @@ 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" +} + +func (e SCIMUserStaleError) Is(target error) bool { + return target == errStale +} + +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 SCIMUserDeletedError struct{} + +func (e SCIMUserDeletedError) Error() string { + return "user was deleted by 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" +} + +func (e SCIMGroupStaleError) Is(target error) bool { + return target == errStale +} + +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/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) + } +} diff --git a/internal/models/linking.go b/internal/models/linking.go index 5f5f2d0cd7..a7aa23ec5a 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" @@ -53,6 +55,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 +73,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) @@ -212,3 +221,30 @@ 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 { + 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.go b/internal/models/scim.go new file mode 100644 index 0000000000..4125fa41b2 --- /dev/null +++ b/internal/models/scim.go @@ -0,0 +1,253 @@ +package models + +import ( + "database/sql" + "fmt" + "slices" + "strings" + "time" + + "github.com/gofrs/uuid" + "github.com/jackc/pgconn" + "github.com/jackc/pgerrcode" + "github.com/pkg/errors" + "github.com/supabase/auth/internal/storage" +) + +const scimVersionClause = "(?::timestamptz IS NULL OR updated_at = ?)" + +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 SCIMTarget struct { + ProviderID uuid.UUID + ID uuid.UUID + UpdatedAt *time.Time +} + +type scimTable struct { + tableName 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) { + 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) + } + 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.tableName), + uuid.Must(uuid.NewV4()), providerID, string(resource), + ).First(row); err != nil { + return nil, table.wrapError(err, "creating") + } + return row, nil +} + +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.tableName, table.targetClause(), lock), + target.ID, target.ProviderID, + ).First(row); err != nil { + return nil, table.wrapError(err, "finding") + } + 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( + 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) { + return nil, nil + } + return nil, table.wrapError(err, "finding") + } + 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( + 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") + } + return row, nil +} + +func (t scimTable) where(providerID uuid.UUID, filter SCIMFilter) (string, []any) { + clauses := []string{"sso_provider_id = ?"} + args := []any{providerID} + if t.liveClause != "" { + clauses = append(clauses, t.liveClause) + } + if filter.Name != nil { + clauses = append(clauses, t.nameColumn+` 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) targetClause() string { + if t.liveClause == "" { + return "id = ? AND sso_provider_id = ?" + } + return "id = ? AND sso_provider_id = ? AND " + t.liveClause +} + +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.tableName, 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.wrapError(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 { + direction = "DESC" + } + switch order.By { + case SCIMSortByID: + return "id " + direction + case SCIMSortByName: + 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) wrapError(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 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..092fd153d0 --- /dev/null +++ b/internal/models/scim_group.go @@ -0,0 +1,259 @@ +package models + +import ( + "bytes" + "fmt" + "slices" + "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{ + tableName: 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) { + return createSCIMRow[SCIMGroup](tx, scimGroupsTable, providerID, resource) +} + +func FindSCIMGroup(tx *storage.Connection, providerID, id uuid.UUID) (*SCIMGroup, error) { + 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, 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 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) +} + +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) { + touched := &SCIMGroup{} + if err := tx.RawQuery( + 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") + } + return touched, nil +} + +func DeleteSCIMGroup(tx *storage.Connection, target SCIMTarget) (*SCIMGroup, error) { + group := &SCIMGroup{} + err := tx.RawQuery( + 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 { + return nil, scimGroupsTable.writeError(tx, target, err, "deleting") + } + return group, nil +} + +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.tableName), + groupIDs, providerID, + ).All(&members) + if err != nil { + return nil, errors.Wrap(err, "error finding SCIM group members") + } + return members, nil +} + +func FindSCIMMembershipsByUser(tx *storage.Connection, providerID uuid.UUID, scimUserIDs []uuid.UUID) ([]SCIMGroupMembership, error) { + memberships := []SCIMGroupMembership{} + if len(scimUserIDs) == 0 { + 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.tableName), + scimUserIDs, providerID, + ).All(&memberships) + if err != nil { + return nil, errors.Wrap(err, "error finding SCIM groups for users") + } + return memberships, nil +} + +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, nil, err + } + 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)) + if err != nil { + return nil, nil, nil, err + } + 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( + 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) { + 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, + ).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 = clock_timestamp() WHERE id = ANY(?::uuid[])", groups), + groupIDs, + ).Exec(); err != nil { + return nil, errors.Wrap(err, "error updating SCIM groups") + } + } + return groupIDs, nil +} + +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.tableName), + 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} + } + return found, nil +} + +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 = ? ORDER BY created_at, scim_user_id", SCIMGroupMember{}.TableName()), + groupID, + ).All(&ids); err != nil { + return nil, errors.Wrap(err, "error finding SCIM group members") + } + return ids, nil +} + +func removeSCIMGroupMembers(tx *storage.Connection, groupID uuid.UUID, scimUserIDs []uuid.UUID) ([]uuid.UUID, error) { + removed := []uuid.UUID{} + if len(scimUserIDs) == 0 { + return removed, nil + } + if err := tx.RawQuery( + fmt.Sprintf("DELETE FROM %q WHERE group_id = ? AND scim_user_id = ANY(?::uuid[]) RETURNING scim_user_id", SCIMGroupMember{}.TableName()), + groupID, scimUserIDs, + ).All(&removed); err != nil { + return nil, errors.Wrap(err, "error removing SCIM group members") + } + return removed, nil +} + +func addSCIMGroupMembers(tx *storage.Connection, group *SCIMGroup, scimUserIDs []uuid.UUID) ([]uuid.UUID, error) { + added := []uuid.UUID{} + if len(scimUserIDs) == 0 { + 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[]) ON CONFLICT DO NOTHING RETURNING scim_user_id", SCIMGroupMember{}.TableName()), + group.ID, locked, + ).All(&added); err != nil { + return nil, errors.Wrap(err, "error adding SCIM group members") + } + 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 new file mode 100644 index 0000000000..acc1c79acd --- /dev/null +++ b/internal/models/scim_group_test.go @@ -0,0 +1,434 @@ +package models + +import ( + "bytes" + "fmt" + "testing" + "time" + + "github.com/gofrs/uuid" + "github.com/stretchr/testify/require" + "github.com/stretchr/testify/suite" + "github.com/supabase/auth/internal/storage" +) + +type SCIMGroupTestSuite struct { + suite.Suite + db *storage.Connection + provider *SSOProvider +} + +func TestSCIMGroup(t *testing.T) { + ts := &SCIMGroupTestSuite{db: setupSCIMTestDB(t)} + defer func() { require.NoError(t, 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) 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, 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, SCIMTarget{ProviderID: ts.provider.ID, ID: group.ID, UpdatedAt: &group.UpdatedAt}, []byte(`{"displayName":"Stale"}`)) + require.ErrorIs(ts.T(), err, SCIMGroupStaleError{}) + + _, err = ReplaceSCIMGroup(ts.db, SCIMTarget{ProviderID: ts.createProvider().ID, ID: group.ID}, []byte(`{"displayName":"Other"}`)) + 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") + _, _, _, 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}) + 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, SCIMTarget{ProviderID: ts.provider.ID, ID: group.ID}) + 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") + 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 := 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}) + + 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 := FindSCIMMembershipsByGroup(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, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) + 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, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) + 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, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) + 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 := 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) +} + +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, SCIMTarget{ProviderID: ts.provider.ID, ID: alice.ID}) + require.NoError(ts.T(), err) + + members, err := FindSCIMMembershipsByGroup(ts.db, ts.provider.ID, []uuid.UUID{group.ID}) + require.NoError(ts.T(), err) + 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) 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") + 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 := FindSCIMMembershipsByUser(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 = FindSCIMMembershipsByUser(ts.db, ts.createProvider().ID, []uuid.UUID{alice.ID}) + require.NoError(ts.T(), err) + require.Empty(ts.T(), groups) +} + +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 +} + +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_settings.go b/internal/models/scim_settings.go new file mode 100644 index 0000000000..e211e42f6f --- /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..26735ea246 --- /dev/null +++ b/internal/models/scim_settings_test.go @@ -0,0 +1,97 @@ +package models + +import ( + "sync" + "testing" + + "github.com/gofrs/uuid" + "github.com/stretchr/testify/require" + "github.com/stretchr/testify/suite" + "github.com/supabase/auth/internal/storage" +) + +type SCIMSettingsTestSuite struct { + suite.Suite + db *storage.Connection + provider *SSOProvider +} + +func TestSCIMSettings(t *testing.T) { + ts := &SCIMSettingsTestSuite{db: setupSCIMTestDB(t)} + defer func() { require.NoError(t, 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_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.go b/internal/models/scim_token.go new file mode 100644 index 0000000000..5de0f2cf6e --- /dev/null +++ b/internal/models/scim_token.go @@ -0,0 +1,175 @@ +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 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 := advisoryXactLock(tx, "scim_tokens|"+providerID.String()); err != nil { + return errors.Wrap(err, "error locking SCIM tokens") + } + return 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 +} + +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 new file mode 100644 index 0000000000..0ca617b262 --- /dev/null +++ b/internal/models/scim_token_test.go @@ -0,0 +1,253 @@ +package models + +import ( + "testing" + "time" + + "github.com/gofrs/uuid" + "github.com/stretchr/testify/require" + "github.com/stretchr/testify/suite" + "github.com/supabase/auth/internal/storage" +) + +type SCIMTokenTestSuite struct { + suite.Suite + db *storage.Connection + provider *SSOProvider +} + +func TestSCIMToken(t *testing.T) { + ts := &SCIMTokenTestSuite{db: setupSCIMTestDB(t)} + defer func() { require.NoError(t, 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) TestCreate() { + token, plaintext := ts.createToken(nil) + + 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) + 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) 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) 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()) +} diff --git a/internal/models/scim_user.go b/internal/models/scim_user.go new file mode 100644 index 0000000000..96d4cfaab1 --- /dev/null +++ b/internal/models/scim_user.go @@ -0,0 +1,307 @@ +package models + +import ( + "encoding/json" + "fmt" + "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" +} + +type SCIMIdentityRename struct { + UserID uuid.UUID + Provider string + From string + To string + Data map[string]any +} + +type SCIMIdentityEmailChange struct { + UserID uuid.UUID + Provider string + Subject string + Email string +} + +var scimUsersTable = scimTable{ + tableName: 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) { + return createSCIMRow[SCIMUser](tx, scimUsersTable, providerID, resource) +} + +func FindSCIMUser(tx *storage.Connection, providerID, id uuid.UUID) (*SCIMUser, error) { + 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, SCIMTarget{ProviderID: providerID, ID: id}, true) +} + +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) { + return findSCIMPage[SCIMUser](tx, scimUsersTable, providerID, query) +} + +func ReplaceSCIMUser(tx *storage.Connection, target SCIMTarget, resource []byte) (*SCIMUser, error) { + return replaceSCIMRow[SCIMUser](tx, scimUsersTable, target, resource) +} + +func DeleteSCIMUser(tx *storage.Connection, target SCIMTarget) (*SCIMUser, error) { + user := &SCIMUser{} + err := tx.RawQuery( + 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 { + return nil, scimUsersTable.writeError(tx, target, err, "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", scimUsersTable.tableName), + 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 LogoutUserForSCIM(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( + 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") + } + return rows, nil +} + +func BanDeprovisionedSCIMUsers(tx *storage.Connection, providerID uuid.UUID, until time.Time) (int, error) { + users, scimUsers := User{}.TableName(), scimUsersTable.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 + } + + 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.tableName, + ), + user.SSOProviderID, userID, user.SSOProviderID, userID, + ).First(&existing); err != nil { + return errors.Wrap(err, "error finding linked SCIM user") + } + if existing.Live { + return SCIMUserLinkedError{} + } + 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, + ).Exec(); err != nil { + return errors.Wrap(err, "error linking SCIM user") + } + user.UserID = &userID + return 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 IsSCIMUserDeprovisionedByProvider(tx *storage.Connection, providerID, userID uuid.UUID) (bool, error) { + result := struct { + AnyRow bool `db:"any_row"` + Live bool `db:"live"` + }{} + if err := tx.RawQuery( + 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.tableName, + ), + providerID, userID, providerID, userID, + ).First(&result); err != nil { + return false, errors.Wrap(err, "error finding SCIM user") + } + if !result.AnyRow { + return false, nil + } + return !result.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 = ?", scimUsersTable.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, 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), + 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), + rename.To, string(encoded), rename.UserID, rename.Provider, rename.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 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 +} 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) +);