From a7e424374188d65a049521165b0843fe8bf13abc Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 27 Sep 2026 21:46:49 +0530 Subject: [PATCH 1/3] feat(aws-cognito): users, admin user management and pool-delete parity (C1) --- docs/coverage/README.md | 2 +- docs/coverage/aws/README.md | 2 +- docs/coverage/aws/cognito.md | 15 +- docs/coverage/coverage.json | 40 +- providers/aws/cognito/cognito.go | 19 +- providers/aws/cognito/cognito_test.go | 30 +- providers/aws/cognito/custom_attributes.go | 83 +++ providers/aws/cognito/errors.go | 37 ++ providers/aws/cognito/passwords.go | 96 ++++ providers/aws/cognito/pool_delete_test.go | 87 ++++ providers/aws/cognito/snapshot.go | 7 + providers/aws/cognito/user_attributes.go | 90 ++++ providers/aws/cognito/user_filter.go | 106 ++++ providers/aws/cognito/user_pools.go | 45 +- providers/aws/cognito/users.go | 545 ++++++++++++++++++++ providers/aws/cognito/users_test.go | 557 +++++++++++++++++++++ server/aws/cognito/handler.go | 16 +- server/aws/cognito/user_ops.go | 250 +++++++++ server/aws/cognito/users_sdk_test.go | 367 ++++++++++++++ services/cognito/driver/driver.go | 45 +- services/cognito/driver/errors.go | 9 + services/cognito/driver/types.go | 51 ++ 22 files changed, 2433 insertions(+), 66 deletions(-) create mode 100644 providers/aws/cognito/custom_attributes.go create mode 100644 providers/aws/cognito/passwords.go create mode 100644 providers/aws/cognito/pool_delete_test.go create mode 100644 providers/aws/cognito/user_attributes.go create mode 100644 providers/aws/cognito/user_filter.go create mode 100644 providers/aws/cognito/users.go create mode 100644 providers/aws/cognito/users_test.go create mode 100644 server/aws/cognito/user_ops.go create mode 100644 server/aws/cognito/users_sdk_test.go diff --git a/docs/coverage/README.md b/docs/coverage/README.md index 9aeff26c3..000e42ddc 100644 --- a/docs/coverage/README.md +++ b/docs/coverage/README.md @@ -53,7 +53,7 @@ code does not implement. Machine-readable: [`coverage.json`](./coverage.json). | `cloudtasks` | - | - | [CloudTasks](./gcp/cloudtasks.md) | - | 11 | | `cloudtrail` | [CloudTrail](./aws/cloudtrail.md) | - | - | - | 60 | | `codeartifact` | [CodeArtifact](./aws/codeartifact.md) | - | - | - | 15 | -| `cognito` | [Cognito](./aws/cognito.md) | - | - | - | 18 | +| `cognito` | [Cognito](./aws/cognito.md) | - | - | - | 29 | | `communication` | - | [Communication](./azure/communication.md) | - | - | 10 | | `composer` | - | - | [Composer](./gcp/composer.md) | - | 6 | | `compute` | [EC2](./aws/ec2.md) | [VirtualMachines](./azure/virtualmachines.md) | [GCE](./gcp/gce.md) | - | 37 | diff --git a/docs/coverage/aws/README.md b/docs/coverage/aws/README.md index a6ac8a635..90b7f050a 100644 --- a/docs/coverage/aws/README.md +++ b/docs/coverage/aws/README.md @@ -25,7 +25,7 @@ Services cloudemu emulates for AWS, by native name. Back to the [cross-provider | [CloudWatch](./cloudwatch.md) | `monitoring` | 12 | | [CloudWatchLogs](./cloudwatchlogs.md) | `logging` | 17 | | [CodeArtifact](./codeartifact.md) | `codeartifact` | 15 | -| [Cognito](./cognito.md) | `cognito` | 18 | +| [Cognito](./cognito.md) | `cognito` | 29 | | [Config](./config.md) | `configservice` | 102 | | [CostExplorer](./costexplorer.md) | (provider-native) | 4 | | [DynamoDB](./dynamodb.md) | `database` | 24 | diff --git a/docs/coverage/aws/cognito.md b/docs/coverage/aws/cognito.md index e112485f1..9c4148c48 100644 --- a/docs/coverage/aws/cognito.md +++ b/docs/coverage/aws/cognito.md @@ -3,14 +3,24 @@ AWS's `cognito` service · portable interface `driver.Cognito` · [AWS index](./README.md) -## Operations (18) +## Operations (29) | Operation | Description | | --- | --- | +| `AddCustomAttributes` | AddCustomAttributes appends custom attributes to a pool's schema. Names get | +| `AdminCreateUser` | AdminCreateUser creates a user in FORCE_CHANGE_PASSWORD with a generated | +| `AdminDeleteUser` | | +| `AdminDeleteUserAttributes` | | +| `AdminDisableUser` | | +| `AdminEnableUser` | | +| `AdminGetUser` | | +| `AdminResetUserPassword` | AdminResetUserPassword moves the user to RESET_REQUIRED. | +| `AdminSetUserPassword` | AdminSetUserPassword sets a password checked against the pool policy. A | +| `AdminUpdateUserAttributes` | | | `CreateUserPool` | CreateUserPool creates a user pool, generating its id and ARN, seeding the | | `CreateUserPoolClient` | CreateUserPoolClient creates an app client, generating its 26-character id | | `CreateUserPoolDomain` | | -| `DeleteUserPool` | DeleteUserPool removes a user pool and its clients, domains, and tags. | +| `DeleteUserPool` | DeleteUserPool removes a user pool with its users, clients and tags. Like | | `DeleteUserPoolClient` | | | `DeleteUserPoolDomain` | | | `DescribeUserPool` | DescribeUserPool returns a deep copy of a user pool, or a | @@ -20,6 +30,7 @@ AWS's `cognito` service · portable interface `driver.Cognito` · [AWS index](./ | `ListTagsForResource` | | | `ListUserPoolClients` | ListUserPoolClients returns client descriptions in a user pool in a | | `ListUserPools` | ListUserPools returns pool descriptions in a deterministic order. | +| `ListUsers` | ListUsers returns users sorted by username, filtered by an optional | | `SetUserPoolMfaConfig` | SetUserPoolMfaConfig replaces a pool's MFA configuration and returns the | | `TagResource` | | | `UntagResource` | | diff --git a/docs/coverage/coverage.json b/docs/coverage/coverage.json index e3674849d..17964e1a6 100644 --- a/docs/coverage/coverage.json +++ b/docs/coverage/coverage.json @@ -3768,6 +3768,40 @@ "service": "cognito", "interface": "Cognito", "operations": [ + { + "name": "AddCustomAttributes", + "doc": "AddCustomAttributes appends custom attributes to a pool's schema. Names get" + }, + { + "name": "AdminCreateUser", + "doc": "AdminCreateUser creates a user in FORCE_CHANGE_PASSWORD with a generated" + }, + { + "name": "AdminDeleteUser" + }, + { + "name": "AdminDeleteUserAttributes" + }, + { + "name": "AdminDisableUser" + }, + { + "name": "AdminEnableUser" + }, + { + "name": "AdminGetUser" + }, + { + "name": "AdminResetUserPassword", + "doc": "AdminResetUserPassword moves the user to RESET_REQUIRED." + }, + { + "name": "AdminSetUserPassword", + "doc": "AdminSetUserPassword sets a password checked against the pool policy. A" + }, + { + "name": "AdminUpdateUserAttributes" + }, { "name": "CreateUserPool", "doc": "CreateUserPool creates a user pool, generating its id and ARN, seeding the" @@ -3781,7 +3815,7 @@ }, { "name": "DeleteUserPool", - "doc": "DeleteUserPool removes a user pool and its clients, domains, and tags." + "doc": "DeleteUserPool removes a user pool with its users, clients and tags. Like" }, { "name": "DeleteUserPoolClient" @@ -3815,6 +3849,10 @@ "name": "ListUserPools", "doc": "ListUserPools returns pool descriptions in a deterministic order." }, + { + "name": "ListUsers", + "doc": "ListUsers returns users sorted by username, filtered by an optional" + }, { "name": "SetUserPoolMfaConfig", "doc": "SetUserPoolMfaConfig replaces a pool's MFA configuration and returns the" diff --git a/providers/aws/cognito/cognito.go b/providers/aws/cognito/cognito.go index 84baac03e..65c563218 100644 --- a/providers/aws/cognito/cognito.go +++ b/providers/aws/cognito/cognito.go @@ -1,12 +1,8 @@ -// Package cognito provides an in-memory mock implementation of the AWS Cognito -// user-pools (cognito-idp) control plane: user pools, their app clients, and -// hosted-UI domains, plus resource tagging. +// Package cognito provides an in-memory mock of AWS Cognito user pools +// (cognito-idp): user pools, their app clients, hosted-UI domains, resource +// tagging, and pool users with the admin user-management operations. // -// This is the configuration control plane only. There is no authentication data -// plane behind the emulator (sign-up, sign-in, token issuance, users, and -// groups are out of scope), so the mock covers provisioning and reading pools, -// clients, and domains and their settings, which is what IaC tools (Terraform, -// CloudFormation) and the console's create flow exercise. +// Sign-up, sign-in and token issuance are not modeled yet. package cognito import ( @@ -29,13 +25,15 @@ const clientKeySep = "/" // plane. type Mock struct { // userPools is keyed by pool id; clients is keyed by "/"; - // domains is keyed by the domain string. + // domains is keyed by the domain string; users is keyed by + // "/". userPools *memstore.Store[driver.UserPool] clients *memstore.Store[driver.UserPoolClient] domains *memstore.Store[driver.UserPoolDomain] + users *memstore.Store[userRecord] // mu serializes compound read-modify-write mutations (pool update, cascading - // pool delete) that span more than one store operation. + // pool delete, user changes) that span more than one store operation. mu sync.Mutex // tagsMu guards the resource-tag side map, keyed by resource ARN. @@ -51,6 +49,7 @@ func New(opts *config.Options) *Mock { userPools: memstore.New[driver.UserPool](), clients: memstore.New[driver.UserPoolClient](), domains: memstore.New[driver.UserPoolDomain](), + users: memstore.New[userRecord](), tags: map[string]map[string]string{}, opts: opts, } diff --git a/providers/aws/cognito/cognito_test.go b/providers/aws/cognito/cognito_test.go index 585b6b1bb..57da6cb88 100644 --- a/providers/aws/cognito/cognito_test.go +++ b/providers/aws/cognito/cognito_test.go @@ -160,8 +160,9 @@ func TestUpdateAndDeleteUserPool(t *testing.T) { t.Fatalf("update not applied: %+v", got) } - requireNoError(t, m.DeleteUserPool(context.Background(), pool.ID), "DeleteUserPool") - assertNotFound(t, m.DeleteUserPool(context.Background(), pool.ID)) + if err := m.DeleteUserPool(context.Background(), pool.ID); err == nil { + t.Fatal("DeleteUserPool succeeded with deletion protection ACTIVE") + } } func TestClientSecretOnlyWithGenerate(t *testing.T) { @@ -282,31 +283,6 @@ func TestUserPoolDomainLifecycle(t *testing.T) { } } -func TestDeletePoolCascades(t *testing.T) { - m := newMock(t) - pool := mustCreatePool(t, m, "cascade-pool") - ctx := context.Background() - - client, err := m.CreateUserPoolClient(ctx, driver.CreateUserPoolClientInput{UserPoolID: pool.ID, ClientName: "c"}) - requireNoError(t, err, "CreateUserPoolClient") - - requireNoError(t, m.CreateUserPoolDomain(ctx, - driver.CreateUserPoolDomainInput{Domain: "d.example", UserPoolID: pool.ID}), "CreateUserPoolDomain") - - requireNoError(t, m.DeleteUserPool(ctx, pool.ID), "DeleteUserPool") - - if _, err := m.DescribeUserPoolClient(ctx, pool.ID, client.ClientID); !cerrors.IsNotFound(err) { - t.Fatal("client not removed on pool delete") - } - - dom, err := m.DescribeUserPoolDomain(ctx, "d.example") - requireNoError(t, err, "DescribeUserPoolDomain") - - if dom.Domain != "" { - t.Fatal("domain not removed on pool delete") - } -} - func TestTagsRoundTripAndReplace(t *testing.T) { m := newMock(t) ctx := context.Background() diff --git a/providers/aws/cognito/custom_attributes.go b/providers/aws/cognito/custom_attributes.go new file mode 100644 index 000000000..f025eebdb --- /dev/null +++ b/providers/aws/cognito/custom_attributes.go @@ -0,0 +1,83 @@ +package cognito + +import ( + "context" + "strings" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// maxCustomAttributes is the per-pool ceiling on custom attributes. +const maxCustomAttributes = 50 + +// AddCustomAttributes appends custom attributes to a pool's schema. Each name +// gets the custom: prefix (dev: for developer-only). Cognito never changes or +// removes an attribute once added, so a name already in the schema, or repeated +// in the request, is rejected and nothing is added. +func (m *Mock) AddCustomAttributes(_ context.Context, userPoolID string, attrs []driver.SchemaAttribute) error { + if len(attrs) == 0 { + return invalidParameter("1 validation error detected: Value null at 'customAttributes' failed to satisfy constraint: " + + "Member must not be null") + } + + m.mu.Lock() + defer m.mu.Unlock() + + pool, ok := m.userPools.Get(userPoolID) + if !ok { + return poolNotFound(userPoolID) + } + + pool = copyUserPool(pool) + schema := pool.SchemaAttributes + + for _, a := range attrs { + if a.Name == "" { + return invalidParameter("1 validation error detected: Value null at 'customAttributes.member.name' " + + "failed to satisfy constraint: Member must not be null") + } + + if a.AttributeDataType == "" { + a.AttributeDataType = driver.AttributeTypeString + } + + added := customAttribute(a) + if _, exists := schemaAttributeIn(schema, added.Name); exists { + return invalidParameter("Existing attribute already has name %s.", added.Name) + } + + schema = append(schema, added) + } + + if countCustom(schema) > maxCustomAttributes { + return invalidParameter("The user pool has reached the limit of %d custom attributes.", maxCustomAttributes) + } + + pool.SchemaAttributes = schema + pool.LastModifiedDate = m.now() + m.userPools.Set(userPoolID, pool) + + return nil +} + +func schemaAttributeIn(schema []driver.SchemaAttribute, name string) (driver.SchemaAttribute, bool) { + for _, a := range schema { + if a.Name == name { + return a, true + } + } + + return driver.SchemaAttribute{}, false +} + +func countCustom(schema []driver.SchemaAttribute) int { + n := 0 + + for _, a := range schema { + if strings.HasPrefix(a.Name, "custom:") || strings.HasPrefix(a.Name, "dev:") { + n++ + } + } + + return n +} diff --git a/providers/aws/cognito/errors.go b/providers/aws/cognito/errors.go index a9aa33ab4..989e0c83b 100644 --- a/providers/aws/cognito/errors.go +++ b/providers/aws/cognito/errors.go @@ -17,3 +17,40 @@ func invalidParameter(format string, args ...any) error { func resourceNotFound(format string, args ...any) error { return &driver.APIError{Exception: driver.ExResourceNotFound, Err: errors.Newf(errors.NotFound, format, args...)} } + +// poolNotFound is the ResourceNotFoundException for a missing user pool. +func poolNotFound(id string) error { + return resourceNotFound("User pool %s does not exist.", id) +} + +// userNotFound is the UserNotFoundException real Cognito returns for an unknown +// username. +func userNotFound() error { + //nolint:revive // exact Cognito message, surfaced verbatim to the SDK + return &driver.APIError{Exception: driver.ExUserNotFound, Err: errors.New(errors.NotFound, "User does not exist.")} +} + +// usernameExists builds a UsernameExistsException for a duplicate user. +func usernameExists(msg string) error { + return &driver.APIError{Exception: driver.ExUsernameExists, Err: errors.New(errors.AlreadyExists, msg)} +} + +// invalidPassword builds an InvalidPasswordException for a password that breaks +// the pool's policy. +func invalidPassword(reason string) error { + return &driver.APIError{ + Exception: driver.ExInvalidPassword, + Err: errors.New(errors.InvalidArgument, "Password did not conform with policy: "+reason), + } +} + +// notAuthorized builds a NotAuthorizedException for an operation the user's +// current state does not allow. +func notAuthorized(msg string) error { + return &driver.APIError{Exception: driver.ExNotAuthorized, Err: errors.New(errors.FailedPrecondition, msg)} +} + +// unsupportedUserState builds an UnsupportedUserStateException. +func unsupportedUserState(format string, args ...any) error { + return &driver.APIError{Exception: driver.ExUnsupportedUserState, Err: errors.Newf(errors.FailedPrecondition, format, args...)} +} diff --git a/providers/aws/cognito/passwords.go b/providers/aws/cognito/passwords.go new file mode 100644 index 000000000..eff358353 --- /dev/null +++ b/providers/aws/cognito/passwords.go @@ -0,0 +1,96 @@ +package cognito + +import ( + "crypto/rand" + "crypto/sha256" + "encoding/hex" + "strings" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// Password handling. Only a salted SHA-256 digest is kept; the plaintext is +// never stored or logged. +const ( + saltLen = 16 + generatedPasswordLen = 16 + maxPasswordLen = 256 +) + +// policySymbols is the set of special characters Cognito counts toward the +// "require symbols" rule. +const policySymbols = "^$*.[]{}()?\"!@#%&/\\,><':;|_~`=+- " + +// Character classes used to build a generated temporary password. +const ( + upperChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZ" + lowerChars = "abcdefghijklmnopqrstuvwxyz" + digitChars = "0123456789" + symbolChars = "!@#%&*?" +) + +// checkPassword validates a password against a pool's policy and returns the +// InvalidPasswordException real Cognito sends for the first rule it breaks. +func checkPassword(pw string, pp driver.PasswordPolicy) error { + if len(pw) < int(pp.MinimumLength) { + return invalidPassword("Password not long enough") + } + + if len(pw) > maxPasswordLen { + return invalidPassword("Password must have length less than or equal to 256") + } + + rules := []struct { + required bool + chars string + reason string + }{ + {pp.RequireUppercase, upperChars, "Password must have uppercase characters"}, + {pp.RequireLowercase, lowerChars, "Password must have lowercase characters"}, + {pp.RequireNumbers, digitChars, "Password must have numeric characters"}, + {pp.RequireSymbols, policySymbols, "Password must have symbol characters"}, + } + + for _, r := range rules { + if r.required && !strings.ContainsAny(pw, r.chars) { + return invalidPassword(r.reason) + } + } + + return nil +} + +// generatePassword returns a temporary password that satisfies any policy with +// a minimum length up to its length: it always carries every character class. +func generatePassword(pp driver.PasswordPolicy) string { + n := max(generatedPasswordLen, int(pp.MinimumLength)) + + var b strings.Builder + + b.WriteString(randString(1, upperChars)) + b.WriteString(randString(1, lowerChars)) + b.WriteString(randString(1, digitChars)) + b.WriteString(randString(1, symbolChars)) + b.WriteString(randString(n-b.Len(), alnumMixed)) + + return b.String() +} + +// hashPassword returns a fresh random salt and the SHA-256 digest of salt+pw, +// both hex encoded. +// +//nolint:gocritic // unnamedResult: (salt, digest) reads clearly at the call sites +func hashPassword(pw string) (string, string) { + raw := make([]byte, saltLen) + _, _ = rand.Read(raw) + + salt := hex.EncodeToString(raw) + + return salt, digest(salt, pw) +} + +func digest(salt, pw string) string { + sum := sha256.Sum256([]byte(salt + pw)) + + return hex.EncodeToString(sum[:]) +} diff --git a/providers/aws/cognito/pool_delete_test.go b/providers/aws/cognito/pool_delete_test.go new file mode 100644 index 000000000..02403bef1 --- /dev/null +++ b/providers/aws/cognito/pool_delete_test.go @@ -0,0 +1,87 @@ +package cognito + +import ( + "context" + "errors" + "testing" + + cerrors "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +func assertInvalidParameter(t *testing.T, err error, wantMsg string) { + t.Helper() + + var apiErr *driver.APIError + if !errors.As(err, &apiErr) || apiErr.Exception != driver.ExInvalidParameter { + t.Fatalf("expected InvalidParameterException, got %v", err) + } + + if wantMsg != "" && cerrors.Message(err) != wantMsg { + t.Fatalf("message = %q, want %q", cerrors.Message(err), wantMsg) + } +} + +func TestDeleteUserPoolBlockedByDomain(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "domain-pool") + + requireNoError(t, m.CreateUserPoolDomain(ctx, + driver.CreateUserPoolDomainInput{Domain: "d.example", UserPoolID: pool.ID}), "CreateUserPoolDomain") + + assertInvalidParameter(t, m.DeleteUserPool(ctx, pool.ID), + "User pool cannot be deleted. It has a domain configured that should be deleted first.") + + // The pool and its domain are both untouched by the refused delete. + if _, err := m.DescribeUserPool(ctx, pool.ID); err != nil { + t.Fatalf("pool gone after refused delete: %v", err) + } + + dom, err := m.DescribeUserPoolDomain(ctx, "d.example") + requireNoError(t, err, "DescribeUserPoolDomain") + + if dom.UserPoolID != pool.ID { + t.Fatalf("domain changed after refused delete: %+v", dom) + } + + requireNoError(t, m.DeleteUserPoolDomain(ctx, "d.example", pool.ID), "DeleteUserPoolDomain") + requireNoError(t, m.DeleteUserPool(ctx, pool.ID), "DeleteUserPool after domain removed") +} + +func TestDeleteUserPoolBlockedByDeletionProtection(t *testing.T) { + m := newMock(t) + ctx := context.Background() + + pool, err := m.CreateUserPool(ctx, driver.CreateUserPoolInput{ + Name: "protected", + DeletionProtection: driver.DeletionProtectionActive, + }) + requireNoError(t, err, "CreateUserPool") + + assertInvalidParameter(t, m.DeleteUserPool(ctx, pool.ID), + "The user pool cannot be deleted because deletion protection is activated. Deletion protection must be inactivated first.") + + requireNoError(t, m.UpdateUserPool(ctx, driver.UpdateUserPoolInput{ + ID: pool.ID, + DeletionProtection: driver.DeletionProtectionInactive, + }), "UpdateUserPool") + + requireNoError(t, m.DeleteUserPool(ctx, pool.ID), "DeleteUserPool after protection off") + assertNotFound(t, m.DeleteUserPool(ctx, pool.ID)) +} + +func TestDeletePoolCascadesClients(t *testing.T) { + m := newMock(t) + pool := mustCreatePool(t, m, "cascade-pool") + ctx := context.Background() + + client, err := m.CreateUserPoolClient(ctx, driver.CreateUserPoolClientInput{UserPoolID: pool.ID, ClientName: "c"}) + requireNoError(t, err, "CreateUserPoolClient") + + requireNoError(t, m.DeleteUserPool(ctx, pool.ID), "DeleteUserPool") + + if _, err := m.DescribeUserPoolClient(ctx, pool.ID, client.ClientID); !cerrors.IsNotFound(err) { + t.Fatal("client not removed on pool delete") + } +} diff --git a/providers/aws/cognito/snapshot.go b/providers/aws/cognito/snapshot.go index 3864edb64..61a4b292f 100644 --- a/providers/aws/cognito/snapshot.go +++ b/providers/aws/cognito/snapshot.go @@ -19,6 +19,7 @@ type cognitoSnapshot struct { UserPools map[string]driver.UserPool `json:"userPools,omitempty"` Clients map[string]driver.UserPoolClient `json:"clients,omitempty"` Domains map[string]driver.UserPoolDomain `json:"domains,omitempty"` + Users map[string]userRecord `json:"users,omitempty"` Tags map[string]map[string]string `json:"tags,omitempty"` } @@ -29,6 +30,7 @@ func (m *Mock) Snapshot(_ context.Context, _ bool) (json.RawMessage, error) { UserPools: deepCopyMap(m.userPools.All(), copyUserPool), Clients: deepCopyMap(m.clients.All(), copyUserPoolClient), Domains: deepCopyMap(m.domains.All(), copyUserPoolDomain), + Users: deepCopyMap(m.users.All(), copyUserRecord), } m.tagsMu.RLock() @@ -48,6 +50,7 @@ func (m *Mock) Restore(_ context.Context, data json.RawMessage) error { m.userPools.Clear() m.clients.Clear() m.domains.Clear() + m.users.Clear() for k := range snap.UserPools { m.userPools.Set(k, copyUserPool(snap.UserPools[k])) @@ -61,6 +64,10 @@ func (m *Mock) Restore(_ context.Context, data json.RawMessage) error { m.domains.Set(k, copyUserPoolDomain(snap.Domains[k])) } + for k := range snap.Users { + m.users.Set(k, copyUserRecord(snap.Users[k])) + } + m.tagsMu.Lock() m.tags = deepCopyTags(snap.Tags) diff --git a/providers/aws/cognito/user_attributes.go b/providers/aws/cognito/user_attributes.go new file mode 100644 index 000000000..5aa0be5db --- /dev/null +++ b/providers/aws/cognito/user_attributes.go @@ -0,0 +1,90 @@ +package cognito + +import ( + "slices" + "strconv" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// schemaError builds the InvalidParameterException Cognito sends when a user +// attribute does not fit the pool schema. +func schemaError(name, reason string) error { + return invalidParameter("Attributes did not conform to the schema: %s: %s", name, reason) +} + +// schemaAttribute looks up an attribute in the pool schema by its full name. +func schemaAttribute(pool *driver.UserPool, name string) (driver.SchemaAttribute, bool) { + return schemaAttributeIn(pool.SchemaAttributes, name) +} + +// validateAttributes checks user attributes against the pool schema: every name +// must exist, sub is never caller-set, an update may not touch an immutable +// attribute, and a string may not pass its maximum length. +func validateAttributes(pool *driver.UserPool, attrs []driver.Attribute, update bool) error { + for _, attr := range attrs { + a, ok := schemaAttribute(pool, attr.Name) + if !ok { + return schemaError(attr.Name, "Attribute does not exist in the schema.") + } + + if attr.Name == attrSub || (update && !a.Mutable) { + return schemaError(attr.Name, "Attribute cannot be updated. (changing an immutable attribute)") + } + + if c := a.StringAttributeConstraints; c != nil && c.MaxLength != "" { + if limit, err := strconv.Atoi(c.MaxLength); err == nil && len(attr.Value) > limit { + return schemaError(attr.Name, "String must be no longer than "+c.MaxLength+" characters") + } + } + } + + return nil +} + +// mergeAttributes returns base with updates applied: an existing name is +// replaced in place and a new one is appended. base is not modified. +func mergeAttributes(base, updates []driver.Attribute) []driver.Attribute { + out := slices.Clone(base) + + for _, u := range updates { + i := slices.IndexFunc(out, func(a driver.Attribute) bool { return a.Name == u.Name }) + if i >= 0 { + out[i].Value = u.Value + + continue + } + + out = append(out, u) + } + + return out +} + +// attrValue returns the value of a named attribute, or "" when absent. +func attrValue(attrs []driver.Attribute, name string) string { + v, _ := lastValue(attrs, name) + + return v +} + +// lastValue returns the last value given for name and whether it was present. +func lastValue(attrs []driver.Attribute, name string) (string, bool) { + for i := len(attrs) - 1; i >= 0; i-- { + if attrs[i].Name == name { + return attrs[i].Value, true + } + } + + return "", false +} + +// selectAttributes applies ListUsers AttributesToGet: nil keeps everything, and +// an empty list keeps nothing. +func selectAttributes(attrs []driver.Attribute, want []string) []driver.Attribute { + if want == nil { + return attrs + } + + return slices.DeleteFunc(attrs, func(a driver.Attribute) bool { return !slices.Contains(want, a.Name) }) +} diff --git a/providers/aws/cognito/user_filter.go b/providers/aws/cognito/user_filter.go new file mode 100644 index 000000000..0d38c4223 --- /dev/null +++ b/providers/aws/cognito/user_filter.go @@ -0,0 +1,106 @@ +package cognito + +import ( + "regexp" + "strings" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// filterPrefix is the ListUsers "starts with" operator; "=" is an exact match. +const filterPrefix = "^=" + +// Pseudo-attributes a ListUsers filter can search besides real attributes. +const ( + searchUsername = "username" + searchUserStatus = "cognito:user_status" + searchStatus = "status" +) + +// Values of the "status" search attribute. +const ( + statusEnabled = "Enabled" + statusDisabled = "Disabled" +) + +// userFilterPattern matches `attr = "value"` or `attr ^= "value"`, with \" and +// \\ escapes inside the quoted value. +var userFilterPattern = regexp.MustCompile(`^\s*([^\s=^]+)\s*(\^=|=)\s*"((?:[^"\\]|\\.)*)"\s*$`) + +// searchableAttributes lists what ListUsers can filter on. Custom attributes +// are not searchable. +// +//nolint:gochecknoglobals // read-only lookup table +var searchableAttributes = map[string]bool{ + searchUsername: true, attrEmail: true, attrPhoneNumber: true, "name": true, + "given_name": true, "family_name": true, attrPreferredUsername: true, + searchUserStatus: true, searchStatus: true, attrSub: true, +} + +// userFilter is a parsed ListUsers filter. The zero value matches every user. +type userFilter struct { + attr string + op string + value string +} + +// parseUserFilter parses a ListUsers Filter. An empty filter matches all users. +func parseUserFilter(s string) (userFilter, error) { + if strings.TrimSpace(s) == "" { + return userFilter{}, nil + } + + m := userFilterPattern.FindStringSubmatch(s) + if m == nil { + return userFilter{}, invalidParameter("Error while parsing filter.") + } + + if !searchableAttributes[m[1]] { + return userFilter{}, invalidParameter("Invalid search attribute: %s", m[1]) + } + + value := strings.NewReplacer(`\"`, `"`, `\\`, `\`).Replace(m[3]) + + return userFilter{attr: m[1], op: m[2], value: value}, nil +} + +// matches reports whether a user passes the filter. cognito:user_status +// compares case-insensitively; everything else is case-sensitive. +func (f userFilter) matches(u *driver.User) bool { + if f.attr == "" { + return true + } + + got, ok := f.lookup(u) + if !ok { + return false + } + + want := f.value + if f.attr == searchUserStatus { + got, want = strings.ToUpper(got), strings.ToUpper(want) + } + + if f.op == filterPrefix { + return strings.HasPrefix(got, want) + } + + return got == want +} + +func (f userFilter) lookup(u *driver.User) (string, bool) { + switch f.attr { + case searchUsername: + return u.Username, true + case searchUserStatus: + return u.UserStatus, true + case searchStatus: + if u.Enabled { + return statusEnabled, true + } + + return statusDisabled, true + default: + return lastValue(u.Attributes, f.attr) + } +} diff --git a/providers/aws/cognito/user_pools.go b/providers/aws/cognito/user_pools.go index 6c7d21abb..b89ead006 100644 --- a/providers/aws/cognito/user_pools.go +++ b/providers/aws/cognito/user_pools.go @@ -67,6 +67,7 @@ func (m *Mock) DescribeUserPool(_ context.Context, id string) (*driver.UserPool, out := copyUserPool(pool) out.Tags = m.currentTags(pool.ARN) + out.EstimatedNumberOfUsers = m.countUsers(id) return &out, nil } @@ -112,28 +113,36 @@ func (m *Mock) UpdateUserPool(_ context.Context, in driver.UpdateUserPoolInput) return nil } -// DeleteUserPool removes a user pool along with its clients, domains, and tags. +// DeleteUserPool removes a user pool along with its users, clients, and tags. +// Real Cognito refuses the delete while deletion protection is ACTIVE or a +// hosted-UI domain is still attached, so the caller must clear those first. func (m *Mock) DeleteUserPool(_ context.Context, id string) error { m.mu.Lock() defer m.mu.Unlock() pool, ok := m.userPools.Get(id) if !ok { - return resourceNotFound("User pool %s does not exist.", id) + return poolNotFound(id) } - for _, key := range m.clients.Keys() { - if c, ok := m.clients.Get(key); ok && c.UserPoolID == id { - m.clients.Delete(key) + if pool.DeletionProtection == driver.DeletionProtectionActive { + return invalidParameter("The user pool cannot be deleted because deletion protection is activated. " + + "Deletion protection must be inactivated first.") + } + + for _, d := range m.domains.All() { + if d.UserPoolID == id { + return invalidParameter("User pool cannot be deleted. It has a domain configured that should be deleted first.") } } - for _, domain := range m.domains.Keys() { - if d, ok := m.domains.Get(domain); ok && d.UserPoolID == id { - m.domains.Delete(domain) + for _, key := range m.clients.Keys() { + if c, ok := m.clients.Get(key); ok && c.UserPoolID == id { + m.clients.Delete(key) } } + m.deletePoolUsers(id) m.userPools.Delete(id) m.deleteTags(pool.ARN) @@ -142,6 +151,11 @@ func (m *Mock) DeleteUserPool(_ context.Context, id string) error { // ListUserPools returns pool descriptions sorted by id. func (m *Mock) ListUserPools(_ context.Context, page driver.Pagination) ([]driver.UserPoolDescription, string, error) { + if page.MaxResults > defaultPageSize { + return nil, "", invalidParameter("1 validation error detected: Value '%d' at 'maxResults' failed to satisfy constraint: "+ + "Member must have value less than or equal to %d", page.MaxResults, defaultPageSize) + } + ids := sortedKeys(m.userPools.Keys()) all := make([]driver.UserPoolDescription, 0, len(ids)) @@ -214,15 +228,22 @@ func mergeSchema(custom []driver.SchemaAttribute) []driver.SchemaAttribute { continue } - a.Name = customPrefix(a.DeveloperOnlyAttribute) + a.Name - a.StringAttributeConstraints = copyStringConstraints(a.StringAttributeConstraints) - a.NumberAttributeConstraints = copyNumberConstraints(a.NumberAttributeConstraints) - attrs = append(attrs, a) + attrs = append(attrs, customAttribute(a)) } return attrs } +// customAttribute returns a caller-supplied attribute with its custom: or dev: +// prefix applied and its constraints deep-copied. +func customAttribute(a driver.SchemaAttribute) driver.SchemaAttribute { + a.Name = customPrefix(a.DeveloperOnlyAttribute) + a.Name + a.StringAttributeConstraints = copyStringConstraints(a.StringAttributeConstraints) + a.NumberAttributeConstraints = copyNumberConstraints(a.NumberAttributeConstraints) + + return a +} + // customPrefix returns the attribute-name prefix Cognito applies to a // non-standard attribute. func customPrefix(developerOnly bool) string { diff --git a/providers/aws/cognito/users.go b/providers/aws/cognito/users.go new file mode 100644 index 000000000..e69b0867a --- /dev/null +++ b/providers/aws/cognito/users.go @@ -0,0 +1,545 @@ +package cognito + +import ( + "context" + "slices" + "sort" + "strings" + + "github.com/stackshy/cloudemu/v2/internal/idgen" + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// Attribute names with special handling. +const ( + attrSub = "sub" + attrEmail = "email" + attrEmailVerified = "email_verified" + attrPhoneNumber = "phone_number" + attrPhoneNumberVerified = "phone_number_verified" + attrPreferredUsername = "preferred_username" + attrTrue = "true" + attrFalse = "false" +) + +// maxUsernameLen is the UsernameType length ceiling. +const maxUsernameLen = 128 + +// userRecord is a stored user: the public view plus the password digest. It +// lives in the users store keyed by userKey(poolID, username). +type userRecord struct { + PoolID string `json:"poolId"` + User driver.User `json:"user"` + PasswordSalt string `json:"passwordSalt,omitempty"` + PasswordHash string `json:"passwordHash,omitempty"` +} + +func userKey(poolID, username string) string { return poolID + clientKeySep + username } + +//nolint:gocritic // hugeParam: value signature required by the func(V) V copy callback +func copyUserRecord(in userRecord) userRecord { + out := in + out.User.Attributes = slices.Clone(in.User.Attributes) + + return out +} + +// AdminCreateUser creates a user in FORCE_CHANGE_PASSWORD, or with MessageAction +// RESEND re-invites an existing one. +// +//nolint:gocritic // hugeParam: taken by value to match the driver interface +func (m *Mock) AdminCreateUser(_ context.Context, in driver.AdminCreateUserInput) (*driver.User, error) { + if err := checkUsername(in.Username); err != nil { + return nil, err + } + + if err := checkMessageAction(in.MessageAction); err != nil { + return nil, err + } + + m.mu.Lock() + defer m.mu.Unlock() + + pool, ok := m.userPools.Get(in.UserPoolID) + if !ok { + return nil, poolNotFound(in.UserPoolID) + } + + if in.MessageAction == driver.MessageActionResend { + return m.resendInvite(&pool, in) + } + + if err := validateAttributes(&pool, in.UserAttributes, false); err != nil { + return nil, err + } + + password := in.TemporaryPassword + if password == "" { + password = generatePassword(pool.Policies.PasswordPolicy) + } else if err := checkPassword(password, pool.Policies.PasswordPolicy); err != nil { + return nil, err + } + + sub := idgen.UUID() + attrs := mergeAttributes([]driver.Attribute{{Name: attrSub, Value: sub}}, in.UserAttributes) + + username, attrs, err := m.newUsername(&pool, in.Username, sub, attrs) + if err != nil { + return nil, err + } + + now := m.now() + rec := userRecord{ + PoolID: pool.ID, + User: driver.User{ + Username: username, + Attributes: attrs, + UserCreateDate: now, + UserLastModifiedDate: now, + Enabled: true, + UserStatus: driver.UserStatusForceChangePassword, + }, + } + rec.PasswordSalt, rec.PasswordHash = hashPassword(password) + + m.users.Set(userKey(pool.ID, username), copyUserRecord(rec)) + + out := copyUserRecord(rec).User + + return &out, nil +} + +// resendInvite handles AdminCreateUser with MessageAction RESEND: only a user +// still in FORCE_CHANGE_PASSWORD can be re-invited, and it gets a fresh +// temporary password. +// +//nolint:gocritic // hugeParam: in is the caller's input, passed through by value +func (m *Mock) resendInvite(pool *driver.UserPool, in driver.AdminCreateUserInput) (*driver.User, error) { + key, rec, ok := m.resolveUser(pool, in.Username) + if !ok { + return nil, userNotFound() + } + + if rec.User.UserStatus != driver.UserStatusForceChangePassword { + return nil, unsupportedUserState("Resend not possible. %s status is not FORCE_CHANGE_PASSWORD", in.Username) + } + + password := in.TemporaryPassword + if password == "" { + password = generatePassword(pool.Policies.PasswordPolicy) + } else if err := checkPassword(password, pool.Policies.PasswordPolicy); err != nil { + return nil, err + } + + rec = copyUserRecord(rec) + rec.PasswordSalt, rec.PasswordHash = hashPassword(password) + rec.User.UserLastModifiedDate = m.now() + m.users.Set(key, rec) + + out := copyUserRecord(rec).User + + return &out, nil +} + +// newUsername decides the stored username for a new user and checks it is +// free. In a pool that signs in with email or phone number, the username must +// be one of those, the stored username is the sub, and the value is copied into +// the matching attribute. +func (m *Mock) newUsername( + pool *driver.UserPool, username, sub string, attrs []driver.Attribute, +) (string, []driver.Attribute, error) { + if len(pool.UsernameAttributes) == 0 { + if err := checkNotAliasFormat(pool, username); err != nil { + return "", nil, err + } + + if m.users.Has(userKey(pool.ID, username)) { + return "", nil, usernameExists("User account already exists") + } + + return username, attrs, nil + } + + attr, err := usernameAttributeFor(pool.UsernameAttributes, username) + if err != nil { + return "", nil, err + } + + users := m.poolUsers(pool.ID) + for i := range users { + if attrValue(users[i].User.Attributes, attr) == username { + return "", nil, usernameExists("An account with the given " + attr + " already exists.") + } + } + + return sub, mergeAttributes(attrs, []driver.Attribute{{Name: attr, Value: username}}), nil +} + +// usernameAttributeFor returns the username attribute (email or phone_number) a +// sign-in name belongs to, or the InvalidParameterException Cognito sends when +// it is neither. +func usernameAttributeFor(allowed []string, username string) (string, error) { + switch { + case isEmailFormat(username) && slices.Contains(allowed, attrEmail): + return attrEmail, nil + case isPhoneFormat(username) && slices.Contains(allowed, attrPhoneNumber): + return attrPhoneNumber, nil + } + + switch { + case slices.Contains(allowed, attrEmail) && slices.Contains(allowed, attrPhoneNumber): + return "", invalidParameter("Username should be either an email or a phone number.") + case slices.Contains(allowed, attrEmail): + return "", invalidParameter("Username should be an email.") + default: + return "", invalidParameter("Username should be a phone number.") + } +} + +// checkNotAliasFormat rejects a username shaped like an email or phone number +// when the pool uses that attribute as a sign-in alias. +func checkNotAliasFormat(pool *driver.UserPool, username string) error { + if slices.Contains(pool.AliasAttributes, attrEmail) && isEmailFormat(username) { + return invalidParameter("Username cannot be of email format, since user pool is configured for email alias.") + } + + if slices.Contains(pool.AliasAttributes, attrPhoneNumber) && isPhoneFormat(username) { + return invalidParameter("Username cannot be of phone number format, since user pool is configured for phone number alias.") + } + + return nil +} + +// AdminGetUser returns one user. +func (m *Mock) AdminGetUser(_ context.Context, userPoolID, username string) (*driver.User, error) { + pool, ok := m.userPools.Get(userPoolID) + if !ok { + return nil, poolNotFound(userPoolID) + } + + _, rec, ok := m.resolveUser(&pool, username) + if !ok { + return nil, userNotFound() + } + + out := copyUserRecord(rec).User + + return &out, nil +} + +// ListUsers returns a page of a pool's users sorted by username. +// +//nolint:gocritic // hugeParam: taken by value to match the driver interface +func (m *Mock) ListUsers(_ context.Context, in driver.ListUsersInput) ([]driver.User, string, error) { + if in.Limit > defaultPageSize { + return nil, "", invalidParameter("1 validation error detected: Value '%d' at 'limit' failed to satisfy constraint: "+ + "Member must have value less than or equal to %d", in.Limit, defaultPageSize) + } + + if !m.userPools.Has(in.UserPoolID) { + return nil, "", poolNotFound(in.UserPoolID) + } + + filter, err := parseUserFilter(in.Filter) + if err != nil { + return nil, "", err + } + + var matched []driver.User + + users := m.poolUsers(in.UserPoolID) + for i := range users { + if filter.matches(&users[i].User) { + u := copyUserRecord(users[i]).User + u.Attributes = selectAttributes(u.Attributes, in.AttributesToGet) + matched = append(matched, u) + } + } + + page, next, err := paginate(matched, driver.Pagination{NextToken: in.PaginationToken, MaxResults: in.Limit}) + if err != nil { + return nil, "", err + } + + return page, next, nil +} + +// AdminDeleteUser removes a user. +func (m *Mock) AdminDeleteUser(_ context.Context, userPoolID, username string) error { + m.mu.Lock() + defer m.mu.Unlock() + + pool, ok := m.userPools.Get(userPoolID) + if !ok { + return poolNotFound(userPoolID) + } + + key, _, ok := m.resolveUser(&pool, username) + if !ok { + return userNotFound() + } + + m.users.Delete(key) + + return nil +} + +// AdminUpdateUserAttributes sets attribute values. Changing email or phone +// number without also setting its verified flag marks it unverified, as Cognito +// does. +func (m *Mock) AdminUpdateUserAttributes(_ context.Context, userPoolID, username string, attrs []driver.Attribute) error { + return m.updateUser(userPoolID, username, func(pool *driver.UserPool, rec *userRecord) error { + if err := validateAttributes(pool, attrs, true); err != nil { + return err + } + + updates := slices.Clone(attrs) + updates = append(updates, unverifyChanged(rec.User.Attributes, attrs)...) + rec.User.Attributes = mergeAttributes(rec.User.Attributes, updates) + + return nil + }) +} + +// unverifyChanged returns the "_verified=false" updates for an email or +// phone number that changes value without its verified flag in the same call. +func unverifyChanged(current, updates []driver.Attribute) []driver.Attribute { + var out []driver.Attribute + + for _, pair := range [][2]string{{attrEmail, attrEmailVerified}, {attrPhoneNumber, attrPhoneNumberVerified}} { + v, changed := lastValue(updates, pair[0]) + if !changed || v == attrValue(current, pair[0]) { + continue + } + + if _, set := lastValue(updates, pair[1]); !set { + out = append(out, driver.Attribute{Name: pair[1], Value: attrFalse}) + } + } + + return out +} + +// AdminDeleteUserAttributes removes attributes from a user. +func (m *Mock) AdminDeleteUserAttributes(_ context.Context, userPoolID, username string, names []string) error { + return m.updateUser(userPoolID, username, func(pool *driver.UserPool, rec *userRecord) error { + for _, name := range names { + a, ok := schemaAttribute(pool, name) + if !ok { + return schemaError(name, "Attribute does not exist in the schema.") + } + + if !a.Mutable { + return schemaError(name, "Attribute cannot be updated. (changing an immutable attribute)") + } + } + + rec.User.Attributes = slices.DeleteFunc(rec.User.Attributes, func(a driver.Attribute) bool { + return slices.Contains(names, a.Name) + }) + + return nil + }) +} + +// AdminSetUserPassword sets a user's password. +func (m *Mock) AdminSetUserPassword(_ context.Context, userPoolID, username, password string, permanent bool) error { + return m.updateUser(userPoolID, username, func(pool *driver.UserPool, rec *userRecord) error { + if err := checkPassword(password, pool.Policies.PasswordPolicy); err != nil { + return err + } + + rec.PasswordSalt, rec.PasswordHash = hashPassword(password) + rec.User.UserStatus = driver.UserStatusForceChangePassword + + if permanent { + rec.User.UserStatus = driver.UserStatusConfirmed + } + + return nil + }) +} + +// AdminEnableUser enables a user. +func (m *Mock) AdminEnableUser(_ context.Context, userPoolID, username string) error { + return m.updateUser(userPoolID, username, func(_ *driver.UserPool, rec *userRecord) error { + rec.User.Enabled = true + + return nil + }) +} + +// AdminDisableUser disables a user. +func (m *Mock) AdminDisableUser(_ context.Context, userPoolID, username string) error { + return m.updateUser(userPoolID, username, func(_ *driver.UserPool, rec *userRecord) error { + rec.User.Enabled = false + + return nil + }) +} + +// AdminResetUserPassword moves a user to RESET_REQUIRED. A user who has not +// yet replaced the temporary password cannot be reset. +func (m *Mock) AdminResetUserPassword(_ context.Context, userPoolID, username string) error { + return m.updateUser(userPoolID, username, func(_ *driver.UserPool, rec *userRecord) error { + if rec.User.UserStatus == driver.UserStatusForceChangePassword { + return notAuthorized("User password cannot be reset in the current state.") + } + + rec.User.UserStatus = driver.UserStatusResetRequired + + return nil + }) +} + +// updateUser runs fn on a copy of the resolved user under the mutation lock and +// stores the result with a fresh last-modified time. +func (m *Mock) updateUser(userPoolID, username string, fn func(*driver.UserPool, *userRecord) error) error { + m.mu.Lock() + defer m.mu.Unlock() + + pool, ok := m.userPools.Get(userPoolID) + if !ok { + return poolNotFound(userPoolID) + } + + key, rec, ok := m.resolveUser(&pool, username) + if !ok { + return userNotFound() + } + + rec = copyUserRecord(rec) + if err := fn(&pool, &rec); err != nil { + return err + } + + rec.User.UserLastModifiedDate = m.now() + m.users.Set(key, rec) + + return nil +} + +// resolveUser finds a user by username, then by the sign-in attributes the pool +// allows: a username attribute (email or phone number), or an alias. Email and +// phone aliases only resolve once verified; preferred_username always does. +func (m *Mock) resolveUser(pool *driver.UserPool, name string) (string, userRecord, bool) { + if rec, ok := m.users.Get(userKey(pool.ID, name)); ok { + return userKey(pool.ID, name), rec, true + } + + users := m.poolUsers(pool.ID) + for i := range users { + if signInMatches(pool, &users[i].User, name) { + return userKey(pool.ID, users[i].User.Username), users[i], true + } + } + + return "", userRecord{}, false +} + +func signInMatches(pool *driver.UserPool, u *driver.User, name string) bool { + for _, attr := range pool.UsernameAttributes { + if attrValue(u.Attributes, attr) == name { + return true + } + } + + for _, attr := range pool.AliasAttributes { + if attrValue(u.Attributes, attr) != name { + continue + } + + switch attr { + case attrEmail: + return attrValue(u.Attributes, attrEmailVerified) == attrTrue + case attrPhoneNumber: + return attrValue(u.Attributes, attrPhoneNumberVerified) == attrTrue + default: + return true + } + } + + return false +} + +// poolUserKeys returns the store keys of a pool's users, sorted. Pool ids never +// contain the key separator, so the "/" prefix selects exactly one pool. +func (m *Mock) poolUserKeys(poolID string) []string { + prefix := userKey(poolID, "") + + keys := slices.DeleteFunc(m.users.Keys(), func(k string) bool { return !strings.HasPrefix(k, prefix) }) + sort.Strings(keys) + + return keys +} + +// poolUsers returns a pool's users sorted by username. +func (m *Mock) poolUsers(poolID string) []userRecord { + keys := m.poolUserKeys(poolID) + out := make([]userRecord, 0, len(keys)) + + for _, k := range keys { + if rec, ok := m.users.Get(k); ok { + out = append(out, rec) + } + } + + return out +} + +// countUsers returns the number of users in a pool. +func (m *Mock) countUsers(poolID string) int32 { + return int32(len(m.poolUserKeys(poolID))) //nolint:gosec // user counts stay far below int32 max +} + +// deletePoolUsers removes every user of a pool. +func (m *Mock) deletePoolUsers(poolID string) { + for _, k := range m.poolUserKeys(poolID) { + m.users.Delete(k) + } +} + +func checkUsername(username string) error { + if username == "" { + return invalidParameter("1 validation error detected: Value null at 'username' failed to satisfy constraint: Member must not be null") + } + + if len(username) > maxUsernameLen || strings.ContainsFunc(username, isSpaceRune) { + return invalidParameter("1 validation error detected: Value at 'username' failed to satisfy constraint: " + + "Member must satisfy regular expression pattern: [\\p{L}\\p{M}\\p{S}\\p{N}\\p{P}]+") + } + + return nil +} + +func isSpaceRune(r rune) bool { return r == ' ' || r == '\t' || r == '\n' || r == '\r' } + +func checkMessageAction(action string) error { + switch action { + case "", driver.MessageActionResend, driver.MessageActionSuppress: + return nil + default: + return invalidParameter("1 validation error detected: Value '%s' at 'messageAction' failed to satisfy constraint: "+ + "Member must satisfy enum value set: [RESEND, SUPPRESS]", action) + } +} + +func isEmailFormat(s string) bool { + at := strings.IndexByte(s, '@') + + return at > 0 && at < len(s)-1 && !strings.Contains(s[at+1:], "@") +} + +func isPhoneFormat(s string) bool { + if len(s) < 2 || s[0] != '+' { + return false + } + + for _, r := range s[1:] { + if r < '0' || r > '9' { + return false + } + } + + return true +} diff --git a/providers/aws/cognito/users_test.go b/providers/aws/cognito/users_test.go new file mode 100644 index 000000000..02d14a6c7 --- /dev/null +++ b/providers/aws/cognito/users_test.go @@ -0,0 +1,557 @@ +package cognito + +import ( + "context" + "errors" + "fmt" + "testing" + + cerrors "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +func assertException(t *testing.T, err error, exception, wantMsg string) { + t.Helper() + + var apiErr *driver.APIError + if !errors.As(err, &apiErr) || apiErr.Exception != exception { + t.Fatalf("expected %s, got %v", exception, err) + } + + if wantMsg != "" && cerrors.Message(err) != wantMsg { + t.Fatalf("message = %q, want %q", cerrors.Message(err), wantMsg) + } +} + +func mustCreateUser(t *testing.T, m *Mock, poolID, username string, attrs ...driver.Attribute) *driver.User { + t.Helper() + + u, err := m.AdminCreateUser(context.Background(), driver.AdminCreateUserInput{ + UserPoolID: poolID, + Username: username, + UserAttributes: attrs, + MessageAction: driver.MessageActionSuppress, + }) + requireNoError(t, err, "AdminCreateUser "+username) + + return u +} + +func TestAdminCreateUserDefaults(t *testing.T) { + m := newMock(t) + pool := mustCreatePool(t, m, "users") + + u := mustCreateUser(t, m, pool.ID, "alice", driver.Attribute{Name: "email", Value: "a@example.com"}) + + if u.Username != "alice" || !u.Enabled || u.UserStatus != driver.UserStatusForceChangePassword { + t.Fatalf("new user = %+v", u) + } + + if len(u.Attributes) != 2 || u.Attributes[0].Name != "sub" || len(u.Attributes[0].Value) != 36 { + t.Fatalf("attributes = %+v, want sub then email", u.Attributes) + } + + if u.UserCreateDate.IsZero() || !u.UserCreateDate.Equal(u.UserLastModifiedDate) { + t.Fatalf("dates = %v / %v", u.UserCreateDate, u.UserLastModifiedDate) + } + + got, err := m.DescribeUserPool(context.Background(), pool.ID) + requireNoError(t, err, "DescribeUserPool") + + if got.EstimatedNumberOfUsers != 1 { + t.Fatalf("EstimatedNumberOfUsers = %d, want 1", got.EstimatedNumberOfUsers) + } +} + +func TestAdminCreateUserErrors(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "errs") + mustCreateUser(t, m, pool.ID, "bob") + + cases := []struct { + name string + in driver.AdminCreateUserInput + exception string + msg string + }{ + {"duplicate", driver.AdminCreateUserInput{UserPoolID: pool.ID, Username: "bob"}, + driver.ExUsernameExists, "User account already exists"}, + {"missing pool", driver.AdminCreateUserInput{UserPoolID: "us-east-1_nope00000", Username: "x"}, + driver.ExResourceNotFound, "User pool us-east-1_nope00000 does not exist."}, + {"unknown attribute", driver.AdminCreateUserInput{UserPoolID: pool.ID, Username: "c", + UserAttributes: []driver.Attribute{{Name: "custom:nope", Value: "1"}}}, + driver.ExInvalidParameter, "Attributes did not conform to the schema: custom:nope: Attribute does not exist in the schema."}, + {"short temp password", driver.AdminCreateUserInput{UserPoolID: pool.ID, Username: "d", TemporaryPassword: "Ab1!"}, + driver.ExInvalidPassword, "Password did not conform with policy: Password not long enough"}, + {"no symbol", driver.AdminCreateUserInput{UserPoolID: pool.ID, Username: "e", TemporaryPassword: "Abcdefg12"}, + driver.ExInvalidPassword, "Password did not conform with policy: Password must have symbol characters"}, + {"resend unknown", driver.AdminCreateUserInput{UserPoolID: pool.ID, Username: "ghost", + MessageAction: driver.MessageActionResend}, driver.ExUserNotFound, "User does not exist."}, + {"bad message action", driver.AdminCreateUserInput{UserPoolID: pool.ID, Username: "f", MessageAction: "LOUD"}, + driver.ExInvalidParameter, ""}, + {"empty username", driver.AdminCreateUserInput{UserPoolID: pool.ID}, driver.ExInvalidParameter, ""}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + _, err := m.AdminCreateUser(ctx, tc.in) + assertException(t, err, tc.exception, tc.msg) + }) + } +} + +func TestAdminCreateUserResend(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "resend") + mustCreateUser(t, m, pool.ID, "carol") + + before, _ := m.users.Get(userKey(pool.ID, "carol")) + + u, err := m.AdminCreateUser(ctx, driver.AdminCreateUserInput{ + UserPoolID: pool.ID, Username: "carol", MessageAction: driver.MessageActionResend, + }) + requireNoError(t, err, "RESEND") + + after, _ := m.users.Get(userKey(pool.ID, "carol")) + if u.UserStatus != driver.UserStatusForceChangePassword || after.PasswordHash == before.PasswordHash { + t.Fatal("RESEND should keep FORCE_CHANGE_PASSWORD and issue a new temporary password") + } + + requireNoError(t, m.AdminSetUserPassword(ctx, pool.ID, "carol", "Perm4nent!pw", true), "AdminSetUserPassword") + + _, err = m.AdminCreateUser(ctx, driver.AdminCreateUserInput{ + UserPoolID: pool.ID, Username: "carol", MessageAction: driver.MessageActionResend, + }) + assertException(t, err, driver.ExUnsupportedUserState, "") +} + +func TestGeneratedTemporaryPasswordMeetsPolicy(t *testing.T) { + pp := driver.PasswordPolicy{ + MinimumLength: 30, RequireUppercase: true, RequireLowercase: true, RequireNumbers: true, RequireSymbols: true, + } + + for range 50 { + if err := checkPassword(generatePassword(pp), pp); err != nil { + t.Fatalf("generated password failed policy: %v", err) + } + } +} + +func TestPasswordIsStoredHashed(t *testing.T) { + m := newMock(t) + pool := mustCreatePool(t, m, "hash") + + _, err := m.AdminCreateUser(context.Background(), driver.AdminCreateUserInput{ + UserPoolID: pool.ID, Username: "dan", TemporaryPassword: "Temp0rary!pw", + }) + requireNoError(t, err, "AdminCreateUser") + + rec, _ := m.users.Get(userKey(pool.ID, "dan")) + if rec.PasswordHash == "" || rec.PasswordHash == "Temp0rary!pw" || rec.PasswordSalt == "" { + t.Fatalf("password not hashed: %+v", rec) + } + + if digest(rec.PasswordSalt, "Temp0rary!pw") != rec.PasswordHash { + t.Fatal("stored digest does not match the password") + } +} + +func TestEmailUsernamePool(t *testing.T) { + m := newMock(t) + ctx := context.Background() + + pool, err := m.CreateUserPool(ctx, driver.CreateUserPoolInput{Name: "email-login", UsernameAttributes: []string{"email"}}) + requireNoError(t, err, "CreateUserPool") + + _, err = m.AdminCreateUser(ctx, driver.AdminCreateUserInput{UserPoolID: pool.ID, Username: "plainname"}) + assertException(t, err, driver.ExInvalidParameter, "Username should be an email.") + + u := mustCreateUser(t, m, pool.ID, "erin@example.com") + if u.Username != attrValue(u.Attributes, "sub") || attrValue(u.Attributes, "email") != "erin@example.com" { + t.Fatalf("email-username user = %+v, want username=sub and email set", u) + } + + got, err := m.AdminGetUser(ctx, pool.ID, "erin@example.com") + requireNoError(t, err, "AdminGetUser by email") + + if got.Username != u.Username { + t.Fatalf("lookup by email found %q, want %q", got.Username, u.Username) + } + + _, err = m.AdminCreateUser(ctx, driver.AdminCreateUserInput{UserPoolID: pool.ID, Username: "erin@example.com"}) + assertException(t, err, driver.ExUsernameExists, "An account with the given email already exists.") +} + +func TestEmailAliasPool(t *testing.T) { + m := newMock(t) + ctx := context.Background() + + pool, err := m.CreateUserPool(ctx, driver.CreateUserPoolInput{Name: "alias", AliasAttributes: []string{"email"}}) + requireNoError(t, err, "CreateUserPool") + + _, err = m.AdminCreateUser(ctx, driver.AdminCreateUserInput{UserPoolID: pool.ID, Username: "x@example.com"}) + assertException(t, err, driver.ExInvalidParameter, + "Username cannot be of email format, since user pool is configured for email alias.") + + mustCreateUser(t, m, pool.ID, "frank", driver.Attribute{Name: "email", Value: "f@example.com"}) + + _, err = m.AdminGetUser(ctx, pool.ID, "f@example.com") + assertException(t, err, driver.ExUserNotFound, "") + + requireNoError(t, m.AdminUpdateUserAttributes(ctx, pool.ID, "frank", + []driver.Attribute{{Name: "email_verified", Value: "true"}}), "verify email") + + got, err := m.AdminGetUser(ctx, pool.ID, "f@example.com") + requireNoError(t, err, "AdminGetUser by verified alias") + + if got.Username != "frank" { + t.Fatalf("alias lookup = %q", got.Username) + } +} + +func TestAdminGetUserNotFound(t *testing.T) { + m := newMock(t) + pool := mustCreatePool(t, m, "get") + + _, err := m.AdminGetUser(context.Background(), pool.ID, "nobody") + assertException(t, err, driver.ExUserNotFound, "User does not exist.") + + _, err = m.AdminGetUser(context.Background(), "us-east-1_nope00000", "nobody") + assertException(t, err, driver.ExResourceNotFound, "") +} + +func TestListUsersFilter(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "list") + + mustCreateUser(t, m, pool.ID, "amy", driver.Attribute{Name: "email", Value: "amy@corp.example"}, + driver.Attribute{Name: "given_name", Value: "Amy"}) + mustCreateUser(t, m, pool.ID, "ben", driver.Attribute{Name: "email", Value: "ben@other.example"}) + mustCreateUser(t, m, pool.ID, "abe", driver.Attribute{Name: "email", Value: "abe@corp.example"}) + requireNoError(t, m.AdminDisableUser(ctx, pool.ID, "ben"), "AdminDisableUser") + + cases := []struct { + filter string + want []string + }{ + {"", []string{"abe", "amy", "ben"}}, + {`username = "amy"`, []string{"amy"}}, + {`username ^= "a"`, []string{"abe", "amy"}}, + {`email = "ben@other.example"`, []string{"ben"}}, + {`given_name ^= "Am"`, []string{"amy"}}, + {`status = "Disabled"`, []string{"ben"}}, + {`status = "Enabled"`, []string{"abe", "amy"}}, + {`cognito:user_status = "force_change_password"`, []string{"abe", "amy", "ben"}}, + {`email = "nobody@x"`, nil}, + } + + for _, tc := range cases { + t.Run(tc.filter, func(t *testing.T) { + users, _, err := m.ListUsers(ctx, driver.ListUsersInput{UserPoolID: pool.ID, Filter: tc.filter}) + requireNoError(t, err, "ListUsers") + + var names []string + for _, u := range users { + names = append(names, u.Username) + } + + if fmt.Sprint(names) != fmt.Sprint(tc.want) { + t.Fatalf("filter %q = %v, want %v", tc.filter, names, tc.want) + } + }) + } +} + +func TestListUsersErrorsAndProjection(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "list-errs") + mustCreateUser(t, m, pool.ID, "gus", driver.Attribute{Name: "email", Value: "g@example.com"}) + + _, _, err := m.ListUsers(ctx, driver.ListUsersInput{UserPoolID: pool.ID, Filter: `custom:tier = "x"`}) + assertException(t, err, driver.ExInvalidParameter, "Invalid search attribute: custom:tier") + + _, _, err = m.ListUsers(ctx, driver.ListUsersInput{UserPoolID: pool.ID, Filter: `email equals x`}) + assertException(t, err, driver.ExInvalidParameter, "Error while parsing filter.") + + _, _, err = m.ListUsers(ctx, driver.ListUsersInput{UserPoolID: pool.ID, Limit: 61}) + assertException(t, err, driver.ExInvalidParameter, "") + + _, _, err = m.ListUsers(ctx, driver.ListUsersInput{UserPoolID: "us-east-1_nope00000"}) + assertException(t, err, driver.ExResourceNotFound, "") + + users, _, err := m.ListUsers(ctx, driver.ListUsersInput{UserPoolID: pool.ID, AttributesToGet: []string{"email"}}) + requireNoError(t, err, "ListUsers AttributesToGet") + + if len(users[0].Attributes) != 1 || users[0].Attributes[0].Name != "email" { + t.Fatalf("AttributesToGet=[email] gave %+v", users[0].Attributes) + } + + users, _, err = m.ListUsers(ctx, driver.ListUsersInput{UserPoolID: pool.ID, AttributesToGet: []string{}}) + requireNoError(t, err, "ListUsers AttributesToGet empty") + + if len(users[0].Attributes) != 0 { + t.Fatalf("empty AttributesToGet gave %+v", users[0].Attributes) + } + + // The stored user is not changed by a projected read. + got, err := m.AdminGetUser(ctx, pool.ID, "gus") + requireNoError(t, err, "AdminGetUser") + + if len(got.Attributes) != 2 { + t.Fatalf("stored attributes changed by projection: %+v", got.Attributes) + } +} + +func TestListUsersPagination(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "paged") + + for i := range 65 { + mustCreateUser(t, m, pool.ID, fmt.Sprintf("user%03d", i)) + } + + seen := map[string]bool{} + token := "" + pages := 0 + + for { + users, next, err := m.ListUsers(ctx, driver.ListUsersInput{UserPoolID: pool.ID, PaginationToken: token}) + requireNoError(t, err, "ListUsers") + + pages++ + + for _, u := range users { + seen[u.Username] = true + } + + if next == "" { + break + } + + token = next + } + + if pages != 2 || len(seen) != 65 { + t.Fatalf("pages=%d users=%d, want 2 pages covering 65 users", pages, len(seen)) + } +} + +func TestListUserPoolsMaxResultsCeiling(t *testing.T) { + m := newMock(t) + ctx := context.Background() + + for i := range 61 { + mustCreatePool(t, m, fmt.Sprintf("pool-%02d", i)) + } + + _, _, err := m.ListUserPools(ctx, driver.Pagination{MaxResults: 61}) + assertException(t, err, driver.ExInvalidParameter, "") + + first, next, err := m.ListUserPools(ctx, driver.Pagination{MaxResults: 60}) + requireNoError(t, err, "ListUserPools page 1") + + second, last, err := m.ListUserPools(ctx, driver.Pagination{MaxResults: 60, NextToken: next}) + requireNoError(t, err, "ListUserPools page 2") + + if len(first) != 60 || next == "" || len(second) != 1 || last != "" { + t.Fatalf("pages = %d (next %q) + %d (next %q)", len(first), next, len(second), last) + } +} + +func TestAdminUpdateUserAttributes(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "update") + mustCreateUser(t, m, pool.ID, "hal", + driver.Attribute{Name: "email", Value: "old@example.com"}, driver.Attribute{Name: "email_verified", Value: "true"}) + + requireNoError(t, m.AdminUpdateUserAttributes(ctx, pool.ID, "hal", []driver.Attribute{ + {Name: "email", Value: "new@example.com"}, {Name: "name", Value: "Hal"}, + }), "AdminUpdateUserAttributes") + + u, err := m.AdminGetUser(ctx, pool.ID, "hal") + requireNoError(t, err, "AdminGetUser") + + if attrValue(u.Attributes, "email") != "new@example.com" || attrValue(u.Attributes, "name") != "Hal" { + t.Fatalf("attributes not updated: %+v", u.Attributes) + } + + if attrValue(u.Attributes, "email_verified") != "false" { + t.Fatalf("changed email should be unverified: %+v", u.Attributes) + } + + err = m.AdminUpdateUserAttributes(ctx, pool.ID, "hal", []driver.Attribute{{Name: "sub", Value: "x"}}) + assertException(t, err, driver.ExInvalidParameter, "") + + err = m.AdminUpdateUserAttributes(ctx, pool.ID, "ghost", []driver.Attribute{{Name: "name", Value: "x"}}) + assertException(t, err, driver.ExUserNotFound, "") +} + +func TestAdminDeleteUserAttributes(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "delattr") + mustCreateUser(t, m, pool.ID, "ivy", driver.Attribute{Name: "name", Value: "Ivy"}) + + requireNoError(t, m.AdminDeleteUserAttributes(ctx, pool.ID, "ivy", []string{"name"}), "AdminDeleteUserAttributes") + + u, err := m.AdminGetUser(ctx, pool.ID, "ivy") + requireNoError(t, err, "AdminGetUser") + + if _, ok := lastValue(u.Attributes, "name"); ok { + t.Fatalf("name not deleted: %+v", u.Attributes) + } + + assertException(t, m.AdminDeleteUserAttributes(ctx, pool.ID, "ivy", []string{"sub"}), driver.ExInvalidParameter, "") + assertException(t, m.AdminDeleteUserAttributes(ctx, pool.ID, "ivy", []string{"custom:x"}), driver.ExInvalidParameter, + "Attributes did not conform to the schema: custom:x: Attribute does not exist in the schema.") +} + +func TestUserStateTransitions(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "states") + mustCreateUser(t, m, pool.ID, "jay") + + err := m.AdminResetUserPassword(ctx, pool.ID, "jay") + assertException(t, err, driver.ExNotAuthorized, "User password cannot be reset in the current state.") + + assertException(t, m.AdminSetUserPassword(ctx, pool.ID, "jay", "weak", true), driver.ExInvalidPassword, "") + + requireNoError(t, m.AdminSetUserPassword(ctx, pool.ID, "jay", "Str0ng!Password", true), "AdminSetUserPassword") + requireStatus(t, m, pool.ID, "jay", driver.UserStatusConfirmed) + + requireNoError(t, m.AdminResetUserPassword(ctx, pool.ID, "jay"), "AdminResetUserPassword") + requireStatus(t, m, pool.ID, "jay", driver.UserStatusResetRequired) + + requireNoError(t, m.AdminSetUserPassword(ctx, pool.ID, "jay", "Str0ng!Password", false), "temporary set") + requireStatus(t, m, pool.ID, "jay", driver.UserStatusForceChangePassword) + + requireNoError(t, m.AdminDisableUser(ctx, pool.ID, "jay"), "AdminDisableUser") + + u, _ := m.AdminGetUser(ctx, pool.ID, "jay") + if u.Enabled { + t.Fatal("user still enabled after AdminDisableUser") + } + + requireNoError(t, m.AdminEnableUser(ctx, pool.ID, "jay"), "AdminEnableUser") + + u, _ = m.AdminGetUser(ctx, pool.ID, "jay") + if !u.Enabled { + t.Fatal("user still disabled after AdminEnableUser") + } + + requireNoError(t, m.AdminDeleteUser(ctx, pool.ID, "jay"), "AdminDeleteUser") + assertException(t, m.AdminDeleteUser(ctx, pool.ID, "jay"), driver.ExUserNotFound, "User does not exist.") + assertException(t, m.AdminEnableUser(ctx, pool.ID, "jay"), driver.ExUserNotFound, "") +} + +func requireStatus(t *testing.T, m *Mock, poolID, username, want string) { + t.Helper() + + u, err := m.AdminGetUser(context.Background(), poolID, username) + requireNoError(t, err, "AdminGetUser") + + if u.UserStatus != want { + t.Fatalf("status = %s, want %s", u.UserStatus, want) + } +} + +func TestAddCustomAttributes(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "custom") + + requireNoError(t, m.AddCustomAttributes(ctx, pool.ID, []driver.SchemaAttribute{ + {Name: "tier", AttributeDataType: driver.AttributeTypeString, Mutable: false}, + }), "AddCustomAttributes") + + got, err := m.DescribeUserPool(ctx, pool.ID) + requireNoError(t, err, "DescribeUserPool") + + a, ok := schemaAttribute(got, "custom:tier") + if !ok || a.Mutable { + t.Fatalf("custom:tier not in schema as immutable: %+v", got.SchemaAttributes) + } + + err = m.AddCustomAttributes(ctx, pool.ID, []driver.SchemaAttribute{{Name: "tier", AttributeDataType: driver.AttributeTypeString}}) + assertException(t, err, driver.ExInvalidParameter, "Existing attribute already has name custom:tier.") + + // The new attribute can be set on create, but is immutable afterwards. + mustCreateUser(t, m, pool.ID, "kim", driver.Attribute{Name: "custom:tier", Value: "gold"}) + + err = m.AdminUpdateUserAttributes(ctx, pool.ID, "kim", []driver.Attribute{{Name: "custom:tier", Value: "silver"}}) + assertException(t, err, driver.ExInvalidParameter, "") + + assertException(t, m.AddCustomAttributes(ctx, "us-east-1_nope00000", + []driver.SchemaAttribute{{Name: "x"}}), driver.ExResourceNotFound, "") +} + +func TestAddCustomAttributesLimit(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "limit") + + attrs := make([]driver.SchemaAttribute, 0, maxCustomAttributes) + for i := range maxCustomAttributes { + attrs = append(attrs, driver.SchemaAttribute{Name: fmt.Sprintf("a%02d", i)}) + } + + requireNoError(t, m.AddCustomAttributes(ctx, pool.ID, attrs), "AddCustomAttributes 50") + + err := m.AddCustomAttributes(ctx, pool.ID, []driver.SchemaAttribute{{Name: "one-more"}}) + assertException(t, err, driver.ExInvalidParameter, "") +} + +func TestDeletePoolCascadesUsers(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "cascade-users") + other := mustCreatePool(t, m, "other") + + mustCreateUser(t, m, pool.ID, "leo") + mustCreateUser(t, m, other.ID, "leo") + + requireNoError(t, m.DeleteUserPool(ctx, pool.ID), "DeleteUserPool") + + if m.users.Has(userKey(pool.ID, "leo")) { + t.Fatal("user survived pool delete") + } + + if _, err := m.AdminGetUser(ctx, other.ID, "leo"); err != nil { + t.Fatalf("user in another pool removed: %v", err) + } +} + +func TestSnapshotRoundTripUsers(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "snap") + u := mustCreateUser(t, m, pool.ID, "mia", driver.Attribute{Name: "email", Value: "m@example.com"}) + requireNoError(t, m.AdminSetUserPassword(ctx, pool.ID, "mia", "Str0ng!Password", true), "AdminSetUserPassword") + + data, err := m.Snapshot(ctx, false) + requireNoError(t, err, "Snapshot") + + restored := newMock(t) + requireNoError(t, restored.Restore(ctx, data), "Restore") + + got, err := restored.AdminGetUser(ctx, pool.ID, "mia") + requireNoError(t, err, "AdminGetUser after restore") + + if got.UserStatus != driver.UserStatusConfirmed || attrValue(got.Attributes, "sub") != attrValue(u.Attributes, "sub") { + t.Fatalf("restored user = %+v", got) + } + + before, _ := m.users.Get(userKey(pool.ID, "mia")) + after, _ := restored.users.Get(userKey(pool.ID, "mia")) + + if before.PasswordHash != after.PasswordHash || before.PasswordSalt != after.PasswordSalt { + t.Fatal("password digest not persisted") + } +} diff --git a/server/aws/cognito/handler.go b/server/aws/cognito/handler.go index 4a881410f..eb9b11a75 100644 --- a/server/aws/cognito/handler.go +++ b/server/aws/cognito/handler.go @@ -2,8 +2,8 @@ // protocol as a server.Handler. Point the real // aws-sdk-go-v2/service/cognitoidentityprovider client (or the `aws cognito-idp` // CLI, or the Terraform AWS provider) at a Server registered with this handler -// and the user-pool, app-client, and hosted-UI-domain control-plane operations -// run against an in-memory Cognito driver. +// and the user-pool, app-client, hosted-UI-domain and admin user-management +// operations run against an in-memory Cognito driver. // // Cognito uses the AWS JSON 1.1 wire shape (POST + JSON body dispatched on the // X-Amz-Target header, prefix "AWSCognitoIdentityProviderService."). @@ -50,6 +50,18 @@ func New(d cognitodriver.Cognito) *Handler { "TagResource": h.tagResource, "UntagResource": h.untagResource, "ListTagsForResource": h.listTagsForResource, + + "AddCustomAttributes": h.addCustomAttributes, + "AdminCreateUser": h.adminCreateUser, + "AdminGetUser": h.adminGetUser, + "ListUsers": h.listUsers, + "AdminDeleteUser": h.adminDeleteUser, + "AdminUpdateUserAttributes": h.adminUpdateUserAttributes, + "AdminDeleteUserAttributes": h.adminDeleteUserAttributes, + "AdminSetUserPassword": h.adminSetUserPassword, + "AdminEnableUser": h.adminEnableUser, + "AdminDisableUser": h.adminDisableUser, + "AdminResetUserPassword": h.adminResetUserPassword, } return h diff --git a/server/aws/cognito/user_ops.go b/server/aws/cognito/user_ops.go new file mode 100644 index 000000000..970c5fb4d --- /dev/null +++ b/server/aws/cognito/user_ops.go @@ -0,0 +1,250 @@ +package cognito + +import ( + "context" + "net/http" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +type attributeJSON struct { + Name string `json:"Name"` + Value string `json:"Value"` +} + +// userTypeJSON is the UserType shape AdminCreateUser and ListUsers return. +type userTypeJSON struct { + Username string `json:"Username"` + Attributes []attributeJSON `json:"Attributes"` + UserCreateDate *float64 `json:"UserCreateDate,omitempty"` + UserLastModifiedDate *float64 `json:"UserLastModifiedDate,omitempty"` + Enabled bool `json:"Enabled"` + UserStatus string `json:"UserStatus,omitempty"` +} + +// adminGetUserResponse is AdminGetUser's flat output; its attribute list is +// named UserAttributes rather than Attributes. +type adminGetUserResponse struct { + Username string `json:"Username"` + UserAttributes []attributeJSON `json:"UserAttributes"` + UserCreateDate *float64 `json:"UserCreateDate,omitempty"` + UserLastModifiedDate *float64 `json:"UserLastModifiedDate,omitempty"` + Enabled bool `json:"Enabled"` + UserStatus string `json:"UserStatus,omitempty"` +} + +func attributesToWire(in []driver.Attribute) []attributeJSON { + out := make([]attributeJSON, len(in)) + for i, a := range in { + out[i] = attributeJSON(a) + } + + return out +} + +func attributesFromWire(in []attributeJSON) []driver.Attribute { + if in == nil { + return nil + } + + out := make([]driver.Attribute, len(in)) + for i, a := range in { + out[i] = driver.Attribute(a) + } + + return out +} + +func userToWire(u *driver.User) userTypeJSON { + return userTypeJSON{ + Username: u.Username, + Attributes: attributesToWire(u.Attributes), + UserCreateDate: epochOrNil(u.UserCreateDate), + UserLastModifiedDate: epochOrNil(u.UserLastModifiedDate), + Enabled: u.Enabled, + UserStatus: u.UserStatus, + } +} + +type adminCreateUserRequest struct { + UserPoolID string `json:"UserPoolId"` + Username string `json:"Username"` + UserAttributes []attributeJSON `json:"UserAttributes"` + TemporaryPassword string `json:"TemporaryPassword"` + MessageAction string `json:"MessageAction"` + DesiredDeliveryMediums []string `json:"DesiredDeliveryMediums"` +} + +type adminCreateUserResponse struct { + User userTypeJSON `json:"User"` +} + +func (h *Handler) adminCreateUser(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *adminCreateUserRequest) (any, error) { + u, err := h.cognito.AdminCreateUser(ctx, driver.AdminCreateUserInput{ + UserPoolID: req.UserPoolID, + Username: req.Username, + UserAttributes: attributesFromWire(req.UserAttributes), + TemporaryPassword: req.TemporaryPassword, + MessageAction: req.MessageAction, + DesiredDeliveryMediums: req.DesiredDeliveryMediums, + }) + if err != nil { + return nil, err + } + + return adminCreateUserResponse{User: userToWire(u)}, nil + }) +} + +// userRef is the {UserPoolId, Username} pair most admin user operations take. +type userRef struct { + UserPoolID string `json:"UserPoolId"` + Username string `json:"Username"` +} + +func (h *Handler) adminGetUser(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *userRef) (any, error) { + u, err := h.cognito.AdminGetUser(ctx, req.UserPoolID, req.Username) + if err != nil { + return nil, err + } + + return adminGetUserResponse{ + Username: u.Username, + UserAttributes: attributesToWire(u.Attributes), + UserCreateDate: epochOrNil(u.UserCreateDate), + UserLastModifiedDate: epochOrNil(u.UserLastModifiedDate), + Enabled: u.Enabled, + UserStatus: u.UserStatus, + }, nil + }) +} + +type listUsersRequest struct { + UserPoolID string `json:"UserPoolId"` + AttributesToGet []string `json:"AttributesToGet"` + Limit int32 `json:"Limit"` + PaginationToken string `json:"PaginationToken"` + Filter string `json:"Filter"` +} + +type listUsersResponse struct { + Users []userTypeJSON `json:"Users"` + PaginationToken string `json:"PaginationToken,omitempty"` +} + +func (h *Handler) listUsers(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *listUsersRequest) (any, error) { + users, next, err := h.cognito.ListUsers(ctx, driver.ListUsersInput{ + UserPoolID: req.UserPoolID, + Filter: req.Filter, + AttributesToGet: req.AttributesToGet, + Limit: req.Limit, + PaginationToken: req.PaginationToken, + }) + if err != nil { + return nil, err + } + + out := make([]userTypeJSON, len(users)) + for i := range users { + out[i] = userToWire(&users[i]) + } + + return listUsersResponse{Users: out, PaginationToken: next}, nil + }) +} + +func (h *Handler) adminDeleteUser(w http.ResponseWriter, r *http.Request) { + h.simpleUserOp(w, r, h.cognito.AdminDeleteUser) +} + +func (h *Handler) adminEnableUser(w http.ResponseWriter, r *http.Request) { + h.simpleUserOp(w, r, h.cognito.AdminEnableUser) +} + +func (h *Handler) adminDisableUser(w http.ResponseWriter, r *http.Request) { + h.simpleUserOp(w, r, h.cognito.AdminDisableUser) +} + +func (h *Handler) adminResetUserPassword(w http.ResponseWriter, r *http.Request) { + h.simpleUserOp(w, r, h.cognito.AdminResetUserPassword) +} + +// simpleUserOp serves an operation that takes only {UserPoolId, Username} and +// returns an empty body. +func (h *Handler) simpleUserOp(w http.ResponseWriter, r *http.Request, call func(context.Context, string, string) error) { + dispatch(h, w, r, func(_ *Handler, ctx context.Context, req *userRef) (any, error) { + if err := call(ctx, req.UserPoolID, req.Username); err != nil { + return nil, err + } + + return struct{}{}, nil + }) +} + +type adminUpdateUserAttributesRequest struct { + UserPoolID string `json:"UserPoolId"` + Username string `json:"Username"` + UserAttributes []attributeJSON `json:"UserAttributes"` +} + +func (h *Handler) adminUpdateUserAttributes(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *adminUpdateUserAttributesRequest) (any, error) { + err := h.cognito.AdminUpdateUserAttributes(ctx, req.UserPoolID, req.Username, attributesFromWire(req.UserAttributes)) + if err != nil { + return nil, err + } + + return struct{}{}, nil + }) +} + +type adminDeleteUserAttributesRequest struct { + UserPoolID string `json:"UserPoolId"` + Username string `json:"Username"` + UserAttributeNames []string `json:"UserAttributeNames"` +} + +func (h *Handler) adminDeleteUserAttributes(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *adminDeleteUserAttributesRequest) (any, error) { + if err := h.cognito.AdminDeleteUserAttributes(ctx, req.UserPoolID, req.Username, req.UserAttributeNames); err != nil { + return nil, err + } + + return struct{}{}, nil + }) +} + +type adminSetUserPasswordRequest struct { + UserPoolID string `json:"UserPoolId"` + Username string `json:"Username"` + Password string `json:"Password"` + Permanent bool `json:"Permanent"` +} + +func (h *Handler) adminSetUserPassword(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *adminSetUserPasswordRequest) (any, error) { + if err := h.cognito.AdminSetUserPassword(ctx, req.UserPoolID, req.Username, req.Password, req.Permanent); err != nil { + return nil, err + } + + return struct{}{}, nil + }) +} + +type addCustomAttributesRequest struct { + UserPoolID string `json:"UserPoolId"` + CustomAttributes []schemaAttributeJSON `json:"CustomAttributes"` +} + +func (h *Handler) addCustomAttributes(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *addCustomAttributesRequest) (any, error) { + if err := h.cognito.AddCustomAttributes(ctx, req.UserPoolID, schemaAttributesFromWire(req.CustomAttributes)); err != nil { + return nil, err + } + + return struct{}{}, nil + }) +} diff --git a/server/aws/cognito/users_sdk_test.go b/server/aws/cognito/users_sdk_test.go new file mode 100644 index 000000000..62dec6961 --- /dev/null +++ b/server/aws/cognito/users_sdk_test.go @@ -0,0 +1,367 @@ +package cognito_test + +import ( + "context" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" + awsconfig "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" + awsct "github.com/aws/aws-sdk-go-v2/service/cloudtrail" + cip "github.com/aws/aws-sdk-go-v2/service/cognitoidentityprovider" + ciptypes "github.com/aws/aws-sdk-go-v2/service/cognitoidentityprovider/types" + "github.com/aws/smithy-go" + + "github.com/stackshy/cloudemu/v2" + awsserver "github.com/stackshy/cloudemu/v2/server/aws" +) + +func requireErrorCode(t *testing.T, err error, code, msg string) { + t.Helper() + + var apiErr smithy.APIError + if !errors.As(err, &apiErr) { + t.Fatalf("expected %s, got %v", code, err) + } + + if apiErr.ErrorCode() != code { + t.Fatalf("error code = %q, want %q (%s)", apiErr.ErrorCode(), code, apiErr.ErrorMessage()) + } + + if msg != "" && apiErr.ErrorMessage() != msg { + t.Fatalf("message = %q, want %q", apiErr.ErrorMessage(), msg) + } +} + +func attr(attrs []ciptypes.AttributeType, name string) string { + for _, a := range attrs { + if aws.ToString(a.Name) == name { + return aws.ToString(a.Value) + } + } + + return "" +} + +func TestSDKAdminUserLifecycle(t *testing.T) { + ctx := context.Background() + c := newCognitoClient(t) + pool := createPool(t, c, "users") + + created, err := c.AdminCreateUser(ctx, &cip.AdminCreateUserInput{ + UserPoolId: aws.String(pool), + Username: aws.String("alice"), + TemporaryPassword: aws.String("Temp0rary!pw"), + MessageAction: ciptypes.MessageActionTypeSuppress, + UserAttributes: []ciptypes.AttributeType{ + {Name: aws.String("email"), Value: aws.String("alice@example.com")}, + }, + }) + if err != nil { + t.Fatalf("AdminCreateUser: %v", err) + } + + u := created.User + if aws.ToString(u.Username) != "alice" || !u.Enabled || u.UserStatus != ciptypes.UserStatusTypeForceChangePassword { + t.Fatalf("created user = %+v", u) + } + + if len(attr(u.Attributes, "sub")) != 36 || u.UserCreateDate == nil { + t.Fatalf("created user missing sub or dates: %+v", u) + } + + _, err = c.AdminCreateUser(ctx, &cip.AdminCreateUserInput{UserPoolId: aws.String(pool), Username: aws.String("alice")}) + requireErrorCode(t, err, "UsernameExistsException", "User account already exists") + + _, err = c.AdminSetUserPassword(ctx, &cip.AdminSetUserPasswordInput{ + UserPoolId: aws.String(pool), Username: aws.String("alice"), Password: aws.String("short"), Permanent: true, + }) + requireErrorCode(t, err, "InvalidPasswordException", "Password did not conform with policy: Password not long enough") + + _, err = c.AdminResetUserPassword(ctx, &cip.AdminResetUserPasswordInput{UserPoolId: aws.String(pool), Username: aws.String("alice")}) + requireErrorCode(t, err, "NotAuthorizedException", "User password cannot be reset in the current state.") + + if _, err = c.AdminSetUserPassword(ctx, &cip.AdminSetUserPasswordInput{ + UserPoolId: aws.String(pool), Username: aws.String("alice"), Password: aws.String("Str0ng!Password"), Permanent: true, + }); err != nil { + t.Fatalf("AdminSetUserPassword: %v", err) + } + + if _, err = c.AdminUpdateUserAttributes(ctx, &cip.AdminUpdateUserAttributesInput{ + UserPoolId: aws.String(pool), Username: aws.String("alice"), + UserAttributes: []ciptypes.AttributeType{{Name: aws.String("name"), Value: aws.String("Alice")}}, + }); err != nil { + t.Fatalf("AdminUpdateUserAttributes: %v", err) + } + + if _, err = c.AdminDisableUser(ctx, &cip.AdminDisableUserInput{UserPoolId: aws.String(pool), Username: aws.String("alice")}); err != nil { + t.Fatalf("AdminDisableUser: %v", err) + } + + got, err := c.AdminGetUser(ctx, &cip.AdminGetUserInput{UserPoolId: aws.String(pool), Username: aws.String("alice")}) + if err != nil { + t.Fatalf("AdminGetUser: %v", err) + } + + if got.Enabled || got.UserStatus != ciptypes.UserStatusTypeConfirmed || attr(got.UserAttributes, "name") != "Alice" { + t.Fatalf("AdminGetUser = enabled %v status %s attrs %+v", got.Enabled, got.UserStatus, got.UserAttributes) + } + + if _, err = c.AdminDeleteUserAttributes(ctx, &cip.AdminDeleteUserAttributesInput{ + UserPoolId: aws.String(pool), Username: aws.String("alice"), UserAttributeNames: []string{"name"}, + }); err != nil { + t.Fatalf("AdminDeleteUserAttributes: %v", err) + } + + if _, err = c.AdminEnableUser(ctx, &cip.AdminEnableUserInput{UserPoolId: aws.String(pool), Username: aws.String("alice")}); err != nil { + t.Fatalf("AdminEnableUser: %v", err) + } + + if _, err = c.AdminResetUserPassword(ctx, &cip.AdminResetUserPasswordInput{ + UserPoolId: aws.String(pool), Username: aws.String("alice"), + }); err != nil { + t.Fatalf("AdminResetUserPassword: %v", err) + } + + listed, err := c.ListUsers(ctx, &cip.ListUsersInput{UserPoolId: aws.String(pool), Filter: aws.String(`email = "alice@example.com"`)}) + if err != nil { + t.Fatalf("ListUsers: %v", err) + } + + if len(listed.Users) != 1 || listed.Users[0].UserStatus != ciptypes.UserStatusTypeResetRequired || + attr(listed.Users[0].Attributes, "name") != "" { + t.Fatalf("ListUsers = %+v", listed.Users) + } + + desc, err := c.DescribeUserPool(ctx, &cip.DescribeUserPoolInput{UserPoolId: aws.String(pool)}) + if err != nil { + t.Fatalf("DescribeUserPool: %v", err) + } + + if desc.UserPool.EstimatedNumberOfUsers != 1 { + t.Fatalf("EstimatedNumberOfUsers = %d, want 1", desc.UserPool.EstimatedNumberOfUsers) + } + + if _, err = c.AdminDeleteUser(ctx, &cip.AdminDeleteUserInput{UserPoolId: aws.String(pool), Username: aws.String("alice")}); err != nil { + t.Fatalf("AdminDeleteUser: %v", err) + } + + _, err = c.AdminGetUser(ctx, &cip.AdminGetUserInput{UserPoolId: aws.String(pool), Username: aws.String("alice")}) + + var unf *ciptypes.UserNotFoundException + if !errors.As(err, &unf) { + t.Fatalf("expected typed UserNotFoundException, got %v", err) + } +} + +func TestSDKListUsersPaginatorAndFilterErrors(t *testing.T) { + ctx := context.Background() + c := newCognitoClient(t) + pool := createPool(t, c, "paged") + + for i := range 62 { + if _, err := c.AdminCreateUser(ctx, &cip.AdminCreateUserInput{ + UserPoolId: aws.String(pool), Username: aws.String(fmt.Sprintf("u%02d", i)), + MessageAction: ciptypes.MessageActionTypeSuppress, + }); err != nil { + t.Fatalf("AdminCreateUser: %v", err) + } + } + + pages, total := 0, 0 + + p := cip.NewListUsersPaginator(c, &cip.ListUsersInput{UserPoolId: aws.String(pool)}) + for p.HasMorePages() { + out, err := p.NextPage(ctx) + if err != nil { + t.Fatalf("ListUsers page: %v", err) + } + + pages++ + total += len(out.Users) + } + + if pages != 2 || total != 62 { + t.Fatalf("paginator saw %d pages, %d users; want 2 pages, 62 users", pages, total) + } + + _, err := c.ListUsers(ctx, &cip.ListUsersInput{UserPoolId: aws.String(pool), Filter: aws.String(`custom:x = "1"`)}) + requireErrorCode(t, err, "InvalidParameterException", "Invalid search attribute: custom:x") +} + +func TestSDKListUserPoolsPastSixty(t *testing.T) { + ctx := context.Background() + c := newCognitoClient(t) + + for i := range 61 { + createPool(t, c, fmt.Sprintf("pool-%02d", i)) + } + + total := 0 + + p := cip.NewListUserPoolsPaginator(c, &cip.ListUserPoolsInput{MaxResults: aws.Int32(60)}) + for p.HasMorePages() { + out, err := p.NextPage(ctx) + if err != nil { + t.Fatalf("ListUserPools page: %v", err) + } + + total += len(out.UserPools) + } + + if total != 61 { + t.Fatalf("paginator saw %d pools, want 61", total) + } + + _, err := c.ListUserPools(ctx, &cip.ListUserPoolsInput{MaxResults: aws.Int32(61)}) + requireErrorCode(t, err, "InvalidParameterException", "") +} + +func TestSDKDeleteUserPoolBlockedByDomain(t *testing.T) { + ctx := context.Background() + c := newCognitoClient(t) + pool := createPool(t, c, "with-domain") + + if _, err := c.CreateUserPoolDomain(ctx, &cip.CreateUserPoolDomainInput{ + UserPoolId: aws.String(pool), Domain: aws.String("my-auth"), + }); err != nil { + t.Fatalf("CreateUserPoolDomain: %v", err) + } + + _, err := c.DeleteUserPool(ctx, &cip.DeleteUserPoolInput{UserPoolId: aws.String(pool)}) + + var ipe *ciptypes.InvalidParameterException + if !errors.As(err, &ipe) || aws.ToString(ipe.Message) != + "User pool cannot be deleted. It has a domain configured that should be deleted first." { + t.Fatalf("expected InvalidParameterException for attached domain, got %v", err) + } + + if _, err = c.DeleteUserPoolDomain(ctx, &cip.DeleteUserPoolDomainInput{ + UserPoolId: aws.String(pool), Domain: aws.String("my-auth"), + }); err != nil { + t.Fatalf("DeleteUserPoolDomain: %v", err) + } + + if _, err = c.DeleteUserPool(ctx, &cip.DeleteUserPoolInput{UserPoolId: aws.String(pool)}); err != nil { + t.Fatalf("DeleteUserPool after domain removed: %v", err) + } +} + +func TestSDKAddCustomAttributes(t *testing.T) { + ctx := context.Background() + c := newCognitoClient(t) + pool := createPool(t, c, "custom") + + in := &cip.AddCustomAttributesInput{ + UserPoolId: aws.String(pool), + CustomAttributes: []ciptypes.SchemaAttributeType{{ + Name: aws.String("tier"), AttributeDataType: ciptypes.AttributeDataTypeString, Mutable: aws.Bool(true), + }}, + } + if _, err := c.AddCustomAttributes(ctx, in); err != nil { + t.Fatalf("AddCustomAttributes: %v", err) + } + + _, err := c.AddCustomAttributes(ctx, in) + requireErrorCode(t, err, "InvalidParameterException", "Existing attribute already has name custom:tier.") + + desc, err := c.DescribeUserPool(ctx, &cip.DescribeUserPoolInput{UserPoolId: aws.String(pool)}) + if err != nil { + t.Fatalf("DescribeUserPool: %v", err) + } + + found := false + + for _, a := range desc.UserPool.SchemaAttributes { + if aws.ToString(a.Name) == "custom:tier" { + found = true + } + } + + if !found { + t.Fatal("custom:tier missing from SchemaAttributes") + } +} + +// TestTagOpsReturnEmptyBody pins the exact TagResource / UntagResource output: +// real Cognito returns an empty JSON object. +func TestTagOpsReturnEmptyBody(t *testing.T) { + cloud := cloudemu.NewAWS() + ts := httptest.NewServer(awsserver.New(awsserver.Drivers{Cognito: cloud.Cognito})) + t.Cleanup(ts.Close) + + arn := "arn:aws:cognito-idp:us-east-1:123456789012:userpool/us-east-1_abcdefghi" + + for op, body := range map[string]string{ + "TagResource": `{"ResourceArn":"` + arn + `","Tags":{"env":"dev"}}`, + "UntagResource": `{"ResourceArn":"` + arn + `","TagKeys":["env"]}`, + } { + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, ts.URL, strings.NewReader(body)) + if err != nil { + t.Fatal(err) + } + + req.Header.Set("X-Amz-Target", "AWSCognitoIdentityProviderService."+op) + req.Header.Set("Content-Type", "application/x-amz-json-1.1") + + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("%s: %v", op, err) + } + + raw, _ := io.ReadAll(resp.Body) + resp.Body.Close() + + if resp.StatusCode != http.StatusOK || strings.TrimSpace(string(raw)) != "{}" { + t.Fatalf("%s = %d %q, want 200 {}", op, resp.StatusCode, raw) + } + } +} + +// TestCloudTrailRecordsCognitoCalls pins that Cognito control-plane and admin +// user calls show up in CloudTrail LookupEvents. +func TestCloudTrailRecordsCognitoCalls(t *testing.T) { + ctx := context.Background() + cloud := cloudemu.NewAWS() + ts := httptest.NewServer(awsserver.New(awsserver.DriversFrom(cloud))) + t.Cleanup(ts.Close) + + cfg, err := awsconfig.LoadDefaultConfig(ctx, + awsconfig.WithRegion("us-east-1"), + awsconfig.WithCredentialsProvider(credentials.NewStaticCredentialsProvider("test", "test", ""))) + if err != nil { + t.Fatalf("aws config: %v", err) + } + + cfg.BaseEndpoint = aws.String(ts.URL) + c := cip.NewFromConfig(cfg) + pool := createPool(t, c, "audited") + + if _, err = c.AdminCreateUser(ctx, &cip.AdminCreateUserInput{ + UserPoolId: aws.String(pool), Username: aws.String("audit-user"), MessageAction: ciptypes.MessageActionTypeSuppress, + }); err != nil { + t.Fatalf("AdminCreateUser: %v", err) + } + + out, err := awsct.NewFromConfig(cfg).LookupEvents(ctx, &awsct.LookupEventsInput{}) + if err != nil { + t.Fatalf("LookupEvents: %v", err) + } + + names := map[string]string{} + for _, e := range out.Events { + names[aws.ToString(e.EventName)] = aws.ToString(e.EventSource) + } + + for _, want := range []string{"CreateUserPool", "AdminCreateUser"} { + if src, ok := names[want]; !ok || src != "cognito-idp.amazonaws.com" { + t.Fatalf("%s not recorded from cognito-idp.amazonaws.com; events = %v", want, names) + } + } +} diff --git a/services/cognito/driver/driver.go b/services/cognito/driver/driver.go index af4293184..6a4d78da9 100644 --- a/services/cognito/driver/driver.go +++ b/services/cognito/driver/driver.go @@ -1,13 +1,9 @@ -// Package driver defines the interface and types for the AWS Cognito user-pools -// (cognito-idp) control plane. It models user pools, their app clients, and -// hosted-UI domains, plus resource tagging. +// Package driver defines the interface and types for AWS Cognito user pools +// (cognito-idp). It models user pools, their app clients, hosted-UI domains, +// resource tagging, and the users of a pool with the admin user-management +// operations. // -// This is the configuration control plane only: creating and reading the pool, -// client, and domain resources and their settings. There is no authentication -// data plane behind the emulator: sign-up, sign-in, token issuance, users, and -// groups are out of scope. A caller that only provisions Cognito resources -// (Terraform, CloudFormation, the console's create flow) behaves as it would -// against real Cognito, while token/user operations are deferred. +// Sign-up, sign-in and token issuance are not modeled yet. package driver import "context" @@ -19,6 +15,7 @@ type Cognito interface { userPoolClientAPI userPoolDomainAPI tagAPI + userAPI } // userPoolAPI covers the user-pool control plane. @@ -34,8 +31,13 @@ type userPoolAPI interface { // UpdateUserPool applies the mutable pool settings. A nil field is left // unchanged; UserPoolTags, when non-nil, replaces the pool's tag set. UpdateUserPool(ctx context.Context, in UpdateUserPoolInput) error - // DeleteUserPool removes a user pool and its clients, domains, and tags. + // DeleteUserPool removes a user pool with its users, clients and tags. Like + // real Cognito it refuses (InvalidParameterException) while deletion + // protection is ACTIVE or a hosted-UI domain is still attached. DeleteUserPool(ctx context.Context, id string) error + // AddCustomAttributes appends custom attributes to a pool's schema. Names get + // the "custom:" prefix, and an existing name is rejected. + AddCustomAttributes(ctx context.Context, userPoolID string, attrs []SchemaAttribute) error // ListUserPools returns pool descriptions in a deterministic order. ListUserPools(ctx context.Context, page Pagination) ([]UserPoolDescription, string, error) // GetUserPoolMfaConfig returns a pool's MFA configuration. The Terraform AWS @@ -80,3 +82,26 @@ type tagAPI interface { UntagResource(ctx context.Context, resourceARN string, tagKeys []string) error ListTagsForResource(ctx context.Context, resourceARN string) (map[string]string, error) } + +// userAPI covers the users of a user pool and the admin user-management +// operations. Usernames are looked up the way Cognito does: by username, or by +// email/phone number when the pool uses them as username attributes or aliases. +type userAPI interface { + // AdminCreateUser creates a user in FORCE_CHANGE_PASSWORD with a generated + // sub. MessageAction RESEND re-invites an existing FORCE_CHANGE_PASSWORD user. + AdminCreateUser(ctx context.Context, in AdminCreateUserInput) (*User, error) + AdminGetUser(ctx context.Context, userPoolID, username string) (*User, error) + // ListUsers returns users sorted by username, filtered by an optional + // `attr = "v"` or `attr ^= "v"` expression on a searchable attribute. + ListUsers(ctx context.Context, in ListUsersInput) ([]User, string, error) + AdminDeleteUser(ctx context.Context, userPoolID, username string) error + AdminUpdateUserAttributes(ctx context.Context, userPoolID, username string, attrs []Attribute) error + AdminDeleteUserAttributes(ctx context.Context, userPoolID, username string, names []string) error + // AdminSetUserPassword sets a password checked against the pool policy. A + // permanent password confirms the user; otherwise it is FORCE_CHANGE_PASSWORD. + AdminSetUserPassword(ctx context.Context, userPoolID, username, password string, permanent bool) error + AdminEnableUser(ctx context.Context, userPoolID, username string) error + AdminDisableUser(ctx context.Context, userPoolID, username string) error + // AdminResetUserPassword moves the user to RESET_REQUIRED. + AdminResetUserPassword(ctx context.Context, userPoolID, username string) error +} diff --git a/services/cognito/driver/errors.go b/services/cognito/driver/errors.go index ac1d5eed1..df6752812 100644 --- a/services/cognito/driver/errors.go +++ b/services/cognito/driver/errors.go @@ -12,6 +12,15 @@ const ( ExLimitExceeded = "LimitExceededException" ) +// Exception names for the user-management operations. +const ( + ExUserNotFound = "UserNotFoundException" + ExUsernameExists = "UsernameExistsException" + ExInvalidPassword = "InvalidPasswordException" + ExNotAuthorized = "NotAuthorizedException" + ExUnsupportedUserState = "UnsupportedUserStateException" +) + // APIError tags a canonical cloudemu error with the Cognito exception name it // concerns, so the server can emit the right __type while GetCode still resolves // the HTTP status through Unwrap. diff --git a/services/cognito/driver/types.go b/services/cognito/driver/types.go index 81f3ec2cd..ce31b0509 100644 --- a/services/cognito/driver/types.go +++ b/services/cognito/driver/types.go @@ -264,3 +264,54 @@ type CreateUserPoolDomainInput struct { Domain string UserPoolID string } + +// User status values reported by AdminGetUser and ListUsers. +const ( + UserStatusUnconfirmed = "UNCONFIRMED" + UserStatusConfirmed = "CONFIRMED" + UserStatusResetRequired = "RESET_REQUIRED" + UserStatusForceChangePassword = "FORCE_CHANGE_PASSWORD" +) + +// AdminCreateUser MessageAction values. +const ( + MessageActionResend = "RESEND" + MessageActionSuppress = "SUPPRESS" +) + +// Attribute is one name/value pair on a user. +type Attribute struct { + Name string + Value string +} + +// User is a user in a user pool as AdminGetUser and ListUsers report it. The +// password is never part of this value. +type User struct { + Username string + Attributes []Attribute + UserCreateDate time.Time + UserLastModifiedDate time.Time + Enabled bool + UserStatus string +} + +// AdminCreateUserInput is the input to AdminCreateUser. +type AdminCreateUserInput struct { + UserPoolID string + Username string + UserAttributes []Attribute + TemporaryPassword string + MessageAction string + DesiredDeliveryMediums []string +} + +// ListUsersInput is the input to ListUsers. A nil AttributesToGet returns every +// attribute; an empty, non-nil one returns none. +type ListUsersInput struct { + UserPoolID string + Filter string + AttributesToGet []string + Limit int32 + PaginationToken string +} From b35bd3211487d8257660ab4ab3ee20a1b93c1de7 Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 27 Sep 2026 22:02:07 +0530 Subject: [PATCH 2/3] fix(aws-cognito): enforce sign-in uniqueness, ForceAliasCreation, custom attribute names --- providers/aws/cognito/custom_attributes.go | 19 +- providers/aws/cognito/errors.go | 9 + providers/aws/cognito/sign_in.go | 160 ++++++++++++++ .../aws/cognito/sign_in_uniqueness_test.go | 204 ++++++++++++++++++ providers/aws/cognito/user_pools.go | 54 +++-- providers/aws/cognito/users.go | 63 ++---- server/aws/cognito/user_ops.go | 2 + server/aws/cognito/users_sdk_test.go | 50 +++++ services/cognito/driver/errors.go | 1 + services/cognito/driver/types.go | 3 + 10 files changed, 502 insertions(+), 63 deletions(-) create mode 100644 providers/aws/cognito/sign_in.go create mode 100644 providers/aws/cognito/sign_in_uniqueness_test.go diff --git a/providers/aws/cognito/custom_attributes.go b/providers/aws/cognito/custom_attributes.go index f025eebdb..8ac0e9b84 100644 --- a/providers/aws/cognito/custom_attributes.go +++ b/providers/aws/cognito/custom_attributes.go @@ -7,13 +7,18 @@ import ( "github.com/stackshy/cloudemu/v2/services/cognito/driver" ) -// maxCustomAttributes is the per-pool ceiling on custom attributes. -const maxCustomAttributes = 50 +// Custom attribute limits: at most 50 per pool, and a name (without its +// custom: prefix) of 1 to 20 characters. +const ( + maxCustomAttributes = 50 + maxCustomNameLen = 20 +) // AddCustomAttributes appends custom attributes to a pool's schema. Each name // gets the custom: prefix (dev: for developer-only). Cognito never changes or // removes an attribute once added, so a name already in the schema, or repeated -// in the request, is rejected and nothing is added. +// in the request, is rejected and nothing is added. A name may be given with or +// without its custom: prefix. func (m *Mock) AddCustomAttributes(_ context.Context, userPoolID string, attrs []driver.SchemaAttribute) error { if len(attrs) == 0 { return invalidParameter("1 validation error detected: Value null at 'customAttributes' failed to satisfy constraint: " + @@ -31,10 +36,10 @@ func (m *Mock) AddCustomAttributes(_ context.Context, userPoolID string, attrs [ pool = copyUserPool(pool) schema := pool.SchemaAttributes - for _, a := range attrs { - if a.Name == "" { - return invalidParameter("1 validation error detected: Value null at 'customAttributes.member.name' " + - "failed to satisfy constraint: Member must not be null") + for i, a := range attrs { + if n := len(bareCustomName(a.Name)); n < 1 || n > maxCustomNameLen { + return invalidParameter("1 validation error detected: Value '%s' at 'customAttributes.%d.member.name' "+ + "failed to satisfy constraint: Member must have length between 1 and %d", a.Name, i+1, maxCustomNameLen) } if a.AttributeDataType == "" { diff --git a/providers/aws/cognito/errors.go b/providers/aws/cognito/errors.go index 989e0c83b..e3c62af8c 100644 --- a/providers/aws/cognito/errors.go +++ b/providers/aws/cognito/errors.go @@ -35,6 +35,15 @@ func usernameExists(msg string) error { return &driver.APIError{Exception: driver.ExUsernameExists, Err: errors.New(errors.AlreadyExists, msg)} } +// aliasExists builds the AliasExistsException for a sign-in value (email, phone +// number or preferred username) that another user already holds. +func aliasExists(attr string) error { + return &driver.APIError{ + Exception: driver.ExAliasExists, + Err: errors.New(errors.AlreadyExists, "An account with the given "+attr+" already exists."), + } +} + // invalidPassword builds an InvalidPasswordException for a password that breaks // the pool's policy. func invalidPassword(reason string) error { diff --git a/providers/aws/cognito/sign_in.go b/providers/aws/cognito/sign_in.go new file mode 100644 index 000000000..af3571121 --- /dev/null +++ b/providers/aws/cognito/sign_in.go @@ -0,0 +1,160 @@ +package cognito + +import ( + "slices" + "strings" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// Sign-in values other than the username. A pool either signs in with email or +// phone number in place of a username (UsernameAttributes), or keeps usernames +// and lets email, phone number and preferred_username act as aliases +// (AliasAttributes). Either way a value can belong to one user only. Email and +// phone aliases only count once verified; an unverified copy is just data. + +// verifiedFlag maps a verifiable alias attribute to its verification flag. +func verifiedFlag(attr string) (string, bool) { + switch attr { + case attrEmail: + return attrEmailVerified, true + case attrPhoneNumber: + return attrPhoneNumberVerified, true + default: + return "", false + } +} + +// sameSignIn compares two sign-in values. Email addresses ignore case. +func sameSignIn(attr, a, b string) bool { + if attr == attrEmail { + return strings.EqualFold(a, b) + } + + return a == b +} + +// activeSignIns returns the attribute values a user can sign in with, keyed by +// attribute name. +func activeSignIns(pool *driver.UserPool, attrs []driver.Attribute) map[string]string { + out := map[string]string{} + + for _, attr := range pool.UsernameAttributes { + if v := attrValue(attrs, attr); v != "" { + out[attr] = v + } + } + + for _, attr := range pool.AliasAttributes { + v := attrValue(attrs, attr) + if v == "" { + continue + } + + if flag, ok := verifiedFlag(attr); ok && attrValue(attrs, flag) != attrTrue { + continue + } + + out[attr] = v + } + + return out +} + +// signInMatches reports whether name is one of the user's sign-in values. +func signInMatches(pool *driver.UserPool, u *driver.User, name string) bool { + for attr, v := range activeSignIns(pool, u.Attributes) { + if sameSignIn(attr, v, name) { + return true + } + } + + return false +} + +// signInClaim is another user holding a sign-in value the caller wants. +type signInClaim struct { + key string + attr string +} + +// claimSignIns checks that the sign-in values in attrs (the user's attributes +// after the change) are free, ignoring the user stored under selfKey. +// +// A taken username attribute fails with UsernameExistsException on create and +// AliasExistsException on update. A taken alias fails with AliasExistsException, +// unless force is set and the alias is an email or phone number: then it moves +// to this user and the previous holder is marked unverified. Nothing changes +// unless every value is free or movable. +func (m *Mock) claimSignIns(pool *driver.UserPool, selfKey string, attrs []driver.Attribute, force, create bool) error { + var moves []signInClaim + + wanted := activeSignIns(pool, attrs) + + for _, key := range m.poolUserKeys(pool.ID) { + if key == selfKey { + continue + } + + other, ok := m.users.Get(key) + if !ok { + continue + } + + for attr, v := range activeSignIns(pool, other.User.Attributes) { + want, ok := wanted[attr] + if !ok || !sameSignIn(attr, want, v) { + continue + } + + if err := claimError(pool, attr, force, create); err != nil { + return err + } + + moves = append(moves, signInClaim{key: key, attr: attr}) + } + } + + for _, mv := range moves { + m.unverify(mv.key, mv.attr) + } + + return nil +} + +// claimError returns the error for a sign-in value held by another user, or nil +// when ForceAliasCreation may move it. +func claimError(pool *driver.UserPool, attr string, force, create bool) error { + if slices.Contains(pool.UsernameAttributes, attr) { + if create { + return usernameExists("An account with the given " + attr + " already exists.") + } + + return aliasExists(attr) + } + + if _, verifiable := verifiedFlag(attr); force && verifiable { + return nil + } + + return aliasExists(attr) +} + +// unverify marks a user's email or phone number unverified, which drops it as +// a sign-in alias. +func (m *Mock) unverify(key, attr string) { + flag, ok := verifiedFlag(attr) + if !ok { + return + } + + rec, ok := m.users.Get(key) + if !ok { + return + } + + rec = copyUserRecord(rec) + rec.User.Attributes = mergeAttributes(rec.User.Attributes, []driver.Attribute{{Name: flag, Value: attrFalse}}) + rec.User.UserLastModifiedDate = m.now() + m.users.Set(key, rec) +} diff --git a/providers/aws/cognito/sign_in_uniqueness_test.go b/providers/aws/cognito/sign_in_uniqueness_test.go new file mode 100644 index 000000000..cbc14b8f9 --- /dev/null +++ b/providers/aws/cognito/sign_in_uniqueness_test.go @@ -0,0 +1,204 @@ +package cognito + +import ( + "context" + "strings" + "testing" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +func createPoolWith(t *testing.T, m *Mock, in driver.CreateUserPoolInput) *driver.UserPool { + t.Helper() + + pool, err := m.CreateUserPool(context.Background(), in) + requireNoError(t, err, "CreateUserPool") + + return pool +} + +func TestUsernameAttributeEmailStaysUnique(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := createPoolWith(t, m, driver.CreateUserPoolInput{Name: "email-login", UsernameAttributes: []string{"email"}}) + + a := mustCreateUser(t, m, pool.ID, "a@x.com") + b := mustCreateUser(t, m, pool.ID, "b@x.com") + + // Email comparison ignores case, on create and on update. + _, err := m.AdminCreateUser(ctx, driver.AdminCreateUserInput{UserPoolID: pool.ID, Username: "A@X.com"}) + assertException(t, err, driver.ExUsernameExists, "An account with the given email already exists.") + + err = m.AdminUpdateUserAttributes(ctx, pool.ID, b.Username, []driver.Attribute{{Name: "email", Value: "A@x.com"}}) + assertException(t, err, driver.ExAliasExists, "An account with the given email already exists.") + + got, err := m.AdminGetUser(ctx, pool.ID, "b@x.com") + requireNoError(t, err, "AdminGetUser b") + + if got.Username != b.Username { + t.Fatalf("b's sign-in changed after the refused update") + } + + got, err = m.AdminGetUser(ctx, pool.ID, "A@X.COM") + requireNoError(t, err, "AdminGetUser mixed case") + + if got.Username != a.Username { + t.Fatalf("mixed-case lookup found %q, want %q", got.Username, a.Username) + } + + // Moving b to a free address still works. + requireNoError(t, m.AdminUpdateUserAttributes(ctx, pool.ID, b.Username, + []driver.Attribute{{Name: "email", Value: "c@x.com"}}), "update to free email") +} + +func verifiedEmail(email string) []driver.Attribute { + return []driver.Attribute{{Name: "email", Value: email}, {Name: "email_verified", Value: "true"}} +} + +func TestEmailAliasCreateConflictAndForce(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := createPoolWith(t, m, driver.CreateUserPoolInput{Name: "alias", AliasAttributes: []string{"email"}}) + + mustCreateUser(t, m, pool.ID, "first", verifiedEmail("shared@x.com")...) + + _, err := m.AdminCreateUser(ctx, driver.AdminCreateUserInput{ + UserPoolID: pool.ID, Username: "second", UserAttributes: verifiedEmail("Shared@x.com"), + }) + assertException(t, err, driver.ExAliasExists, "An account with the given email already exists.") + + if m.users.Has(userKey(pool.ID, "second")) { + t.Fatal("refused create still stored the user") + } + + // An unverified copy of the address is not an alias, so it never conflicts. + mustCreateUser(t, m, pool.ID, "unverified", driver.Attribute{Name: "email", Value: "shared@x.com"}) + + _, err = m.AdminCreateUser(ctx, driver.AdminCreateUserInput{ + UserPoolID: pool.ID, Username: "second", UserAttributes: verifiedEmail("shared@x.com"), + ForceAliasCreation: true, MessageAction: driver.MessageActionSuppress, + }) + requireNoError(t, err, "AdminCreateUser ForceAliasCreation") + + first, err := m.AdminGetUser(ctx, pool.ID, "first") + requireNoError(t, err, "AdminGetUser first") + + if attrValue(first.Attributes, "email_verified") != "false" || attrValue(first.Attributes, "email") != "shared@x.com" { + t.Fatalf("old owner should keep the email but lose verification: %+v", first.Attributes) + } + + owner, err := m.AdminGetUser(ctx, pool.ID, "shared@x.com") + requireNoError(t, err, "AdminGetUser by alias") + + if owner.Username != "second" { + t.Fatalf("alias resolves to %q, want second", owner.Username) + } +} + +func TestEmailAliasUpdateConflict(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := createPoolWith(t, m, driver.CreateUserPoolInput{Name: "alias-upd", AliasAttributes: []string{"email", "phone_number"}}) + + mustCreateUser(t, m, pool.ID, "owner", verifiedEmail("taken@x.com")...) + mustCreateUser(t, m, pool.ID, "other", driver.Attribute{Name: "email", Value: "taken@x.com"}) + + // Flipping email_verified to true on a value another user holds verified. + err := m.AdminUpdateUserAttributes(ctx, pool.ID, "other", []driver.Attribute{{Name: "email_verified", Value: "true"}}) + assertException(t, err, driver.ExAliasExists, "An account with the given email already exists.") + + // Setting a verified email that is already someone's alias. + mustCreateUser(t, m, pool.ID, "third") + + err = m.AdminUpdateUserAttributes(ctx, pool.ID, "third", verifiedEmail("TAKEN@x.com")) + assertException(t, err, driver.ExAliasExists, "") + + // Phone numbers follow the same rule. + mustCreateUser(t, m, pool.ID, "caller", driver.Attribute{Name: "phone_number", Value: "+15550100"}, + driver.Attribute{Name: "phone_number_verified", Value: "true"}) + + err = m.AdminUpdateUserAttributes(ctx, pool.ID, "third", []driver.Attribute{ + {Name: "phone_number", Value: "+15550100"}, {Name: "phone_number_verified", Value: "true"}, + }) + assertException(t, err, driver.ExAliasExists, "An account with the given phone_number already exists.") + + // An unverified value is fine. + requireNoError(t, m.AdminUpdateUserAttributes(ctx, pool.ID, "third", + []driver.Attribute{{Name: "email", Value: "taken@x.com"}}), "unverified duplicate") +} + +func TestPreferredUsernameAliasUnique(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := createPoolWith(t, m, driver.CreateUserPoolInput{Name: "pref", AliasAttributes: []string{"preferred_username"}}) + + mustCreateUser(t, m, pool.ID, "u1", driver.Attribute{Name: "preferred_username", Value: "neo"}) + mustCreateUser(t, m, pool.ID, "u2") + + _, err := m.AdminCreateUser(ctx, driver.AdminCreateUserInput{ + UserPoolID: pool.ID, Username: "u3", UserAttributes: []driver.Attribute{{Name: "preferred_username", Value: "neo"}}, + }) + assertException(t, err, driver.ExAliasExists, "An account with the given preferred_username already exists.") + + err = m.AdminUpdateUserAttributes(ctx, pool.ID, "u2", []driver.Attribute{{Name: "preferred_username", Value: "neo"}}) + assertException(t, err, driver.ExAliasExists, "") + + got, err := m.AdminGetUser(ctx, pool.ID, "neo") + requireNoError(t, err, "AdminGetUser by preferred_username") + + if got.Username != "u1" { + t.Fatalf("preferred_username resolves to %q", got.Username) + } +} + +func TestCustomAttributeNameForms(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "names") + + requireNoError(t, m.AddCustomAttributes(ctx, pool.ID, []driver.SchemaAttribute{{Name: "custom:t3"}}), "prefixed name") + + got, err := m.DescribeUserPool(ctx, pool.ID) + requireNoError(t, err, "DescribeUserPool") + + if _, ok := schemaAttribute(got, "custom:t3"); !ok { + t.Fatalf("custom:t3 not stored under its own name: %+v", got.SchemaAttributes) + } + + if _, ok := schemaAttribute(got, "custom:custom:t3"); ok { + t.Fatal("prefix doubled") + } + + err = m.AddCustomAttributes(ctx, pool.ID, []driver.SchemaAttribute{{Name: "t3"}}) + assertException(t, err, driver.ExInvalidParameter, "Existing attribute already has name custom:t3.") + + err = m.AddCustomAttributes(ctx, pool.ID, []driver.SchemaAttribute{{Name: strings.Repeat("n", 21)}}) + assertException(t, err, driver.ExInvalidParameter, "") + + requireNoError(t, m.AddCustomAttributes(ctx, pool.ID, + []driver.SchemaAttribute{{Name: strings.Repeat("n", 20)}}), "20-character name") +} + +func TestStandardAttributeOverrides(t *testing.T) { + m := newMock(t) + pool := createPoolWith(t, m, driver.CreateUserPoolInput{ + Name: "required-email", + SchemaAttributes: []driver.SchemaAttribute{{ + Name: "email", AttributeDataType: driver.AttributeTypeString, Required: true, Mutable: true, + StringAttributeConstraints: &driver.StringAttributeConstraints{MinLength: "5", MaxLength: "100"}, + }}, + }) + + email, ok := schemaAttribute(pool, "email") + if !ok || !email.Required || !email.Mutable || email.StringAttributeConstraints.MaxLength != "100" { + t.Fatalf("email override not applied: %+v", email) + } + + if len(pool.SchemaAttributes) != 20 { + t.Fatalf("override added a duplicate attribute: %d attributes", len(pool.SchemaAttributes)) + } + + // AdminCreateUser may leave required attributes empty (AWS docs, "Working + // with user attributes"), so no required check applies here. + mustCreateUser(t, m, pool.ID, "no-email") +} diff --git a/providers/aws/cognito/user_pools.go b/providers/aws/cognito/user_pools.go index b89ead006..512ce0124 100644 --- a/providers/aws/cognito/user_pools.go +++ b/providers/aws/cognito/user_pools.go @@ -2,6 +2,8 @@ package cognito import ( "context" + "slices" + "strings" "github.com/stackshy/cloudemu/v2/services/cognito/driver" ) @@ -219,12 +221,16 @@ func orDefault(v, def string) string { // mergeSchema returns the 20 default attributes followed by any caller-supplied // custom attributes (prefixed "dev:" for developer-only, else "custom:"), -// matching how Cognito names non-standard attributes. +// matching how Cognito names non-standard attributes. An entry naming a +// standard attribute adjusts that attribute instead of adding one. func mergeSchema(custom []driver.SchemaAttribute) []driver.SchemaAttribute { attrs := defaultSchemaAttributes() + n := len(attrs) for _, a := range custom { - if isStandardAttribute(a.Name) { + if i := slices.IndexFunc(attrs[:n], func(s driver.SchemaAttribute) bool { return s.Name == a.Name }); i >= 0 { + overrideStandard(&attrs[i], &a) + continue } @@ -234,16 +240,47 @@ func mergeSchema(custom []driver.SchemaAttribute) []driver.SchemaAttribute { return attrs } +// overrideStandard applies the caller's Required flag and length or value +// constraints to a standard attribute. sub is fixed. Mutable keeps its default: +// the wire cannot tell an omitted Mutable from false. +func overrideStandard(std, in *driver.SchemaAttribute) { + if std.Name == attrSub { + return + } + + std.Required = in.Required + + if c := in.StringAttributeConstraints; c != nil && std.StringAttributeConstraints != nil { + std.StringAttributeConstraints = copyStringConstraints(c) + } + + if c := in.NumberAttributeConstraints; c != nil && std.NumberAttributeConstraints != nil { + std.NumberAttributeConstraints = copyNumberConstraints(c) + } +} + // customAttribute returns a caller-supplied attribute with its custom: or dev: // prefix applied and its constraints deep-copied. func customAttribute(a driver.SchemaAttribute) driver.SchemaAttribute { - a.Name = customPrefix(a.DeveloperOnlyAttribute) + a.Name + a.Name = customPrefix(a.DeveloperOnlyAttribute) + bareCustomName(a.Name) a.StringAttributeConstraints = copyStringConstraints(a.StringAttributeConstraints) a.NumberAttributeConstraints = copyNumberConstraints(a.NumberAttributeConstraints) return a } +// bareCustomName strips a custom: or dev: prefix the caller already supplied, +// so "custom:tier" and "tier" name the same attribute. +func bareCustomName(name string) string { + for _, p := range []string{customPrefix(false), customPrefix(true)} { + if rest, ok := strings.CutPrefix(name, p); ok { + return rest + } + } + + return name +} + // customPrefix returns the attribute-name prefix Cognito applies to a // non-standard attribute. func customPrefix(developerOnly bool) string { @@ -253,14 +290,3 @@ func customPrefix(developerOnly bool) string { return "custom:" } - -// isStandardAttribute reports whether name is one of the 20 default attributes. -func isStandardAttribute(name string) bool { - for _, a := range defaultSchemaAttributes() { - if a.Name == name { - return true - } - } - - return false -} diff --git a/providers/aws/cognito/users.go b/providers/aws/cognito/users.go index e69b0867a..c0a492da3 100644 --- a/providers/aws/cognito/users.go +++ b/providers/aws/cognito/users.go @@ -88,6 +88,10 @@ func (m *Mock) AdminCreateUser(_ context.Context, in driver.AdminCreateUserInput return nil, err } + if err := m.claimSignIns(&pool, userKey(pool.ID, username), attrs, in.ForceAliasCreation, true); err != nil { + return nil, err + } + now := m.now() rec := userRecord{ PoolID: pool.ID, @@ -165,13 +169,6 @@ func (m *Mock) newUsername( return "", nil, err } - users := m.poolUsers(pool.ID) - for i := range users { - if attrValue(users[i].User.Attributes, attr) == username { - return "", nil, usernameExists("An account with the given " + attr + " already exists.") - } - } - return sub, mergeAttributes(attrs, []driver.Attribute{{Name: attr, Value: username}}), nil } @@ -286,16 +283,23 @@ func (m *Mock) AdminDeleteUser(_ context.Context, userPoolID, username string) e // AdminUpdateUserAttributes sets attribute values. Changing email or phone // number without also setting its verified flag marks it unverified, as Cognito -// does. +// does. A sign-in value another user already holds fails with +// AliasExistsException. func (m *Mock) AdminUpdateUserAttributes(_ context.Context, userPoolID, username string, attrs []driver.Attribute) error { - return m.updateUser(userPoolID, username, func(pool *driver.UserPool, rec *userRecord) error { + return m.updateUser(userPoolID, username, func(pool *driver.UserPool, key string, rec *userRecord) error { if err := validateAttributes(pool, attrs, true); err != nil { return err } updates := slices.Clone(attrs) updates = append(updates, unverifyChanged(rec.User.Attributes, attrs)...) - rec.User.Attributes = mergeAttributes(rec.User.Attributes, updates) + merged := mergeAttributes(rec.User.Attributes, updates) + + if err := m.claimSignIns(pool, key, merged, false, false); err != nil { + return err + } + + rec.User.Attributes = merged return nil }) @@ -322,7 +326,7 @@ func unverifyChanged(current, updates []driver.Attribute) []driver.Attribute { // AdminDeleteUserAttributes removes attributes from a user. func (m *Mock) AdminDeleteUserAttributes(_ context.Context, userPoolID, username string, names []string) error { - return m.updateUser(userPoolID, username, func(pool *driver.UserPool, rec *userRecord) error { + return m.updateUser(userPoolID, username, func(pool *driver.UserPool, _ string, rec *userRecord) error { for _, name := range names { a, ok := schemaAttribute(pool, name) if !ok { @@ -344,7 +348,7 @@ func (m *Mock) AdminDeleteUserAttributes(_ context.Context, userPoolID, username // AdminSetUserPassword sets a user's password. func (m *Mock) AdminSetUserPassword(_ context.Context, userPoolID, username, password string, permanent bool) error { - return m.updateUser(userPoolID, username, func(pool *driver.UserPool, rec *userRecord) error { + return m.updateUser(userPoolID, username, func(pool *driver.UserPool, _ string, rec *userRecord) error { if err := checkPassword(password, pool.Policies.PasswordPolicy); err != nil { return err } @@ -362,7 +366,7 @@ func (m *Mock) AdminSetUserPassword(_ context.Context, userPoolID, username, pas // AdminEnableUser enables a user. func (m *Mock) AdminEnableUser(_ context.Context, userPoolID, username string) error { - return m.updateUser(userPoolID, username, func(_ *driver.UserPool, rec *userRecord) error { + return m.updateUser(userPoolID, username, func(_ *driver.UserPool, _ string, rec *userRecord) error { rec.User.Enabled = true return nil @@ -371,7 +375,7 @@ func (m *Mock) AdminEnableUser(_ context.Context, userPoolID, username string) e // AdminDisableUser disables a user. func (m *Mock) AdminDisableUser(_ context.Context, userPoolID, username string) error { - return m.updateUser(userPoolID, username, func(_ *driver.UserPool, rec *userRecord) error { + return m.updateUser(userPoolID, username, func(_ *driver.UserPool, _ string, rec *userRecord) error { rec.User.Enabled = false return nil @@ -381,7 +385,7 @@ func (m *Mock) AdminDisableUser(_ context.Context, userPoolID, username string) // AdminResetUserPassword moves a user to RESET_REQUIRED. A user who has not // yet replaced the temporary password cannot be reset. func (m *Mock) AdminResetUserPassword(_ context.Context, userPoolID, username string) error { - return m.updateUser(userPoolID, username, func(_ *driver.UserPool, rec *userRecord) error { + return m.updateUser(userPoolID, username, func(_ *driver.UserPool, _ string, rec *userRecord) error { if rec.User.UserStatus == driver.UserStatusForceChangePassword { return notAuthorized("User password cannot be reset in the current state.") } @@ -394,7 +398,7 @@ func (m *Mock) AdminResetUserPassword(_ context.Context, userPoolID, username st // updateUser runs fn on a copy of the resolved user under the mutation lock and // stores the result with a fresh last-modified time. -func (m *Mock) updateUser(userPoolID, username string, fn func(*driver.UserPool, *userRecord) error) error { +func (m *Mock) updateUser(userPoolID, username string, fn func(pool *driver.UserPool, key string, rec *userRecord) error) error { m.mu.Lock() defer m.mu.Unlock() @@ -409,7 +413,7 @@ func (m *Mock) updateUser(userPoolID, username string, fn func(*driver.UserPool, } rec = copyUserRecord(rec) - if err := fn(&pool, &rec); err != nil { + if err := fn(&pool, key, &rec); err != nil { return err } @@ -437,31 +441,6 @@ func (m *Mock) resolveUser(pool *driver.UserPool, name string) (string, userReco return "", userRecord{}, false } -func signInMatches(pool *driver.UserPool, u *driver.User, name string) bool { - for _, attr := range pool.UsernameAttributes { - if attrValue(u.Attributes, attr) == name { - return true - } - } - - for _, attr := range pool.AliasAttributes { - if attrValue(u.Attributes, attr) != name { - continue - } - - switch attr { - case attrEmail: - return attrValue(u.Attributes, attrEmailVerified) == attrTrue - case attrPhoneNumber: - return attrValue(u.Attributes, attrPhoneNumberVerified) == attrTrue - default: - return true - } - } - - return false -} - // poolUserKeys returns the store keys of a pool's users, sorted. Pool ids never // contain the key separator, so the "/" prefix selects exactly one pool. func (m *Mock) poolUserKeys(poolID string) []string { diff --git a/server/aws/cognito/user_ops.go b/server/aws/cognito/user_ops.go index 970c5fb4d..7b1deefb9 100644 --- a/server/aws/cognito/user_ops.go +++ b/server/aws/cognito/user_ops.go @@ -73,6 +73,7 @@ type adminCreateUserRequest struct { TemporaryPassword string `json:"TemporaryPassword"` MessageAction string `json:"MessageAction"` DesiredDeliveryMediums []string `json:"DesiredDeliveryMediums"` + ForceAliasCreation bool `json:"ForceAliasCreation"` } type adminCreateUserResponse struct { @@ -88,6 +89,7 @@ func (h *Handler) adminCreateUser(w http.ResponseWriter, r *http.Request) { TemporaryPassword: req.TemporaryPassword, MessageAction: req.MessageAction, DesiredDeliveryMediums: req.DesiredDeliveryMediums, + ForceAliasCreation: req.ForceAliasCreation, }) if err != nil { return nil, err diff --git a/server/aws/cognito/users_sdk_test.go b/server/aws/cognito/users_sdk_test.go index 62dec6961..7c8deffd5 100644 --- a/server/aws/cognito/users_sdk_test.go +++ b/server/aws/cognito/users_sdk_test.go @@ -160,6 +160,56 @@ func TestSDKAdminUserLifecycle(t *testing.T) { } } +func TestSDKAliasExistsAndForceAliasCreation(t *testing.T) { + ctx := context.Background() + c := newCognitoClient(t) + + created, err := c.CreateUserPool(ctx, &cip.CreateUserPoolInput{ + PoolName: aws.String("alias"), + AliasAttributes: []ciptypes.AliasAttributeType{ciptypes.AliasAttributeTypeEmail}, + }) + if err != nil { + t.Fatalf("CreateUserPool: %v", err) + } + + pool := created.UserPool.Id + verified := []ciptypes.AttributeType{ + {Name: aws.String("email"), Value: aws.String("dup@example.com")}, + {Name: aws.String("email_verified"), Value: aws.String("true")}, + } + + newUser := func(name string, force bool) error { + _, err := c.AdminCreateUser(ctx, &cip.AdminCreateUserInput{ + UserPoolId: pool, Username: aws.String(name), UserAttributes: verified, + MessageAction: ciptypes.MessageActionTypeSuppress, ForceAliasCreation: force, + }) + + return err + } + + if err := newUser("first", false); err != nil { + t.Fatalf("AdminCreateUser first: %v", err) + } + + var aee *ciptypes.AliasExistsException + if err := newUser("second", false); !errors.As(err, &aee) { + t.Fatalf("expected typed AliasExistsException, got %v", err) + } + + if err := newUser("second", true); err != nil { + t.Fatalf("AdminCreateUser ForceAliasCreation: %v", err) + } + + first, err := c.AdminGetUser(ctx, &cip.AdminGetUserInput{UserPoolId: pool, Username: aws.String("first")}) + if err != nil { + t.Fatalf("AdminGetUser: %v", err) + } + + if attr(first.UserAttributes, "email_verified") != "false" { + t.Fatalf("old alias owner still verified: %+v", first.UserAttributes) + } +} + func TestSDKListUsersPaginatorAndFilterErrors(t *testing.T) { ctx := context.Background() c := newCognitoClient(t) diff --git a/services/cognito/driver/errors.go b/services/cognito/driver/errors.go index df6752812..4500ccd80 100644 --- a/services/cognito/driver/errors.go +++ b/services/cognito/driver/errors.go @@ -16,6 +16,7 @@ const ( const ( ExUserNotFound = "UserNotFoundException" ExUsernameExists = "UsernameExistsException" + ExAliasExists = "AliasExistsException" ExInvalidPassword = "InvalidPasswordException" ExNotAuthorized = "NotAuthorizedException" ExUnsupportedUserState = "UnsupportedUserStateException" diff --git a/services/cognito/driver/types.go b/services/cognito/driver/types.go index ce31b0509..9daa68684 100644 --- a/services/cognito/driver/types.go +++ b/services/cognito/driver/types.go @@ -304,6 +304,9 @@ type AdminCreateUserInput struct { TemporaryPassword string MessageAction string DesiredDeliveryMediums []string + // ForceAliasCreation moves a verified email or phone alias that another + // user already holds to the new user instead of failing. + ForceAliasCreation bool } // ListUsersInput is the input to ListUsers. A nil AttributesToGet returns every From 0f2ead6075bb31b04af8a75fb9819e5e42d85a94 Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 27 Sep 2026 22:07:30 +0530 Subject: [PATCH 3/3] test(aws-cognito): AliasExistsException is untyped for AdminCreateUser in the SDK --- server/aws/cognito/users_sdk_test.go | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/server/aws/cognito/users_sdk_test.go b/server/aws/cognito/users_sdk_test.go index 7c8deffd5..696b99c68 100644 --- a/server/aws/cognito/users_sdk_test.go +++ b/server/aws/cognito/users_sdk_test.go @@ -191,10 +191,9 @@ func TestSDKAliasExistsAndForceAliasCreation(t *testing.T) { t.Fatalf("AdminCreateUser first: %v", err) } - var aee *ciptypes.AliasExistsException - if err := newUser("second", false); !errors.As(err, &aee) { - t.Fatalf("expected typed AliasExistsException, got %v", err) - } + // AdminCreateUser does not model AliasExistsException, so the SDK surfaces + // it as a generic API error carrying the code. + requireErrorCode(t, newUser("second", false), "AliasExistsException", "An account with the given email already exists.") if err := newUser("second", true); err != nil { t.Fatalf("AdminCreateUser ForceAliasCreation: %v", err)