From 286f8c0f70f564a0f8497a9f2c694bd3d1df4fa2 Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 4 Oct 2026 16:55:27 +0530 Subject: [PATCH 1/3] feat(aws-cognito): groups, sign-up and password sign-in with RS256 tokens (C2) --- docs/coverage/README.md | 2 +- docs/coverage/aws/README.md | 2 +- docs/coverage/aws/cognito.md | 43 +- docs/coverage/coverage.json | 95 ++++ providers/aws/cognito/auth.go | 590 ++++++++++++++++++++ providers/aws/cognito/auth_test.go | 654 +++++++++++++++++++++++ providers/aws/cognito/cognito.go | 33 +- providers/aws/cognito/groups.go | 304 +++++++++++ providers/aws/cognito/groups_test.go | 137 +++++ providers/aws/cognito/keys.go | 89 +++ providers/aws/cognito/passwords.go | 31 ++ providers/aws/cognito/passwords_test.go | 31 -- providers/aws/cognito/sign_up.go | 438 +++++++++++++++ providers/aws/cognito/sign_up_test.go | 267 +++++++++ providers/aws/cognito/snapshot.go | 95 +++- providers/aws/cognito/tokens.go | 403 ++++++++++++++ providers/aws/cognito/user_pools.go | 3 + providers/aws/cognito/users.go | 22 +- server/aws/authbypass_test.go | 1 - server/aws/authz_completeness_test.go | 2 +- server/aws/aws.go | 6 + server/aws/cognito/auth_ops.go | 259 +++++++++ server/aws/cognito/auth_sdk_test.go | 378 +++++++++++++ server/aws/cognito/group_ops.go | 191 +++++++ server/aws/cognito/handler.go | 24 + server/aws/cognito/wellknown.go | 115 ++++ server/aws/cognito_auth_enforced_test.go | 88 +++ server/aws/publicauth_test.go | 4 +- server/serverkit/cognito_codes.go | 98 ++++ server/serverkit/cognito_codes_test.go | 71 +++ server/serverkit/serverkit.go | 2 + services/cognito/driver/auth_types.go | 143 +++++ services/cognito/driver/driver.go | 79 ++- services/cognito/driver/errors.go | 12 + 34 files changed, 4655 insertions(+), 57 deletions(-) create mode 100644 providers/aws/cognito/auth.go create mode 100644 providers/aws/cognito/auth_test.go create mode 100644 providers/aws/cognito/groups.go create mode 100644 providers/aws/cognito/groups_test.go create mode 100644 providers/aws/cognito/keys.go create mode 100644 providers/aws/cognito/sign_up.go create mode 100644 providers/aws/cognito/sign_up_test.go create mode 100644 providers/aws/cognito/tokens.go create mode 100644 server/aws/cognito/auth_ops.go create mode 100644 server/aws/cognito/auth_sdk_test.go create mode 100644 server/aws/cognito/group_ops.go create mode 100644 server/aws/cognito/wellknown.go create mode 100644 server/aws/cognito_auth_enforced_test.go create mode 100644 server/serverkit/cognito_codes.go create mode 100644 server/serverkit/cognito_codes_test.go create mode 100644 services/cognito/driver/auth_types.go diff --git a/docs/coverage/README.md b/docs/coverage/README.md index b9a759762..114a89e10 100644 --- a/docs/coverage/README.md +++ b/docs/coverage/README.md @@ -55,7 +55,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) | - | - | - | 29 | +| `cognito` | [Cognito](./aws/cognito.md) | - | - | - | 50 | | `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 216629258..851163efb 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` | 29 | +| [Cognito](./cognito.md) | `cognito` | 50 | | [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 9c4148c48..6d10eec93 100644 --- a/docs/coverage/aws/cognito.md +++ b/docs/coverage/aws/cognito.md @@ -3,40 +3,81 @@ AWS's `cognito` service · portable interface `driver.Cognito` · [AWS index](./README.md) -## Operations (29) +## Operations (50) | Operation | Description | | --- | --- | | `AddCustomAttributes` | AddCustomAttributes appends custom attributes to a pool's schema. Names get | +| `AdminAddUserToGroup` | | +| `AdminConfirmSignUp` | | | `AdminCreateUser` | AdminCreateUser creates a user in FORCE_CHANGE_PASSWORD with a generated | | `AdminDeleteUser` | | | `AdminDeleteUserAttributes` | | | `AdminDisableUser` | | | `AdminEnableUser` | | | `AdminGetUser` | | +| `AdminInitiateAuth` | AdminInitiateAuth runs ADMIN_USER_PASSWORD_AUTH, ADMIN_NO_SRP_AUTH or | +| `AdminListGroupsForUser` | | +| `AdminRemoveUserFromGroup` | | | `AdminResetUserPassword` | AdminResetUserPassword moves the user to RESET_REQUIRED. | +| `AdminRespondToAuthChallenge` | | | `AdminSetUserPassword` | AdminSetUserPassword sets a password checked against the pool policy. A | | `AdminUpdateUserAttributes` | | +| `AdminUserGlobalSignOut` | | +| `ConfirmSignUp` | ConfirmSignUp confirms a user with the code SignUp or | +| `CreateGroup` | CreateGroup creates a group. A name already in the pool fails with | | `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` | | +| `DeleteGroup` | DeleteGroup removes a group and every membership in it. | | `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 | | `DescribeUserPoolClient` | | | `DescribeUserPoolDomain` | DescribeUserPoolDomain returns a domain's description. An unknown domain | +| `GetGroup` | | +| `GetUser` | GetUser returns the user an access token was issued to. | | `GetUserPoolMfaConfig` | GetUserPoolMfaConfig returns a pool's MFA configuration. The Terraform AWS | +| `GlobalSignOut` | GlobalSignOut revokes every token issued to the access token's user. | +| `InitiateAuth` | InitiateAuth runs USER_PASSWORD_AUTH or REFRESH_TOKEN_AUTH for a client. | +| `ListGroups` | | | `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 | +| `ListUsersInGroup` | | +| `ResendConfirmationCode` | | +| `RespondToAuthChallenge` | RespondToAuthChallenge answers the NEW_PASSWORD_REQUIRED challenge. | +| `RevokeToken` | RevokeToken revokes a refresh token and the access and ID tokens minted | | `SetUserPoolMfaConfig` | SetUserPoolMfaConfig replaces a pool's MFA configuration and returns the | +| `SignUp` | SignUp registers an UNCONFIRMED user, checks the password against the | | `TagResource` | | | `UntagResource` | | +| `UpdateGroup` | UpdateGroup changes the fields the input sets and returns the group. | | `UpdateUserPool` | UpdateUserPool applies the mutable pool settings. A nil field is left | | `UpdateUserPoolClient` | UpdateUserPoolClient replaces the client's settings and returns the result. | +## Optional capabilities + +Discovered by type assertion; only some providers implement these. + +### CodeInspector + +CodeInspector exposes the confirmation codes the emulator would have sent by + +| Operation | Description | +| --- | --- | +| `ConfirmationCode` | | + +### KeySetProvider + +KeySetProvider publishes a user pool's token-signing keys and issuer for the + +| Operation | Description | +| --- | --- | +| `SigningKeys` | SigningKeys returns the pool's issuer URL and its public JSON Web Key Set. | + ## Not in scope _Not documented yet. See the [emulator boundary](../../../README.md) for cloudemu-wide non-goals._ diff --git a/docs/coverage/coverage.json b/docs/coverage/coverage.json index 8c3fc5a3a..4265e9981 100644 --- a/docs/coverage/coverage.json +++ b/docs/coverage/coverage.json @@ -4053,6 +4053,12 @@ "name": "AddCustomAttributes", "doc": "AddCustomAttributes appends custom attributes to a pool's schema. Names get" }, + { + "name": "AdminAddUserToGroup" + }, + { + "name": "AdminConfirmSignUp" + }, { "name": "AdminCreateUser", "doc": "AdminCreateUser creates a user in FORCE_CHANGE_PASSWORD with a generated" @@ -4072,10 +4078,23 @@ { "name": "AdminGetUser" }, + { + "name": "AdminInitiateAuth", + "doc": "AdminInitiateAuth runs ADMIN_USER_PASSWORD_AUTH, ADMIN_NO_SRP_AUTH or" + }, + { + "name": "AdminListGroupsForUser" + }, + { + "name": "AdminRemoveUserFromGroup" + }, { "name": "AdminResetUserPassword", "doc": "AdminResetUserPassword moves the user to RESET_REQUIRED." }, + { + "name": "AdminRespondToAuthChallenge" + }, { "name": "AdminSetUserPassword", "doc": "AdminSetUserPassword sets a password checked against the pool policy. A" @@ -4083,6 +4102,17 @@ { "name": "AdminUpdateUserAttributes" }, + { + "name": "AdminUserGlobalSignOut" + }, + { + "name": "ConfirmSignUp", + "doc": "ConfirmSignUp confirms a user with the code SignUp or" + }, + { + "name": "CreateGroup", + "doc": "CreateGroup creates a group. A name already in the pool fails with" + }, { "name": "CreateUserPool", "doc": "CreateUserPool creates a user pool, generating its id and ARN, seeding the" @@ -4094,6 +4124,10 @@ { "name": "CreateUserPoolDomain" }, + { + "name": "DeleteGroup", + "doc": "DeleteGroup removes a group and every membership in it." + }, { "name": "DeleteUserPool", "doc": "DeleteUserPool removes a user pool with its users, clients and tags. Like" @@ -4115,10 +4149,28 @@ "name": "DescribeUserPoolDomain", "doc": "DescribeUserPoolDomain returns a domain's description. An unknown domain" }, + { + "name": "GetGroup" + }, + { + "name": "GetUser", + "doc": "GetUser returns the user an access token was issued to." + }, { "name": "GetUserPoolMfaConfig", "doc": "GetUserPoolMfaConfig returns a pool's MFA configuration. The Terraform AWS" }, + { + "name": "GlobalSignOut", + "doc": "GlobalSignOut revokes every token issued to the access token's user." + }, + { + "name": "InitiateAuth", + "doc": "InitiateAuth runs USER_PASSWORD_AUTH or REFRESH_TOKEN_AUTH for a client." + }, + { + "name": "ListGroups" + }, { "name": "ListTagsForResource" }, @@ -4134,16 +4186,38 @@ "name": "ListUsers", "doc": "ListUsers returns users sorted by username, filtered by an optional" }, + { + "name": "ListUsersInGroup" + }, + { + "name": "ResendConfirmationCode" + }, + { + "name": "RespondToAuthChallenge", + "doc": "RespondToAuthChallenge answers the NEW_PASSWORD_REQUIRED challenge." + }, + { + "name": "RevokeToken", + "doc": "RevokeToken revokes a refresh token and the access and ID tokens minted" + }, { "name": "SetUserPoolMfaConfig", "doc": "SetUserPoolMfaConfig replaces a pool's MFA configuration and returns the" }, + { + "name": "SignUp", + "doc": "SignUp registers an UNCONFIRMED user, checks the password against the" + }, { "name": "TagResource" }, { "name": "UntagResource" }, + { + "name": "UpdateGroup", + "doc": "UpdateGroup changes the fields the input sets and returns the group." + }, { "name": "UpdateUserPool", "doc": "UpdateUserPool applies the mutable pool settings. A nil field is left" @@ -4153,6 +4227,27 @@ "doc": "UpdateUserPoolClient replaces the client's settings and returns the result." } ], + "optionalCapabilities": [ + { + "name": "CodeInspector", + "doc": "CodeInspector exposes the confirmation codes the emulator would have sent by", + "operations": [ + { + "name": "ConfirmationCode" + } + ] + }, + { + "name": "KeySetProvider", + "doc": "KeySetProvider publishes a user pool's token-signing keys and issuer for the", + "operations": [ + { + "name": "SigningKeys", + "doc": "SigningKeys returns the pool's issuer URL and its public JSON Web Key Set." + } + ] + } + ], "providers": { "aws": "Cognito" } diff --git a/providers/aws/cognito/auth.go b/providers/aws/cognito/auth.go new file mode 100644 index 000000000..c9631c063 --- /dev/null +++ b/providers/aws/cognito/auth.go @@ -0,0 +1,590 @@ +package cognito + +import ( + "context" + "crypto/hmac" + "encoding/json" + "slices" + "strings" + "time" + + "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// Auth parameter and challenge response names. +const ( + paramUsername = "USERNAME" + paramPassword = "PASSWORD" + paramSecretHash = "SECRET_HASH" + paramRefreshToken = "REFRESH_TOKEN" + paramNewPassword = "NEW_PASSWORD" + userAttrPrefix = "userAttributes." + sessionBytes = 96 + existenceEnabled = "ENABLED" + allowFlowPrefix = "ALLOW_" + legacyAdminNoSRP = "ADMIN_NO_SRP_AUTH" + allowAdminUserPwd = "ALLOW_ADMIN_USER_PASSWORD_AUTH" + allowUserPwd = "ALLOW_USER_PASSWORD_AUTH" + allowRefresh = "ALLOW_REFRESH_TOKEN_AUTH" +) + +// authFlows is the AuthFlowType enum, in the order the validation error lists it. +// +//nolint:gochecknoglobals // static protocol enum +var authFlows = []string{ + driver.AuthFlowUserSRP, driver.AuthFlowRefreshTokenAuth, driver.AuthFlowRefreshToken, driver.AuthFlowCustom, + driver.AuthFlowAdminNoSRP, driver.AuthFlowUserPassword, driver.AuthFlowAdminUserPassword, driver.AuthFlowUser, +} + +// challengeSession is an outstanding NEW_PASSWORD_REQUIRED challenge. +type challengeSession struct { + poolID string + clientID string + username string + challenge string + expires time.Time +} + +func missingParam(name string) error { return invalidParameter("Missing required parameter %s", name) } + +func incorrectCredentials() error { return notAuthorized("Incorrect username or password.") } + +func invalidSession() error { return notAuthorized("Invalid session for the user.") } + +func invalidRefreshToken() error { return notAuthorized("Invalid Refresh Token") } + +// unknownUser is the error for a sign-in naming no user. A client with +// PreventUserExistenceErrors ENABLED hides that the user does not exist. +func unknownUser(client *driver.UserPoolClient) error { + if client.PreventUserExistenceErrors == existenceEnabled { + return incorrectCredentials() + } + + return userNotFound() +} + +func notSupported() error { return invalidParameter("Initiate Auth method not supported.") } + +func noSecretReceived(clientID string) string { + return "Client " + clientID + " is configured for secret but secret was not received" +} + +func apiErr(exception string, code errors.Code, msg string) error { + return &driver.APIError{Exception: exception, Err: errors.New(code, msg)} +} + +// InitiateAuth starts a client-side sign-in. +func (m *Mock) InitiateAuth(_ context.Context, in driver.InitiateAuthInput) (*driver.AuthResult, error) { + return m.initiate(&in, false) +} + +// AdminInitiateAuth starts a server-side sign-in in a named pool. +func (m *Mock) AdminInitiateAuth(_ context.Context, in driver.InitiateAuthInput) (*driver.AuthResult, error) { + return m.initiate(&in, true) +} + +// authClient resolves the app client of a sign-in. The admin operations also +// name the pool, which must exist and own the client. +func (m *Mock) authClient(poolID, clientID string, admin bool) (driver.UserPoolClient, driver.UserPool, error) { + if admin && !m.userPools.Has(poolID) { + return driver.UserPoolClient{}, driver.UserPool{}, poolNotFound(poolID) + } + + client, err := m.clientByID(clientID) + if err != nil { + return driver.UserPoolClient{}, driver.UserPool{}, err + } + + if admin && client.UserPoolID != poolID { + return driver.UserPoolClient{}, driver.UserPool{}, resourceNotFound("User pool client %s does not exist.", clientID) + } + + pool, ok := m.userPools.Get(client.UserPoolID) + if !ok { + return driver.UserPoolClient{}, driver.UserPool{}, poolNotFound(client.UserPoolID) + } + + return client, pool, nil +} + +func (m *Mock) initiate(in *driver.InitiateAuthInput, admin bool) (*driver.AuthResult, error) { + if !slices.Contains(authFlows, in.AuthFlow) { + return nil, invalidParameter("1 validation error detected: Value '%s' at 'authFlow' failed to satisfy constraint: "+ + "Member must satisfy enum value set: [%s]", in.AuthFlow, strings.Join(authFlows, ", ")) + } + + m.mu.Lock() + defer m.mu.Unlock() + + client, pool, err := m.authClient(in.UserPoolID, in.ClientID, admin) + if err != nil { + return nil, err + } + + if err := checkFlow(&client, in.AuthFlow, admin); err != nil { + return nil, err + } + + switch in.AuthFlow { + case driver.AuthFlowUserPassword, driver.AuthFlowAdminUserPassword, driver.AuthFlowAdminNoSRP: + return m.passwordAuth(&pool, client, in.AuthParameters) + case driver.AuthFlowRefreshToken, driver.AuthFlowRefreshTokenAuth: + return m.refreshAuth(client, in.AuthParameters) + default: + return nil, invalidParameter("%s is not supported by cloudemu yet.", in.AuthFlow) + } +} + +// checkFlow applies the client's ExplicitAuthFlows. The legacy names +// USER_PASSWORD_AUTH and ADMIN_NO_SRP_AUTH enable the matching ALLOW_ flows, +// and a client configured only with legacy names can always refresh. +func checkFlow(client *driver.UserPoolClient, flow string, admin bool) error { + flows := client.ExplicitAuthFlows + has := func(names ...string) bool { + return slices.ContainsFunc(names, func(n string) bool { return slices.Contains(flows, n) }) + } + + if err := checkFlowSide(flow, admin); err != nil { + return err + } + + switch flow { + case driver.AuthFlowUserPassword: + if !has(allowUserPwd, driver.AuthFlowUserPassword) { + return invalidParameter("USER_PASSWORD_AUTH flow not enabled for this client") + } + case driver.AuthFlowAdminUserPassword, driver.AuthFlowAdminNoSRP: + if !has(allowAdminUserPwd, legacyAdminNoSRP) { + return invalidParameter("Auth flow not enabled for this client") + } + case driver.AuthFlowRefreshToken, driver.AuthFlowRefreshTokenAuth: + if !refreshAllowed(flows) { + return invalidParameter("Refresh Token flow not enabled for this client") + } + case driver.AuthFlowUserSRP: + if !has("ALLOW_USER_SRP_AUTH") { + return invalidParameter("USER_SRP_AUTH is not enabled for the client.") + } + } + + return nil +} + +// refreshAllowed reports whether a client may refresh: it lists +// ALLOW_REFRESH_TOKEN_AUTH, or uses only legacy flow names. +func refreshAllowed(flows []string) bool { + return slices.Contains(flows, allowRefresh) || + !slices.ContainsFunc(flows, func(f string) bool { return strings.HasPrefix(f, allowFlowPrefix) }) +} + +// checkFlowSide rejects a password flow sent to the wrong API: the ADMIN_ +// flows only work through AdminInitiateAuth and USER_PASSWORD_AUTH only through +// InitiateAuth. +func checkFlowSide(flow string, admin bool) error { + adminFlow := flow == driver.AuthFlowAdminUserPassword || flow == driver.AuthFlowAdminNoSRP + if (adminFlow && !admin) || (flow == driver.AuthFlowUserPassword && admin) { + return notSupported() + } + + return nil +} + +// passwordAuth checks a username and password and either issues tokens or +// returns the NEW_PASSWORD_REQUIRED challenge. +// +//nolint:gocritic // hugeParam: stored client copy passed by value +func (m *Mock) passwordAuth(pool *driver.UserPool, client driver.UserPoolClient, params map[string]string) (*driver.AuthResult, error) { + username := params[paramUsername] + if username == "" { + return nil, missingParam(paramUsername) + } + + password := params[paramPassword] + if password == "" { + return nil, missingParam(paramPassword) + } + + _, rec, found := m.resolveUser(pool, username) + + names := []string{username} + if found { + names = append(names, rec.User.Username) + } + + if err := checkSecretHash(client, params[paramSecretHash], noSecretReceived(client.ClientID), names...); err != nil { + return nil, err + } + + if !found { + return nil, unknownUser(&client) + } + + if !verifyPassword(rec.PasswordSalt, rec.PasswordHash, password) { + return nil, incorrectCredentials() + } + + if err := checkCanSignIn(&rec); err != nil { + return nil, err + } + + if rec.User.UserStatus == driver.UserStatusForceChangePassword { + return m.newPasswordChallenge(pool, &client, &rec), nil + } + + return m.signIn(client, &rec) +} + +// checkCanSignIn rejects a correct password for a user who may not sign in +// yet: disabled, unconfirmed, or due a password reset. +func checkCanSignIn(rec *userRecord) error { + if !rec.User.Enabled { + return notAuthorized("User is disabled.") + } + + switch rec.User.UserStatus { + case driver.UserStatusUnconfirmed: + return apiErr(driver.ExUserNotConfirmed, errors.FailedPrecondition, "User is not confirmed.") + case driver.UserStatusResetRequired: + return apiErr(driver.ExPasswordResetRequired, errors.FailedPrecondition, "Password reset required for the user") + } + + return nil +} + +// signIn records a login for a user and returns its tokens. +// +//nolint:gocritic // hugeParam: stored client copy passed by value +func (m *Mock) signIn(client driver.UserPoolClient, rec *userRecord) (*driver.AuthResult, error) { + l, refresh := m.startLogin(client, rec) + + res, err := m.mintTokens(client, rec, &l) + if err != nil { + m.logins.Delete(l.OriginJTI) + + return nil, err + } + + res.RefreshToken = refresh + + return &driver.AuthResult{AuthenticationResult: res, ChallengeParameters: map[string]string{}}, nil +} + +// newPasswordChallenge opens a NEW_PASSWORD_REQUIRED session. +func (m *Mock) newPasswordChallenge(pool *driver.UserPool, client *driver.UserPoolClient, rec *userRecord) *driver.AuthResult { + attrs := map[string]string{} + + for _, a := range rec.User.Attributes { + if a.Name != attrSub { + attrs[a.Name] = a.Value + } + } + + required := []string{} + + for _, a := range pool.SchemaAttributes { + if a.Required && a.Name != attrSub && attrValue(rec.User.Attributes, a.Name) == "" { + required = append(required, userAttrPrefix+a.Name) + } + } + + attrJSON, _ := json.Marshal(attrs) + requiredJSON, _ := json.Marshal(required) + + session := randomB64(sessionBytes) + + m.sessionsMu.Lock() + m.pruneSessions() + m.sessions[session] = challengeSession{ + poolID: pool.ID, + clientID: client.ClientID, + username: rec.User.Username, + challenge: driver.ChallengeNewPasswordRequired, + expires: m.now().Add(time.Duration(client.AuthSessionValidity) * time.Minute), + } + m.sessionsMu.Unlock() + + return &driver.AuthResult{ + ChallengeName: driver.ChallengeNewPasswordRequired, + Session: session, + ChallengeParameters: map[string]string{ + "USER_ID_FOR_SRP": rec.User.Username, + "requiredAttributes": string(requiredJSON), + "userAttributes": string(attrJSON), + }, + } +} + +// pruneSessions drops expired sessions. It must run under sessionsMu. +func (m *Mock) pruneSessions() { + now := m.now() + + for k, s := range m.sessions { + if !now.Before(s.expires) { + delete(m.sessions, k) + } + } +} + +// refreshAuth mints new access and ID tokens from a refresh token. The login +// keeps its origin_jti and auth_time, and no new refresh token is issued. +// +//nolint:gocritic // hugeParam: stored client copy passed by value +func (m *Mock) refreshAuth(client driver.UserPoolClient, params map[string]string) (*driver.AuthResult, error) { + token := params[paramRefreshToken] + if token == "" { + return nil, missingParam(paramRefreshToken) + } + + l, ok := m.lookupRefresh(token) + if !ok || l.ClientID != client.ClientID { + return nil, invalidRefreshToken() + } + + if err := checkSecretHash(client, params[paramSecretHash], noSecretReceived(client.ClientID), + params[paramUsername], l.Username, l.Sub); err != nil { + return nil, err + } + + rec, err := m.refreshUser(&l) + if err != nil { + return nil, err + } + + res, err := m.mintTokens(client, &rec, &l) + if err != nil { + return nil, err + } + + return &driver.AuthResult{AuthenticationResult: res, ChallengeParameters: map[string]string{}}, nil +} + +// refreshUser checks a login can still refresh and returns its user. +func (m *Mock) refreshUser(l *loginRecord) (userRecord, error) { + if l.Revoked { + return userRecord{}, notAuthorized("Refresh Token has been revoked") + } + + if !m.now().Before(l.RefreshExpires) { + return userRecord{}, notAuthorized("Refresh Token has expired") + } + + rec, ok := m.users.Get(userKey(l.PoolID, l.Username)) + if !ok || attrValue(rec.User.Attributes, attrSub) != l.Sub { + return userRecord{}, invalidRefreshToken() + } + + if !rec.User.Enabled { + return userRecord{}, notAuthorized("User is disabled.") + } + + return rec, nil +} + +// RespondToAuthChallenge answers a challenge from InitiateAuth. +func (m *Mock) RespondToAuthChallenge(_ context.Context, in driver.RespondToAuthChallengeInput) (*driver.AuthResult, error) { + return m.respond(&in, false) +} + +// AdminRespondToAuthChallenge answers a challenge from AdminInitiateAuth. +func (m *Mock) AdminRespondToAuthChallenge(_ context.Context, in driver.RespondToAuthChallengeInput) (*driver.AuthResult, error) { + return m.respond(&in, true) +} + +func (m *Mock) respond(in *driver.RespondToAuthChallengeInput, admin bool) (*driver.AuthResult, error) { + if in.ChallengeName != driver.ChallengeNewPasswordRequired { + return nil, invalidParameter("%s is not supported by cloudemu yet.", in.ChallengeName) + } + + m.mu.Lock() + defer m.mu.Unlock() + + client, pool, err := m.authClient(in.UserPoolID, in.ClientID, admin) + if err != nil { + return nil, err + } + + username := in.ChallengeResponses[paramUsername] + if username == "" { + return nil, missingParam(paramUsername) + } + + sess, err := m.takeSession(in.Session, client.ClientID, in.ChallengeName) + if err != nil { + return nil, err + } + + key, rec, found := m.resolveUser(&pool, username) + if !found || rec.User.Username != sess.username { + return nil, invalidSession() + } + + if err := checkSecretHash(client, in.ChallengeResponses[paramSecretHash], noSecretReceived(client.ClientID), + username, rec.User.Username); err != nil { + return nil, err + } + + rec = copyUserRecord(rec) + if err := m.applyNewPassword(&pool, key, &rec, in.ChallengeResponses); err != nil { + return nil, err + } + + m.endSession(in.Session) + m.users.Set(key, rec) + + return m.signIn(client, &rec) +} + +// takeSession returns a live session for the client and challenge. An expired +// session is dropped. +func (m *Mock) takeSession(id, clientID, challenge string) (challengeSession, error) { + m.sessionsMu.Lock() + defer m.sessionsMu.Unlock() + + sess, ok := m.sessions[id] + if !ok || sess.clientID != clientID || sess.challenge != challenge { + return challengeSession{}, invalidSession() + } + + if !m.now().Before(sess.expires) { + delete(m.sessions, id) + + return challengeSession{}, notAuthorized("Invalid session for the user, session is expired.") + } + + return sess, nil +} + +func (m *Mock) endSession(id string) { + m.sessionsMu.Lock() + delete(m.sessions, id) + m.sessionsMu.Unlock() +} + +// applyNewPassword sets the new password and any userAttributes.* responses, +// and confirms the user. +func (m *Mock) applyNewPassword(pool *driver.UserPool, key string, rec *userRecord, responses map[string]string) error { + newPassword := responses[paramNewPassword] + if newPassword == "" { + return missingParam(paramNewPassword) + } + + var updates []driver.Attribute + + for k, v := range responses { + if name, ok := strings.CutPrefix(k, userAttrPrefix); ok { + updates = append(updates, driver.Attribute{Name: name, Value: v}) + } + } + + slices.SortFunc(updates, func(a, b driver.Attribute) int { return strings.Compare(a.Name, b.Name) }) + + if err := validateAttributes(pool, updates, false); err != nil { + return err + } + + attrs := mergeAttributes(rec.User.Attributes, updates) + if err := checkRequired(pool, attrs); err != nil { + return err + } + + if err := checkPassword(newPassword, pool.Policies.PasswordPolicy); err != nil { + return err + } + + if err := m.claimSignIns(pool, key, attrs, false, false); err != nil { + return err + } + + rec.User.Attributes = attrs + rec.PasswordSalt, rec.PasswordHash = hashPassword(newPassword) + rec.User.UserStatus = driver.UserStatusConfirmed + rec.User.UserLastModifiedDate = m.now() + + return nil +} + +// GetUser returns the user an access token belongs to. +func (m *Mock) GetUser(_ context.Context, accessToken string) (*driver.User, error) { + rec, err := m.verifyAccessToken(accessToken) + if err != nil { + return nil, err + } + + out := copyUserRecord(rec).User + + return &out, nil +} + +// GlobalSignOut revokes every login of the access token's user. +func (m *Mock) GlobalSignOut(_ context.Context, accessToken string) error { + m.mu.Lock() + defer m.mu.Unlock() + + rec, err := m.verifyAccessToken(accessToken) + if err != nil { + return err + } + + m.revokeLogins(rec.PoolID, attrValue(rec.User.Attributes, attrSub)) + + return nil +} + +// AdminUserGlobalSignOut revokes every login of a user. +func (m *Mock) AdminUserGlobalSignOut(_ context.Context, userPoolID, username string) error { + m.mu.Lock() + defer m.mu.Unlock() + + pool, ok := m.userPools.Get(userPoolID) + if !ok { + return poolNotFound(userPoolID) + } + + _, rec, ok := m.resolveUser(&pool, username) + if !ok { + return userNotFound() + } + + m.revokeLogins(pool.ID, attrValue(rec.User.Attributes, attrSub)) + + return nil +} + +// RevokeToken revokes a refresh token and every token minted from the same +// login. An unknown refresh token is accepted silently, as RFC 7009 allows. +func (m *Mock) RevokeToken(_ context.Context, in driver.RevokeTokenInput) error { + m.mu.Lock() + defer m.mu.Unlock() + + client, err := m.clientByID(in.ClientID) + if err != nil { + return apiErr(driver.ExUnauthorized, errors.PermissionDenied, "Invalid client id "+in.ClientID) + } + + if client.ClientSecret != "" && !hmac.Equal([]byte(in.ClientSecret), []byte(client.ClientSecret)) { + return apiErr(driver.ExUnauthorized, errors.PermissionDenied, "Invalid client secret for client "+client.ClientID) + } + + if !client.EnableTokenRevocation { + return apiErr(driver.ExUnsupportedOperation, errors.FailedPrecondition, "Token revocation is not enabled for this app client.") + } + + if len(strings.Split(in.Token, ".")) != refreshSegments { + return apiErr(driver.ExUnsupportedTokenType, errors.InvalidArgument, "Unsupported token type. Only refresh tokens can be revoked.") + } + + l, ok := m.lookupRefresh(in.Token) + if !ok { + return nil + } + + if l.ClientID != client.ClientID { + return apiErr(driver.ExUnauthorized, errors.PermissionDenied, "The refresh token was not issued to client "+client.ClientID) + } + + l.Revoked = true + m.logins.Set(l.OriginJTI, l) + + return nil +} diff --git a/providers/aws/cognito/auth_test.go b/providers/aws/cognito/auth_test.go new file mode 100644 index 000000000..71c5db17c --- /dev/null +++ b/providers/aws/cognito/auth_test.go @@ -0,0 +1,654 @@ +package cognito + +import ( + "context" + "crypto" + "crypto/rsa" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "math/big" + "strings" + "testing" + "time" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// verifyJWT checks an RS256 token against the pool's published JWKS with +// crypto/rsa only, independent of the signer, and returns the header kid and +// the claims. +func verifyJWT(t *testing.T, m *Mock, poolID, token string) (string, map[string]any) { + t.Helper() + + _, set, err := m.SigningKeys(context.Background(), poolID) + requireNoError(t, err, "SigningKeys") + + parts := strings.Split(token, ".") + if len(parts) != 3 { + t.Fatalf("token has %d segments", len(parts)) + } + + var header struct { + Alg string `json:"alg"` + Kid string `json:"kid"` + } + + decodeSegment(t, parts[0], &header) + + if header.Alg != "RS256" { + t.Fatalf("alg = %q", header.Alg) + } + + var pub *rsa.PublicKey + + for _, k := range set.Keys { + if k.Kid == header.Kid { + n, _ := base64.RawURLEncoding.DecodeString(k.N) + e, _ := base64.RawURLEncoding.DecodeString(k.E) + pub = &rsa.PublicKey{N: new(big.Int).SetBytes(n), E: int(new(big.Int).SetBytes(e).Int64())} + } + } + + if pub == nil { + t.Fatalf("kid %q not in JWKS", header.Kid) + } + + sig, err := base64.RawURLEncoding.DecodeString(parts[2]) + requireNoError(t, err, "signature encoding") + + sum := sha256.Sum256([]byte(parts[0] + "." + parts[1])) + if err := rsa.VerifyPKCS1v15(pub, crypto.SHA256, sum[:], sig); err != nil { + t.Fatalf("signature does not verify against JWKS: %v", err) + } + + var claims map[string]any + + decodeSegment(t, parts[1], &claims) + + return header.Kid, claims +} + +func decodeSegment(t *testing.T, seg string, v any) { + t.Helper() + + raw, err := base64.RawURLEncoding.DecodeString(seg) + requireNoError(t, err, "segment encoding") + requireNoError(t, json.Unmarshal(raw, v), "segment json") +} + +// confirmedUser signs up and confirms a user with an email, returning its sub. +func confirmedUser(t *testing.T, m *Mock, poolID, clientID, username string) string { + t.Helper() + + ctx := context.Background() + + out, err := m.SignUp(ctx, signUpInput(clientID, username, testPassword, emailAttr(username+"@example.com"))) + requireNoError(t, err, "SignUp") + + requireNoError(t, m.ConfirmSignUp(ctx, driver.ConfirmSignUpInput{ + ClientUserInput: driver.ClientUserInput{ClientID: clientID, Username: username}, + ConfirmationCode: mustCode(t, m, poolID, username), + }), "ConfirmSignUp") + + return out.UserSub +} + +func passwordAuth(clientID, username, password string) driver.InitiateAuthInput { + return driver.InitiateAuthInput{ + ClientID: clientID, AuthFlow: driver.AuthFlowUserPassword, + AuthParameters: map[string]string{"USERNAME": username, "PASSWORD": password}, + } +} + +func TestInitiateAuthIssuesSignedTokens(t *testing.T) { + m, fc := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + sub := confirmedUser(t, m, pool.ID, client.ClientID, "alice") + + _, err := m.CreateGroup(ctx, driver.CreateGroupInput{UserPoolID: pool.ID, GroupName: "readers", Precedence: int32Ptr(5)}) + requireNoError(t, err, "CreateGroup") + _, err = m.CreateGroup(ctx, driver.CreateGroupInput{UserPoolID: pool.ID, GroupName: "admins", Precedence: int32Ptr(1)}) + requireNoError(t, err, "CreateGroup") + requireNoError(t, m.AdminAddUserToGroup(ctx, pool.ID, "alice", "readers"), "add") + requireNoError(t, m.AdminAddUserToGroup(ctx, pool.ID, "alice", "admins"), "add") + + res, err := m.InitiateAuth(ctx, passwordAuth(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "InitiateAuth") + + ar := res.AuthenticationResult + if res.ChallengeName != "" || ar == nil || ar.TokenType != "Bearer" || ar.ExpiresIn != 3600 || ar.RefreshToken == "" { + t.Fatalf("result = %+v / %+v", res, ar) + } + + if n := len(strings.Split(ar.RefreshToken, ".")); n != 5 { + t.Fatalf("refresh token has %d segments, want 5 (JWE shape)", n) + } + + iss := "https://cognito-idp.us-east-1.amazonaws.com/" + pool.ID + now := fc.Now().Unix() + + idKid, id := verifyJWT(t, m, pool.ID, ar.IDToken) + accessKid, access := verifyJWT(t, m, pool.ID, ar.AccessToken) + + if idKid == accessKid { + t.Fatal("ID and access tokens must be signed with different keys") + } + + wantID := map[string]any{ + "iss": iss, "sub": sub, "aud": client.ClientID, "token_use": "id", "cognito:username": "alice", + "email": "alice@example.com", "email_verified": true, + "auth_time": float64(now), "iat": float64(now), "exp": float64(now + 3600), + } + assertClaims(t, "id", id, wantID) + + wantAccess := map[string]any{ + "iss": iss, "sub": sub, "client_id": client.ClientID, "token_use": "access", "username": "alice", + "scope": "aws.cognito.signin.user.admin", "auth_time": float64(now), "exp": float64(now + 3600), + } + assertClaims(t, "access", access, wantAccess) + + for _, c := range []map[string]any{id, access} { + groups, _ := c["cognito:groups"].([]any) + if len(groups) != 2 || groups[0] != "admins" || groups[1] != "readers" { + t.Fatalf("cognito:groups = %v, want precedence order [admins readers]", c["cognito:groups"]) + } + + if c["origin_jti"] == "" || c["jti"] == "" || c["event_id"] == "" { + t.Fatalf("missing jti/origin_jti/event_id in %v", c) + } + } + + if id["origin_jti"] != access["origin_jti"] || id["jti"] == access["jti"] { + t.Fatal("ID and access tokens share origin_jti but not jti") + } + + if _, ok := access["aud"]; ok { + t.Fatal("access token carries no aud") + } + + u, err := m.GetUser(ctx, ar.AccessToken) + requireNoError(t, err, "GetUser") + + if u.Username != "alice" || attrValue(u.Attributes, "sub") != sub { + t.Fatalf("GetUser = %+v", u) + } + + _, err = m.GetUser(ctx, ar.IDToken) + assertException(t, err, driver.ExNotAuthorized, "Invalid Access Token") +} + +func assertClaims(t *testing.T, name string, got, want map[string]any) { + t.Helper() + + for k, v := range want { + if got[k] != v { + t.Fatalf("%s token %s = %v (%T), want %v", name, k, got[k], got[k], v) + } + } +} + +func TestInitiateAuthErrors(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + confirmedUser(t, m, pool.ID, client.ClientID, "alice") + + _, err := m.SignUp(ctx, signUpInput(client.ClientID, "pending", testPassword)) + requireNoError(t, err, "SignUp pending") + + defaultClient := mustCreateClient(t, m, pool.ID, false) + + cases := []struct { + name string + in driver.InitiateAuthInput + exception string + msg string + }{ + {"wrong password", passwordAuth(client.ClientID, "alice", "Wr0ng!pass"), driver.ExNotAuthorized, "Incorrect username or password."}, + {"unknown user legacy", passwordAuth(client.ClientID, "ghost", testPassword), driver.ExUserNotFound, "User does not exist."}, + {"unconfirmed", passwordAuth(client.ClientID, "pending", testPassword), driver.ExUserNotConfirmed, "User is not confirmed."}, + {"flow not enabled", passwordAuth(defaultClient.ClientID, "alice", testPassword), + driver.ExInvalidParameter, "USER_PASSWORD_AUTH flow not enabled for this client"}, + {"admin flow on public api", driver.InitiateAuthInput{ClientID: client.ClientID, AuthFlow: driver.AuthFlowAdminUserPassword}, + driver.ExInvalidParameter, "Initiate Auth method not supported."}, + {"unknown client", passwordAuth("nosuchclient", "alice", testPassword), driver.ExResourceNotFound, ""}, + {"missing password", driver.InitiateAuthInput{ + ClientID: client.ClientID, AuthFlow: driver.AuthFlowUserPassword, AuthParameters: map[string]string{"USERNAME": "alice"}, + }, driver.ExInvalidParameter, "Missing required parameter PASSWORD"}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + _, err := m.InitiateAuth(ctx, tc.in) + assertException(t, err, tc.exception, tc.msg) + }) + } + + requireNoError(t, m.AdminDisableUser(ctx, pool.ID, "alice"), "AdminDisableUser") + + _, err = m.InitiateAuth(ctx, passwordAuth(client.ClientID, "alice", testPassword)) + assertException(t, err, driver.ExNotAuthorized, "User is disabled.") +} + +func TestPreventUserExistenceErrorsEnabled(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + + client, err := m.CreateUserPoolClient(ctx, driver.CreateUserPoolClientInput{ + UserPoolID: pool.ID, ClientName: "hidden", ExplicitAuthFlows: passwordFlows, PreventUserExistenceErrors: "ENABLED", + }) + requireNoError(t, err, "CreateUserPoolClient") + + _, err = m.InitiateAuth(ctx, passwordAuth(client.ClientID, "ghost", testPassword)) + assertException(t, err, driver.ExNotAuthorized, "Incorrect username or password.") +} + +func TestInitiateAuthSecretHash(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, true, passwordFlows...) + + in := signUpInput(client.ClientID, "alice", testPassword) + in.SecretHash = testSecretHash(client.ClientSecret, "alice", client.ClientID) + _, err := m.SignUp(ctx, in) + requireNoError(t, err, "SignUp") + requireNoError(t, m.AdminConfirmSignUp(ctx, pool.ID, "alice"), "AdminConfirmSignUp") + + auth := passwordAuth(client.ClientID, "alice", testPassword) + + _, err = m.InitiateAuth(ctx, auth) + assertException(t, err, driver.ExNotAuthorized, "Client "+client.ClientID+" is configured for secret but secret was not received") + + auth.AuthParameters["SECRET_HASH"] = "bm9wZQ==" + _, err = m.InitiateAuth(ctx, auth) + assertException(t, err, driver.ExNotAuthorized, "Unable to verify secret hash for client "+client.ClientID) + + auth.AuthParameters["SECRET_HASH"] = testSecretHash(client.ClientSecret, "alice", client.ClientID) + res, err := m.InitiateAuth(ctx, auth) + requireNoError(t, err, "InitiateAuth with SECRET_HASH") + + refresh := driver.InitiateAuthInput{ + ClientID: client.ClientID, AuthFlow: driver.AuthFlowRefreshTokenAuth, + AuthParameters: map[string]string{"REFRESH_TOKEN": res.AuthenticationResult.RefreshToken}, + } + + _, err = m.InitiateAuth(ctx, refresh) + assertException(t, err, driver.ExNotAuthorized, "Client "+client.ClientID+" is configured for secret but secret was not received") + + refresh.AuthParameters["SECRET_HASH"] = testSecretHash(client.ClientSecret, "alice", client.ClientID) + _, err = m.InitiateAuth(ctx, refresh) + requireNoError(t, err, "refresh with SECRET_HASH") +} + +func TestRefreshTokenAuth(t *testing.T) { + m, fc := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + confirmedUser(t, m, pool.ID, client.ClientID, "alice") + + first, err := m.InitiateAuth(ctx, passwordAuth(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "InitiateAuth") + + authTime := fc.Now().Unix() + + fc.Advance(10 * time.Minute) + + for _, flow := range []string{driver.AuthFlowRefreshTokenAuth, driver.AuthFlowRefreshToken} { + res, err := m.InitiateAuth(ctx, driver.InitiateAuthInput{ + ClientID: client.ClientID, AuthFlow: flow, + AuthParameters: map[string]string{"REFRESH_TOKEN": first.AuthenticationResult.RefreshToken}, + }) + requireNoError(t, err, "refresh "+flow) + + ar := res.AuthenticationResult + if ar.RefreshToken != "" || ar.AccessToken == "" || ar.IDToken == "" { + t.Fatalf("refresh result = %+v, want new access/id and no refresh token", ar) + } + + _, claims := verifyJWT(t, m, pool.ID, ar.AccessToken) + _, orig := verifyJWT(t, m, pool.ID, first.AuthenticationResult.AccessToken) + + if claims["auth_time"] != float64(authTime) || claims["origin_jti"] != orig["origin_jti"] { + t.Fatalf("refreshed claims keep auth_time and origin_jti: %v vs %v", claims, orig) + } + } + + _, err = m.InitiateAuth(ctx, driver.InitiateAuthInput{ + ClientID: client.ClientID, AuthFlow: driver.AuthFlowRefreshTokenAuth, + AuthParameters: map[string]string{"REFRESH_TOKEN": "garbage"}, + }) + assertException(t, err, driver.ExNotAuthorized, "Invalid Refresh Token") + + other := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + _, err = m.InitiateAuth(ctx, driver.InitiateAuthInput{ + ClientID: other.ClientID, AuthFlow: driver.AuthFlowRefreshTokenAuth, + AuthParameters: map[string]string{"REFRESH_TOKEN": first.AuthenticationResult.RefreshToken}, + }) + assertException(t, err, driver.ExNotAuthorized, "Invalid Refresh Token") + + fc.Advance(31 * 24 * time.Hour) + + _, err = m.InitiateAuth(ctx, driver.InitiateAuthInput{ + ClientID: client.ClientID, AuthFlow: driver.AuthFlowRefreshTokenAuth, + AuthParameters: map[string]string{"REFRESH_TOKEN": first.AuthenticationResult.RefreshToken}, + }) + assertException(t, err, driver.ExNotAuthorized, "Refresh Token has expired") +} + +func TestAccessTokenExpiry(t *testing.T) { + m, fc := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + + client, err := m.CreateUserPoolClient(ctx, driver.CreateUserPoolClientInput{ + UserPoolID: pool.ID, ClientName: "short", ExplicitAuthFlows: passwordFlows, + AccessTokenValidity: int32Ptr(5), IDTokenValidity: int32Ptr(10), + TokenValidityUnits: &driver.TokenValidityUnits{AccessToken: "minutes", IDToken: "minutes", RefreshToken: "days"}, + }) + requireNoError(t, err, "CreateUserPoolClient") + confirmedUser(t, m, pool.ID, client.ClientID, "alice") + + res, err := m.InitiateAuth(ctx, passwordAuth(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "InitiateAuth") + + if res.AuthenticationResult.ExpiresIn != 300 { + t.Fatalf("ExpiresIn = %d, want 300", res.AuthenticationResult.ExpiresIn) + } + + _, id := verifyJWT(t, m, pool.ID, res.AuthenticationResult.IDToken) + if id["exp"] != float64(fc.Now().Unix()+600) { + t.Fatalf("id exp = %v", id["exp"]) + } + + fc.Advance(6 * time.Minute) + + _, err = m.GetUser(ctx, res.AuthenticationResult.AccessToken) + assertException(t, err, driver.ExNotAuthorized, "Access Token has expired") +} + +func TestAccessTokenTamperRejected(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + confirmedUser(t, m, pool.ID, client.ClientID, "alice") + confirmedUser(t, m, pool.ID, client.ClientID, "bob") + + res, err := m.InitiateAuth(ctx, passwordAuth(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "InitiateAuth") + + parts := strings.Split(res.AuthenticationResult.AccessToken, ".") + payload, _ := base64.RawURLEncoding.DecodeString(parts[1]) + forged := strings.Replace(string(payload), `"username":"alice"`, `"username":"bob"`, 1) + parts[1] = base64.RawURLEncoding.EncodeToString([]byte(forged)) + + for _, tok := range []string{strings.Join(parts, "."), "not-a-jwt", ""} { + _, err = m.GetUser(ctx, tok) + assertException(t, err, driver.ExNotAuthorized, "Invalid Access Token") + } +} + +func TestNewPasswordRequiredChallenge(t *testing.T) { + m, fc := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + + _, err := m.AdminCreateUser(ctx, driver.AdminCreateUserInput{ + UserPoolID: pool.ID, Username: "temp", TemporaryPassword: "Temp0rary!", MessageAction: driver.MessageActionSuppress, + UserAttributes: []driver.Attribute{emailAttr("temp@example.com")}, + }) + requireNoError(t, err, "AdminCreateUser") + + res, err := m.AdminInitiateAuth(ctx, driver.InitiateAuthInput{ + UserPoolID: pool.ID, ClientID: client.ClientID, AuthFlow: driver.AuthFlowAdminUserPassword, + AuthParameters: map[string]string{"USERNAME": "temp", "PASSWORD": "Temp0rary!"}, + }) + requireNoError(t, err, "AdminInitiateAuth") + + if res.ChallengeName != driver.ChallengeNewPasswordRequired || res.Session == "" || res.AuthenticationResult != nil { + t.Fatalf("result = %+v", res) + } + + if res.ChallengeParameters["USER_ID_FOR_SRP"] != "temp" || res.ChallengeParameters["requiredAttributes"] != "[]" { + t.Fatalf("challenge parameters = %v", res.ChallengeParameters) + } + + var userAttrs map[string]string + requireNoError(t, json.Unmarshal([]byte(res.ChallengeParameters["userAttributes"]), &userAttrs), "userAttributes json") + + if userAttrs["email"] != "temp@example.com" { + t.Fatalf("userAttributes = %v", userAttrs) + } + + respond := driver.RespondToAuthChallengeInput{ + ClientID: client.ClientID, ChallengeName: driver.ChallengeNewPasswordRequired, Session: res.Session, + ChallengeResponses: map[string]string{"USERNAME": "temp", "NEW_PASSWORD": "short"}, + } + + _, err = m.RespondToAuthChallenge(ctx, respond) + assertException(t, err, driver.ExInvalidPassword, "Password did not conform with policy: Password not long enough") + + _, err = m.RespondToAuthChallenge(ctx, driver.RespondToAuthChallengeInput{ + ClientID: client.ClientID, ChallengeName: driver.ChallengeNewPasswordRequired, Session: "bogus", + ChallengeResponses: map[string]string{"USERNAME": "temp", "NEW_PASSWORD": "N3wPassword!"}, + }) + assertException(t, err, driver.ExNotAuthorized, "Invalid session for the user.") + + respond.ChallengeResponses["NEW_PASSWORD"] = "N3wPassword!" + out, err := m.RespondToAuthChallenge(ctx, respond) + requireNoError(t, err, "RespondToAuthChallenge") + + if out.AuthenticationResult == nil || out.AuthenticationResult.RefreshToken == "" { + t.Fatalf("respond result = %+v", out) + } + + u, _ := m.AdminGetUser(ctx, pool.ID, "temp") + if u.UserStatus != driver.UserStatusConfirmed { + t.Fatalf("status = %s, want CONFIRMED", u.UserStatus) + } + + _, err = m.InitiateAuth(ctx, passwordAuth(client.ClientID, "temp", "N3wPassword!")) + requireNoError(t, err, "sign in with the new password") + + // A session cannot be reused, and expires after AuthSessionValidity minutes. + _, err = m.RespondToAuthChallenge(ctx, respond) + assertException(t, err, driver.ExNotAuthorized, "Invalid session for the user.") + + requireNoError(t, m.AdminSetUserPassword(ctx, pool.ID, "temp", "Temp0rary!", false), "AdminSetUserPassword") + + res, err = m.InitiateAuth(ctx, passwordAuth(client.ClientID, "temp", "Temp0rary!")) + requireNoError(t, err, "InitiateAuth after reset") + + fc.Advance(4 * time.Minute) + + respond.Session = res.Session + _, err = m.RespondToAuthChallenge(ctx, respond) + assertException(t, err, driver.ExNotAuthorized, "Invalid session for the user, session is expired.") +} + +func TestAdminAuthFlows(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, "ALLOW_USER_PASSWORD_AUTH", "ALLOW_REFRESH_TOKEN_AUTH") + adminClient := mustCreateClient(t, m, pool.ID, false, "ADMIN_NO_SRP_AUTH") + confirmedUser(t, m, pool.ID, client.ClientID, "alice") + + admin := func(clientID, flow string) driver.InitiateAuthInput { + return driver.InitiateAuthInput{ + UserPoolID: pool.ID, ClientID: clientID, AuthFlow: flow, + AuthParameters: map[string]string{"USERNAME": "alice", "PASSWORD": testPassword}, + } + } + + _, err := m.AdminInitiateAuth(ctx, admin(client.ClientID, driver.AuthFlowAdminUserPassword)) + assertException(t, err, driver.ExInvalidParameter, "Auth flow not enabled for this client") + + for _, flow := range []string{driver.AuthFlowAdminUserPassword, driver.AuthFlowAdminNoSRP} { + _, err = m.AdminInitiateAuth(ctx, admin(adminClient.ClientID, flow)) + requireNoError(t, err, "AdminInitiateAuth "+flow) + } + + _, err = m.AdminInitiateAuth(ctx, admin(adminClient.ClientID, driver.AuthFlowUserPassword)) + assertException(t, err, driver.ExInvalidParameter, "Initiate Auth method not supported.") + + in := admin(adminClient.ClientID, driver.AuthFlowAdminUserPassword) + in.UserPoolID = "us-east-1_other0000" + _, err = m.AdminInitiateAuth(ctx, in) + assertException(t, err, driver.ExResourceNotFound, "") +} + +func TestGlobalSignOutRevokesTokens(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + confirmedUser(t, m, pool.ID, client.ClientID, "alice") + + a, err := m.InitiateAuth(ctx, passwordAuth(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "sign-in 1") + b, err := m.InitiateAuth(ctx, passwordAuth(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "sign-in 2") + + requireNoError(t, m.GlobalSignOut(ctx, a.AuthenticationResult.AccessToken), "GlobalSignOut") + + for _, tok := range []string{a.AuthenticationResult.AccessToken, b.AuthenticationResult.AccessToken} { + _, err = m.GetUser(ctx, tok) + assertException(t, err, driver.ExNotAuthorized, "Access Token has been revoked") + } + + _, err = m.InitiateAuth(ctx, driver.InitiateAuthInput{ + ClientID: client.ClientID, AuthFlow: driver.AuthFlowRefreshTokenAuth, + AuthParameters: map[string]string{"REFRESH_TOKEN": b.AuthenticationResult.RefreshToken}, + }) + assertException(t, err, driver.ExNotAuthorized, "Refresh Token has been revoked") + + c, err := m.InitiateAuth(ctx, passwordAuth(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "sign-in after sign-out") + + _, err = m.GetUser(ctx, c.AuthenticationResult.AccessToken) + requireNoError(t, err, "a fresh sign-in works") + + requireNoError(t, m.AdminUserGlobalSignOut(ctx, pool.ID, "alice"), "AdminUserGlobalSignOut") + + _, err = m.GetUser(ctx, c.AuthenticationResult.AccessToken) + assertException(t, err, driver.ExNotAuthorized, "Access Token has been revoked") +} + +func TestRevokeToken(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + confirmedUser(t, m, pool.ID, client.ClientID, "alice") + + a, err := m.InitiateAuth(ctx, passwordAuth(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "sign-in 1") + b, err := m.InitiateAuth(ctx, passwordAuth(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "sign-in 2") + + err = m.RevokeToken(ctx, driver.RevokeTokenInput{Token: a.AuthenticationResult.AccessToken, ClientID: client.ClientID}) + assertException(t, err, driver.ExUnsupportedTokenType, "") + + requireNoError(t, m.RevokeToken(ctx, driver.RevokeTokenInput{ + Token: a.AuthenticationResult.RefreshToken, ClientID: client.ClientID, + }), "RevokeToken") + + _, err = m.GetUser(ctx, a.AuthenticationResult.AccessToken) + assertException(t, err, driver.ExNotAuthorized, "Access Token has been revoked") + + _, err = m.GetUser(ctx, b.AuthenticationResult.AccessToken) + requireNoError(t, err, "the other session is untouched") + + noRevoke := false + locked, err := m.CreateUserPoolClient(ctx, driver.CreateUserPoolClientInput{ + UserPoolID: pool.ID, ClientName: "norevoke", ExplicitAuthFlows: passwordFlows, EnableTokenRevocation: &noRevoke, + }) + requireNoError(t, err, "CreateUserPoolClient") + + c, err := m.InitiateAuth(ctx, passwordAuth(locked.ClientID, "alice", testPassword)) + requireNoError(t, err, "sign-in on norevoke client") + + err = m.RevokeToken(ctx, driver.RevokeTokenInput{Token: c.AuthenticationResult.RefreshToken, ClientID: locked.ClientID}) + assertException(t, err, driver.ExUnsupportedOperation, "") +} + +func TestSigningKeysSurviveSnapshot(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + confirmedUser(t, m, pool.ID, client.ClientID, "alice") + + res, err := m.InitiateAuth(ctx, passwordAuth(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "InitiateAuth") + + issuer, before, err := m.SigningKeys(ctx, pool.ID) + requireNoError(t, err, "SigningKeys") + + if issuer != "https://cognito-idp.us-east-1.amazonaws.com/"+pool.ID || len(before.Keys) != 2 { + t.Fatalf("issuer=%q keys=%d", issuer, len(before.Keys)) + } + + snap, err := m.Snapshot(ctx, false) + requireNoError(t, err, "Snapshot") + + if strings.Contains(string(snap), testPassword) { + t.Fatal("snapshot carries a plaintext password") + } + + restored, _ := newClockMock(t) + requireNoError(t, restored.Restore(ctx, snap), "Restore") + + _, after, err := restored.SigningKeys(ctx, pool.ID) + requireNoError(t, err, "SigningKeys after restore") + + b1, _ := json.Marshal(before) + b2, _ := json.Marshal(after) + + if string(b1) != string(b2) { + t.Fatal("JWKS changed across snapshot/restore") + } + + _, err = restored.GetUser(ctx, res.AuthenticationResult.AccessToken) + requireNoError(t, err, "access token still valid after restore") + + _, err = restored.InitiateAuth(ctx, driver.InitiateAuthInput{ + ClientID: client.ClientID, AuthFlow: driver.AuthFlowRefreshTokenAuth, + AuthParameters: map[string]string{"REFRESH_TOKEN": res.AuthenticationResult.RefreshToken}, + }) + requireNoError(t, err, "refresh token still valid after restore") + + _, _, err = m.SigningKeys(ctx, "us-east-1_missing00") + assertException(t, err, driver.ExResourceNotFound, "") +} + +func TestDeleteUserInvalidatesTokens(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + confirmedUser(t, m, pool.ID, client.ClientID, "alice") + + res, err := m.InitiateAuth(ctx, passwordAuth(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "InitiateAuth") + + requireNoError(t, m.AdminDeleteUser(ctx, pool.ID, "alice"), "AdminDeleteUser") + + _, err = m.GetUser(ctx, res.AuthenticationResult.AccessToken) + assertException(t, err, driver.ExNotAuthorized, "Access Token has been revoked") + + // A new user with the same name must not inherit the old user's tokens. + confirmedUser(t, m, pool.ID, client.ClientID, "alice") + + _, err = m.GetUser(ctx, res.AuthenticationResult.AccessToken) + assertException(t, err, driver.ExNotAuthorized, "Access Token has been revoked") +} diff --git a/providers/aws/cognito/cognito.go b/providers/aws/cognito/cognito.go index 65c563218..644cd7deb 100644 --- a/providers/aws/cognito/cognito.go +++ b/providers/aws/cognito/cognito.go @@ -1,8 +1,8 @@ // 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. -// -// Sign-up, sign-in and token issuance are not modeled yet. +// tagging, pool users with the admin user-management operations, groups, self +// sign-up with confirmation codes, and password sign-in that issues RS256 +// tokens a real JWT library verifies against the pool's JWKS. package cognito import ( @@ -15,8 +15,13 @@ import ( "github.com/stackshy/cloudemu/v2/services/cognito/driver" ) -// Compile-time check that Mock implements driver.Cognito. -var _ driver.Cognito = (*Mock)(nil) +// Compile-time checks that Mock implements driver.Cognito and the non-API +// interfaces the wire layer uses. +var ( + _ driver.Cognito = (*Mock)(nil) + _ driver.KeySetProvider = (*Mock)(nil) + _ driver.CodeInspector = (*Mock)(nil) +) // clientKeySep separates the user-pool id and client id in the clients store. const clientKeySep = "/" @@ -26,11 +31,14 @@ const clientKeySep = "/" type Mock struct { // userPools is keyed by pool id; clients is keyed by "/"; // domains is keyed by the domain string; users is keyed by - // "/". + // "/"; groups is keyed by "/"; logins + // holds one record per sign-in, keyed by its origin_jti. userPools *memstore.Store[driver.UserPool] clients *memstore.Store[driver.UserPoolClient] domains *memstore.Store[driver.UserPoolDomain] users *memstore.Store[userRecord] + groups *memstore.Store[driver.Group] + logins *memstore.Store[loginRecord] // mu serializes compound read-modify-write mutations (pool update, cascading // pool delete, user changes) that span more than one store operation. @@ -40,6 +48,15 @@ type Mock struct { tagsMu sync.RWMutex tags map[string]map[string]string + // keysMu guards the per-pool signing keys, created on first use. + keysMu sync.Mutex + keys map[string]*poolKeys + + // sessionsMu guards the challenge sessions. They last minutes, so they + // are not persisted. + sessionsMu sync.Mutex + sessions map[string]challengeSession + opts *config.Options } @@ -50,7 +67,11 @@ func New(opts *config.Options) *Mock { clients: memstore.New[driver.UserPoolClient](), domains: memstore.New[driver.UserPoolDomain](), users: memstore.New[userRecord](), + groups: memstore.New[driver.Group](), + logins: memstore.New[loginRecord](), tags: map[string]map[string]string{}, + keys: map[string]*poolKeys{}, + sessions: map[string]challengeSession{}, opts: opts, } } diff --git a/providers/aws/cognito/groups.go b/providers/aws/cognito/groups.go new file mode 100644 index 000000000..d81c8e512 --- /dev/null +++ b/providers/aws/cognito/groups.go @@ -0,0 +1,304 @@ +package cognito + +import ( + "context" + "slices" + "sort" + "strings" + + "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// maxGroupNameLen is the GroupNameType length ceiling. +const maxGroupNameLen = 128 + +func groupKey(poolID, name string) string { return poolID + clientKeySep + name } + +func groupNotFound() error { return resourceNotFound("Group not found.") } + +//nolint:gocritic // hugeParam: value signature required by the func(V) V copy callback +func copyGroup(in driver.Group) driver.Group { + out := in + out.Precedence = copyInt32Ptr(in.Precedence) + + return out +} + +// checkLimit rejects a list Limit above the cognito-idp ceiling of 60. +func checkLimit(n int32) error { + if n > defaultPageSize { + return invalidParameter("1 validation error detected: Value '%d' at 'limit' failed to satisfy constraint: "+ + "Member must have value less than or equal to %d", n, defaultPageSize) + } + + return nil +} + +func checkGroupInput(name string, precedence *int32) error { + if name == "" || len(name) > maxGroupNameLen { + return invalidParameter("1 validation error detected: Value at 'groupName' failed to satisfy constraint: " + + "Member must have length between 1 and 128") + } + + if precedence != nil && *precedence < 0 { + return invalidParameter("1 validation error detected: Value '%d' at 'precedence' failed to satisfy constraint: "+ + "Member must have value greater than or equal to 0", *precedence) + } + + return nil +} + +// CreateGroup creates a group in a user pool. +func (m *Mock) CreateGroup(_ context.Context, in driver.CreateGroupInput) (*driver.Group, error) { + if err := checkGroupInput(in.GroupName, in.Precedence); err != nil { + return nil, err + } + + m.mu.Lock() + defer m.mu.Unlock() + + if !m.userPools.Has(in.UserPoolID) { + return nil, poolNotFound(in.UserPoolID) + } + + key := groupKey(in.UserPoolID, in.GroupName) + if m.groups.Has(key) { + return nil, &driver.APIError{ + Exception: driver.ExGroupExists, + Err: errors.New(errors.AlreadyExists, "A group with the name "+in.GroupName+" already exists."), + } + } + + now := m.now() + g := driver.Group{ + GroupName: in.GroupName, + UserPoolID: in.UserPoolID, + Description: in.Description, + RoleARN: in.RoleARN, + Precedence: copyInt32Ptr(in.Precedence), + CreationDate: now, + LastModifiedDate: now, + } + m.groups.Set(key, copyGroup(g)) + + return &g, nil +} + +// GetGroup returns one group. +func (m *Mock) GetGroup(_ context.Context, userPoolID, groupName string) (*driver.Group, error) { + if !m.userPools.Has(userPoolID) { + return nil, poolNotFound(userPoolID) + } + + g, ok := m.groups.Get(groupKey(userPoolID, groupName)) + if !ok { + return nil, groupNotFound() + } + + out := copyGroup(g) + + return &out, nil +} + +// UpdateGroup changes the description, role and precedence the input sets. +func (m *Mock) UpdateGroup(_ context.Context, in driver.UpdateGroupInput) (*driver.Group, error) { + if err := checkGroupInput(in.GroupName, in.Precedence); err != nil { + return nil, err + } + + m.mu.Lock() + defer m.mu.Unlock() + + if !m.userPools.Has(in.UserPoolID) { + return nil, poolNotFound(in.UserPoolID) + } + + key := groupKey(in.UserPoolID, in.GroupName) + + g, ok := m.groups.Get(key) + if !ok { + return nil, groupNotFound() + } + + g = copyGroup(g) + + if in.Description != nil { + g.Description = *in.Description + } + + if in.RoleARN != nil { + g.RoleARN = *in.RoleARN + } + + if in.Precedence != nil { + g.Precedence = copyInt32Ptr(in.Precedence) + } + + g.LastModifiedDate = m.now() + m.groups.Set(key, copyGroup(g)) + + return &g, nil +} + +// DeleteGroup removes a group and its memberships. +func (m *Mock) DeleteGroup(_ context.Context, userPoolID, groupName string) error { + m.mu.Lock() + defer m.mu.Unlock() + + if !m.userPools.Has(userPoolID) { + return poolNotFound(userPoolID) + } + + if !m.groups.Delete(groupKey(userPoolID, groupName)) { + return groupNotFound() + } + + for _, key := range m.poolUserKeys(userPoolID) { + rec, ok := m.users.Get(key) + if !ok || !slices.Contains(rec.Groups, groupName) { + continue + } + + rec = copyUserRecord(rec) + rec.Groups = slices.DeleteFunc(rec.Groups, func(g string) bool { return g == groupName }) + m.users.Set(key, rec) + } + + return nil +} + +// ListGroups returns a page of a pool's groups sorted by name. +func (m *Mock) ListGroups(_ context.Context, userPoolID string, page driver.Pagination) ([]driver.Group, string, error) { + if err := checkLimit(page.MaxResults); err != nil { + return nil, "", err + } + + if !m.userPools.Has(userPoolID) { + return nil, "", poolNotFound(userPoolID) + } + + return paginate(m.poolGroups(userPoolID), page) +} + +// AdminAddUserToGroup adds a user to a group. Adding a member again is a no-op. +func (m *Mock) AdminAddUserToGroup(_ context.Context, userPoolID, username, groupName string) error { + return m.changeMembership(userPoolID, username, groupName, func(groups []string) []string { + if slices.Contains(groups, groupName) { + return groups + } + + return append(groups, groupName) + }) +} + +// AdminRemoveUserFromGroup removes a user from a group. +func (m *Mock) AdminRemoveUserFromGroup(_ context.Context, userPoolID, username, groupName string) error { + return m.changeMembership(userPoolID, username, groupName, func(groups []string) []string { + return slices.DeleteFunc(groups, func(g string) bool { return g == groupName }) + }) +} + +func (m *Mock) changeMembership(userPoolID, username, groupName string, fn func([]string) []string) error { + return m.updateUser(userPoolID, username, func(pool *driver.UserPool, _ string, rec *userRecord) error { + if !m.groups.Has(groupKey(pool.ID, groupName)) { + return groupNotFound() + } + + rec.Groups = fn(rec.Groups) + + return nil + }) +} + +// AdminListGroupsForUser returns a page of the groups a user belongs to. +func (m *Mock) AdminListGroupsForUser( + _ context.Context, userPoolID, username string, page driver.Pagination, +) ([]driver.Group, string, error) { + if err := checkLimit(page.MaxResults); err != nil { + return nil, "", err + } + + pool, ok := m.userPools.Get(userPoolID) + if !ok { + return nil, "", poolNotFound(userPoolID) + } + + _, rec, ok := m.resolveUser(&pool, username) + if !ok { + return nil, "", userNotFound() + } + + return paginate(m.userGroups(userPoolID, rec.Groups), page) +} + +// ListUsersInGroup returns a page of a group's members sorted by username. +func (m *Mock) ListUsersInGroup(_ context.Context, userPoolID, groupName string, page driver.Pagination) ([]driver.User, string, error) { + if err := checkLimit(page.MaxResults); err != nil { + return nil, "", err + } + + if !m.userPools.Has(userPoolID) { + return nil, "", poolNotFound(userPoolID) + } + + if !m.groups.Has(groupKey(userPoolID, groupName)) { + return nil, "", groupNotFound() + } + + var members []driver.User + + users := m.poolUsers(userPoolID) + for i := range users { + if slices.Contains(users[i].Groups, groupName) { + members = append(members, copyUserRecord(users[i]).User) + } + } + + return paginate(members, page) +} + +// poolGroups returns a pool's groups sorted by name. +func (m *Mock) poolGroups(poolID string) []driver.Group { + prefix := groupKey(poolID, "") + + keys := slices.DeleteFunc(m.groups.Keys(), func(k string) bool { return !strings.HasPrefix(k, prefix) }) + sort.Strings(keys) + + out := make([]driver.Group, 0, len(keys)) + + for _, k := range keys { + if g, ok := m.groups.Get(k); ok { + out = append(out, copyGroup(g)) + } + } + + return out +} + +// userGroups resolves a user's group names to groups sorted by name, skipping +// any that no longer exist. +func (m *Mock) userGroups(poolID string, names []string) []driver.Group { + out := make([]driver.Group, 0, len(names)) + + for _, name := range names { + if g, ok := m.groups.Get(groupKey(poolID, name)); ok { + out = append(out, copyGroup(g)) + } + } + + sort.Slice(out, func(i, j int) bool { return out[i].GroupName < out[j].GroupName }) + + return out +} + +// deletePoolGroups removes every group of a pool. +func (m *Mock) deletePoolGroups(poolID string) { + prefix := groupKey(poolID, "") + + for _, k := range m.groups.Keys() { + if strings.HasPrefix(k, prefix) { + m.groups.Delete(k) + } + } +} diff --git a/providers/aws/cognito/groups_test.go b/providers/aws/cognito/groups_test.go new file mode 100644 index 000000000..29ecb9398 --- /dev/null +++ b/providers/aws/cognito/groups_test.go @@ -0,0 +1,137 @@ +package cognito + +import ( + "context" + "testing" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +func int32Ptr(v int32) *int32 { return &v } + +func TestGroupLifecycle(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "groups") + + g, err := m.CreateGroup(ctx, driver.CreateGroupInput{ + UserPoolID: pool.ID, GroupName: "admins", Description: "Admins", + RoleARN: "arn:aws:iam::123456789012:role/admin", Precedence: int32Ptr(1), + }) + requireNoError(t, err, "CreateGroup") + + if g.GroupName != "admins" || g.UserPoolID != pool.ID || *g.Precedence != 1 || g.CreationDate.IsZero() { + t.Fatalf("group = %+v", g) + } + + _, err = m.CreateGroup(ctx, driver.CreateGroupInput{UserPoolID: pool.ID, GroupName: "admins"}) + assertException(t, err, driver.ExGroupExists, "A group with the name admins already exists.") + + _, err = m.GetGroup(ctx, pool.ID, "nope") + assertException(t, err, driver.ExResourceNotFound, "Group not found.") + + desc := "Administrators" + up, err := m.UpdateGroup(ctx, driver.UpdateGroupInput{UserPoolID: pool.ID, GroupName: "admins", Description: &desc}) + requireNoError(t, err, "UpdateGroup") + + if up.Description != desc || up.RoleARN == "" || *up.Precedence != 1 { + t.Fatalf("updated group = %+v, want description changed and the rest kept", up) + } + + _, err = m.CreateGroup(ctx, driver.CreateGroupInput{UserPoolID: pool.ID, GroupName: "readers"}) + requireNoError(t, err, "CreateGroup readers") + + groups, next, err := m.ListGroups(ctx, pool.ID, driver.Pagination{MaxResults: 1}) + requireNoError(t, err, "ListGroups") + + if len(groups) != 1 || groups[0].GroupName != "admins" || next == "" { + t.Fatalf("page 1 = %+v next=%q", groups, next) + } + + groups, next, err = m.ListGroups(ctx, pool.ID, driver.Pagination{MaxResults: 1, NextToken: next}) + requireNoError(t, err, "ListGroups page 2") + + if len(groups) != 1 || groups[0].GroupName != "readers" || next != "" { + t.Fatalf("page 2 = %+v next=%q", groups, next) + } + + requireNoError(t, m.DeleteGroup(ctx, pool.ID, "readers"), "DeleteGroup") + assertException(t, m.DeleteGroup(ctx, pool.ID, "readers"), driver.ExResourceNotFound, "Group not found.") + + _, err = m.CreateGroup(ctx, driver.CreateGroupInput{UserPoolID: "us-east-1_missing00", GroupName: "x"}) + assertException(t, err, driver.ExResourceNotFound, "") +} + +func TestGroupMembership(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "members") + mustCreateUser(t, m, pool.ID, "alice") + mustCreateUser(t, m, pool.ID, "bob") + + for _, name := range []string{"admins", "readers"} { + _, err := m.CreateGroup(ctx, driver.CreateGroupInput{UserPoolID: pool.ID, GroupName: name}) + requireNoError(t, err, "CreateGroup "+name) + } + + requireNoError(t, m.AdminAddUserToGroup(ctx, pool.ID, "alice", "admins"), "add alice admins") + requireNoError(t, m.AdminAddUserToGroup(ctx, pool.ID, "alice", "admins"), "re-add is a no-op") + requireNoError(t, m.AdminAddUserToGroup(ctx, pool.ID, "alice", "readers"), "add alice readers") + requireNoError(t, m.AdminAddUserToGroup(ctx, pool.ID, "bob", "readers"), "add bob readers") + + assertException(t, m.AdminAddUserToGroup(ctx, pool.ID, "carol", "admins"), driver.ExUserNotFound, "User does not exist.") + assertException(t, m.AdminAddUserToGroup(ctx, pool.ID, "alice", "ghosts"), driver.ExResourceNotFound, "Group not found.") + + groups, _, err := m.AdminListGroupsForUser(ctx, pool.ID, "alice", driver.Pagination{}) + requireNoError(t, err, "AdminListGroupsForUser") + + if len(groups) != 2 { + t.Fatalf("alice groups = %+v", groups) + } + + users, _, err := m.ListUsersInGroup(ctx, pool.ID, "readers", driver.Pagination{}) + requireNoError(t, err, "ListUsersInGroup") + + if len(users) != 2 || users[0].Username != "alice" || users[1].Username != "bob" { + t.Fatalf("readers = %+v", users) + } + + requireNoError(t, m.AdminRemoveUserFromGroup(ctx, pool.ID, "alice", "readers"), "remove") + + users, _, _ = m.ListUsersInGroup(ctx, pool.ID, "readers", driver.Pagination{}) + if len(users) != 1 || users[0].Username != "bob" { + t.Fatalf("readers after remove = %+v", users) + } + + requireNoError(t, m.DeleteGroup(ctx, pool.ID, "admins"), "DeleteGroup admins") + + groups, _, _ = m.AdminListGroupsForUser(ctx, pool.ID, "alice", driver.Pagination{}) + if len(groups) != 0 { + t.Fatalf("alice groups after group delete = %+v", groups) + } + + _, err = m.CreateGroup(ctx, driver.CreateGroupInput{UserPoolID: pool.ID, GroupName: "admins"}) + requireNoError(t, err, "recreate admins") + + groups, _, _ = m.AdminListGroupsForUser(ctx, pool.ID, "alice", driver.Pagination{}) + if len(groups) != 0 { + t.Fatalf("recreated group must start empty, alice has %+v", groups) + } + + _, _, err = m.ListUsersInGroup(ctx, pool.ID, "ghosts", driver.Pagination{}) + assertException(t, err, driver.ExResourceNotFound, "Group not found.") +} + +func TestDeleteUserPoolCascadesGroups(t *testing.T) { + m := newMock(t) + ctx := context.Background() + pool := mustCreatePool(t, m, "cascade") + + _, err := m.CreateGroup(ctx, driver.CreateGroupInput{UserPoolID: pool.ID, GroupName: "g"}) + requireNoError(t, err, "CreateGroup") + requireNoError(t, m.DeleteUserPool(ctx, pool.ID), "DeleteUserPool") + + if n := len(m.groups.Keys()); n != 0 { + t.Fatalf("%d groups left after pool delete", n) + } +} diff --git a/providers/aws/cognito/keys.go b/providers/aws/cognito/keys.go new file mode 100644 index 000000000..aa6ce3c69 --- /dev/null +++ b/providers/aws/cognito/keys.go @@ -0,0 +1,89 @@ +package cognito + +import ( + "context" + "fmt" + "strings" + + "github.com/stackshy/cloudemu/v2/internal/jwtsign" +) + +// poolKeys are a pool's two RS256 signing keys. Cognito signs ID and access +// tokens with different keys, so the two token kinds have different kids. +type poolKeys struct { + id *jwtsign.Key + access *jwtsign.Key +} + +// keysFor returns a pool's signing keys, creating them on first use. +func (m *Mock) keysFor(poolID string) (*poolKeys, error) { + m.keysMu.Lock() + defer m.keysMu.Unlock() + + if k, ok := m.keys[poolID]; ok { + return k, nil + } + + id, err := jwtsign.NewRSAKey() + if err != nil { + return nil, fmt.Errorf("cognito: id token key: %w", err) + } + + access, err := jwtsign.NewRSAKey() + if err != nil { + return nil, fmt.Errorf("cognito: access token key: %w", err) + } + + k := &poolKeys{id: id, access: access} + m.keys[poolID] = k + + return k, nil +} + +// existingKeys returns a pool's keys without creating them. +func (m *Mock) existingKeys(poolID string) (*poolKeys, bool) { + m.keysMu.Lock() + defer m.keysMu.Unlock() + + k, ok := m.keys[poolID] + + return k, ok +} + +func (m *Mock) deletePoolKeys(poolID string) { + m.keysMu.Lock() + delete(m.keys, poolID) + m.keysMu.Unlock() +} + +// issuer is the iss claim of a pool's tokens. The region comes from the pool +// id, which is "_". +func issuer(poolID string) string { + region, _, _ := strings.Cut(poolID, "_") + + return "https://cognito-idp." + region + ".amazonaws.com/" + poolID +} + +// poolFromIssuer returns the pool id an iss claim names. +func poolFromIssuer(iss string) string { + i := strings.LastIndexByte(iss, '/') + if i < 0 { + return "" + } + + return iss[i+1:] +} + +// SigningKeys returns a pool's issuer and the JWKS that verifies its tokens. +func (m *Mock) SigningKeys(_ context.Context, userPoolID string) (string, jwtsign.JWKSet, error) { + if !m.userPools.Has(userPoolID) { + return "", jwtsign.JWKSet{}, poolNotFound(userPoolID) + } + + k, err := m.keysFor(userPoolID) + if err != nil { + return "", jwtsign.JWKSet{}, err + } + + return issuer(userPoolID), jwtsign.JWKS(k.id, k.access), nil +} diff --git a/providers/aws/cognito/passwords.go b/providers/aws/cognito/passwords.go index 5d9dde511..f8abd59fa 100644 --- a/providers/aws/cognito/passwords.go +++ b/providers/aws/cognito/passwords.go @@ -4,6 +4,7 @@ import ( "crypto/pbkdf2" "crypto/rand" "crypto/sha256" + "crypto/subtle" "encoding/hex" "strconv" "strings" @@ -28,6 +29,9 @@ const ( pbkdf2Prefix = "pbkdf2-sha256$" pbkdf2Iterations = 10000 pbkdf2KeyLen = 32 + // maxPBKDF2Iterations caps the iteration count a stored hash may claim, so + // a corrupt snapshot cannot make one sign-in burn CPU. + maxPBKDF2Iterations = 1_000_000 ) // policySymbols is the set of special characters Cognito counts toward the @@ -110,3 +114,30 @@ func pbkdf2Hash(salt, pw string, iter int) string { return pbkdf2Prefix + strconv.Itoa(iter) + "$" + hex.EncodeToString(key) } + +// verifyPassword reports whether pw matches a stored salt and hash, in either +// the PBKDF2 format or the legacy bare SHA-256 format older snapshots carry. +func verifyPassword(salt, stored, pw string) bool { + if stored == "" { + return false + } + + rest, ok := strings.CutPrefix(stored, pbkdf2Prefix) + if !ok { + sum := sha256.Sum256([]byte(salt + pw)) + + return subtle.ConstantTimeCompare([]byte(hex.EncodeToString(sum[:])), []byte(stored)) == 1 + } + + iterText, _, ok := strings.Cut(rest, "$") + if !ok { + return false + } + + iter, err := strconv.Atoi(iterText) + if err != nil || iter <= 0 || iter > maxPBKDF2Iterations { + return false + } + + return subtle.ConstantTimeCompare([]byte(pbkdf2Hash(salt, pw, iter)), []byte(stored)) == 1 +} diff --git a/providers/aws/cognito/passwords_test.go b/providers/aws/cognito/passwords_test.go index 1aa069903..40e8ecb7a 100644 --- a/providers/aws/cognito/passwords_test.go +++ b/providers/aws/cognito/passwords_test.go @@ -2,42 +2,11 @@ package cognito import ( "crypto/sha256" - "crypto/subtle" "encoding/hex" - "strconv" "strings" "testing" ) -const maxTestIterations = 1_000_000 - -// verifyPassword reports whether pw matches a stored salt and hash, in either -// the PBKDF2 format or the legacy bare SHA-256 format older snapshots carry. -func verifyPassword(salt, stored, pw string) bool { - if stored == "" { - return false - } - - rest, ok := strings.CutPrefix(stored, pbkdf2Prefix) - if !ok { - sum := sha256.Sum256([]byte(salt + pw)) - - return subtle.ConstantTimeCompare([]byte(hex.EncodeToString(sum[:])), []byte(stored)) == 1 - } - - iterText, _, ok := strings.Cut(rest, "$") - if !ok { - return false - } - - iter, err := strconv.Atoi(iterText) - if err != nil || iter <= 0 || iter > maxTestIterations { - return false - } - - return subtle.ConstantTimeCompare([]byte(pbkdf2Hash(salt, pw, iter)), []byte(stored)) == 1 -} - func TestVerifyPassword(t *testing.T) { const pw = "Corr3ct!horse" diff --git a/providers/aws/cognito/sign_up.go b/providers/aws/cognito/sign_up.go new file mode 100644 index 000000000..159825b15 --- /dev/null +++ b/providers/aws/cognito/sign_up.go @@ -0,0 +1,438 @@ +package cognito + +import ( + "context" + "crypto/hmac" + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "math/big" + "slices" + "strings" + "time" + + "github.com/stackshy/cloudemu/v2/errors" + "github.com/stackshy/cloudemu/v2/internal/idgen" + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// Confirmation codes are six digits and valid for 24 hours, like the codes +// Cognito sends by email or SMS. +const ( + codeDigits = 6 + codeSpace = 1_000_000 + codeValidity = 24 * time.Hour + phoneTailLen = 4 +) + +// pendingCode is the confirmation code outstanding for a user. Delivery is +// emulated: nothing is sent, and the code is read back through +// ConfirmationCode (the /_cloudemu/cognito/codes endpoint). +type pendingCode struct { + Code string `json:"code"` + ExpiresAt time.Time `json:"expiresAt"` + Delivery *driver.CodeDeliveryDetails `json:"delivery,omitempty"` +} + +func (c *pendingCode) clone() *pendingCode { + if c == nil { + return nil + } + + out := *c + + if c.Delivery != nil { + d := *c.Delivery + out.Delivery = &d + } + + return &out +} + +func codeMismatch() error { + //nolint:revive // exact Cognito message, surfaced verbatim to the SDK + return &driver.APIError{ + Exception: driver.ExCodeMismatch, + Err: errors.New(errors.InvalidArgument, "Invalid verification code provided, please try again."), + } +} + +func expiredCode() error { + //nolint:revive // exact Cognito message, surfaced verbatim to the SDK + return &driver.APIError{ + Exception: driver.ExExpiredCode, + Err: errors.New(errors.InvalidArgument, "Invalid code provided, please request a code again."), + } +} + +// clientUserNotFound is the UserNotFoundException the client-side operations +// return for an unknown username. +func clientUserNotFound() error { + //nolint:revive // exact Cognito message, surfaced verbatim to the SDK + return &driver.APIError{ + Exception: driver.ExUserNotFound, + Err: errors.New(errors.NotFound, "Username/client id combination not found."), + } +} + +func cannotConfirm(status string) error { + return notAuthorized("User cannot be confirmed. Current status is " + status) +} + +// clientByID finds an app client by id alone, as the client-side operations +// address it. Client ids are unique across pools. +func (m *Mock) clientByID(clientID string) (driver.UserPoolClient, error) { + if clientID != "" { + suffix := clientKeySep + clientID + + for _, key := range m.clients.Keys() { + if !strings.HasSuffix(key, suffix) { + continue + } + + if c, ok := m.clients.Get(key); ok { + return copyUserPoolClient(c), nil + } + } + } + + return driver.UserPoolClient{}, resourceNotFound("User pool client %s does not exist.", clientID) +} + +// secretHash is Base64(HMAC-SHA256(clientSecret, username+clientID)), the +// SECRET_HASH a client with a secret must send. +func secretHash(secret, username, clientID string) string { + mac := hmac.New(sha256.New, []byte(secret)) + mac.Write([]byte(username + clientID)) + + return base64.StdEncoding.EncodeToString(mac.Sum(nil)) +} + +// checkSecretHash verifies SECRET_HASH for a client that has a secret. The +// hash may be computed over any of the names the user is known by. missing is +// the message for an absent hash, which differs between operations. +// +//nolint:gocritic // hugeParam: the client is a stored copy passed by value +func checkSecretHash(client driver.UserPoolClient, got, missing string, names ...string) error { + if client.ClientSecret == "" { + return nil + } + + if got == "" { + return notAuthorized(missing) + } + + for _, name := range names { + if name == "" { + continue + } + + if hmac.Equal([]byte(got), []byte(secretHash(client.ClientSecret, name, client.ClientID))) { + return nil + } + } + + return badSecretHash(client.ClientID) +} + +func badSecretHash(clientID string) error { + return notAuthorized("Unable to verify secret hash for client " + clientID) +} + +// newCode returns a random six-digit code. +func newCode() string { + n, err := rand.Int(rand.Reader, big.NewInt(codeSpace)) + if err != nil { + return "000000" + } + + s := n.String() + + return strings.Repeat("0", codeDigits-len(s)) + s +} + +// codeDelivery picks where a confirmation code goes: the phone number when it +// is auto-verified and set, else the email. Nil means the pool verifies +// neither, so no code is delivered. +func codeDelivery(pool *driver.UserPool, attrs []driver.Attribute) *driver.CodeDeliveryDetails { + if slices.Contains(pool.AutoVerifiedAttributes, attrPhoneNumber) { + if v := attrValue(attrs, attrPhoneNumber); v != "" { + return &driver.CodeDeliveryDetails{ + Destination: maskPhone(v), DeliveryMedium: driver.DeliveryMediumSMS, AttributeName: attrPhoneNumber, + } + } + } + + if slices.Contains(pool.AutoVerifiedAttributes, attrEmail) { + if v := attrValue(attrs, attrEmail); v != "" { + return &driver.CodeDeliveryDetails{ + Destination: maskEmail(v), DeliveryMedium: driver.DeliveryMediumEmail, AttributeName: attrEmail, + } + } + } + + return nil +} + +// maskEmail renders alice@example.com as a***@e***. +func maskEmail(v string) string { + local, domain, ok := strings.Cut(v, "@") + if !ok || local == "" || domain == "" { + return "***" + } + + return local[:1] + "***@" + domain[:1] + "***" +} + +// maskPhone keeps the plus sign and the last four digits. +func maskPhone(v string) string { + if len(v) <= phoneTailLen+1 { + return v + } + + return "+" + strings.Repeat("*", len(v)-phoneTailLen-1) + v[len(v)-phoneTailLen:] +} + +func (m *Mock) issueCode(pool *driver.UserPool, attrs []driver.Attribute) *pendingCode { + return &pendingCode{Code: newCode(), ExpiresAt: m.now().Add(codeValidity), Delivery: codeDelivery(pool, attrs)} +} + +// checkRequired rejects a sign-up that leaves out a schema attribute the pool +// marks required. +func checkRequired(pool *driver.UserPool, attrs []driver.Attribute) error { + for _, a := range pool.SchemaAttributes { + if a.Required && a.Name != attrSub && attrValue(attrs, a.Name) == "" { + return schemaError(a.Name, "The attribute is required") + } + } + + return nil +} + +// SignUp registers an UNCONFIRMED user through an app client and issues a +// confirmation code. +// +//nolint:gocritic // hugeParam: taken by value to match the driver interface +func (m *Mock) SignUp(_ context.Context, in driver.SignUpInput) (*driver.SignUpOutput, error) { + if err := checkUsername(in.Username); err != nil { + return nil, err + } + + if in.Password == "" { + return nil, invalidParameter("1 validation error detected: Value null at 'password' failed to satisfy constraint: " + + "Member must not be null") + } + + m.mu.Lock() + defer m.mu.Unlock() + + client, err := m.clientByID(in.ClientID) + if err != nil { + return nil, err + } + + if err = checkSecretHash(client, in.SecretHash, "Unable to verify secret hash for client "+client.ClientID, in.Username); err != nil { + return nil, err + } + + pool, ok := m.userPools.Get(client.UserPoolID) + if !ok { + return nil, poolNotFound(client.UserPoolID) + } + + rec, err := m.newSignUpRecord(&pool, in) + if err != nil { + return nil, err + } + + m.users.Set(userKey(pool.ID, rec.User.Username), copyUserRecord(rec)) + + return &driver.SignUpOutput{ + UserConfirmed: false, + UserSub: attrValue(rec.User.Attributes, attrSub), + CodeDeliveryDetails: rec.Code.clone().deliveryOrNil(), + }, nil +} + +func (c *pendingCode) deliveryOrNil() *driver.CodeDeliveryDetails { + if c == nil { + return nil + } + + return c.Delivery +} + +//nolint:gocritic // hugeParam: in is the caller's input, passed through by value +func (m *Mock) newSignUpRecord(pool *driver.UserPool, in driver.SignUpInput) (userRecord, error) { + if err := validateAttributes(pool, in.UserAttributes, false); err != nil { + return userRecord{}, err + } + + if len(pool.UsernameAttributes) == 0 && m.users.Has(userKey(pool.ID, in.Username)) { + return userRecord{}, usernameExists("User already exists") + } + + 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 userRecord{}, err + } + + if err := checkRequired(pool, attrs); err != nil { + return userRecord{}, err + } + + if err := checkPassword(in.Password, pool.Policies.PasswordPolicy); err != nil { + return userRecord{}, err + } + + if err := m.claimSignIns(pool, userKey(pool.ID, username), attrs, false, true); err != nil { + return userRecord{}, err + } + + now := m.now() + rec := userRecord{ + PoolID: pool.ID, + User: driver.User{ + Username: username, + Attributes: attrs, + UserCreateDate: now, + UserLastModifiedDate: now, + Enabled: true, + UserStatus: driver.UserStatusUnconfirmed, + }, + Code: m.issueCode(pool, attrs), + } + rec.PasswordSalt, rec.PasswordHash = hashPassword(in.Password) + + return rec, nil +} + +// clientUser resolves the client, verifies SECRET_HASH, and finds the user a +// client-side operation names. It must run under m.mu. +func (m *Mock) clientUser(in driver.ClientUserInput) (driver.UserPool, string, userRecord, error) { + client, err := m.clientByID(in.ClientID) + if err != nil { + return driver.UserPool{}, "", userRecord{}, err + } + + pool, ok := m.userPools.Get(client.UserPoolID) + if !ok { + return driver.UserPool{}, "", userRecord{}, poolNotFound(client.UserPoolID) + } + + key, rec, found := m.resolveUser(&pool, in.Username) + + names := []string{in.Username} + if found { + names = append(names, rec.User.Username) + } + + if err := checkSecretHash(client, in.SecretHash, "Unable to verify secret hash for client "+client.ClientID, names...); err != nil { + return driver.UserPool{}, "", userRecord{}, err + } + + if !found { + return driver.UserPool{}, "", userRecord{}, clientUserNotFound() + } + + return pool, key, copyUserRecord(rec), nil +} + +// ConfirmSignUp confirms an UNCONFIRMED user with the code that was sent. The +// attribute the code went to becomes verified. +func (m *Mock) ConfirmSignUp(_ context.Context, in driver.ConfirmSignUpInput) error { + m.mu.Lock() + defer m.mu.Unlock() + + pool, key, rec, err := m.clientUser(in.ClientUserInput) + if err != nil { + return err + } + + if rec.User.UserStatus != driver.UserStatusUnconfirmed { + return cannotConfirm(rec.User.UserStatus) + } + + if rec.Code == nil || !hmac.Equal([]byte(rec.Code.Code), []byte(in.ConfirmationCode)) { + return codeMismatch() + } + + if !m.now().Before(rec.Code.ExpiresAt) { + return expiredCode() + } + + if d := rec.Code.Delivery; d != nil { + if flag, ok := verifiedFlag(d.AttributeName); ok { + attrs := mergeAttributes(rec.User.Attributes, []driver.Attribute{{Name: flag, Value: attrTrue}}) + if err := m.claimSignIns(&pool, key, attrs, in.ForceAliasCreation, false); err != nil { + return err + } + + rec.User.Attributes = attrs + } + } + + rec.Code = nil + rec.User.UserStatus = driver.UserStatusConfirmed + rec.User.UserLastModifiedDate = m.now() + m.users.Set(key, rec) + + return nil +} + +// ResendConfirmationCode issues a fresh code to an UNCONFIRMED user. +func (m *Mock) ResendConfirmationCode(_ context.Context, in driver.ClientUserInput) (*driver.CodeDeliveryDetails, error) { + m.mu.Lock() + defer m.mu.Unlock() + + pool, key, rec, err := m.clientUser(in) + if err != nil { + return nil, err + } + + if rec.User.UserStatus != driver.UserStatusUnconfirmed { + return nil, invalidParameter("User is already confirmed.") + } + + rec.Code = m.issueCode(&pool, rec.User.Attributes) + m.users.Set(key, rec) + + return rec.Code.clone().deliveryOrNil(), nil +} + +// AdminConfirmSignUp confirms an UNCONFIRMED user without a code. +func (m *Mock) AdminConfirmSignUp(_ context.Context, userPoolID, username string) error { + return m.updateUser(userPoolID, username, func(_ *driver.UserPool, _ string, rec *userRecord) error { + if rec.User.UserStatus != driver.UserStatusUnconfirmed { + return cannotConfirm(rec.User.UserStatus) + } + + rec.Code = nil + rec.User.UserStatus = driver.UserStatusConfirmed + + return nil + }) +} + +// ConfirmationCode returns the code outstanding for a user. +func (m *Mock) ConfirmationCode(_ context.Context, userPoolID, username string) (*driver.IssuedCode, error) { + pool, ok := m.userPools.Get(userPoolID) + if !ok { + return nil, poolNotFound(userPoolID) + } + + _, rec, ok := m.resolveUser(&pool, username) + if !ok { + return nil, userNotFound() + } + + c := rec.Code.clone() + if c == nil { + return nil, resourceNotFound("No confirmation code is outstanding for user %s.", username) + } + + return &driver.IssuedCode{Code: c.Code, ExpiresAt: c.ExpiresAt, Delivery: c.Delivery}, nil +} diff --git a/providers/aws/cognito/sign_up_test.go b/providers/aws/cognito/sign_up_test.go new file mode 100644 index 000000000..1efc8762b --- /dev/null +++ b/providers/aws/cognito/sign_up_test.go @@ -0,0 +1,267 @@ +package cognito + +import ( + "context" + "crypto/hmac" + "crypto/sha256" + "encoding/base64" + "testing" + "time" + + "github.com/stackshy/cloudemu/v2/config" + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// Flows a test client enables for password sign-in. +// +//nolint:gochecknoglobals // test fixture +var passwordFlows = []string{"ALLOW_USER_PASSWORD_AUTH", "ALLOW_ADMIN_USER_PASSWORD_AUTH", "ALLOW_REFRESH_TOKEN_AUTH"} + +const testPassword = "Passw0rd!" + +func newClockMock(t *testing.T) (*Mock, *config.FakeClock) { + t.Helper() + + fc := config.NewFakeClock(time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC)) + + return New(config.NewOptions(config.WithClock(fc))), fc +} + +func mustCreateEmailPool(t *testing.T, m *Mock) *driver.UserPool { + t.Helper() + + pool, err := m.CreateUserPool(context.Background(), driver.CreateUserPoolInput{ + Name: "signup", AutoVerifiedAttributes: []string{"email"}, + }) + requireNoError(t, err, "CreateUserPool") + + return pool +} + +func mustCreateClient(t *testing.T, m *Mock, poolID string, secret bool, flows ...string) *driver.UserPoolClient { + t.Helper() + + c, err := m.CreateUserPoolClient(context.Background(), driver.CreateUserPoolClientInput{ + UserPoolID: poolID, ClientName: "app", GenerateSecret: secret, ExplicitAuthFlows: flows, + }) + requireNoError(t, err, "CreateUserPoolClient") + + return c +} + +// testSecretHash computes SECRET_HASH the way the AWS docs specify it, without +// the provider's helper. +func testSecretHash(secret, username, clientID string) string { + mac := hmac.New(sha256.New, []byte(secret)) + mac.Write([]byte(username + clientID)) + + return base64.StdEncoding.EncodeToString(mac.Sum(nil)) +} + +func signUpInput(clientID, username, password string, attrs ...driver.Attribute) driver.SignUpInput { + return driver.SignUpInput{ + ClientUserInput: driver.ClientUserInput{ClientID: clientID, Username: username}, + Password: password, + UserAttributes: attrs, + } +} + +func emailAttr(v string) driver.Attribute { return driver.Attribute{Name: "email", Value: v} } + +func mustCode(t *testing.T, m *Mock, poolID, username string) string { + t.Helper() + + c, err := m.ConfirmationCode(context.Background(), poolID, username) + requireNoError(t, err, "ConfirmationCode") + + return c.Code +} + +func TestSignUpAndConfirm(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + + out, err := m.SignUp(ctx, signUpInput(client.ClientID, "alice", testPassword, emailAttr("alice@example.com"))) + requireNoError(t, err, "SignUp") + + if out.UserConfirmed || len(out.UserSub) != 36 { + t.Fatalf("SignUp = %+v", out) + } + + d := out.CodeDeliveryDetails + if d == nil || d.DeliveryMedium != driver.DeliveryMediumEmail || d.AttributeName != "email" || d.Destination != "a***@e***" { + t.Fatalf("CodeDeliveryDetails = %+v", d) + } + + u, err := m.AdminGetUser(ctx, pool.ID, "alice") + requireNoError(t, err, "AdminGetUser") + + if u.UserStatus != driver.UserStatusUnconfirmed || attrValue(u.Attributes, "sub") != out.UserSub { + t.Fatalf("user = %+v", u) + } + + code := mustCode(t, m, pool.ID, "alice") + if len(code) != 6 { + t.Fatalf("code %q is not 6 digits", code) + } + + bad := driver.ConfirmSignUpInput{ClientUserInput: driver.ClientUserInput{ClientID: client.ClientID, Username: "alice"}} + + bad.ConfirmationCode = wrongCode(code) + assertException(t, m.ConfirmSignUp(ctx, bad), driver.ExCodeMismatch, "Invalid verification code provided, please try again.") + + bad.ConfirmationCode = code + requireNoError(t, m.ConfirmSignUp(ctx, bad), "ConfirmSignUp") + + u, _ = m.AdminGetUser(ctx, pool.ID, "alice") + if u.UserStatus != driver.UserStatusConfirmed || attrValue(u.Attributes, "email_verified") != "true" { + t.Fatalf("confirmed user = %+v", u) + } + + assertException(t, m.ConfirmSignUp(ctx, bad), driver.ExNotAuthorized, "User cannot be confirmed. Current status is CONFIRMED") + + _, err = m.ResendConfirmationCode(ctx, bad.ClientUserInput) + assertException(t, err, driver.ExInvalidParameter, "User is already confirmed.") +} + +func wrongCode(code string) string { + if code == "000000" { + return "111111" + } + + return "000000" +} + +func TestSignUpErrors(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false) + + _, err := m.SignUp(ctx, signUpInput(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "SignUp") + + cases := []struct { + name string + in driver.SignUpInput + exception string + msg string + }{ + {"duplicate", signUpInput(client.ClientID, "alice", testPassword), driver.ExUsernameExists, "User already exists"}, + {"short password", signUpInput(client.ClientID, "bob", "Sh0rt!"), + driver.ExInvalidPassword, "Password did not conform with policy: Password not long enough"}, + {"no symbol", signUpInput(client.ClientID, "bob", "Passw0rdxx"), + driver.ExInvalidPassword, "Password did not conform with policy: Password must have symbol characters"}, + {"missing password", signUpInput(client.ClientID, "bob", ""), driver.ExInvalidParameter, ""}, + {"unknown client", signUpInput("nosuchclient", "bob", testPassword), driver.ExResourceNotFound, ""}, + {"bad attribute", signUpInput(client.ClientID, "bob", testPassword, driver.Attribute{Name: "custom:nope", Value: "x"}), + driver.ExInvalidParameter, "Attributes did not conform to the schema: custom:nope: Attribute does not exist in the schema."}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + _, err := m.SignUp(ctx, tc.in) + assertException(t, err, tc.exception, tc.msg) + }) + } +} + +func TestSignUpSecretHash(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, true) + + in := signUpInput(client.ClientID, "alice", testPassword) + + _, err := m.SignUp(ctx, in) + assertException(t, err, driver.ExNotAuthorized, "Unable to verify secret hash for client "+client.ClientID) + + in.SecretHash = testSecretHash(client.ClientSecret, "someone-else", client.ClientID) + _, err = m.SignUp(ctx, in) + assertException(t, err, driver.ExNotAuthorized, "Unable to verify secret hash for client "+client.ClientID) + + in.SecretHash = testSecretHash(client.ClientSecret, "alice", client.ClientID) + _, err = m.SignUp(ctx, in) + requireNoError(t, err, "SignUp with SECRET_HASH") +} + +func TestConfirmationCodeExpiresAndResends(t *testing.T) { + m, fc := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false) + + _, err := m.SignUp(ctx, signUpInput(client.ClientID, "alice", testPassword, emailAttr("alice@example.com"))) + requireNoError(t, err, "SignUp") + + code := mustCode(t, m, pool.ID, "alice") + + fc.Advance(25 * time.Hour) + + in := driver.ConfirmSignUpInput{ + ClientUserInput: driver.ClientUserInput{ClientID: client.ClientID, Username: "alice"}, + ConfirmationCode: code, + } + assertException(t, m.ConfirmSignUp(ctx, in), driver.ExExpiredCode, "Invalid code provided, please request a code again.") + + d, err := m.ResendConfirmationCode(ctx, in.ClientUserInput) + requireNoError(t, err, "ResendConfirmationCode") + + if d.Destination != "a***@e***" { + t.Fatalf("resend delivery = %+v", d) + } + + in.ConfirmationCode = mustCode(t, m, pool.ID, "alice") + requireNoError(t, m.ConfirmSignUp(ctx, in), "ConfirmSignUp with resent code") + + _, err = m.ResendConfirmationCode(ctx, driver.ClientUserInput{ClientID: client.ClientID, Username: "ghost"}) + assertException(t, err, driver.ExUserNotFound, "Username/client id combination not found.") +} + +func TestAdminConfirmSignUp(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false) + + _, err := m.SignUp(ctx, signUpInput(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "SignUp") + requireNoError(t, m.AdminConfirmSignUp(ctx, pool.ID, "alice"), "AdminConfirmSignUp") + + u, _ := m.AdminGetUser(ctx, pool.ID, "alice") + if u.UserStatus != driver.UserStatusConfirmed { + t.Fatalf("status = %s", u.UserStatus) + } + + assertException(t, m.AdminConfirmSignUp(ctx, pool.ID, "alice"), driver.ExNotAuthorized, + "User cannot be confirmed. Current status is CONFIRMED") + assertException(t, m.AdminConfirmSignUp(ctx, pool.ID, "ghost"), driver.ExUserNotFound, "User does not exist.") +} + +func TestSignUpEmailUsernamePool(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + + pool, err := m.CreateUserPool(ctx, driver.CreateUserPoolInput{ + Name: "email-login", UsernameAttributes: []string{"email"}, AutoVerifiedAttributes: []string{"email"}, + }) + requireNoError(t, err, "CreateUserPool") + + client := mustCreateClient(t, m, pool.ID, false) + + out, err := m.SignUp(ctx, signUpInput(client.ClientID, "carol@example.com", testPassword)) + requireNoError(t, err, "SignUp") + + u, err := m.AdminGetUser(ctx, pool.ID, "carol@example.com") + requireNoError(t, err, "AdminGetUser by email") + + if u.Username != out.UserSub || attrValue(u.Attributes, "email") != "carol@example.com" { + t.Fatalf("user = %+v, want username = sub and email copied", u) + } + + _, err = m.SignUp(ctx, signUpInput(client.ClientID, "carol@example.com", testPassword)) + assertException(t, err, driver.ExUsernameExists, "An account with the given email already exists.") +} diff --git a/providers/aws/cognito/snapshot.go b/providers/aws/cognito/snapshot.go index 61a4b292f..d8b1969c4 100644 --- a/providers/aws/cognito/snapshot.go +++ b/providers/aws/cognito/snapshot.go @@ -5,6 +5,7 @@ import ( "encoding/json" "fmt" + "github.com/stackshy/cloudemu/v2/internal/jwtsign" "github.com/stackshy/cloudemu/v2/internal/snapshot" "github.com/stackshy/cloudemu/v2/services/cognito/driver" ) @@ -14,15 +15,25 @@ var _ snapshot.Snapshottable = (*Mock)(nil) // cognitoSnapshot is the full serialized state of the Cognito mock. The stores // hold plain driver values keyed by their resource id (clients by the composite // poolID/clientID key), so each map lifts straight out. Tags live only in the -// ARN-keyed side map. +// ARN-keyed side map. Challenge sessions last minutes and are not kept. 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"` + Groups map[string]driver.Group `json:"groups,omitempty"` + Logins map[string]loginRecord `json:"logins,omitempty"` + Keys map[string]keySnapshot `json:"keys,omitempty"` Tags map[string]map[string]string `json:"tags,omitempty"` } +// keySnapshot holds a pool's two signing keys as PKCS#8 DER, so tokens issued +// before a restart still verify after it. +type keySnapshot struct { + ID []byte `json:"id"` + Access []byte `json:"access"` +} + // Snapshot captures the mock's entire state as JSON. includeAssets is unused. Cognito holds no // bulk object bodies. func (m *Mock) Snapshot(_ context.Context, _ bool) (json.RawMessage, error) { @@ -31,8 +42,17 @@ func (m *Mock) Snapshot(_ context.Context, _ bool) (json.RawMessage, error) { Clients: deepCopyMap(m.clients.All(), copyUserPoolClient), Domains: deepCopyMap(m.domains.All(), copyUserPoolDomain), Users: deepCopyMap(m.users.All(), copyUserRecord), + Groups: deepCopyMap(m.groups.All(), copyGroup), + Logins: deepCopyMap(m.logins.All(), copyLoginRecord), + } + + keys, err := m.snapshotKeys() + if err != nil { + return nil, err } + snap.Keys = keys + m.tagsMu.RLock() snap.Tags = deepCopyTags(m.tags) m.tagsMu.RUnlock() @@ -47,10 +67,17 @@ func (m *Mock) Restore(_ context.Context, data json.RawMessage) error { return fmt.Errorf("cognito: parse snapshot: %w", err) } + keys, err := parseKeys(snap.Keys) + if err != nil { + return err + } + m.userPools.Clear() m.clients.Clear() m.domains.Clear() m.users.Clear() + m.groups.Clear() + m.logins.Clear() for k := range snap.UserPools { m.userPools.Set(k, copyUserPool(snap.UserPools[k])) @@ -68,6 +95,22 @@ func (m *Mock) Restore(_ context.Context, data json.RawMessage) error { m.users.Set(k, copyUserRecord(snap.Users[k])) } + for k := range snap.Groups { + m.groups.Set(k, copyGroup(snap.Groups[k])) + } + + for k := range snap.Logins { + m.logins.Set(k, snap.Logins[k]) + } + + m.keysMu.Lock() + m.keys = keys + m.keysMu.Unlock() + + m.sessionsMu.Lock() + m.sessions = map[string]challengeSession{} + m.sessionsMu.Unlock() + m.tagsMu.Lock() m.tags = deepCopyTags(snap.Tags) @@ -79,6 +122,56 @@ func (m *Mock) Restore(_ context.Context, data json.RawMessage) error { return nil } +// snapshotKeys encodes every pool's signing keys as PKCS#8. +func (m *Mock) snapshotKeys() (map[string]keySnapshot, error) { + m.keysMu.Lock() + defer m.keysMu.Unlock() + + if len(m.keys) == 0 { + return nil, nil + } + + out := make(map[string]keySnapshot, len(m.keys)) + + for poolID, k := range m.keys { + id, err := jwtsign.MarshalPKCS8(k.id) + if err != nil { + return nil, err + } + + access, err := jwtsign.MarshalPKCS8(k.access) + if err != nil { + return nil, err + } + + out[poolID] = keySnapshot{ID: id, Access: access} + } + + return out, nil +} + +// parseKeys decodes snapshotted signing keys. The kids are recomputed from the +// keys, so they match the ones published before the snapshot. +func parseKeys(in map[string]keySnapshot) (map[string]*poolKeys, error) { + keys := make(map[string]*poolKeys, len(in)) + + for poolID, ks := range in { + id, err := jwtsign.ParsePKCS8(ks.ID) + if err != nil { + return nil, fmt.Errorf("cognito: restore keys of %s: %w", poolID, err) + } + + access, err := jwtsign.ParsePKCS8(ks.Access) + if err != nil { + return nil, fmt.Errorf("cognito: restore keys of %s: %w", poolID, err) + } + + keys[poolID] = &poolKeys{id: id, access: access} + } + + return keys, nil +} + // deepCopyTags copies a nested ARN->tags map. func deepCopyTags(in map[string]map[string]string) map[string]map[string]string { if len(in) == 0 { diff --git a/providers/aws/cognito/tokens.go b/providers/aws/cognito/tokens.go new file mode 100644 index 000000000..9b479b1ee --- /dev/null +++ b/providers/aws/cognito/tokens.go @@ -0,0 +1,403 @@ +package cognito + +import ( + "crypto/hmac" + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "encoding/hex" + "encoding/json" + stderrors "errors" + "slices" + "sort" + "strings" + "time" + + "github.com/stackshy/cloudemu/v2/internal/idgen" + "github.com/stackshy/cloudemu/v2/internal/jwtsign" + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// Token defaults: access and ID tokens last an hour and refresh tokens 30 +// days unless the app client says otherwise. +const ( + defaultTokenValidity = time.Hour + tokenTypeBearer = "Bearer" + scopeSignInAdmin = "aws.cognito.signin.user.admin" + tokenUseID = "id" + tokenUseAccess = "access" + refreshSegments = 5 + refreshSecretLen = 32 + refreshKeyLen = 256 + refreshIVLen = 12 + refreshTagLen = 16 + jwtSegments = 3 +) + +// refreshHeader is the JWE protected header Cognito refresh tokens carry. The +// emulator's refresh tokens are opaque, like the real ones, but keep the same +// five-segment shape so clients that sniff the format accept them. +const refreshHeader = "eyJjdHkiOiJKV1QiLCJlbmMiOiJBMjU2R0NNIiwiYWxnIjoiUlNBLU9BRVAifQ" + +// loginRecord is one sign-in. Every token minted from it, including those a +// refresh mints later, carries its OriginJTI, so revoking the record revokes +// them all. RefreshHash is the SHA-256 of the refresh token; the token itself +// is never stored. +type loginRecord struct { + OriginJTI string `json:"originJti"` + PoolID string `json:"poolId"` + ClientID string `json:"clientId"` + Username string `json:"username"` + Sub string `json:"sub"` + AuthTime time.Time `json:"authTime"` + RefreshExpires time.Time `json:"refreshExpires"` + RefreshHash string `json:"refreshHash"` + Revoked bool `json:"revoked,omitempty"` +} + +//nolint:gocritic // hugeParam: value signature required by the func(V) V copy callback +func copyLoginRecord(in loginRecord) loginRecord { return in } + +func tokenHash(token string) string { + sum := sha256.Sum256([]byte(token)) + + return hex.EncodeToString(sum[:]) +} + +func randomB64(n int) string { + b := make([]byte, n) + _, _ = rand.Read(b) + + return base64.RawURLEncoding.EncodeToString(b) +} + +// newRefreshToken returns an opaque refresh token that names its login. +func newRefreshToken(originJTI string) string { + secret := randomB64(refreshSecretLen) + body := base64.RawURLEncoding.EncodeToString([]byte(originJTI + "." + secret)) + + return strings.Join([]string{refreshHeader, randomB64(refreshKeyLen), randomB64(refreshIVLen), body, randomB64(refreshTagLen)}, ".") +} + +// lookupRefresh returns the login a refresh token belongs to. +func (m *Mock) lookupRefresh(token string) (loginRecord, bool) { + parts := strings.Split(token, ".") + if len(parts) != refreshSegments { + return loginRecord{}, false + } + + raw, err := base64.RawURLEncoding.DecodeString(parts[3]) + if err != nil { + return loginRecord{}, false + } + + originJTI, _, ok := strings.Cut(string(raw), ".") + if !ok { + return loginRecord{}, false + } + + l, ok := m.logins.Get(originJTI) + if !ok || !hmac.Equal([]byte(l.RefreshHash), []byte(tokenHash(token))) { + return loginRecord{}, false + } + + return l, true +} + +// deleteLogins removes the logins match selects. +func (m *Mock) deleteLogins(match func(*loginRecord) bool) { + for _, key := range m.logins.Keys() { + if l, ok := m.logins.Get(key); ok && match(&l) { + m.logins.Delete(key) + } + } +} + +// revokeLogins marks every login of a user revoked. +func (m *Mock) revokeLogins(poolID, sub string) { + for _, key := range m.logins.Keys() { + l, ok := m.logins.Get(key) + if !ok || l.PoolID != poolID || l.Sub != sub || l.Revoked { + continue + } + + l.Revoked = true + m.logins.Set(key, l) + } +} + +// validity converts a client token-validity value and unit to a duration. +func validity(value *int32, unit string, def time.Duration) time.Duration { + if value == nil { + return def + } + + return unitDuration(*value, unit, driver.TimeUnitHours) +} + +func unitDuration(value int32, unit, defUnit string) time.Duration { + if unit == "" { + unit = defUnit + } + + d := time.Duration(value) + + switch unit { + case driver.TimeUnitSeconds: + return d * time.Second + case driver.TimeUnitMinutes: + return d * time.Minute + case driver.TimeUnitDays: + return d * 24 * time.Hour + default: + return d * time.Hour + } +} + +// tokenLifetimes returns the access, ID and refresh token lifetimes of a client. +// +//nolint:gocritic // hugeParam, unnamedResult: stored client copy; the order matches the names +func tokenLifetimes(c driver.UserPoolClient) (time.Duration, time.Duration, time.Duration) { + units := driver.TokenValidityUnits{} + if c.TokenValidityUnits != nil { + units = *c.TokenValidityUnits + } + + access := validity(c.AccessTokenValidity, units.AccessToken, defaultTokenValidity) + id := validity(c.IDTokenValidity, units.IDToken, defaultTokenValidity) + refresh := unitDuration(c.RefreshTokenValidity, units.RefreshToken, driver.TimeUnitDays) + + return access, id, refresh +} + +// startLogin records a new sign-in and returns it with its refresh token. +// +//nolint:gocritic // hugeParam: stored client copy passed by value +func (m *Mock) startLogin(client driver.UserPoolClient, rec *userRecord) (loginRecord, string) { + _, _, refresh := tokenLifetimes(client) + now := m.now() + originJTI := idgen.UUID() + token := newRefreshToken(originJTI) + + l := loginRecord{ + OriginJTI: originJTI, + PoolID: client.UserPoolID, + ClientID: client.ClientID, + Username: rec.User.Username, + Sub: attrValue(rec.User.Attributes, attrSub), + AuthTime: now, + RefreshExpires: now.Add(refresh), + RefreshHash: tokenHash(token), + } + m.logins.Set(originJTI, l) + + return l, token +} + +// mintTokens signs a fresh access and ID token for a login. +// +//nolint:gocritic // hugeParam: stored client copy passed by value +func (m *Mock) mintTokens(client driver.UserPoolClient, rec *userRecord, l *loginRecord) (*driver.AuthenticationResult, error) { + keys, err := m.keysFor(l.PoolID) + if err != nil { + return nil, err + } + + accessLife, idLife, _ := tokenLifetimes(client) + now := m.now() + eventID := idgen.UUID() + groups := m.userGroupsByPrecedence(l.PoolID, rec.Groups) + + common := func(use string, life time.Duration) map[string]any { + c := map[string]any{ + attrSub: l.Sub, + "iss": issuer(l.PoolID), + "origin_jti": l.OriginJTI, + "event_id": eventID, + "token_use": use, + "auth_time": l.AuthTime.Unix(), + "iat": now.Unix(), + "exp": now.Add(life).Unix(), + "jti": idgen.UUID(), + } + + if len(groups) > 0 { + c["cognito:groups"] = groupNames(groups) + } + + return c + } + + access := common(tokenUseAccess, accessLife) + access["client_id"] = client.ClientID + access["scope"] = scopeSignInAdmin + access["username"] = rec.User.Username + + id := common(tokenUseID, idLife) + id["aud"] = client.ClientID + id["cognito:username"] = rec.User.Username + addRoleClaims(id, groups) + addAttributeClaims(id, rec.User.Attributes, client.ReadAttributes) + + accessToken, err := jwtsign.Sign(keys.access, access) + if err != nil { + return nil, err + } + + idToken, err := jwtsign.Sign(keys.id, id) + if err != nil { + return nil, err + } + + return &driver.AuthenticationResult{ + AccessToken: accessToken, + IDToken: idToken, + ExpiresIn: int32(accessLife / time.Second), //nolint:gosec // validity is capped far below int32 seconds + TokenType: tokenTypeBearer, + }, nil +} + +// userGroupsByPrecedence orders a user's groups the way cognito:groups lists +// them: lowest precedence first, groups without a precedence last. +func (m *Mock) userGroupsByPrecedence(poolID string, names []string) []driver.Group { + groups := m.userGroups(poolID, names) + + sort.SliceStable(groups, func(i, j int) bool { + pi, pj := groups[i].Precedence, groups[j].Precedence + + switch { + case pi == nil || pj == nil: + return pi != nil && pj == nil + default: + return *pi < *pj + } + }) + + return groups +} + +func groupNames(groups []driver.Group) []string { + out := make([]string, len(groups)) + for i := range groups { + out[i] = groups[i].GroupName + } + + return out +} + +// addRoleClaims adds cognito:roles and cognito:preferred_role to an ID token +// when any of the user's groups has an IAM role. +func addRoleClaims(claims map[string]any, groups []driver.Group) { + var roles []string + + for i := range groups { + if groups[i].RoleARN != "" { + roles = append(roles, groups[i].RoleARN) + } + } + + if len(roles) == 0 { + return + } + + claims["cognito:roles"] = roles + claims["cognito:preferred_role"] = roles[0] +} + +// addAttributeClaims copies user attributes into an ID token. The two +// verification flags are booleans; everything else, custom attributes +// included, stays a string. A client with ReadAttributes only sees those. +func addAttributeClaims(claims map[string]any, attrs []driver.Attribute, readable []string) { + for _, a := range attrs { + if a.Name == attrSub || (len(readable) > 0 && !slices.Contains(readable, a.Name)) { + continue + } + + if a.Name == attrEmailVerified || a.Name == attrPhoneNumberVerified { + claims[a.Name] = a.Value == attrTrue + + continue + } + + claims[a.Name] = a.Value + } +} + +func invalidAccessToken() error { return notAuthorized("Invalid Access Token") } + +// verifyAccessToken checks an access token's signature, expiry and use, then +// that its login is still live and its user still exists. It returns the +// user's record. +func (m *Mock) verifyAccessToken(token string) (userRecord, error) { + poolID, ok := unverifiedPool(token) + if !ok { + return userRecord{}, invalidAccessToken() + } + + keys, ok := m.existingKeys(poolID) + if !ok { + return userRecord{}, invalidAccessToken() + } + + claims, err := jwtsign.Verify(token, []*jwtsign.Key{keys.access}, m.opts.Clock) + + switch { + case stderrors.Is(err, jwtsign.ErrExpired): + return userRecord{}, notAuthorized("Access Token has expired") + case err != nil: + return userRecord{}, invalidAccessToken() + } + + if claims["token_use"] != tokenUseAccess || claims["iss"] != issuer(poolID) { + return userRecord{}, invalidAccessToken() + } + + return m.liveTokenUser(poolID, claims) +} + +// liveTokenUser checks that a verified token's login is not revoked and that +// its user still exists and is enabled. +func (m *Mock) liveTokenUser(poolID string, claims map[string]any) (userRecord, error) { + originJTI, _ := claims["origin_jti"].(string) + sub, _ := claims[attrSub].(string) + + l, ok := m.logins.Get(originJTI) + if !ok || l.Revoked || l.Sub != sub || l.PoolID != poolID { + return userRecord{}, notAuthorized("Access Token has been revoked") + } + + rec, ok := m.users.Get(userKey(poolID, l.Username)) + if !ok || attrValue(rec.User.Attributes, attrSub) != sub { + return userRecord{}, notAuthorized("Access Token has been revoked") + } + + if !rec.User.Enabled { + return userRecord{}, notAuthorized("User is disabled.") + } + + return rec, nil +} + +// unverifiedPool reads the pool id from a JWT's iss claim before the signature +// is checked, to pick the key set to check it with. +func unverifiedPool(token string) (string, bool) { + parts := strings.Split(token, ".") + if len(parts) != jwtSegments { + return "", false + } + + raw, err := base64.RawURLEncoding.DecodeString(parts[1]) + if err != nil { + return "", false + } + + var claims struct { + Iss string `json:"iss"` + } + + if json.Unmarshal(raw, &claims) != nil || claims.Iss == "" { + return "", false + } + + return poolFromIssuer(claims.Iss), true +} diff --git a/providers/aws/cognito/user_pools.go b/providers/aws/cognito/user_pools.go index 512ce0124..32eefe5de 100644 --- a/providers/aws/cognito/user_pools.go +++ b/providers/aws/cognito/user_pools.go @@ -145,6 +145,9 @@ func (m *Mock) DeleteUserPool(_ context.Context, id string) error { } m.deletePoolUsers(id) + m.deletePoolGroups(id) + m.deleteLogins(func(l *loginRecord) bool { return l.PoolID == id }) + m.deletePoolKeys(id) m.userPools.Delete(id) m.deleteTags(pool.ARN) diff --git a/providers/aws/cognito/users.go b/providers/aws/cognito/users.go index c0a492da3..c83ba8f81 100644 --- a/providers/aws/cognito/users.go +++ b/providers/aws/cognito/users.go @@ -25,13 +25,16 @@ const ( // 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). +// userRecord is a stored user: the public view plus the password digest, its +// group memberships, and any outstanding confirmation code. 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"` + PoolID string `json:"poolId"` + User driver.User `json:"user"` + PasswordSalt string `json:"passwordSalt,omitempty"` + PasswordHash string `json:"passwordHash,omitempty"` + Groups []string `json:"groups,omitempty"` + Code *pendingCode `json:"code,omitempty"` } func userKey(poolID, username string) string { return poolID + clientKeySep + username } @@ -40,6 +43,8 @@ func userKey(poolID, username string) string { return poolID + clientKeySep + us func copyUserRecord(in userRecord) userRecord { out := in out.User.Attributes = slices.Clone(in.User.Attributes) + out.Groups = slices.Clone(in.Groups) + out.Code = in.Code.clone() return out } @@ -271,12 +276,15 @@ func (m *Mock) AdminDeleteUser(_ context.Context, userPoolID, username string) e return poolNotFound(userPoolID) } - key, _, ok := m.resolveUser(&pool, username) + key, rec, ok := m.resolveUser(&pool, username) if !ok { return userNotFound() } m.users.Delete(key) + m.deleteLogins(func(l *loginRecord) bool { + return l.PoolID == pool.ID && l.Sub == attrValue(rec.User.Attributes, attrSub) + }) return nil } diff --git a/server/aws/authbypass_test.go b/server/aws/authbypass_test.go index 1173a68d7..70143971a 100644 --- a/server/aws/authbypass_test.go +++ b/server/aws/authbypass_test.go @@ -60,7 +60,6 @@ func TestEnforcedGateBypassAttempts(t *testing.T) { {"/_cognito path, public target", cognitoOn("", "/_cognito/x", "InitiateAuth")}, {"well-known path, ListUserPools", rawReq{method: http.MethodGet, path: "/us-east-1_abcDEF123/.well-known/jwks.json", header: map[string]string{"X-Amz-Target": idpTarget + "ListUserPools"}}}, - {"jwks GET before Cognito serves it", rawReq{method: http.MethodGet, path: "/us-east-1_abcDEF123/.well-known/jwks.json"}}, {"cognito CreateUserPool", cognitoOn("", "/", "CreateUserPool")}, {"cognito AdminInitiateAuth", cognitoOn("", "/", "AdminInitiateAuth")}, {"cognito lower-case op", cognitoOn("", "/", "initiateAuth")}, diff --git a/server/aws/authz_completeness_test.go b/server/aws/authz_completeness_test.go index 6d20f7648..4d7748501 100644 --- a/server/aws/authz_completeness_test.go +++ b/server/aws/authz_completeness_test.go @@ -90,7 +90,7 @@ func TestHandlerIAMServicesMatchTable(t *testing.T) { "*dynamodb.Handler": "dynamodb", "*dynamodb.StreamsHandler": "dynamodb", "*sqs.Handler": "sqs", "*ssm.Handler": "ssm", "*kms.Handler": "kms", "*acm.Handler": "acm", "*sfn.Handler": "states", "*kinesis.Handler": "kinesis", "*cloudtrail.Handler": "cloudtrail", "*glue.Handler": "glue", "*aoss.Handler": "aoss", "*kendra.Handler": "kendra", - "*athena.Handler": "athena", "*cognito.Handler": "cognito-idp", "*configservice.Handler": "config", + "*athena.Handler": "athena", "*cognito.Handler": "cognito-idp", "*cognito.WellKnown": "cognito-idp", "*configservice.Handler": "config", "*wafv2.Handler": "wafv2", "*ecs.Handler": "ecs", "*ecr.Handler": "ecr", "*route53resolver.Handler": "route53resolver", "*eventbridge.Handler": "events", "*cloudwatchlogs.Handler": "logs", "*secretsmanager.Handler": "secretsmanager", "*keyspaces.Handler": "cassandra", "*memorydb.Handler": "memorydb", "*networkfirewall.Handler": "network-firewall", diff --git a/server/aws/aws.go b/server/aws/aws.go index b2c1e42f4..d7bda7a0b 100644 --- a/server/aws/aws.go +++ b/server/aws/aws.go @@ -765,6 +765,12 @@ func newServer(d Drivers) (*server.Server, authzSets) { // services, so registration order is unconstrained. if d.Cognito != nil { srv.Register(rpc(cognitosrv.New(d.Cognito))) + + // The pool's GET /{poolId}/.well-known/* documents must register before + // S3, the permissive REST fallback that would otherwise claim the path. + if keys, ok := d.Cognito.(cognitodriver.KeySetProvider); ok { + srv.Register(cognitosrv.NewWellKnown(keys)) + } } // AWS Config matches the X-Amz-Target prefix "StarlingDoveService.", diff --git a/server/aws/cognito/auth_ops.go b/server/aws/cognito/auth_ops.go new file mode 100644 index 000000000..3996465c1 --- /dev/null +++ b/server/aws/cognito/auth_ops.go @@ -0,0 +1,259 @@ +package cognito + +import ( + "context" + "net/http" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +type codeDeliveryJSON struct { + Destination string `json:"Destination,omitempty"` + DeliveryMedium string `json:"DeliveryMedium,omitempty"` + AttributeName string `json:"AttributeName,omitempty"` +} + +func deliveryToWire(d *driver.CodeDeliveryDetails) *codeDeliveryJSON { + if d == nil { + return nil + } + + out := codeDeliveryJSON(*d) + + return &out +} + +type signUpRequest struct { + ClientID string `json:"ClientId"` + SecretHash string `json:"SecretHash"` + Username string `json:"Username"` + Password string `json:"Password"` + UserAttributes []attributeJSON `json:"UserAttributes"` +} + +type signUpResponse struct { + UserConfirmed bool `json:"UserConfirmed"` + UserSub string `json:"UserSub"` + CodeDeliveryDetails *codeDeliveryJSON `json:"CodeDeliveryDetails,omitempty"` +} + +func (h *Handler) signUp(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *signUpRequest) (any, error) { + out, err := h.cognito.SignUp(ctx, driver.SignUpInput{ + ClientUserInput: driver.ClientUserInput{ClientID: req.ClientID, SecretHash: req.SecretHash, Username: req.Username}, + Password: req.Password, + UserAttributes: attributesFromWire(req.UserAttributes), + }) + if err != nil { + return nil, err + } + + return signUpResponse{ + UserConfirmed: out.UserConfirmed, + UserSub: out.UserSub, + CodeDeliveryDetails: deliveryToWire(out.CodeDeliveryDetails), + }, nil + }) +} + +type clientUserRequest struct { + ClientID string `json:"ClientId"` + SecretHash string `json:"SecretHash"` + Username string `json:"Username"` + ConfirmationCode string `json:"ConfirmationCode"` + ForceAliasCreation bool `json:"ForceAliasCreation"` +} + +func (req *clientUserRequest) ref() driver.ClientUserInput { + return driver.ClientUserInput{ClientID: req.ClientID, SecretHash: req.SecretHash, Username: req.Username} +} + +func (h *Handler) confirmSignUp(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *clientUserRequest) (any, error) { + err := h.cognito.ConfirmSignUp(ctx, driver.ConfirmSignUpInput{ + ClientUserInput: req.ref(), ConfirmationCode: req.ConfirmationCode, ForceAliasCreation: req.ForceAliasCreation, + }) + if err != nil { + return nil, err + } + + return struct{}{}, nil + }) +} + +type resendResponse struct { + CodeDeliveryDetails *codeDeliveryJSON `json:"CodeDeliveryDetails,omitempty"` +} + +func (h *Handler) resendConfirmationCode(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *clientUserRequest) (any, error) { + d, err := h.cognito.ResendConfirmationCode(ctx, req.ref()) + if err != nil { + return nil, err + } + + return resendResponse{CodeDeliveryDetails: deliveryToWire(d)}, nil + }) +} + +func (h *Handler) adminConfirmSignUp(w http.ResponseWriter, r *http.Request) { + h.simpleUserOp(w, r, h.cognito.AdminConfirmSignUp) +} + +func (h *Handler) adminUserGlobalSignOut(w http.ResponseWriter, r *http.Request) { + h.simpleUserOp(w, r, h.cognito.AdminUserGlobalSignOut) +} + +type authenticationResultJSON struct { + AccessToken string `json:"AccessToken,omitempty"` + ExpiresIn int32 `json:"ExpiresIn"` + TokenType string `json:"TokenType,omitempty"` + RefreshToken string `json:"RefreshToken,omitempty"` + IDToken string `json:"IdToken,omitempty"` +} + +type authResponse struct { + ChallengeName string `json:"ChallengeName,omitempty"` + Session string `json:"Session,omitempty"` + ChallengeParameters map[string]string `json:"ChallengeParameters"` + AuthenticationResult *authenticationResultJSON `json:"AuthenticationResult,omitempty"` +} + +func authToWire(res *driver.AuthResult) authResponse { + out := authResponse{ + ChallengeName: res.ChallengeName, + Session: res.Session, + ChallengeParameters: res.ChallengeParameters, + } + + if out.ChallengeParameters == nil { + out.ChallengeParameters = map[string]string{} + } + + if a := res.AuthenticationResult; a != nil { + out.AuthenticationResult = &authenticationResultJSON{ + AccessToken: a.AccessToken, + ExpiresIn: a.ExpiresIn, + TokenType: a.TokenType, + RefreshToken: a.RefreshToken, + IDToken: a.IDToken, + } + } + + return out +} + +type initiateAuthRequest struct { + UserPoolID string `json:"UserPoolId"` + ClientID string `json:"ClientId"` + AuthFlow string `json:"AuthFlow"` + AuthParameters map[string]string `json:"AuthParameters"` +} + +func (req *initiateAuthRequest) input() driver.InitiateAuthInput { + return driver.InitiateAuthInput{ + UserPoolID: req.UserPoolID, ClientID: req.ClientID, AuthFlow: req.AuthFlow, AuthParameters: req.AuthParameters, + } +} + +func (h *Handler) initiateAuth(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *initiateAuthRequest) (any, error) { + in := req.input() + in.UserPoolID = "" + + return authResult(h.cognito.InitiateAuth(ctx, in)) + }) +} + +func (h *Handler) adminInitiateAuth(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *initiateAuthRequest) (any, error) { + return authResult(h.cognito.AdminInitiateAuth(ctx, req.input())) + }) +} + +func authResult(res *driver.AuthResult, err error) (any, error) { + if err != nil { + return nil, err + } + + return authToWire(res), nil +} + +type respondRequest struct { + UserPoolID string `json:"UserPoolId"` + ClientID string `json:"ClientId"` + ChallengeName string `json:"ChallengeName"` + Session string `json:"Session"` + ChallengeResponses map[string]string `json:"ChallengeResponses"` +} + +func (req *respondRequest) input() driver.RespondToAuthChallengeInput { + return driver.RespondToAuthChallengeInput{ + UserPoolID: req.UserPoolID, ClientID: req.ClientID, ChallengeName: req.ChallengeName, + Session: req.Session, ChallengeResponses: req.ChallengeResponses, + } +} + +func (h *Handler) respondToAuthChallenge(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *respondRequest) (any, error) { + in := req.input() + in.UserPoolID = "" + + return authResult(h.cognito.RespondToAuthChallenge(ctx, in)) + }) +} + +func (h *Handler) adminRespondToAuthChallenge(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *respondRequest) (any, error) { + return authResult(h.cognito.AdminRespondToAuthChallenge(ctx, req.input())) + }) +} + +type accessTokenRequest struct { + AccessToken string `json:"AccessToken"` +} + +type getUserResponse struct { + Username string `json:"Username"` + UserAttributes []attributeJSON `json:"UserAttributes"` +} + +func (h *Handler) getUser(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *accessTokenRequest) (any, error) { + u, err := h.cognito.GetUser(ctx, req.AccessToken) + if err != nil { + return nil, err + } + + return getUserResponse{Username: u.Username, UserAttributes: attributesToWire(u.Attributes)}, nil + }) +} + +func (h *Handler) globalSignOut(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *accessTokenRequest) (any, error) { + if err := h.cognito.GlobalSignOut(ctx, req.AccessToken); err != nil { + return nil, err + } + + return struct{}{}, nil + }) +} + +type revokeTokenRequest struct { + Token string `json:"Token"` + ClientID string `json:"ClientId"` + ClientSecret string `json:"ClientSecret"` +} + +func (h *Handler) revokeToken(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *revokeTokenRequest) (any, error) { + err := h.cognito.RevokeToken(ctx, driver.RevokeTokenInput{ + Token: req.Token, ClientID: req.ClientID, ClientSecret: req.ClientSecret, + }) + if err != nil { + return nil, err + } + + return struct{}{}, nil + }) +} diff --git a/server/aws/cognito/auth_sdk_test.go b/server/aws/cognito/auth_sdk_test.go new file mode 100644 index 000000000..5cca62246 --- /dev/null +++ b/server/aws/cognito/auth_sdk_test.go @@ -0,0 +1,378 @@ +package cognito_test + +import ( + "context" + "crypto" + "crypto/rsa" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "math/big" + "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" + cip "github.com/aws/aws-sdk-go-v2/service/cognitoidentityprovider" + ciptypes "github.com/aws/aws-sdk-go-v2/service/cognitoidentityprovider/types" + + "github.com/stackshy/cloudemu/v2" + awsprovider "github.com/stackshy/cloudemu/v2/providers/aws" + awsserver "github.com/stackshy/cloudemu/v2/server/aws" +) + +// authEnv is a full AWS wire server with an admin (signed) and a public +// (anonymous) Cognito client. +type authEnv struct { + url string + admin *cip.Client + public *cip.Client + cloud *awsprovider.Provider +} + +func newAuthEnv(t *testing.T) *authEnv { + t.Helper() + + cloud := cloudemu.NewAWS() + ts := httptest.NewServer(awsserver.New(awsserver.DriversFrom(cloud))) + t.Cleanup(ts.Close) + + client := func(creds aws.CredentialsProvider) *cip.Client { + cfg, err := awsconfig.LoadDefaultConfig(context.Background(), + awsconfig.WithRegion("us-east-1"), awsconfig.WithCredentialsProvider(creds)) + if err != nil { + t.Fatalf("aws config: %v", err) + } + + return cip.NewFromConfig(cfg, func(o *cip.Options) { o.BaseEndpoint = aws.String(ts.URL) }) + } + + return &authEnv{ + url: ts.URL, + admin: client(credentials.NewStaticCredentialsProvider("test", "test", "")), + public: client(aws.AnonymousCredentials{}), + cloud: cloud, + } +} + +// jwks fetches the pool's key set over HTTP, as a JWT library would. +func (e *authEnv) jwks(t *testing.T, poolID string) map[string]*rsa.PublicKey { + t.Helper() + + req, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, e.url+"/"+poolID+"/.well-known/jwks.json", http.NoBody) + + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("GET jwks: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + t.Fatalf("GET jwks = %d", resp.StatusCode) + } + + var set struct { + Keys []struct { + Kid, Kty, Alg, Use, N, E string + } `json:"keys"` + } + + if err := json.NewDecoder(resp.Body).Decode(&set); err != nil { + t.Fatalf("decode jwks: %v", err) + } + + out := map[string]*rsa.PublicKey{} + + for _, k := range set.Keys { + if k.Kty != "RSA" || k.Alg != "RS256" || k.Use != "sig" { + t.Fatalf("unexpected JWK %+v", k) + } + + n, _ := base64.RawURLEncoding.DecodeString(k.N) + ex, _ := base64.RawURLEncoding.DecodeString(k.E) + out[k.Kid] = &rsa.PublicKey{N: new(big.Int).SetBytes(n), E: int(new(big.Int).SetBytes(ex).Int64())} + } + + return out +} + +// verify checks an RS256 JWT against a key set with crypto/rsa and returns +// its claims. +func verify(t *testing.T, keys map[string]*rsa.PublicKey, token string) map[string]any { + t.Helper() + + parts := strings.Split(token, ".") + if len(parts) != 3 { + t.Fatalf("token has %d segments", len(parts)) + } + + var header struct{ Alg, Kid string } + + raw, _ := base64.RawURLEncoding.DecodeString(parts[0]) + _ = json.Unmarshal(raw, &header) + + pub, ok := keys[header.Kid] + if !ok || header.Alg != "RS256" { + t.Fatalf("header %+v not in JWKS", header) + } + + sig, _ := base64.RawURLEncoding.DecodeString(parts[2]) + sum := sha256.Sum256([]byte(parts[0] + "." + parts[1])) + + if err := rsa.VerifyPKCS1v15(pub, crypto.SHA256, sum[:], sig); err != nil { + t.Fatalf("signature: %v", err) + } + + var claims map[string]any + + raw, _ = base64.RawURLEncoding.DecodeString(parts[1]) + _ = json.Unmarshal(raw, &claims) + + return claims +} + +func TestSDKSignUpSignInFlow(t *testing.T) { + ctx := context.Background() + e := newAuthEnv(t) + + pool, err := e.admin.CreateUserPool(ctx, &cip.CreateUserPoolInput{ + PoolName: aws.String("app"), + AutoVerifiedAttributes: []ciptypes.VerifiedAttributeType{ciptypes.VerifiedAttributeTypeEmail}, + }) + if err != nil { + t.Fatalf("CreateUserPool: %v", err) + } + + poolID := aws.ToString(pool.UserPool.Id) + + client, err := e.admin.CreateUserPoolClient(ctx, &cip.CreateUserPoolClientInput{ + UserPoolId: aws.String(poolID), ClientName: aws.String("web"), + ExplicitAuthFlows: []ciptypes.ExplicitAuthFlowsType{ + ciptypes.ExplicitAuthFlowsTypeAllowUserPasswordAuth, ciptypes.ExplicitAuthFlowsTypeAllowRefreshTokenAuth, + }, + }) + if err != nil { + t.Fatalf("CreateUserPoolClient: %v", err) + } + + clientID := client.UserPoolClient.ClientId + + _, err = e.public.SignUp(ctx, &cip.SignUpInput{ClientId: clientID, Username: aws.String("alice"), Password: aws.String("short")}) + requireErrorCode(t, err, "InvalidPasswordException", "Password did not conform with policy: Password not long enough") + + su, err := e.public.SignUp(ctx, &cip.SignUpInput{ + ClientId: clientID, Username: aws.String("alice"), Password: aws.String("Passw0rd!"), + UserAttributes: []ciptypes.AttributeType{{Name: aws.String("email"), Value: aws.String("alice@example.com")}}, + }) + if err != nil { + t.Fatalf("SignUp: %v", err) + } + + if su.UserConfirmed || su.CodeDeliveryDetails == nil || su.CodeDeliveryDetails.DeliveryMedium != ciptypes.DeliveryMediumTypeEmail { + t.Fatalf("SignUp = %+v", su) + } + + _, err = e.public.SignUp(ctx, &cip.SignUpInput{ClientId: clientID, Username: aws.String("alice"), Password: aws.String("Passw0rd!")}) + requireErrorCode(t, err, "UsernameExistsException", "User already exists") + + auth := &cip.InitiateAuthInput{ + ClientId: clientID, AuthFlow: ciptypes.AuthFlowTypeUserPasswordAuth, + AuthParameters: map[string]string{"USERNAME": "alice", "PASSWORD": "Passw0rd!"}, + } + + _, err = e.public.InitiateAuth(ctx, auth) + requireErrorCode(t, err, "UserNotConfirmedException", "User is not confirmed.") + + _, err = e.public.ConfirmSignUp(ctx, &cip.ConfirmSignUpInput{ClientId: clientID, Username: aws.String("alice"), ConfirmationCode: aws.String("x")}) + requireErrorCode(t, err, "CodeMismatchException", "") + + code, err := e.cloud.Cognito.ConfirmationCode(ctx, poolID, "alice") + if err != nil { + t.Fatalf("ConfirmationCode: %v", err) + } + + if _, err = e.public.ConfirmSignUp(ctx, &cip.ConfirmSignUpInput{ + ClientId: clientID, Username: aws.String("alice"), ConfirmationCode: aws.String(code.Code), + }); err != nil { + t.Fatalf("ConfirmSignUp: %v", err) + } + + if _, err = e.admin.CreateGroup(ctx, &cip.CreateGroupInput{UserPoolId: aws.String(poolID), GroupName: aws.String("admins")}); err != nil { + t.Fatalf("CreateGroup: %v", err) + } + + if _, err = e.admin.AdminAddUserToGroup(ctx, &cip.AdminAddUserToGroupInput{ + UserPoolId: aws.String(poolID), Username: aws.String("alice"), GroupName: aws.String("admins"), + }); err != nil { + t.Fatalf("AdminAddUserToGroup: %v", err) + } + + res, err := e.public.InitiateAuth(ctx, auth) + if err != nil { + t.Fatalf("InitiateAuth: %v", err) + } + + ar := res.AuthenticationResult + keys := e.jwks(t, poolID) + id := verify(t, keys, aws.ToString(ar.IdToken)) + access := verify(t, keys, aws.ToString(ar.AccessToken)) + + if id["iss"] != "https://cognito-idp.us-east-1.amazonaws.com/"+poolID || id["aud"] != aws.ToString(clientID) || + id["token_use"] != "id" || id["email"] != "alice@example.com" || access["token_use"] != "access" { + t.Fatalf("claims id=%v access=%v", id, access) + } + + if groups, _ := access["cognito:groups"].([]any); len(groups) != 1 || groups[0] != "admins" { + t.Fatalf("cognito:groups = %v", access["cognito:groups"]) + } + + u, err := e.public.GetUser(ctx, &cip.GetUserInput{AccessToken: ar.AccessToken}) + if err != nil || aws.ToString(u.Username) != "alice" { + t.Fatalf("GetUser = %+v, %v", u, err) + } + + ref, err := e.public.InitiateAuth(ctx, &cip.InitiateAuthInput{ + ClientId: clientID, AuthFlow: ciptypes.AuthFlowTypeRefreshTokenAuth, + AuthParameters: map[string]string{"REFRESH_TOKEN": aws.ToString(ar.RefreshToken)}, + }) + if err != nil || ref.AuthenticationResult.RefreshToken != nil { + t.Fatalf("refresh = %+v, %v", ref, err) + } + + if _, err = e.public.GlobalSignOut(ctx, &cip.GlobalSignOutInput{AccessToken: ar.AccessToken}); err != nil { + t.Fatalf("GlobalSignOut: %v", err) + } + + _, err = e.public.GetUser(ctx, &cip.GetUserInput{AccessToken: ar.AccessToken}) + requireErrorCode(t, err, "NotAuthorizedException", "Access Token has been revoked") + + _, err = e.public.InitiateAuth(ctx, &cip.InitiateAuthInput{ + ClientId: clientID, AuthFlow: ciptypes.AuthFlowTypeUserPasswordAuth, + AuthParameters: map[string]string{"USERNAME": "alice", "PASSWORD": "Wr0ng!pass"}, + }) + requireErrorCode(t, err, "NotAuthorizedException", "Incorrect username or password.") +} + +func TestSDKGroupsAndAdminAuth(t *testing.T) { + ctx := context.Background() + e := newAuthEnv(t) + poolID := createPool(t, e.admin, "groups") + + g, err := e.admin.CreateGroup(ctx, &cip.CreateGroupInput{ + UserPoolId: aws.String(poolID), GroupName: aws.String("ops"), Description: aws.String("Ops"), Precedence: aws.Int32(2), + }) + if err != nil || aws.ToInt32(g.Group.Precedence) != 2 || g.Group.CreationDate == nil { + t.Fatalf("CreateGroup = %+v, %v", g, err) + } + + _, err = e.admin.CreateGroup(ctx, &cip.CreateGroupInput{UserPoolId: aws.String(poolID), GroupName: aws.String("ops")}) + requireErrorCode(t, err, "GroupExistsException", "A group with the name ops already exists.") + + if _, err = e.admin.UpdateGroup(ctx, &cip.UpdateGroupInput{ + UserPoolId: aws.String(poolID), GroupName: aws.String("ops"), Description: aws.String("Operations"), + }); err != nil { + t.Fatalf("UpdateGroup: %v", err) + } + + got, err := e.admin.GetGroup(ctx, &cip.GetGroupInput{UserPoolId: aws.String(poolID), GroupName: aws.String("ops")}) + if err != nil || aws.ToString(got.Group.Description) != "Operations" || aws.ToInt32(got.Group.Precedence) != 2 { + t.Fatalf("GetGroup = %+v, %v", got, err) + } + + client, err := e.admin.CreateUserPoolClient(ctx, &cip.CreateUserPoolClientInput{ + UserPoolId: aws.String(poolID), ClientName: aws.String("srv"), + ExplicitAuthFlows: []ciptypes.ExplicitAuthFlowsType{ciptypes.ExplicitAuthFlowsTypeAllowAdminUserPasswordAuth}, + }) + if err != nil { + t.Fatalf("CreateUserPoolClient: %v", err) + } + + if _, err = e.admin.AdminCreateUser(ctx, &cip.AdminCreateUserInput{ + UserPoolId: aws.String(poolID), Username: aws.String("bob"), TemporaryPassword: aws.String("Temp0rary!"), + MessageAction: ciptypes.MessageActionTypeSuppress, + }); err != nil { + t.Fatalf("AdminCreateUser: %v", err) + } + + if _, err = e.admin.AdminAddUserToGroup(ctx, &cip.AdminAddUserToGroupInput{ + UserPoolId: aws.String(poolID), Username: aws.String("bob"), GroupName: aws.String("ops"), + }); err != nil { + t.Fatalf("AdminAddUserToGroup: %v", err) + } + + members, err := e.admin.ListUsersInGroup(ctx, &cip.ListUsersInGroupInput{UserPoolId: aws.String(poolID), GroupName: aws.String("ops")}) + if err != nil || len(members.Users) != 1 || aws.ToString(members.Users[0].Username) != "bob" { + t.Fatalf("ListUsersInGroup = %+v, %v", members, err) + } + + start, err := e.admin.AdminInitiateAuth(ctx, &cip.AdminInitiateAuthInput{ + UserPoolId: aws.String(poolID), ClientId: client.UserPoolClient.ClientId, AuthFlow: ciptypes.AuthFlowTypeAdminUserPasswordAuth, + AuthParameters: map[string]string{"USERNAME": "bob", "PASSWORD": "Temp0rary!"}, + }) + if err != nil || start.ChallengeName != ciptypes.ChallengeNameTypeNewPasswordRequired { + t.Fatalf("AdminInitiateAuth = %+v, %v", start, err) + } + + done, err := e.admin.AdminRespondToAuthChallenge(ctx, &cip.AdminRespondToAuthChallengeInput{ + UserPoolId: aws.String(poolID), ClientId: client.UserPoolClient.ClientId, + ChallengeName: ciptypes.ChallengeNameTypeNewPasswordRequired, Session: start.Session, + ChallengeResponses: map[string]string{"USERNAME": "bob", "NEW_PASSWORD": "N3wPassword!"}, + }) + if err != nil || done.AuthenticationResult == nil { + t.Fatalf("AdminRespondToAuthChallenge = %+v, %v", done, err) + } + + groups, err := e.admin.AdminListGroupsForUser(ctx, &cip.AdminListGroupsForUserInput{ + UserPoolId: aws.String(poolID), Username: aws.String("bob"), + }) + if err != nil || len(groups.Groups) != 1 { + t.Fatalf("AdminListGroupsForUser = %+v, %v", groups, err) + } + + if _, err = e.admin.DeleteGroup(ctx, &cip.DeleteGroupInput{UserPoolId: aws.String(poolID), GroupName: aws.String("ops")}); err != nil { + t.Fatalf("DeleteGroup: %v", err) + } + + list, err := e.admin.ListGroups(ctx, &cip.ListGroupsInput{UserPoolId: aws.String(poolID)}) + if err != nil || len(list.Groups) != 0 { + t.Fatalf("ListGroups = %+v, %v", list, err) + } +} + +func TestOpenIDConfiguration(t *testing.T) { + e := newAuthEnv(t) + poolID := createPool(t, e.admin, "oidc") + + req, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, + e.url+"/"+poolID+"/.well-known/openid-configuration", http.NoBody) + + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("GET: %v", err) + } + defer resp.Body.Close() + + var doc map[string]any + _ = json.NewDecoder(resp.Body).Decode(&doc) + + if doc["issuer"] != "https://cognito-idp.us-east-1.amazonaws.com/"+poolID || + doc["jwks_uri"] != e.url+"/"+poolID+"/.well-known/jwks.json" { + t.Fatalf("openid-configuration = %v", doc) + } + + req, _ = http.NewRequestWithContext(context.Background(), http.MethodGet, + e.url+"/us-east-1_nope12345/.well-known/jwks.json", http.NoBody) + + resp2, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("GET: %v", err) + } + resp2.Body.Close() + + if resp2.StatusCode != http.StatusNotFound { + t.Fatalf("unknown pool jwks = %d, want 404", resp2.StatusCode) + } +} diff --git a/server/aws/cognito/group_ops.go b/server/aws/cognito/group_ops.go new file mode 100644 index 000000000..f0d90b773 --- /dev/null +++ b/server/aws/cognito/group_ops.go @@ -0,0 +1,191 @@ +package cognito + +import ( + "context" + "net/http" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// groupJSON is the GroupType shape. +type groupJSON struct { + GroupName string `json:"GroupName"` + UserPoolID string `json:"UserPoolId"` + Description string `json:"Description,omitempty"` + RoleArn string `json:"RoleArn,omitempty"` + Precedence *int32 `json:"Precedence,omitempty"` + LastModifiedDate *float64 `json:"LastModifiedDate,omitempty"` + CreationDate *float64 `json:"CreationDate,omitempty"` +} + +func groupToWire(g *driver.Group) groupJSON { + return groupJSON{ + GroupName: g.GroupName, + UserPoolID: g.UserPoolID, + Description: g.Description, + RoleArn: g.RoleARN, + Precedence: g.Precedence, + LastModifiedDate: epochOrNil(g.LastModifiedDate), + CreationDate: epochOrNil(g.CreationDate), + } +} + +func groupsToWire(in []driver.Group) []groupJSON { + out := make([]groupJSON, len(in)) + for i := range in { + out[i] = groupToWire(&in[i]) + } + + return out +} + +type groupResponse struct { + Group groupJSON `json:"Group"` +} + +func groupResult(g *driver.Group, err error) (any, error) { + if err != nil { + return nil, err + } + + return groupResponse{Group: groupToWire(g)}, nil +} + +type createGroupRequest struct { + UserPoolID string `json:"UserPoolId"` + GroupName string `json:"GroupName"` + Description string `json:"Description"` + RoleArn string `json:"RoleArn"` + Precedence *int32 `json:"Precedence"` +} + +func (h *Handler) createGroup(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *createGroupRequest) (any, error) { + return groupResult(h.cognito.CreateGroup(ctx, driver.CreateGroupInput{ + UserPoolID: req.UserPoolID, GroupName: req.GroupName, Description: req.Description, + RoleARN: req.RoleArn, Precedence: req.Precedence, + })) + }) +} + +type groupRef struct { + UserPoolID string `json:"UserPoolId"` + GroupName string `json:"GroupName"` +} + +func (h *Handler) getGroup(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *groupRef) (any, error) { + return groupResult(h.cognito.GetGroup(ctx, req.UserPoolID, req.GroupName)) + }) +} + +type updateGroupRequest struct { + UserPoolID string `json:"UserPoolId"` + GroupName string `json:"GroupName"` + Description *string `json:"Description"` + RoleArn *string `json:"RoleArn"` + Precedence *int32 `json:"Precedence"` +} + +func (h *Handler) updateGroup(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *updateGroupRequest) (any, error) { + return groupResult(h.cognito.UpdateGroup(ctx, driver.UpdateGroupInput{ + UserPoolID: req.UserPoolID, GroupName: req.GroupName, Description: req.Description, + RoleARN: req.RoleArn, Precedence: req.Precedence, + })) + }) +} + +func (h *Handler) deleteGroup(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *groupRef) (any, error) { + if err := h.cognito.DeleteGroup(ctx, req.UserPoolID, req.GroupName); err != nil { + return nil, err + } + + return struct{}{}, nil + }) +} + +type listGroupsRequest struct { + UserPoolID string `json:"UserPoolId"` + Username string `json:"Username"` + GroupName string `json:"GroupName"` + Limit int32 `json:"Limit"` + NextToken string `json:"NextToken"` +} + +func (req *listGroupsRequest) page() driver.Pagination { + return driver.Pagination{NextToken: req.NextToken, MaxResults: req.Limit} +} + +type listGroupsResponse struct { + Groups []groupJSON `json:"Groups"` + NextToken string `json:"NextToken,omitempty"` +} + +func (h *Handler) listGroups(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *listGroupsRequest) (any, error) { + groups, next, err := h.cognito.ListGroups(ctx, req.UserPoolID, req.page()) + if err != nil { + return nil, err + } + + return listGroupsResponse{Groups: groupsToWire(groups), NextToken: next}, nil + }) +} + +func (h *Handler) adminListGroupsForUser(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *listGroupsRequest) (any, error) { + groups, next, err := h.cognito.AdminListGroupsForUser(ctx, req.UserPoolID, req.Username, req.page()) + if err != nil { + return nil, err + } + + return listGroupsResponse{Groups: groupsToWire(groups), NextToken: next}, nil + }) +} + +type listUsersInGroupResponse struct { + Users []userTypeJSON `json:"Users"` + NextToken string `json:"NextToken,omitempty"` +} + +func (h *Handler) listUsersInGroup(w http.ResponseWriter, r *http.Request) { + dispatch(h, w, r, func(h *Handler, ctx context.Context, req *listGroupsRequest) (any, error) { + users, next, err := h.cognito.ListUsersInGroup(ctx, req.UserPoolID, req.GroupName, req.page()) + if err != nil { + return nil, err + } + + out := make([]userTypeJSON, len(users)) + for i := range users { + out[i] = userToWire(&users[i]) + } + + return listUsersInGroupResponse{Users: out, NextToken: next}, nil + }) +} + +type userGroupRequest struct { + UserPoolID string `json:"UserPoolId"` + Username string `json:"Username"` + GroupName string `json:"GroupName"` +} + +func (h *Handler) adminAddUserToGroup(w http.ResponseWriter, r *http.Request) { + h.membershipOp(w, r, h.cognito.AdminAddUserToGroup) +} + +func (h *Handler) adminRemoveUserFromGroup(w http.ResponseWriter, r *http.Request) { + h.membershipOp(w, r, h.cognito.AdminRemoveUserFromGroup) +} + +func (h *Handler) membershipOp(w http.ResponseWriter, r *http.Request, call func(context.Context, string, string, string) error) { + dispatch(h, w, r, func(_ *Handler, ctx context.Context, req *userGroupRequest) (any, error) { + if err := call(ctx, req.UserPoolID, req.Username, req.GroupName); err != nil { + return nil, err + } + + return struct{}{}, nil + }) +} diff --git a/server/aws/cognito/handler.go b/server/aws/cognito/handler.go index 84ac5d441..97f36c481 100644 --- a/server/aws/cognito/handler.go +++ b/server/aws/cognito/handler.go @@ -31,6 +31,7 @@ type Handler struct { // New returns a Cognito handler backed by d. func New(d cognitodriver.Cognito) *Handler { h := &Handler{cognito: d} + //nolint:goconst // operation names, which publicOps also lists h.routes = map[string]http.HandlerFunc{ "CreateUserPool": h.createUserPool, "DescribeUserPool": h.describeUserPool, @@ -62,6 +63,29 @@ func New(d cognitodriver.Cognito) *Handler { "AdminEnableUser": h.adminEnableUser, "AdminDisableUser": h.adminDisableUser, "AdminResetUserPassword": h.adminResetUserPassword, + + "CreateGroup": h.createGroup, + "GetGroup": h.getGroup, + "UpdateGroup": h.updateGroup, + "DeleteGroup": h.deleteGroup, + "ListGroups": h.listGroups, + "AdminAddUserToGroup": h.adminAddUserToGroup, + "AdminRemoveUserFromGroup": h.adminRemoveUserFromGroup, + "AdminListGroupsForUser": h.adminListGroupsForUser, + "ListUsersInGroup": h.listUsersInGroup, + + "SignUp": h.signUp, + "ConfirmSignUp": h.confirmSignUp, + "ResendConfirmationCode": h.resendConfirmationCode, + "AdminConfirmSignUp": h.adminConfirmSignUp, + "InitiateAuth": h.initiateAuth, + "AdminInitiateAuth": h.adminInitiateAuth, + "RespondToAuthChallenge": h.respondToAuthChallenge, + "AdminRespondToAuthChallenge": h.adminRespondToAuthChallenge, + "GetUser": h.getUser, + "GlobalSignOut": h.globalSignOut, + "AdminUserGlobalSignOut": h.adminUserGlobalSignOut, + "RevokeToken": h.revokeToken, } return h diff --git a/server/aws/cognito/wellknown.go b/server/aws/cognito/wellknown.go new file mode 100644 index 000000000..373e92b2e --- /dev/null +++ b/server/aws/cognito/wellknown.go @@ -0,0 +1,115 @@ +package cognito + +import ( + "encoding/json" + "errors" + "net/http" + "regexp" + + cerrors "github.com/stackshy/cloudemu/v2/errors" + cognitodriver "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// wellKnownPath matches GET /{poolId}/.well-known/{jwks.json|openid-configuration}. +// A pool id such as us-east-1_AbC123 contains an underscore and uppercase +// letters, so it is never a valid S3 bucket name and the route cannot shadow a +// bucket. +var wellKnownPath = regexp.MustCompile(`^/([a-z]{2}(?:-[a-z]+)+-\d_[0-9A-Za-z]+)/\.well-known/(jwks\.json|openid-configuration)$`) + +const jwksDoc = "jwks.json" + +// WellKnown serves a user pool's public signing keys and OpenID Connect +// discovery document, the endpoints a JWT library fetches to verify Cognito +// tokens. Both are unauthenticated GETs on the real service. +type WellKnown struct { + keys cognitodriver.KeySetProvider +} + +// NewWellKnown returns the .well-known handler backed by keys. +func NewWellKnown(keys cognitodriver.KeySetProvider) *WellKnown { + return &WellKnown{keys: keys} +} + +// Matches claims GET requests for a pool's .well-known documents. +func (*WellKnown) Matches(r *http.Request) bool { + return r.Method == http.MethodGet && wellKnownPath.MatchString(r.URL.Path) +} + +// PublicRequest reports that the .well-known documents need no credentials. +// It answers true only for the exact routes Matches claims. +func (w *WellKnown) PublicRequest(r *http.Request) bool { return w.Matches(r) } + +// IAMService returns the IAM service prefix of the user-pool documents. They +// are public, so no request reaches IAM authorization. +func (*WellKnown) IAMService() string { return "cognito-idp" } + +// openIDConfiguration is the discovery document Cognito publishes for a pool. +type openIDConfiguration struct { + Issuer string `json:"issuer"` + JWKSURI string `json:"jwks_uri"` + IDTokenSigningAlgValuesSupported []string `json:"id_token_signing_alg_values_supported"` + ResponseTypesSupported []string `json:"response_types_supported"` + ScopesSupported []string `json:"scopes_supported"` + SubjectTypesSupported []string `json:"subject_types_supported"` + TokenEndpointAuthMethods []string `json:"token_endpoint_auth_methods_supported"` +} + +// ServeHTTP writes the JWKS or the discovery document. +func (w *WellKnown) ServeHTTP(rw http.ResponseWriter, r *http.Request) { + m := wellKnownPath.FindStringSubmatch(r.URL.Path) + if m == nil { + http.NotFound(rw, r) + + return + } + + poolID, doc := m[1], m[2] + + iss, set, err := w.keys.SigningKeys(r.Context(), poolID) + if err != nil { + writeWellKnownError(rw, err) + + return + } + + var body any = set + + if doc != jwksDoc { + body = openIDConfiguration{ + Issuer: iss, + JWKSURI: requestBase(r) + "/" + poolID + "/.well-known/" + jwksDoc, + IDTokenSigningAlgValuesSupported: []string{"RS256"}, + ResponseTypesSupported: []string{"code", "token"}, + ScopesSupported: []string{"openid", "email", "phone", "profile"}, + SubjectTypesSupported: []string{"public"}, + TokenEndpointAuthMethods: []string{"client_secret_basic", "client_secret_post"}, + } + } + + rw.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(rw).Encode(body) +} + +// requestBase is the scheme and host the client reached the emulator on, so +// jwks_uri points back at this server rather than at AWS. +func requestBase(r *http.Request) string { + scheme := "http" + if r.TLS != nil { + scheme = "https" + } + + return scheme + "://" + r.Host +} + +func writeWellKnownError(rw http.ResponseWriter, err error) { + status := http.StatusInternalServerError + + var apiErr *cognitodriver.APIError + if errors.As(err, &apiErr) && apiErr.Exception == cognitodriver.ExResourceNotFound { + status = http.StatusNotFound + } + + rw.Header().Set("Content-Type", "application/json") + rw.WriteHeader(status) + _ = json.NewEncoder(rw).Encode(map[string]string{"message": cerrors.Message(err)}) +} diff --git a/server/aws/cognito_auth_enforced_test.go b/server/aws/cognito_auth_enforced_test.go new file mode 100644 index 000000000..345402c25 --- /dev/null +++ b/server/aws/cognito_auth_enforced_test.go @@ -0,0 +1,88 @@ +package aws + +import ( + "context" + "encoding/json" + "net/http" + "strings" + "testing" + + cognitodriver "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// TestEnforcedCognitoSignInIsPublic drives sign-up, confirmation, sign-in and +// the JWKS fetch with no SigV4 at all under --enforce-auth, and checks that the +// admin operations and other .well-known paths still need credentials. +func TestEnforcedCognitoSignInIsPublic(t *testing.T) { + ts, cloud := enforcedServer(t) + ctx := context.Background() + + pool, err := cloud.Cognito.CreateUserPool(ctx, cognitodriver.CreateUserPoolInput{ + Name: "enforced", AutoVerifiedAttributes: []string{"email"}, + }) + if err != nil { + t.Fatalf("CreateUserPool: %v", err) + } + + client, err := cloud.Cognito.CreateUserPoolClient(ctx, cognitodriver.CreateUserPoolClientInput{ + UserPoolID: pool.ID, ClientName: "web", ExplicitAuthFlows: []string{"ALLOW_USER_PASSWORD_AUTH", "ALLOW_REFRESH_TOKEN_AUTH"}, + }) + if err != nil { + t.Fatalf("CreateUserPoolClient: %v", err) + } + + call := func(op, body string) (int, string) { + return doRaw(t, ts, jsonRPC(idpTarget+op, body)) + } + + if status, body := call("SignUp", `{"ClientId":"`+client.ClientID+`","Username":"alice","Password":"Passw0rd!",`+ + `"UserAttributes":[{"Name":"email","Value":"alice@example.com"}]}`); status != http.StatusOK { + t.Fatalf("SignUp = %d %s", status, body) + } + + code, err := cloud.Cognito.ConfirmationCode(ctx, pool.ID, "alice") + if err != nil { + t.Fatalf("ConfirmationCode: %v", err) + } + + if status, body := call("ConfirmSignUp", `{"ClientId":"`+client.ClientID+`","Username":"alice","ConfirmationCode":"`+ + code.Code+`"}`); status != http.StatusOK { + t.Fatalf("ConfirmSignUp = %d %s", status, body) + } + + status, body := call("InitiateAuth", `{"ClientId":"`+client.ClientID+`","AuthFlow":"USER_PASSWORD_AUTH",`+ + `"AuthParameters":{"USERNAME":"alice","PASSWORD":"Passw0rd!"}}`) + if status != http.StatusOK { + t.Fatalf("InitiateAuth = %d %s", status, body) + } + + var auth struct { + AuthenticationResult struct{ AccessToken string } `json:"AuthenticationResult"` + } + _ = json.Unmarshal([]byte(body), &auth) + + if status, body = call("GetUser", `{"AccessToken":"`+auth.AuthenticationResult.AccessToken+`"}`); status != http.StatusOK { + t.Fatalf("GetUser = %d %s", status, body) + } + + if status, body = doRaw(t, ts, rawReq{method: http.MethodGet, path: "/" + pool.ID + "/.well-known/jwks.json"}); status != http.StatusOK || + !strings.Contains(body, `"keys"`) { + t.Fatalf("jwks = %d %s", status, body) + } + + denied := []struct { + name string + req rawReq + }{ + {"CreateGroup", jsonRPC(idpTarget+"CreateGroup", `{"UserPoolId":"`+pool.ID+`","GroupName":"g"}`)}, + {"AdminInitiateAuth", jsonRPC(idpTarget+"AdminInitiateAuth", `{}`)}, + {"other well-known path", rawReq{method: http.MethodGet, path: "/" + pool.ID + "/.well-known/other"}}, + {"POST jwks", rawReq{method: http.MethodPost, path: "/" + pool.ID + "/.well-known/jwks.json"}}, + } + + for _, tc := range denied { + if status, body := doRaw(t, ts, tc.req); status != http.StatusForbidden || !strings.Contains(body, missingTok) { + t.Fatalf("%s unsigned = %d %s, want 403 %s", tc.name, status, body, missingTok) + } + } +} diff --git a/server/aws/publicauth_test.go b/server/aws/publicauth_test.go index cc880eb2c..91495e762 100644 --- a/server/aws/publicauth_test.go +++ b/server/aws/publicauth_test.go @@ -105,8 +105,8 @@ func TestEnforcedGateAdmitsUnsignedPublicOps(t *testing.T) { req rawReq want int }{ - // Cognito user-pool public ops are not routed yet, so the handler answers - // UnknownOperationException (400): the point is the gate let them through. + // The empty bodies fail Cognito's own validation (400): the point is the + // gate let them through. {"cognito-idp InitiateAuth", jsonRPC(idpTarget+"InitiateAuth", `{}`), http.StatusBadRequest}, {"cognito-idp SignUp", jsonRPC(idpTarget+"SignUp", `{}`), http.StatusBadRequest}, {"cognito-idp RespondToAuthChallenge", jsonRPC(idpTarget+"RespondToAuthChallenge", `{}`), http.StatusBadRequest}, diff --git a/server/serverkit/cognito_codes.go b/server/serverkit/cognito_codes.go new file mode 100644 index 000000000..b19a50418 --- /dev/null +++ b/server/serverkit/cognito_codes.go @@ -0,0 +1,98 @@ +package serverkit + +import ( + "encoding/json" + "errors" + "net/http" + "strings" + "time" + + cerrors "github.com/stackshy/cloudemu/v2/errors" + awsprovider "github.com/stackshy/cloudemu/v2/providers/aws" + cognitodriver "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// cognitoCodesPath is the admin endpoint that reads back the confirmation code +// the emulator would have emailed or texted to a Cognito user: +// GET /_cloudemu/cognito/codes?userPoolId=&username=. +const cognitoCodesPath = "cognito/codes" + +type cognitoCodeJSON struct { + UserPoolID string `json:"userPoolId"` + Username string `json:"username"` + Code string `json:"code"` + ExpiresAt string `json:"expiresAt"` + DeliveryMedium string `json:"deliveryMedium,omitempty"` + Destination string `json:"destination,omitempty"` + AttributeName string `json:"attributeName,omitempty"` +} + +// serveCognitoCodes answers the codes endpoint from the live AWS regions. +func (a *App) serveCognitoCodes(w http.ResponseWriter, r *http.Request) { + a.rebuildMu.Lock() + mux := a.awsMux + a.rebuildMu.Unlock() + + if mux == nil { + writeNetErr(w, http.StatusServiceUnavailable, "cognito codes require the aws provider") + + return + } + + serveCognitoCode(w, r, mux.LiveProviders()) +} + +// serveCognitoCode looks the code up in the region the pool id names +// ("_"). +func serveCognitoCode(w http.ResponseWriter, r *http.Request, regions map[string]*awsprovider.Provider) { + if r.Method != http.MethodGet { + writeNetErr(w, http.StatusMethodNotAllowed, "cognito/codes requires GET") + + return + } + + poolID := r.URL.Query().Get("userPoolId") + username := r.URL.Query().Get("username") + + region, _, ok := strings.Cut(poolID, "_") + if !ok || region == "" || username == "" { + writeNetErr(w, http.StatusBadRequest, "userPoolId and username are required") + + return + } + + prov, ok := regions[region] + if !ok { + writeNetErr(w, http.StatusNotFound, "User pool "+poolID+" does not exist.") + + return + } + + code, err := prov.Cognito.ConfirmationCode(r.Context(), poolID, username) + if err != nil { + status := http.StatusInternalServerError + + var apiErr *cognitodriver.APIError + if errors.As(err, &apiErr) { + status = http.StatusNotFound + } + + writeNetErr(w, status, cerrors.Message(err)) + + return + } + + out := cognitoCodeJSON{ + UserPoolID: poolID, + Username: username, + Code: code.Code, + ExpiresAt: code.ExpiresAt.UTC().Format(time.RFC3339), + } + + if d := code.Delivery; d != nil { + out.DeliveryMedium, out.Destination, out.AttributeName = d.DeliveryMedium, d.Destination, d.AttributeName + } + + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(out) +} diff --git a/server/serverkit/cognito_codes_test.go b/server/serverkit/cognito_codes_test.go new file mode 100644 index 000000000..a9026b04c --- /dev/null +++ b/server/serverkit/cognito_codes_test.go @@ -0,0 +1,71 @@ +package serverkit + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + cloudemu "github.com/stackshy/cloudemu/v2" + awsprovider "github.com/stackshy/cloudemu/v2/providers/aws" + cognitodriver "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +func TestServeCognitoCode(t *testing.T) { + ctx := context.Background() + cloud := cloudemu.NewAWS() + + pool, err := cloud.Cognito.CreateUserPool(ctx, cognitodriver.CreateUserPoolInput{ + Name: "codes", AutoVerifiedAttributes: []string{"email"}, + }) + if err != nil { + t.Fatalf("CreateUserPool: %v", err) + } + + client, err := cloud.Cognito.CreateUserPoolClient(ctx, cognitodriver.CreateUserPoolClientInput{UserPoolID: pool.ID, ClientName: "c"}) + if err != nil { + t.Fatalf("CreateUserPoolClient: %v", err) + } + + if _, err = cloud.Cognito.SignUp(ctx, cognitodriver.SignUpInput{ + ClientUserInput: cognitodriver.ClientUserInput{ClientID: client.ClientID, Username: "alice"}, + Password: "Passw0rd!", + UserAttributes: []cognitodriver.Attribute{{Name: "email", Value: "alice@example.com"}}, + }); err != nil { + t.Fatalf("SignUp: %v", err) + } + + regions := map[string]*awsprovider.Provider{"us-east-1": cloud} + get := func(query string) *httptest.ResponseRecorder { + rec := httptest.NewRecorder() + serveCognitoCode(rec, httptest.NewRequest(http.MethodGet, "/_cloudemu/cognito/codes?"+query, nil), regions) + + return rec + } + + rec := get("userPoolId=" + pool.ID + "&username=alice") + if rec.Code != http.StatusOK { + t.Fatalf("status = %d (%s)", rec.Code, rec.Body.String()) + } + + var out cognitoCodeJSON + if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil { + t.Fatalf("decode: %v", err) + } + + want, _ := cloud.Cognito.ConfirmationCode(ctx, pool.ID, "alice") + if out.Code != want.Code || out.DeliveryMedium != "EMAIL" || out.Destination != "a***@e***" { + t.Fatalf("codes = %+v", out) + } + + for query, status := range map[string]int{ + "userPoolId=" + pool.ID + "&username=ghost": http.StatusNotFound, + "userPoolId=eu-west-1_abc&username=alice": http.StatusNotFound, + "username=alice": http.StatusBadRequest, + } { + if rec := get(query); rec.Code != status { + t.Fatalf("%s: status = %d, want %d", query, rec.Code, status) + } + } +} diff --git a/server/serverkit/serverkit.go b/server/serverkit/serverkit.go index 267cc79cc..76e5c78e9 100644 --- a/server/serverkit/serverkit.go +++ b/server/serverkit/serverkit.go @@ -964,6 +964,8 @@ func (a *App) extraHandler() http.Handler { a.rebuildMu.Unlock() serveCost(w, r, ds) + case cognitoCodesPath: + a.serveCognitoCodes(w, r) default: if strings.HasPrefix(strings.TrimPrefix(r.URL.Path, admin.Prefix), timeTravelPrefix) { a.serveTimeTravel(w, r) diff --git a/services/cognito/driver/auth_types.go b/services/cognito/driver/auth_types.go new file mode 100644 index 000000000..568582c23 --- /dev/null +++ b/services/cognito/driver/auth_types.go @@ -0,0 +1,143 @@ +package driver + +import "time" + +// Auth flows accepted by InitiateAuth and AdminInitiateAuth. +const ( + AuthFlowUserPassword = "USER_PASSWORD_AUTH" + AuthFlowAdminUserPassword = "ADMIN_USER_PASSWORD_AUTH" + AuthFlowAdminNoSRP = "ADMIN_NO_SRP_AUTH" + AuthFlowRefreshToken = "REFRESH_TOKEN" + AuthFlowRefreshTokenAuth = "REFRESH_TOKEN_AUTH" + AuthFlowUserSRP = "USER_SRP_AUTH" + AuthFlowCustom = "CUSTOM_AUTH" + AuthFlowUser = "USER_AUTH" +) + +// ChallengeNewPasswordRequired is the challenge a FORCE_CHANGE_PASSWORD user +// gets on sign-in. +const ChallengeNewPasswordRequired = "NEW_PASSWORD_REQUIRED" + +// Delivery media reported in CodeDeliveryDetails. +const ( + DeliveryMediumEmail = "EMAIL" + DeliveryMediumSMS = "SMS" +) + +// Group is a user-pool group. +type Group struct { + GroupName string + UserPoolID string + Description string + RoleARN string + Precedence *int32 + CreationDate time.Time + LastModifiedDate time.Time +} + +// CreateGroupInput is the input to CreateGroup. +type CreateGroupInput struct { + UserPoolID string + GroupName string + Description string + RoleARN string + Precedence *int32 +} + +// UpdateGroupInput is the input to UpdateGroup. A nil field is left unchanged. +type UpdateGroupInput struct { + UserPoolID string + GroupName string + Description *string + RoleARN *string + Precedence *int32 +} + +// ClientUserInput names a user through an app client, with the SECRET_HASH the +// client needs when it has a secret. +type ClientUserInput struct { + ClientID string + SecretHash string + Username string +} + +// SignUpInput is the input to SignUp. +type SignUpInput struct { + ClientUserInput + Password string + UserAttributes []Attribute +} + +// SignUpOutput is the result of SignUp. +type SignUpOutput struct { + UserConfirmed bool + UserSub string + CodeDeliveryDetails *CodeDeliveryDetails +} + +// CodeDeliveryDetails describes where a confirmation code was sent. The +// destination is masked the way Cognito masks it. +type CodeDeliveryDetails struct { + Destination string + DeliveryMedium string + AttributeName string +} + +// ConfirmSignUpInput is the input to ConfirmSignUp. +type ConfirmSignUpInput struct { + ClientUserInput + ConfirmationCode string + ForceAliasCreation bool +} + +// IssuedCode is the confirmation code currently outstanding for a user. +type IssuedCode struct { + Code string + ExpiresAt time.Time + Delivery *CodeDeliveryDetails +} + +// InitiateAuthInput is the input to InitiateAuth and AdminInitiateAuth. +// UserPoolID is set only by AdminInitiateAuth. +type InitiateAuthInput struct { + UserPoolID string + ClientID string + AuthFlow string + AuthParameters map[string]string +} + +// RespondToAuthChallengeInput is the input to RespondToAuthChallenge and +// AdminRespondToAuthChallenge. UserPoolID is set only by the admin variant. +type RespondToAuthChallengeInput struct { + UserPoolID string + ClientID string + ChallengeName string + Session string + ChallengeResponses map[string]string +} + +// AuthResult is a sign-in step: either a challenge to answer or the issued +// tokens. +type AuthResult struct { + ChallengeName string + Session string + ChallengeParameters map[string]string + AuthenticationResult *AuthenticationResult +} + +// AuthenticationResult carries the issued tokens. RefreshToken is empty on a +// refresh. +type AuthenticationResult struct { + AccessToken string + IDToken string + RefreshToken string + ExpiresIn int32 + TokenType string +} + +// RevokeTokenInput is the input to RevokeToken. +type RevokeTokenInput struct { + Token string + ClientID string + ClientSecret string +} diff --git a/services/cognito/driver/driver.go b/services/cognito/driver/driver.go index 6a4d78da9..5580e4ce3 100644 --- a/services/cognito/driver/driver.go +++ b/services/cognito/driver/driver.go @@ -1,12 +1,14 @@ // 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. -// -// Sign-up, sign-in and token issuance are not modeled yet. +// resource tagging, the users of a pool with the admin user-management +// operations, groups, self sign-up, and password sign-in with RS256 tokens. package driver -import "context" +import ( + "context" + + "github.com/stackshy/cloudemu/v2/internal/jwtsign" +) // Cognito is the interface an AWS Cognito user-pools backend implements. It // covers user pools, app clients, hosted-UI domains, and resource tagging. @@ -16,6 +18,9 @@ type Cognito interface { userPoolDomainAPI tagAPI userAPI + groupAPI + signUpAPI + authAPI } // userPoolAPI covers the user-pool control plane. @@ -105,3 +110,67 @@ type userAPI interface { // AdminResetUserPassword moves the user to RESET_REQUIRED. AdminResetUserPassword(ctx context.Context, userPoolID, username string) error } + +// groupAPI covers user-pool groups and group membership. +type groupAPI interface { + // CreateGroup creates a group. A name already in the pool fails with + // GroupExistsException. + CreateGroup(ctx context.Context, in CreateGroupInput) (*Group, error) + GetGroup(ctx context.Context, userPoolID, groupName string) (*Group, error) + // UpdateGroup changes the fields the input sets and returns the group. + UpdateGroup(ctx context.Context, in UpdateGroupInput) (*Group, error) + // DeleteGroup removes a group and every membership in it. + DeleteGroup(ctx context.Context, userPoolID, groupName string) error + ListGroups(ctx context.Context, userPoolID string, page Pagination) ([]Group, string, error) + AdminAddUserToGroup(ctx context.Context, userPoolID, username, groupName string) error + AdminRemoveUserFromGroup(ctx context.Context, userPoolID, username, groupName string) error + AdminListGroupsForUser(ctx context.Context, userPoolID, username string, page Pagination) ([]Group, string, error) + ListUsersInGroup(ctx context.Context, userPoolID, groupName string, page Pagination) ([]User, string, error) +} + +// signUpAPI covers self-service registration through an app client. +type signUpAPI interface { + // SignUp registers an UNCONFIRMED user, checks the password against the + // pool policy, and issues a confirmation code. + SignUp(ctx context.Context, in SignUpInput) (*SignUpOutput, error) + // ConfirmSignUp confirms a user with the code SignUp or + // ResendConfirmationCode issued. + ConfirmSignUp(ctx context.Context, in ConfirmSignUpInput) error + ResendConfirmationCode(ctx context.Context, in ClientUserInput) (*CodeDeliveryDetails, error) + AdminConfirmSignUp(ctx context.Context, userPoolID, username string) error +} + +// authAPI covers password sign-in, challenges, and the token-authenticated +// user operations. Tokens are RS256 JWTs signed with per-pool keys. +type authAPI interface { + // InitiateAuth runs USER_PASSWORD_AUTH or REFRESH_TOKEN_AUTH for a client. + InitiateAuth(ctx context.Context, in InitiateAuthInput) (*AuthResult, error) + // AdminInitiateAuth runs ADMIN_USER_PASSWORD_AUTH, ADMIN_NO_SRP_AUTH or + // REFRESH_TOKEN_AUTH for a client of in.UserPoolID. + AdminInitiateAuth(ctx context.Context, in InitiateAuthInput) (*AuthResult, error) + // RespondToAuthChallenge answers the NEW_PASSWORD_REQUIRED challenge. + RespondToAuthChallenge(ctx context.Context, in RespondToAuthChallengeInput) (*AuthResult, error) + AdminRespondToAuthChallenge(ctx context.Context, in RespondToAuthChallengeInput) (*AuthResult, error) + // GetUser returns the user an access token was issued to. + GetUser(ctx context.Context, accessToken string) (*User, error) + // GlobalSignOut revokes every token issued to the access token's user. + GlobalSignOut(ctx context.Context, accessToken string) error + AdminUserGlobalSignOut(ctx context.Context, userPoolID, username string) error + // RevokeToken revokes a refresh token and the access and ID tokens minted + // from the same authentication. + RevokeToken(ctx context.Context, in RevokeTokenInput) error +} + +// KeySetProvider publishes a user pool's token-signing keys and issuer for the +// /{poolId}/.well-known endpoints. It is not a cognito-idp API operation. +type KeySetProvider interface { + // SigningKeys returns the pool's issuer URL and its public JSON Web Key Set. + SigningKeys(ctx context.Context, userPoolID string) (issuer string, keys jwtsign.JWKSet, err error) +} + +// CodeInspector exposes the confirmation codes the emulator would have sent by +// email or SMS, so a test can complete ConfirmSignUp. It is not a cognito-idp +// API operation. +type CodeInspector interface { + ConfirmationCode(ctx context.Context, userPoolID, username string) (*IssuedCode, error) +} diff --git a/services/cognito/driver/errors.go b/services/cognito/driver/errors.go index 4500ccd80..f39dc5efb 100644 --- a/services/cognito/driver/errors.go +++ b/services/cognito/driver/errors.go @@ -32,3 +32,15 @@ type APIError struct { func (e *APIError) Error() string { return e.Err.Error() } func (e *APIError) Unwrap() error { return e.Err } + +// Exception names for groups, sign-up and sign-in. +const ( + ExGroupExists = "GroupExistsException" + ExCodeMismatch = "CodeMismatchException" + ExExpiredCode = "ExpiredCodeException" + ExUserNotConfirmed = "UserNotConfirmedException" + ExPasswordResetRequired = "PasswordResetRequiredException" + ExUnsupportedOperation = "UnsupportedOperationException" + ExUnsupportedTokenType = "UnsupportedTokenTypeException" + ExUnauthorized = "UnauthorizedException" +) From 0a6bfea0dfea055674c795080ef0248cedbb9b28 Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 4 Oct 2026 17:25:32 +0530 Subject: [PATCH 2/3] fix(aws-cognito): recheck user state on challenge and refresh, hide user existence, validate token validity --- providers/aws/cognito/auth.go | 47 ++++- providers/aws/cognito/auth_state_test.go | 223 +++++++++++++++++++++ providers/aws/cognito/sign_up.go | 94 +++++++-- providers/aws/cognito/token_validity.go | 77 +++++++ providers/aws/cognito/tokens.go | 6 +- providers/aws/cognito/user_pool_clients.go | 14 +- server/aws/cognito/handler.go | 5 +- server/aws/cognito/sdk_roundtrip_test.go | 6 +- server/aws/cognito/wellknown.go | 2 +- 9 files changed, 436 insertions(+), 38 deletions(-) create mode 100644 providers/aws/cognito/auth_state_test.go create mode 100644 providers/aws/cognito/token_validity.go diff --git a/providers/aws/cognito/auth.go b/providers/aws/cognito/auth.go index c9631c063..2fd948b82 100644 --- a/providers/aws/cognito/auth.go +++ b/providers/aws/cognito/auth.go @@ -22,6 +22,7 @@ const ( userAttrPrefix = "userAttributes." sessionBytes = 96 existenceEnabled = "ENABLED" + dummySalt = "00000000000000000000000000000000" allowFlowPrefix = "ALLOW_" legacyAdminNoSRP = "ADMIN_NO_SRP_AUTH" allowAdminUserPwd = "ALLOW_ADMIN_USER_PASSWORD_AUTH" @@ -42,6 +43,7 @@ type challengeSession struct { poolID string clientID string username string + sub string challenge string expires time.Time } @@ -217,6 +219,10 @@ func (m *Mock) passwordAuth(pool *driver.UserPool, client driver.UserPoolClient, } if !found { + // Spend the same hashing work as a real check so the response time + // does not tell an unknown username from a wrong password. + _ = pbkdf2Hash(dummySalt, password, pbkdf2Iterations) + return nil, unknownUser(&client) } @@ -299,6 +305,7 @@ func (m *Mock) newPasswordChallenge(pool *driver.UserPool, client *driver.UserPo poolID: pool.ID, clientID: client.ClientID, username: rec.User.Username, + sub: attrValue(rec.User.Attributes, attrSub), challenge: driver.ChallengeNewPasswordRequired, expires: m.now().Add(time.Duration(client.AuthSessionValidity) * time.Minute), } @@ -374,8 +381,8 @@ func (m *Mock) refreshUser(l *loginRecord) (userRecord, error) { return userRecord{}, invalidRefreshToken() } - if !rec.User.Enabled { - return userRecord{}, notAuthorized("User is disabled.") + if err := checkCanSignIn(&rec); err != nil { + return userRecord{}, err } return rec, nil @@ -414,17 +421,11 @@ func (m *Mock) respond(in *driver.RespondToAuthChallengeInput, admin bool) (*dri return nil, err } - key, rec, found := m.resolveUser(&pool, username) - if !found || rec.User.Username != sess.username { - return nil, invalidSession() - } - - if err := checkSecretHash(client, in.ChallengeResponses[paramSecretHash], noSecretReceived(client.ClientID), - username, rec.User.Username); err != nil { + key, rec, err := m.challengeUser(&pool, &client, &sess, username, in.ChallengeResponses[paramSecretHash]) + if err != nil { return nil, err } - rec = copyUserRecord(rec) if err := m.applyNewPassword(&pool, key, &rec, in.ChallengeResponses); err != nil { return nil, err } @@ -435,6 +436,32 @@ func (m *Mock) respond(in *driver.RespondToAuthChallengeInput, admin bool) (*dri return m.signIn(client, &rec) } +// challengeUser finds the user a challenge answer names and checks it is still +// the user the session was issued to and still in the challenge's state. An +// admin can disable, reset or confirm the user between the two calls. +func (m *Mock) challengeUser( + pool *driver.UserPool, client *driver.UserPoolClient, sess *challengeSession, username, hash string, +) (string, userRecord, error) { + key, rec, found := m.resolveUser(pool, username) + if !found || rec.User.Username != sess.username || attrValue(rec.User.Attributes, attrSub) != sess.sub { + return "", userRecord{}, invalidSession() + } + + if err := checkSecretHash(*client, hash, noSecretReceived(client.ClientID), username, rec.User.Username); err != nil { + return "", userRecord{}, err + } + + if err := checkCanSignIn(&rec); err != nil { + return "", userRecord{}, err + } + + if rec.User.UserStatus != driver.UserStatusForceChangePassword { + return "", userRecord{}, invalidSession() + } + + return key, copyUserRecord(rec), nil +} + // takeSession returns a live session for the client and challenge. An expired // session is dropped. func (m *Mock) takeSession(id, clientID, challenge string) (challengeSession, error) { diff --git a/providers/aws/cognito/auth_state_test.go b/providers/aws/cognito/auth_state_test.go new file mode 100644 index 000000000..6d6e01756 --- /dev/null +++ b/providers/aws/cognito/auth_state_test.go @@ -0,0 +1,223 @@ +package cognito + +import ( + "context" + "testing" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// startNewPasswordChallenge creates a FORCE_CHANGE_PASSWORD user and opens its +// NEW_PASSWORD_REQUIRED challenge. +func startNewPasswordChallenge(t *testing.T, m *Mock, poolID, clientID string) driver.RespondToAuthChallengeInput { + t.Helper() + + ctx := context.Background() + + _, err := m.AdminCreateUser(ctx, driver.AdminCreateUserInput{ + UserPoolID: poolID, Username: "temp", TemporaryPassword: "Temp0rary!", MessageAction: driver.MessageActionSuppress, + }) + requireNoError(t, err, "AdminCreateUser") + + res, err := m.InitiateAuth(ctx, passwordAuth(clientID, "temp", "Temp0rary!")) + requireNoError(t, err, "InitiateAuth") + + if res.ChallengeName != driver.ChallengeNewPasswordRequired { + t.Fatalf("challenge = %q", res.ChallengeName) + } + + return driver.RespondToAuthChallengeInput{ + ClientID: clientID, ChallengeName: driver.ChallengeNewPasswordRequired, Session: res.Session, + ChallengeResponses: map[string]string{"USERNAME": "temp", "NEW_PASSWORD": "N3wPassword!"}, + } +} + +func TestRespondRejectsUserDisabledMidChallenge(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + + respond := startNewPasswordChallenge(t, m, pool.ID, client.ClientID) + requireNoError(t, m.AdminDisableUser(ctx, pool.ID, "temp"), "AdminDisableUser") + + _, err := m.RespondToAuthChallenge(ctx, respond) + assertException(t, err, driver.ExNotAuthorized, "User is disabled.") + + u, _ := m.AdminGetUser(ctx, pool.ID, "temp") + if u.Enabled || u.UserStatus != driver.UserStatusForceChangePassword { + t.Fatalf("user = enabled %v status %s, want disabled and still FORCE_CHANGE_PASSWORD", u.Enabled, u.UserStatus) + } + + if n := len(m.logins.Keys()); n != 0 { + t.Fatalf("%d logins recorded for a disabled user", n) + } +} + +func TestRespondRejectsUserNoLongerInChallengeState(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + + respond := startNewPasswordChallenge(t, m, pool.ID, client.ClientID) + requireNoError(t, m.AdminSetUserPassword(ctx, pool.ID, "temp", "Perm4nent!pw", true), "AdminSetUserPassword") + + _, err := m.RespondToAuthChallenge(ctx, respond) + assertException(t, err, driver.ExNotAuthorized, "Invalid session for the user.") + + // A user deleted and re-created under the same name cannot use the old + // user's session. + respond = startNewPasswordChallengeAgain(t, m, pool.ID, client.ClientID) + requireNoError(t, m.AdminDeleteUser(ctx, pool.ID, "temp"), "AdminDeleteUser") + _, err = m.AdminCreateUser(ctx, driver.AdminCreateUserInput{ + UserPoolID: pool.ID, Username: "temp", TemporaryPassword: "Temp0rary!", MessageAction: driver.MessageActionSuppress, + }) + requireNoError(t, err, "AdminCreateUser again") + + _, err = m.RespondToAuthChallenge(ctx, respond) + assertException(t, err, driver.ExNotAuthorized, "Invalid session for the user.") +} + +// startNewPasswordChallengeAgain reopens the challenge for the existing temp +// user after resetting it to a temporary password. +func startNewPasswordChallengeAgain(t *testing.T, m *Mock, poolID, clientID string) driver.RespondToAuthChallengeInput { + t.Helper() + + ctx := context.Background() + requireNoError(t, m.AdminSetUserPassword(ctx, poolID, "temp", "Temp0rary!", false), "AdminSetUserPassword temp") + + res, err := m.InitiateAuth(ctx, passwordAuth(clientID, "temp", "Temp0rary!")) + requireNoError(t, err, "InitiateAuth") + + return driver.RespondToAuthChallengeInput{ + ClientID: clientID, ChallengeName: driver.ChallengeNewPasswordRequired, Session: res.Session, + ChallengeResponses: map[string]string{"USERNAME": "temp", "NEW_PASSWORD": "N3wPassword!"}, + } +} + +func TestRefreshAndTokensRecheckUserState(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + client := mustCreateClient(t, m, pool.ID, false, passwordFlows...) + confirmedUser(t, m, pool.ID, client.ClientID, "alice") + + res, err := m.InitiateAuth(ctx, passwordAuth(client.ClientID, "alice", testPassword)) + requireNoError(t, err, "InitiateAuth") + + refresh := driver.InitiateAuthInput{ + ClientID: client.ClientID, AuthFlow: driver.AuthFlowRefreshTokenAuth, + AuthParameters: map[string]string{"REFRESH_TOKEN": res.AuthenticationResult.RefreshToken}, + } + + requireNoError(t, m.AdminDisableUser(ctx, pool.ID, "alice"), "AdminDisableUser") + + _, err = m.InitiateAuth(ctx, refresh) + assertException(t, err, driver.ExNotAuthorized, "User is disabled.") + + _, err = m.GetUser(ctx, res.AuthenticationResult.AccessToken) + assertException(t, err, driver.ExNotAuthorized, "User is disabled.") + + requireNoError(t, m.AdminEnableUser(ctx, pool.ID, "alice"), "AdminEnableUser") + requireNoError(t, m.AdminResetUserPassword(ctx, pool.ID, "alice"), "AdminResetUserPassword") + + _, err = m.InitiateAuth(ctx, refresh) + assertException(t, err, driver.ExPasswordResetRequired, "Password reset required for the user") + + _, err = m.GetUser(ctx, res.AuthenticationResult.AccessToken) + assertException(t, err, driver.ExPasswordResetRequired, "Password reset required for the user") +} + +func TestUserExistenceHiddenOnConfirmAndResend(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + + client, err := m.CreateUserPoolClient(ctx, driver.CreateUserPoolClientInput{ + UserPoolID: pool.ID, ClientName: "hidden", PreventUserExistenceErrors: "ENABLED", + }) + requireNoError(t, err, "CreateUserPoolClient") + + ghost := driver.ClientUserInput{ClientID: client.ClientID, Username: "ghost@example.com"} + + err = m.ConfirmSignUp(ctx, driver.ConfirmSignUpInput{ClientUserInput: ghost, ConfirmationCode: "123456"}) + assertException(t, err, driver.ExCodeMismatch, "Invalid verification code provided, please try again.") + + d, err := m.ResendConfirmationCode(ctx, ghost) + requireNoError(t, err, "ResendConfirmationCode for an unknown user") + + if d == nil || d.DeliveryMedium != driver.DeliveryMediumEmail || d.Destination != "g***@e***" { + t.Fatalf("simulated delivery = %+v", d) + } + + if n := len(m.users.Keys()); n != 0 { + t.Fatalf("%d users created by the simulated paths", n) + } + + legacy := mustCreateClient(t, m, pool.ID, false) + + _, err = m.ResendConfirmationCode(ctx, driver.ClientUserInput{ClientID: legacy.ClientID, Username: "ghost"}) + assertException(t, err, driver.ExUserNotFound, "Username/client id combination not found.") +} + +func TestClientTokenValidityRange(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + + units := func(access, id, refresh string) *driver.TokenValidityUnits { + return &driver.TokenValidityUnits{AccessToken: access, IDToken: id, RefreshToken: refresh} + } + + cases := []struct { + name string + access, id, refresh *int32 + units *driver.TokenValidityUnits + ok bool + }{ + {"defaults", nil, nil, nil, nil, true}, + {"access 5 minutes", int32Ptr(5), nil, nil, units("minutes", "", ""), true}, + {"access 4 minutes", int32Ptr(4), nil, nil, units("minutes", "", ""), false}, + {"access 24 hours", int32Ptr(24), nil, nil, nil, true}, + {"access 25 hours", int32Ptr(25), nil, nil, nil, false}, + {"access 1 day", int32Ptr(1), nil, nil, units("days", "", ""), true}, + {"access 2 days", int32Ptr(2), nil, nil, units("days", "", ""), false}, + {"id 299 seconds", nil, int32Ptr(299), nil, units("", "seconds", ""), false}, + {"id 86400 seconds", nil, int32Ptr(86400), nil, units("", "seconds", ""), true}, + {"refresh 60 minutes", nil, nil, int32Ptr(60), units("", "", "minutes"), true}, + {"refresh 59 minutes", nil, nil, int32Ptr(59), units("", "", "minutes"), false}, + {"refresh 3650 days", nil, nil, int32Ptr(3650), nil, true}, + {"refresh 3651 days", nil, nil, int32Ptr(3651), nil, false}, + {"bad unit", int32Ptr(1), nil, nil, units("weeks", "", ""), false}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + in := driver.CreateUserPoolClientInput{ + UserPoolID: pool.ID, ClientName: "v", AccessTokenValidity: tc.access, IDTokenValidity: tc.id, + RefreshTokenValidity: tc.refresh, TokenValidityUnits: tc.units, + } + + c, err := m.CreateUserPoolClient(ctx, in) + if tc.ok { + requireNoError(t, err, "CreateUserPoolClient") + + in.ClientID = c.ClientID + _, err = m.UpdateUserPoolClient(ctx, in) + requireNoError(t, err, "UpdateUserPoolClient") + + return + } + + assertException(t, err, driver.ExInvalidParameter, "") + }) + } + + c := mustCreateClient(t, m, pool.ID, false) + _, err := m.UpdateUserPoolClient(ctx, driver.CreateUserPoolClientInput{ + UserPoolID: pool.ID, ClientID: c.ClientID, AccessTokenValidity: int32Ptr(2), + TokenValidityUnits: units("minutes", "", ""), + }) + assertException(t, err, driver.ExInvalidParameter, "") +} diff --git a/providers/aws/cognito/sign_up.go b/providers/aws/cognito/sign_up.go index 159825b15..e4caacab6 100644 --- a/providers/aws/cognito/sign_up.go +++ b/providers/aws/cognito/sign_up.go @@ -310,17 +310,35 @@ func (m *Mock) newSignUpRecord(pool *driver.UserPool, in driver.SignUpInput) (us return rec, nil } -// clientUser resolves the client, verifies SECRET_HASH, and finds the user a -// client-side operation names. It must run under m.mu. -func (m *Mock) clientUser(in driver.ClientUserInput) (driver.UserPool, string, userRecord, error) { +// clientTarget is the client, pool and (when found) user a client-side +// operation names. +type clientTarget struct { + client driver.UserPoolClient + pool driver.UserPool + key string + rec userRecord + found bool +} + +// hidesUsers reports whether the client hides whether a user exists +// (PreventUserExistenceErrors ENABLED). +func (t *clientTarget) hidesUsers() bool { + return t.client.PreventUserExistenceErrors == existenceEnabled +} + +// clientUser resolves the client, verifies SECRET_HASH, and looks up the user a +// client-side operation names. A missing user is reported through found, so +// the caller can answer the way the client's existence-error setting asks. It +// must run under m.mu. +func (m *Mock) clientUser(in driver.ClientUserInput) (clientTarget, error) { client, err := m.clientByID(in.ClientID) if err != nil { - return driver.UserPool{}, "", userRecord{}, err + return clientTarget{}, err } pool, ok := m.userPools.Get(client.UserPoolID) if !ok { - return driver.UserPool{}, "", userRecord{}, poolNotFound(client.UserPoolID) + return clientTarget{}, poolNotFound(client.UserPoolID) } key, rec, found := m.resolveUser(&pool, in.Username) @@ -331,14 +349,27 @@ func (m *Mock) clientUser(in driver.ClientUserInput) (driver.UserPool, string, u } if err := checkSecretHash(client, in.SecretHash, "Unable to verify secret hash for client "+client.ClientID, names...); err != nil { - return driver.UserPool{}, "", userRecord{}, err + return clientTarget{}, err + } + + return clientTarget{client: client, pool: pool, key: key, rec: copyUserRecord(rec), found: found}, nil +} + +// simulatedDelivery is the CodeDeliveryDetails a client that hides user +// existence returns for an unknown username, shaped like a real delivery to +// the attribute the pool verifies. +func simulatedDelivery(pool *driver.UserPool, username string) *driver.CodeDeliveryDetails { + email := username + if !isEmailFormat(email) { + email = username + "@example.com" } - if !found { - return driver.UserPool{}, "", userRecord{}, clientUserNotFound() + attrs := []driver.Attribute{{Name: attrEmail, Value: email}} + if isPhoneFormat(username) { + attrs = append(attrs, driver.Attribute{Name: attrPhoneNumber, Value: username}) } - return pool, key, copyUserRecord(rec), nil + return codeDelivery(pool, attrs) } // ConfirmSignUp confirms an UNCONFIRMED user with the code that was sent. The @@ -347,21 +378,27 @@ func (m *Mock) ConfirmSignUp(_ context.Context, in driver.ConfirmSignUpInput) er m.mu.Lock() defer m.mu.Unlock() - pool, key, rec, err := m.clientUser(in.ClientUserInput) + t, err := m.clientUser(in.ClientUserInput) if err != nil { return err } - if rec.User.UserStatus != driver.UserStatusUnconfirmed { - return cannotConfirm(rec.User.UserStatus) + if !t.found { + if t.hidesUsers() { + return codeMismatch() + } + + return clientUserNotFound() } - if rec.Code == nil || !hmac.Equal([]byte(rec.Code.Code), []byte(in.ConfirmationCode)) { - return codeMismatch() + pool, key, rec := t.pool, t.key, t.rec + + if rec.User.UserStatus != driver.UserStatusUnconfirmed { + return cannotConfirm(rec.User.UserStatus) } - if !m.now().Before(rec.Code.ExpiresAt) { - return expiredCode() + if err := m.checkCode(rec.Code, in.ConfirmationCode); err != nil { + return err } if d := rec.Code.Delivery; d != nil { @@ -383,16 +420,39 @@ func (m *Mock) ConfirmSignUp(_ context.Context, in driver.ConfirmSignUpInput) er return nil } +// checkCode compares a confirmation code with the outstanding one. +func (m *Mock) checkCode(want *pendingCode, got string) error { + if want == nil || !hmac.Equal([]byte(want.Code), []byte(got)) { + return codeMismatch() + } + + if !m.now().Before(want.ExpiresAt) { + return expiredCode() + } + + return nil +} + // ResendConfirmationCode issues a fresh code to an UNCONFIRMED user. func (m *Mock) ResendConfirmationCode(_ context.Context, in driver.ClientUserInput) (*driver.CodeDeliveryDetails, error) { m.mu.Lock() defer m.mu.Unlock() - pool, key, rec, err := m.clientUser(in) + t, err := m.clientUser(in) if err != nil { return nil, err } + if !t.found { + if t.hidesUsers() { + return simulatedDelivery(&t.pool, in.Username), nil + } + + return nil, clientUserNotFound() + } + + pool, key, rec := t.pool, t.key, t.rec + if rec.User.UserStatus != driver.UserStatusUnconfirmed { return nil, invalidParameter("User is already confirmed.") } diff --git a/providers/aws/cognito/token_validity.go b/providers/aws/cognito/token_validity.go new file mode 100644 index 000000000..3043f9b89 --- /dev/null +++ b/providers/aws/cognito/token_validity.go @@ -0,0 +1,77 @@ +package cognito + +import ( + "time" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// Token-validity bounds Cognito enforces on an app client: access and ID +// tokens last between 5 minutes and 1 day, refresh tokens between 60 minutes +// and 10 years. +const ( + minAccessValidity = 5 * time.Minute + maxAccessValidity = 24 * time.Hour + minRefreshValidity = 60 * time.Minute + maxRefreshValidity = 3650 * 24 * time.Hour +) + +// checkTokenValidity validates an app client's token-validity values and +// units. A zero or unset value means the default and is not range checked. +func checkTokenValidity(in *driver.CreateUserPoolClientInput) error { + units := driver.TokenValidityUnits{} + if in.TokenValidityUnits != nil { + units = *in.TokenValidityUnits + } + + for _, u := range []struct{ field, value string }{ + {"accessToken", units.AccessToken}, {"idToken", units.IDToken}, {"refreshToken", units.RefreshToken}, + } { + if !validUnit(u.value) { + return invalidParameter("1 validation error detected: Value '%s' at 'tokenValidityUnits.%s' failed to satisfy constraint: "+ + "Member must satisfy enum value set: [seconds, minutes, hours, days]", u.value, u.field) + } + } + + checks := []struct { + value *int32 + unit string + defUnit string + min, max time.Duration + }{ + {in.AccessTokenValidity, units.AccessToken, driver.TimeUnitHours, minAccessValidity, maxAccessValidity}, + {in.IDTokenValidity, units.IDToken, driver.TimeUnitHours, minAccessValidity, maxAccessValidity}, + {in.RefreshTokenValidity, units.RefreshToken, driver.TimeUnitDays, minRefreshValidity, maxRefreshValidity}, + } + + for _, c := range checks { + if c.value == nil || *c.value == 0 { + continue + } + + d := unitDuration(*c.value, c.unit, c.defUnit) + if *c.value < 0 || d < c.min || d > c.max { + return invalidParameter("Invalid range for token validity.") + } + } + + return nil +} + +func validUnit(u string) bool { + switch u { + case "", driver.TimeUnitSeconds, driver.TimeUnitMinutes, driver.TimeUnitHours, driver.TimeUnitDays: + return true + default: + return false + } +} + +// setValidity returns a client token-validity value, treating zero as unset. +func setValidity(v *int32) *int32 { + if v == nil || *v == 0 { + return nil + } + + return copyInt32Ptr(v) +} diff --git a/providers/aws/cognito/tokens.go b/providers/aws/cognito/tokens.go index 9b479b1ee..b934f91f7 100644 --- a/providers/aws/cognito/tokens.go +++ b/providers/aws/cognito/tokens.go @@ -356,7 +356,7 @@ func (m *Mock) verifyAccessToken(token string) (userRecord, error) { } // liveTokenUser checks that a verified token's login is not revoked and that -// its user still exists and is enabled. +// its user still exists and may still sign in. func (m *Mock) liveTokenUser(poolID string, claims map[string]any) (userRecord, error) { originJTI, _ := claims["origin_jti"].(string) sub, _ := claims[attrSub].(string) @@ -371,8 +371,8 @@ func (m *Mock) liveTokenUser(poolID string, claims map[string]any) (userRecord, return userRecord{}, notAuthorized("Access Token has been revoked") } - if !rec.User.Enabled { - return userRecord{}, notAuthorized("User is disabled.") + if err := checkCanSignIn(&rec); err != nil { + return userRecord{}, err } return rec, nil diff --git a/providers/aws/cognito/user_pool_clients.go b/providers/aws/cognito/user_pool_clients.go index 9e80b397c..e0ac788da 100644 --- a/providers/aws/cognito/user_pool_clients.go +++ b/providers/aws/cognito/user_pool_clients.go @@ -32,6 +32,10 @@ func (m *Mock) CreateUserPoolClient(_ context.Context, in driver.CreateUserPoolC return nil, invalidParameter("ClientName is required") } + if err := checkTokenValidity(&in); err != nil { + return nil, err + } + m.mu.Lock() defer m.mu.Unlock() @@ -73,6 +77,10 @@ func (m *Mock) DescribeUserPoolClient(_ context.Context, userPoolID, clientID st // //nolint:gocritic // hugeParam: taken by value to match the driver interface func (m *Mock) UpdateUserPoolClient(_ context.Context, in driver.CreateUserPoolClientInput) (*driver.UserPoolClient, error) { + if err := checkTokenValidity(&in); err != nil { + return nil, err + } + m.mu.Lock() defer m.mu.Unlock() @@ -141,9 +149,9 @@ func buildClient(in driver.CreateUserPoolClientInput) driver.UserPoolClient { return driver.UserPoolClient{ ClientName: in.ClientName, UserPoolID: in.UserPoolID, - RefreshTokenValidity: int32OrDefault(in.RefreshTokenValidity, defaultRefreshTokenValidity), - AccessTokenValidity: copyInt32Ptr(in.AccessTokenValidity), - IDTokenValidity: copyInt32Ptr(in.IDTokenValidity), + RefreshTokenValidity: int32OrDefault(setValidity(in.RefreshTokenValidity), defaultRefreshTokenValidity), + AccessTokenValidity: setValidity(in.AccessTokenValidity), + IDTokenValidity: setValidity(in.IDTokenValidity), TokenValidityUnits: copyTokenValidityUnits(in.TokenValidityUnits), ExplicitAuthFlows: resolveExplicitAuthFlows(in.ExplicitAuthFlows), AuthSessionValidity: int32OrDefault(in.AuthSessionValidity, defaultAuthSessionValidity), diff --git a/server/aws/cognito/handler.go b/server/aws/cognito/handler.go index 97f36c481..b3712688a 100644 --- a/server/aws/cognito/handler.go +++ b/server/aws/cognito/handler.go @@ -22,6 +22,9 @@ import ( const targetPrefix = "AWSCognitoIdentityProviderService." +// iamService is the IAM service prefix of every user-pool operation. +const iamService = "cognito-idp" + // Handler serves Cognito user-pools JSON-RPC requests against a Cognito driver. type Handler struct { cognito cognitodriver.Cognito @@ -160,4 +163,4 @@ func writeErr(w http.ResponseWriter, err error) { // IAMService returns the IAM service prefix of the operations this handler // serves. -func (*Handler) IAMService() string { return "cognito-idp" } +func (*Handler) IAMService() string { return iamService } diff --git a/server/aws/cognito/sdk_roundtrip_test.go b/server/aws/cognito/sdk_roundtrip_test.go index 2e48186c4..57df113e5 100644 --- a/server/aws/cognito/sdk_roundtrip_test.go +++ b/server/aws/cognito/sdk_roundtrip_test.go @@ -234,7 +234,7 @@ func TestSDKUserPoolClientExplicitTokenUnits(t *testing.T) { out, err := c.CreateUserPoolClient(ctx, &cip.CreateUserPoolClientInput{ UserPoolId: aws.String(id), ClientName: aws.String("units-client"), - AccessTokenValidity: aws.Int32(2), + AccessTokenValidity: aws.Int32(15), TokenValidityUnits: &ciptypes.TokenValidityUnitsType{ AccessToken: ciptypes.TimeUnitsTypeMinutes, }, @@ -244,8 +244,8 @@ func TestSDKUserPoolClientExplicitTokenUnits(t *testing.T) { } client := out.UserPoolClient - if aws.ToInt32(client.AccessTokenValidity) != 2 { - t.Fatalf("AccessTokenValidity = %d, want 2", aws.ToInt32(client.AccessTokenValidity)) + if aws.ToInt32(client.AccessTokenValidity) != 15 { + t.Fatalf("AccessTokenValidity = %d, want 15", aws.ToInt32(client.AccessTokenValidity)) } if client.TokenValidityUnits == nil || client.TokenValidityUnits.AccessToken != ciptypes.TimeUnitsTypeMinutes { diff --git a/server/aws/cognito/wellknown.go b/server/aws/cognito/wellknown.go index 373e92b2e..ffd0eba10 100644 --- a/server/aws/cognito/wellknown.go +++ b/server/aws/cognito/wellknown.go @@ -41,7 +41,7 @@ func (w *WellKnown) PublicRequest(r *http.Request) bool { return w.Matches(r) } // IAMService returns the IAM service prefix of the user-pool documents. They // are public, so no request reaches IAM authorization. -func (*WellKnown) IAMService() string { return "cognito-idp" } +func (*WellKnown) IAMService() string { return iamService } // openIDConfiguration is the discovery document Cognito publishes for a pool. type openIDConfiguration struct { From d6fa5c2b0ae317493db6e25569fc9c0767cbea09 Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 4 Oct 2026 17:39:40 +0530 Subject: [PATCH 3/3] fix(aws-cognito): saturate token validity overflow, shape simulated delivery by pool attribute --- providers/aws/cognito/sign_up.go | 42 ++++++-- providers/aws/cognito/tokens.go | 25 ++++- .../aws/cognito/validity_overflow_test.go | 96 +++++++++++++++++++ 3 files changed, 151 insertions(+), 12 deletions(-) create mode 100644 providers/aws/cognito/validity_overflow_test.go diff --git a/providers/aws/cognito/sign_up.go b/providers/aws/cognito/sign_up.go index e4caacab6..00eb208b4 100644 --- a/providers/aws/cognito/sign_up.go +++ b/providers/aws/cognito/sign_up.go @@ -6,6 +6,7 @@ import ( "crypto/rand" "crypto/sha256" "encoding/base64" + "fmt" "math/big" "slices" "strings" @@ -23,6 +24,8 @@ const ( codeSpace = 1_000_000 codeValidity = 24 * time.Hour phoneTailLen = 4 + // fakePhoneTail bounds the last four digits of a simulated phone number. + fakePhoneTail = 10000 ) // pendingCode is the confirmation code outstanding for a user. Delivery is @@ -359,17 +362,42 @@ func (m *Mock) clientUser(in driver.ClientUserInput) (clientTarget, error) { // existence returns for an unknown username, shaped like a real delivery to // the attribute the pool verifies. func simulatedDelivery(pool *driver.UserPool, username string) *driver.CodeDeliveryDetails { - email := username - if !isEmailFormat(email) { - email = username + "@example.com" + switch { + case slices.Contains(pool.AutoVerifiedAttributes, attrEmail): + dest := firstChar(username) + "***@e***" + if isEmailFormat(username) { + dest = maskEmail(username) + } + + return &driver.CodeDeliveryDetails{Destination: dest, DeliveryMedium: driver.DeliveryMediumEmail, AttributeName: attrEmail} + case slices.Contains(pool.AutoVerifiedAttributes, attrPhoneNumber): + dest := maskPhone(fakePhone(username)) + if isPhoneFormat(username) { + dest = maskPhone(username) + } + + return &driver.CodeDeliveryDetails{Destination: dest, DeliveryMedium: driver.DeliveryMediumSMS, AttributeName: attrPhoneNumber} + default: + return nil } +} - attrs := []driver.Attribute{{Name: attrEmail, Value: email}} - if isPhoneFormat(username) { - attrs = append(attrs, driver.Attribute{Name: attrPhoneNumber, Value: username}) +// firstChar returns the first character of s, or "u" for an empty string. +func firstChar(s string) string { + for _, r := range s { + return string(r) } - return codeDelivery(pool, attrs) + return "u" +} + +// fakePhone derives a stable phone number from a username, so repeated +// resends for the same unknown user report the same masked destination. +func fakePhone(username string) string { + sum := sha256.Sum256([]byte(username)) + n := (int(sum[0])<<8 | int(sum[1])) % fakePhoneTail + + return fmt.Sprintf("+1555555%04d", n) } // ConfirmSignUp confirms an UNCONFIRMED user with the code that was sent. The diff --git a/providers/aws/cognito/tokens.go b/providers/aws/cognito/tokens.go index b934f91f7..23cb3868a 100644 --- a/providers/aws/cognito/tokens.go +++ b/providers/aws/cognito/tokens.go @@ -8,6 +8,7 @@ import ( "encoding/hex" "encoding/json" stderrors "errors" + "math" "slices" "sort" "strings" @@ -32,6 +33,7 @@ const ( refreshIVLen = 12 refreshTagLen = 16 jwtSegments = 3 + maxDurationSeconds = math.MaxInt64 / int64(time.Second) ) // refreshHeader is the JWE protected header Cognito refresh tokens carry. The @@ -135,22 +137,35 @@ func validity(value *int32, unit string, def time.Duration) time.Duration { return unitDuration(*value, unit, driver.TimeUnitHours) } +// unitDuration converts a validity value in a unit to a duration. The product +// is computed in seconds, which cannot overflow for an int32 value, and a +// result past the largest time.Duration saturates rather than wrapping, so a +// huge value is rejected by the range check instead of becoming a short one. func unitDuration(value int32, unit, defUnit string) time.Duration { if unit == "" { unit = defUnit } - d := time.Duration(value) + per := int64(time.Hour / time.Second) switch unit { case driver.TimeUnitSeconds: - return d * time.Second + per = 1 case driver.TimeUnitMinutes: - return d * time.Minute + per = int64(time.Minute / time.Second) case driver.TimeUnitDays: - return d * 24 * time.Hour + per = int64(24 * time.Hour / time.Second) + } + + secs := int64(value) * per + + switch { + case secs > maxDurationSeconds: + return time.Duration(math.MaxInt64) + case secs < -maxDurationSeconds: + return time.Duration(math.MinInt64) default: - return d * time.Hour + return time.Duration(secs) * time.Second } } diff --git a/providers/aws/cognito/validity_overflow_test.go b/providers/aws/cognito/validity_overflow_test.go new file mode 100644 index 000000000..782c6134b --- /dev/null +++ b/providers/aws/cognito/validity_overflow_test.go @@ -0,0 +1,96 @@ +package cognito + +import ( + "context" + "regexp" + "testing" + + "github.com/stackshy/cloudemu/v2/services/cognito/driver" +) + +// TestTokenValidityHugeValuesRejected pins that a value whose nanosecond +// duration overflows int64 is rejected instead of wrapping to a short token. +// 213504 days is just past the largest time.Duration. +func TestTokenValidityHugeValuesRejected(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + pool := mustCreateEmailPool(t, m) + + for _, tc := range []struct { + name string + in driver.CreateUserPoolClientInput + units *driver.TokenValidityUnits + }{ + {"access days", driver.CreateUserPoolClientInput{AccessTokenValidity: int32Ptr(213504)}, + &driver.TokenValidityUnits{AccessToken: driver.TimeUnitDays}}, + {"id max int32 days", driver.CreateUserPoolClientInput{IDTokenValidity: int32Ptr(2147483647)}, + &driver.TokenValidityUnits{IDToken: driver.TimeUnitDays}}, + {"refresh max int32 days", driver.CreateUserPoolClientInput{RefreshTokenValidity: int32Ptr(2147483647)}, nil}, + {"access min int32 days", driver.CreateUserPoolClientInput{AccessTokenValidity: int32Ptr(-2147483648)}, + &driver.TokenValidityUnits{AccessToken: driver.TimeUnitDays}}, + } { + t.Run(tc.name, func(t *testing.T) { + in := tc.in + in.UserPoolID, in.ClientName, in.TokenValidityUnits = pool.ID, "huge", tc.units + + _, err := m.CreateUserPoolClient(ctx, in) + assertException(t, err, driver.ExInvalidParameter, "Invalid range for token validity.") + }) + } + + if d := unitDuration(213504, driver.TimeUnitDays, ""); d < maxAccessValidity { + t.Fatalf("unitDuration(213504 days) = %v, want a saturated (huge) duration", d) + } +} + +func TestSimulatedDeliveryFollowsPoolAttribute(t *testing.T) { + m, _ := newClockMock(t) + ctx := context.Background() + + resend := func(verified []string, username string) *driver.CodeDeliveryDetails { + t.Helper() + + pool, err := m.CreateUserPool(ctx, driver.CreateUserPoolInput{Name: "sim", AutoVerifiedAttributes: verified}) + requireNoError(t, err, "CreateUserPool") + + client, err := m.CreateUserPoolClient(ctx, driver.CreateUserPoolClientInput{ + UserPoolID: pool.ID, ClientName: "hidden", PreventUserExistenceErrors: "ENABLED", + }) + requireNoError(t, err, "CreateUserPoolClient") + + d, err := m.ResendConfirmationCode(ctx, driver.ClientUserInput{ClientID: client.ClientID, Username: username}) + requireNoError(t, err, "ResendConfirmationCode") + + return d + } + + d := resend([]string{"email"}, "ghost") + if d.DeliveryMedium != driver.DeliveryMediumEmail || d.AttributeName != "email" || d.Destination != "g***@e***" { + t.Fatalf("email pool, plain username = %+v", d) + } + + d = resend([]string{"email"}, "casper@example.org") + if d.Destination != "c***@e***" { + t.Fatalf("email pool, email username = %+v", d) + } + + phone := regexp.MustCompile(`^\+\*+\d{4}$`) + + d = resend([]string{"phone_number"}, "ghost") + if d.DeliveryMedium != driver.DeliveryMediumSMS || d.AttributeName != "phone_number" || !phone.MatchString(d.Destination) { + t.Fatalf("phone pool, plain username = %+v", d) + } + + if again := resend([]string{"phone_number"}, "ghost"); again.Destination != d.Destination { + t.Fatalf("simulated phone not stable: %q then %q", d.Destination, again.Destination) + } + + d = resend([]string{"phone_number"}, "+15551231234") + if d.Destination != "+*******1234" { + t.Fatalf("phone pool, phone username = %+v", d) + } + + if d = resend(nil, "ghost"); d != nil { + t.Fatalf("pool without verification = %+v, want no delivery", d) + } +}