@nakedev/go-scaffold 0.4.0 → 0.5.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/README.md +598 -306
- package/dist/commands/auth.js +65 -23
- package/dist/commands/check.js +281 -0
- package/dist/commands/config.js +50 -0
- package/dist/commands/create.js +33 -2
- package/dist/commands/generate.js +29 -3
- package/dist/commands/method.js +74 -63
- package/dist/commands/migration.js +2 -2
- package/dist/commands/observability.js +4 -53
- package/dist/commands/rbac.js +21 -10
- package/dist/commands/undo.js +11 -3
- package/dist/commands/worker.js +15 -5
- package/dist/index.js +198 -59
- package/dist/prompts/auth-wizard.js +40 -6
- package/dist/prompts/create-wizard.js +42 -1
- package/dist/prompts/generate-wizard.js +89 -9
- package/dist/templates/auth-manifest.js +50 -19
- package/dist/templates/create-manifest.js +8 -0
- package/dist/templates/module-manifest.js +84 -26
- package/dist/templates/rbac-manifest.js +16 -11
- package/dist/templates/worker-manifest.js +4 -1
- package/dist/types.js +8 -0
- package/dist/utils/auth-patcher.js +124 -33
- package/dist/utils/config.js +167 -4
- package/dist/utils/docs-patcher.js +68 -0
- package/dist/utils/hexagonal-method-patcher.js +334 -0
- package/dist/utils/main-patcher.js +32 -30
- package/dist/utils/marker-patch.js +7 -1
- package/dist/utils/module-location.js +17 -11
- package/dist/utils/module-profile.js +32 -0
- package/dist/utils/platform-patcher.js +56 -7
- package/dist/utils/rbac-patcher.js +89 -210
- package/package.json +7 -2
- package/templates/add/auth/cmd/seed/main.go.hbs +15 -3
- package/templates/add/auth/docs/login.yaml.hbs +11 -1
- package/templates/add/auth/docs/mfa-verify.yaml.hbs +19 -0
- package/templates/add/auth/docs/provider-exchange.yaml.hbs +40 -0
- package/templates/add/auth/docs/provider-login.yaml.hbs +31 -0
- package/templates/add/auth/docs/refresh.yaml.hbs +7 -0
- package/templates/add/auth/docs/register.yaml.hbs +7 -0
- package/templates/add/auth/docs/reset-password.yaml.hbs +1 -1
- package/templates/add/auth/docs/schemas.yaml.hbs +59 -1
- package/templates/add/auth/docs/users-me-mfa-confirm.yaml.hbs +19 -0
- package/templates/add/auth/docs/users-me-mfa-disable.yaml.hbs +15 -0
- package/templates/add/auth/docs/users-me-mfa-setup.yaml.hbs +14 -0
- package/templates/add/auth/docs/users-me-mfa.yaml.hbs +12 -0
- package/templates/add/auth/internal/app/user/adapters/inbound/http/browser_policy.go.hbs +98 -0
- package/templates/add/auth/internal/app/user/adapters/inbound/http/dto.go.hbs +159 -0
- package/templates/add/auth/internal/app/user/adapters/inbound/http/handler.go.hbs +228 -0
- package/templates/add/auth/internal/app/user/adapters/inbound/http/handler_local.go.hbs +76 -0
- package/templates/add/auth/internal/app/user/adapters/inbound/http/handler_mfa.go.hbs +83 -0
- package/templates/add/auth/internal/app/user/adapters/inbound/http/handler_oauth.go.hbs +70 -0
- package/templates/add/auth/internal/app/user/adapters/inbound/http/handler_recovery.go.hbs +49 -0
- package/templates/add/auth/internal/app/user/adapters/inbound/http/handler_test.go.hbs +311 -0
- package/templates/add/auth/internal/app/user/adapters/inbound/http/handler_user.go.hbs +41 -0
- package/templates/add/auth/internal/app/user/adapters/inbound/http/session_cookie.go.hbs +35 -0
- package/templates/add/auth/internal/app/user/adapters/outbound/password/bcrypt.go.hbs +35 -0
- package/templates/add/auth/internal/app/user/adapters/outbound/password/bcrypt_test.go.hbs +20 -0
- package/templates/add/auth/internal/app/user/adapters/outbound/postgres/mfa_store.go.hbs +129 -0
- package/templates/add/auth/internal/app/user/adapters/outbound/postgres/mfa_store_test.go.hbs +174 -0
- package/templates/add/auth/internal/app/user/adapters/outbound/postgres/model.go.hbs +84 -0
- package/templates/add/auth/internal/app/user/adapters/outbound/postgres/repository.go.hbs +211 -0
- package/templates/add/auth/internal/app/user/{repository_test.go.hbs → adapters/outbound/postgres/repository_test.go.hbs} +18 -19
- package/templates/add/auth/internal/app/user/adapters/outbound/postgres/tokenstore_pg.go.hbs +213 -0
- package/templates/add/auth/internal/app/user/adapters/outbound/postgres/tokenstore_pg_test.go.hbs +103 -0
- package/templates/add/auth/internal/app/user/adapters/outbound/postgres/tokenstore_recovery.go.hbs +84 -0
- package/templates/add/auth/internal/app/user/adapters/outbound/redis/tokenstore.go.hbs +228 -0
- package/templates/add/auth/internal/app/user/adapters/outbound/redis/tokenstore_test.go.hbs +196 -0
- package/templates/add/auth/internal/app/user/application/contracts.go.hbs +52 -0
- package/templates/add/auth/internal/app/user/application/dto.go.hbs +75 -0
- package/templates/add/auth/internal/app/user/application/errors.go.hbs +62 -0
- package/templates/add/auth/internal/app/user/application/external_login.go.hbs +198 -0
- package/templates/add/auth/internal/app/user/application/jwt.go.hbs +58 -0
- package/templates/add/auth/internal/app/user/application/local_auth.go.hbs +96 -0
- package/templates/add/auth/internal/app/user/application/mfa_service.go.hbs +449 -0
- package/templates/add/auth/internal/app/user/application/mfa_service_test.go.hbs +200 -0
- package/templates/add/auth/internal/app/user/application/oauth.go.hbs +132 -0
- package/templates/add/auth/internal/app/user/application/provider_test.go.hbs +285 -0
- package/templates/add/auth/internal/app/user/application/recovery.go.hbs +82 -0
- package/templates/add/auth/internal/app/user/application/recovery_service.go.hbs +112 -0
- package/templates/add/auth/internal/app/user/application/service.go.hbs +145 -0
- package/templates/add/auth/internal/app/user/application/service_test.go.hbs +891 -0
- package/templates/add/auth/internal/app/user/application/sessions.go.hbs +99 -0
- package/templates/add/auth/internal/app/user/application/tokenstore_ports.go.hbs +14 -0
- package/templates/add/auth/internal/app/user/application/user_query.go.hbs +65 -0
- package/templates/add/auth/internal/app/user/composition.go.hbs +168 -0
- package/templates/add/auth/internal/app/user/domain/entity.go.hbs +41 -0
- package/templates/add/auth/internal/app/user/domain/errors.go.hbs +32 -0
- package/templates/add/auth/internal/app/user/ports/password.go.hbs +9 -0
- package/templates/add/auth/internal/app/user/ports/repository.go.hbs +90 -0
- package/templates/add/auth/internal/platform/authprovider/google/google.go.hbs +389 -0
- package/templates/add/auth/internal/platform/authprovider/google/google_test.go.hbs +312 -0
- package/templates/add/auth/migrations/create_auth_tokens.up.sql.hbs +10 -5
- package/templates/add/auth/migrations/create_identities.up.sql.hbs +1 -1
- package/templates/add/auth/migrations/create_login_throttle.up.sql.hbs +1 -1
- package/templates/add/auth/migrations/create_mfa.down.sql.hbs +3 -0
- package/templates/add/auth/migrations/create_mfa.up.sql.hbs +29 -0
- package/templates/add/auth/migrations/create_users.up.sql.hbs +4 -3
- package/templates/add/rbac/internal/app/role/adapters/inbound/http/handler.go.hbs +142 -0
- package/templates/add/rbac/internal/app/role/adapters/inbound/http/handler_test.go.hbs +19 -0
- package/templates/add/rbac/internal/app/role/adapters/outbound/postgres/model.go.hbs +48 -0
- package/templates/add/rbac/internal/app/role/adapters/outbound/postgres/repository.go.hbs +127 -0
- package/templates/add/rbac/internal/app/role/{repository_test.go.hbs → adapters/outbound/postgres/repository_test.go.hbs} +8 -8
- package/templates/add/rbac/internal/app/role/application/dto.go.hbs +47 -0
- package/templates/add/rbac/internal/app/role/application/errors.go.hbs +19 -0
- package/templates/add/rbac/internal/app/role/application/service.go.hbs +157 -0
- package/templates/add/rbac/internal/app/role/{service_test.go.hbs → application/service_test.go.hbs} +26 -19
- package/templates/add/rbac/internal/app/role/composition.go.hbs +48 -0
- package/templates/add/rbac/internal/app/role/domain/entity.go.hbs +23 -0
- package/templates/add/rbac/internal/app/role/domain/errors.go.hbs +26 -0
- package/templates/add/rbac/internal/app/role/ports/repository.go.hbs +25 -0
- package/templates/add/rbac/migrations/add_roles.down.sql.hbs +3 -11
- package/templates/add/rbac/migrations/add_roles.up.sql.hbs +17 -6
- package/templates/add/worker/internal/platform/queue/river_test.go.hbs +84 -0
- package/templates/create/base/.claude/skills/go-scaffold/SKILL.md.hbs +358 -121
- package/templates/create/base/.env.example.hbs +0 -1
- package/templates/create/base/.golangci.yml.hbs +2 -2
- package/templates/create/base/AGENTS.md.hbs +279 -67
- package/templates/create/base/Makefile.hbs +2 -1
- package/templates/create/base/README.md.hbs +115 -32
- package/templates/create/base/cmd/api/wiring.go.hbs +13 -9
- package/templates/create/base/internal/composition/doc.go.hbs +7 -0
- package/templates/create/base/internal/platform/database/database.go.hbs +3 -3
- package/templates/create/base/internal/shared/apperror/apperror.go.hbs +15 -2
- package/templates/create/base/internal/shared/config/config.go.hbs +0 -8
- package/templates/create/base/internal/shared/middleware/cors_test.go.hbs +40 -0
- package/templates/create/base/internal/shared/middleware/error.go.hbs +15 -5
- package/templates/create/features/docs/architecture.md.hbs +92 -32
- package/templates/create/features/docs/patterns.md.hbs +137 -91
- package/templates/create/features/docs/techstack.md.hbs +18 -3
- package/templates/generate/module/hexagonal/adapters/inbound/http/dto.go.hbs +45 -0
- package/templates/generate/module/hexagonal/adapters/inbound/http/dto.minimal.go.hbs +28 -0
- package/templates/generate/module/hexagonal/adapters/inbound/http/handler.go.hbs +182 -0
- package/templates/generate/module/hexagonal/adapters/inbound/http/handler.minimal.go.hbs +83 -0
- package/templates/generate/module/hexagonal/adapters/inbound/http/handler_crud_test.go.hbs +18 -0
- package/templates/generate/module/hexagonal/adapters/inbound/http/handler_test.go.hbs +30 -0
- package/templates/generate/module/hexagonal/adapters/outbound/postgres/model.go.hbs +37 -0
- package/templates/generate/module/hexagonal/adapters/outbound/postgres/repository.go.hbs +95 -0
- package/templates/generate/module/{repository_test.go.hbs → hexagonal/adapters/outbound/postgres/repository_test.go.hbs} +8 -8
- package/templates/generate/module/hexagonal/application/commands.crud.go.hbs +54 -0
- package/templates/generate/module/hexagonal/application/commands.go.hbs +25 -0
- package/templates/generate/module/hexagonal/application/cqrs_test.go.hbs +66 -0
- package/templates/generate/module/hexagonal/application/dto.go.hbs +35 -0
- package/templates/generate/module/hexagonal/application/dto.minimal.go.hbs +25 -0
- package/templates/generate/module/hexagonal/application/queries.crud.go.hbs +33 -0
- package/templates/generate/module/hexagonal/application/queries.go.hbs +25 -0
- package/templates/generate/module/hexagonal/application/service.crud.go.hbs +73 -0
- package/templates/generate/module/hexagonal/application/service.go.hbs +29 -0
- package/templates/generate/module/hexagonal/application/service_test.go.hbs +62 -0
- package/templates/generate/module/hexagonal/composition.go.hbs +27 -0
- package/templates/generate/module/hexagonal/domain/entity.go.hbs +20 -0
- package/templates/generate/module/hexagonal/domain/errors.go.hbs +11 -0
- package/templates/generate/module/hexagonal/ports/repository.go.hbs +38 -0
- package/templates/generate/module/migration.up.sql.hbs +1 -1
- package/dist/utils/method-patcher.js +0 -357
- package/templates/add/auth/docs/google-callback.yaml.hbs +0 -22
- package/templates/add/auth/docs/google-login.yaml.hbs +0 -7
- package/templates/add/auth/internal/app/user/dto.go.hbs +0 -77
- package/templates/add/auth/internal/app/user/errors.go.hbs +0 -43
- package/templates/add/auth/internal/app/user/handler.go.hbs +0 -276
- package/templates/add/auth/internal/app/user/jwt.go.hbs +0 -108
- package/templates/add/auth/internal/app/user/model/authtoken.go.hbs +0 -39
- package/templates/add/auth/internal/app/user/model/identity.go.hbs +0 -31
- package/templates/add/auth/internal/app/user/model/loginthrottle.go.hbs +0 -26
- package/templates/add/auth/internal/app/user/model/user.go.hbs +0 -30
- package/templates/add/auth/internal/app/user/repository.go.hbs +0 -137
- package/templates/add/auth/internal/app/user/service.go.hbs +0 -531
- package/templates/add/auth/internal/app/user/service_test.go.hbs +0 -316
- package/templates/add/auth/internal/app/user/tokenstore.go.hbs +0 -30
- package/templates/add/auth/internal/app/user/tokenstore_pg.go.hbs +0 -144
- package/templates/add/auth/internal/app/user/tokenstore_redis.go.hbs +0 -147
- package/templates/add/rbac/internal/app/role/dto.go.hbs +0 -45
- package/templates/add/rbac/internal/app/role/errors.go.hbs +0 -39
- package/templates/add/rbac/internal/app/role/handler.go.hbs +0 -104
- package/templates/add/rbac/internal/app/role/model/permission.go.hbs +0 -12
- package/templates/add/rbac/internal/app/role/model/role.go.hbs +0 -22
- package/templates/add/rbac/internal/app/role/model/role_permission.go.hbs +0 -11
- package/templates/add/rbac/internal/app/role/repository.go.hbs +0 -97
- package/templates/add/rbac/internal/app/role/service.go.hbs +0 -217
- package/templates/generate/module/dto.go.hbs +0 -36
- package/templates/generate/module/errors.go.hbs +0 -33
- package/templates/generate/module/handler.go.hbs +0 -134
- package/templates/generate/module/handler_test.go.hbs +0 -174
- package/templates/generate/module/minimal/dto.go.hbs +0 -28
- package/templates/generate/module/minimal/handler.go.hbs +0 -48
- package/templates/generate/module/minimal/handler_test.go.hbs +0 -10
- package/templates/generate/module/minimal/service.go.hbs +0 -45
- package/templates/generate/module/minimal/service_test.go.hbs +0 -77
- package/templates/generate/module/model/model.go.hbs +0 -36
- package/templates/generate/module/repository.go.hbs +0 -103
- package/templates/generate/module/service.go.hbs +0 -108
- package/templates/generate/module/service_test.go.hbs +0 -161
|
@@ -0,0 +1,200 @@
|
|
|
1
|
+
package application
|
|
2
|
+
|
|
3
|
+
import (
|
|
4
|
+
"context"
|
|
5
|
+
"errors"
|
|
6
|
+
"strings"
|
|
7
|
+
"testing"
|
|
8
|
+
"time"
|
|
9
|
+
|
|
10
|
+
"{{goModule}}/internal/app/user/domain"
|
|
11
|
+
|
|
12
|
+
"github.com/google/uuid"
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
const testMFAEncryptionKey = "MDEyMzQ1Njc4OWFiY2RlZjAxMjM0NTY3ODlhYmNkZWY="
|
|
16
|
+
|
|
17
|
+
func configuredMFATestService(t *testing.T) (*Service, *fakeMFAStore, uuid.UUID, time.Time) {
|
|
18
|
+
t.Helper()
|
|
19
|
+
userID := uuid.New()
|
|
20
|
+
mfa := newFakeMFAStore()
|
|
21
|
+
svc := newTestService(&fakeRepo{user: &domain.User{ID: userID, Email: "mfa@example.com"}}, newFakeTokenStore())
|
|
22
|
+
now := time.Now()
|
|
23
|
+
svc.mfa = mfa
|
|
24
|
+
svc.now = func() time.Time { return now }
|
|
25
|
+
svc.config.MFA = MFASettings{
|
|
26
|
+
Enabled: true,
|
|
27
|
+
Issuer: "Example",
|
|
28
|
+
EncryptionKey: testMFAEncryptionKey,
|
|
29
|
+
ChallengeTTL: 5 * time.Minute,
|
|
30
|
+
TOTPWindow: 1,
|
|
31
|
+
RecoveryCodeCount: 10,
|
|
32
|
+
}
|
|
33
|
+
return svc, mfa, userID, now
|
|
34
|
+
}
|
|
35
|
+
|
|
36
|
+
func TestTOTPCodeMatchesRFC6238SHA1Vector(t *testing.T) {
|
|
37
|
+
secret := "GEZDGNBVGY3TQOJQGEZDGNBVGY3TQOJQ"
|
|
38
|
+
code, err := totpCode(secret, time.Unix(59, 0))
|
|
39
|
+
if err != nil {
|
|
40
|
+
t.Fatalf("generate TOTP: %v", err)
|
|
41
|
+
}
|
|
42
|
+
if code != "287082" {
|
|
43
|
+
t.Fatalf("expected RFC 6238 six-digit code 287082, got %s", code)
|
|
44
|
+
}
|
|
45
|
+
}
|
|
46
|
+
|
|
47
|
+
func TestMFADefaultIsDisabled(t *testing.T) {
|
|
48
|
+
userID := uuid.New()
|
|
49
|
+
svc := newTestService(&fakeRepo{user: &domain.User{ID: userID, Email: "default@example.com"}}, newFakeTokenStore())
|
|
50
|
+
|
|
51
|
+
status, err := svc.MFAStatus(context.Background(), userID)
|
|
52
|
+
if err != nil {
|
|
53
|
+
t.Fatalf("read default MFA status: %v", err)
|
|
54
|
+
}
|
|
55
|
+
if status.Available || status.Enabled {
|
|
56
|
+
t.Fatalf("MFA must be disabled and unavailable by default: %+v", status)
|
|
57
|
+
}
|
|
58
|
+
}
|
|
59
|
+
|
|
60
|
+
func TestValidateMFASettingsRejectsUnsafeEnabledConfiguration(t *testing.T) {
|
|
61
|
+
valid := MFASettings{
|
|
62
|
+
Enabled: true,
|
|
63
|
+
EncryptionKey: testMFAEncryptionKey,
|
|
64
|
+
ChallengeTTL: 5 * time.Minute,
|
|
65
|
+
TOTPWindow: 1,
|
|
66
|
+
RecoveryCodeCount: 10,
|
|
67
|
+
}
|
|
68
|
+
if err := ValidateMFASettings(valid); err != nil {
|
|
69
|
+
t.Fatalf("valid MFA settings rejected: %v", err)
|
|
70
|
+
}
|
|
71
|
+
|
|
72
|
+
for name, settings := range map[string]MFASettings{
|
|
73
|
+
"missing encryption key": validWithMFAOverride(valid, func(s *MFASettings) { s.EncryptionKey = "" }),
|
|
74
|
+
"non-positive challenge": validWithMFAOverride(valid, func(s *MFASettings) { s.ChallengeTTL = 0 }),
|
|
75
|
+
"wide TOTP window": validWithMFAOverride(valid, func(s *MFASettings) { s.TOTPWindow = 4 }),
|
|
76
|
+
"too few recovery codes": validWithMFAOverride(valid, func(s *MFASettings) { s.RecoveryCodeCount = 4 }),
|
|
77
|
+
"too many recovery codes": validWithMFAOverride(valid, func(s *MFASettings) { s.RecoveryCodeCount = 21 }),
|
|
78
|
+
} {
|
|
79
|
+
if err := ValidateMFASettings(settings); err == nil {
|
|
80
|
+
t.Errorf("%s: expected invalid enabled configuration to fail", name)
|
|
81
|
+
}
|
|
82
|
+
}
|
|
83
|
+
if err := ValidateMFASettings(MFASettings{Enabled: false}); err != nil {
|
|
84
|
+
t.Fatalf("disabled MFA should not require an encryption key: %v", err)
|
|
85
|
+
}
|
|
86
|
+
}
|
|
87
|
+
|
|
88
|
+
func validWithMFAOverride(base MFASettings, override func(*MFASettings)) MFASettings {
|
|
89
|
+
copy := base
|
|
90
|
+
override(©)
|
|
91
|
+
return copy
|
|
92
|
+
}
|
|
93
|
+
|
|
94
|
+
func TestMFAEnrollmentIsPendingUntilTOTPConfirmation(t *testing.T) {
|
|
95
|
+
svc, store, userID, now := configuredMFATestService(t)
|
|
96
|
+
ctx := context.Background()
|
|
97
|
+
|
|
98
|
+
setup, err := svc.SetupMFA(ctx, userID)
|
|
99
|
+
if err != nil {
|
|
100
|
+
t.Fatalf("setup MFA: %v", err)
|
|
101
|
+
}
|
|
102
|
+
if setup.Secret == "" || !strings.HasPrefix(setup.OTPAuthURI, "otpauth://totp/") {
|
|
103
|
+
t.Fatalf("unexpected setup response: %+v", setup)
|
|
104
|
+
}
|
|
105
|
+
if strings.Contains(store.enrollments[userID].EncryptedSecret, setup.Secret) {
|
|
106
|
+
t.Fatal("MFA secret must not be stored in plaintext")
|
|
107
|
+
}
|
|
108
|
+
|
|
109
|
+
status, err := svc.MFAStatus(ctx, userID)
|
|
110
|
+
if err != nil || !status.Available || status.Enabled {
|
|
111
|
+
t.Fatalf("expected available but disabled pending enrollment: %+v, %v", status, err)
|
|
112
|
+
}
|
|
113
|
+
code, _ := totpCode(setup.Secret, now)
|
|
114
|
+
codes, err := svc.ConfirmMFA(ctx, userID, code)
|
|
115
|
+
if err != nil {
|
|
116
|
+
t.Fatalf("confirm MFA: %v", err)
|
|
117
|
+
}
|
|
118
|
+
if len(codes) != 10 {
|
|
119
|
+
t.Fatalf("expected 10 recovery codes, got %d", len(codes))
|
|
120
|
+
}
|
|
121
|
+
status, err = svc.MFAStatus(ctx, userID)
|
|
122
|
+
if err != nil || !status.Enabled {
|
|
123
|
+
t.Fatalf("expected enabled MFA status: %+v, %v", status, err)
|
|
124
|
+
}
|
|
125
|
+
}
|
|
126
|
+
|
|
127
|
+
func TestMFAChallengeIsRequiredAndOneUse(t *testing.T) {
|
|
128
|
+
svc, _, userID, now := configuredMFATestService(t)
|
|
129
|
+
ctx := context.Background()
|
|
130
|
+
secret, err := newTOTPSecret()
|
|
131
|
+
if err != nil {
|
|
132
|
+
t.Fatalf("secret: %v", err)
|
|
133
|
+
}
|
|
134
|
+
encrypted, err := encryptMFASecret(testMFAEncryptionKey, secret)
|
|
135
|
+
if err != nil {
|
|
136
|
+
t.Fatalf("encrypt: %v", err)
|
|
137
|
+
}
|
|
138
|
+
if err := svc.mfa.ConfirmEnrollment(ctx, userID, encrypted, []string{}); err == nil {
|
|
139
|
+
t.Fatal("expected store to reject an enrollment without recovery codes")
|
|
140
|
+
}
|
|
141
|
+
// Use the service path so the exact recovery-code policy is exercised.
|
|
142
|
+
if err := svc.mfa.PutPendingEnrollment(ctx, userID, encrypted); err != nil {
|
|
143
|
+
t.Fatalf("pending: %v", err)
|
|
144
|
+
}
|
|
145
|
+
code, _ := totpCode(secret, now)
|
|
146
|
+
if _, err := svc.ConfirmMFA(ctx, userID, code); err != nil {
|
|
147
|
+
t.Fatalf("confirm: %v", err)
|
|
148
|
+
}
|
|
149
|
+
|
|
150
|
+
result, err := svc.completeLogin(ctx, &domain.User{ID: userID, Email: "mfa@example.com"})
|
|
151
|
+
if err != nil || result.MFAChallenge == "" || result.AuthResponse != nil {
|
|
152
|
+
t.Fatalf("expected pre-session MFA challenge: %+v, %v", result, err)
|
|
153
|
+
}
|
|
154
|
+
auth, err := svc.VerifyMFA(ctx, result.MFAChallenge, code)
|
|
155
|
+
if err != nil || auth == nil || auth.AccessToken == "" || auth.RefreshToken == "" {
|
|
156
|
+
t.Fatalf("verify MFA: %+v, %v", auth, err)
|
|
157
|
+
}
|
|
158
|
+
if _, err := svc.VerifyMFA(ctx, result.MFAChallenge, code); codeOfMFAError(err) != "AUTH_MFA_INVALID" {
|
|
159
|
+
t.Fatalf("expected one-use challenge rejection, got %v", err)
|
|
160
|
+
}
|
|
161
|
+
}
|
|
162
|
+
|
|
163
|
+
func TestMFARecoveryCodeIsOneUse(t *testing.T) {
|
|
164
|
+
svc, _, userID, now := configuredMFATestService(t)
|
|
165
|
+
ctx := context.Background()
|
|
166
|
+
setup, err := svc.SetupMFA(ctx, userID)
|
|
167
|
+
if err != nil {
|
|
168
|
+
t.Fatalf("setup: %v", err)
|
|
169
|
+
}
|
|
170
|
+
code, _ := totpCode(setup.Secret, now)
|
|
171
|
+
recoveryCodes, err := svc.ConfirmMFA(ctx, userID, code)
|
|
172
|
+
if err != nil {
|
|
173
|
+
t.Fatalf("confirm: %v", err)
|
|
174
|
+
}
|
|
175
|
+
result, err := svc.completeLogin(ctx, &domain.User{ID: userID, Email: "mfa@example.com"})
|
|
176
|
+
if err != nil {
|
|
177
|
+
t.Fatalf("login challenge: %v", err)
|
|
178
|
+
}
|
|
179
|
+
if _, err := svc.VerifyMFA(ctx, result.MFAChallenge, recoveryCodes[0]); err != nil {
|
|
180
|
+
t.Fatalf("verify with recovery code: %v", err)
|
|
181
|
+
}
|
|
182
|
+
result, err = svc.completeLogin(ctx, &domain.User{ID: userID, Email: "mfa@example.com"})
|
|
183
|
+
if err != nil {
|
|
184
|
+
t.Fatalf("second login challenge: %v", err)
|
|
185
|
+
}
|
|
186
|
+
if _, err := svc.VerifyMFA(ctx, result.MFAChallenge, recoveryCodes[0]); codeOfMFAError(err) != "AUTH_MFA_INVALID" {
|
|
187
|
+
t.Fatalf("expected consumed recovery code rejection, got %v", err)
|
|
188
|
+
}
|
|
189
|
+
}
|
|
190
|
+
|
|
191
|
+
func codeOfMFAError(err error) string {
|
|
192
|
+
if err == nil {
|
|
193
|
+
return ""
|
|
194
|
+
}
|
|
195
|
+
var ruleErr *domain.RuleError
|
|
196
|
+
if errors.As(err, &ruleErr) {
|
|
197
|
+
return ruleErr.Code
|
|
198
|
+
}
|
|
199
|
+
return ""
|
|
200
|
+
}
|
|
@@ -0,0 +1,132 @@
|
|
|
1
|
+
// Package application contains auth use cases and ports that are independent
|
|
2
|
+
// of Gin, GORM models, OAuth SDKs, and HTTP error payloads.
|
|
3
|
+
package application
|
|
4
|
+
|
|
5
|
+
import (
|
|
6
|
+
"context"
|
|
7
|
+
"errors"
|
|
8
|
+
"fmt"
|
|
9
|
+
)
|
|
10
|
+
|
|
11
|
+
// LoginProvider is the outbound port for an external authorization-code
|
|
12
|
+
// provider. Implementations own provider SDKs, token exchange, claims/userinfo
|
|
13
|
+
// validation, and provider-specific HTTP details. The user application only
|
|
14
|
+
// sees these small, provider-neutral values.
|
|
15
|
+
type LoginProvider interface {
|
|
16
|
+
Name() string
|
|
17
|
+
Begin(context.Context, LoginStartInput) (Authorization, error)
|
|
18
|
+
Complete(context.Context, LoginCompleteInput) (ExternalIdentity, error)
|
|
19
|
+
}
|
|
20
|
+
|
|
21
|
+
type LoginStartInput struct {
|
|
22
|
+
State string
|
|
23
|
+
CodeChallenge string
|
|
24
|
+
CodeChallengeMethod string
|
|
25
|
+
Nonce string
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
type Authorization struct {
|
|
29
|
+
URL string
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
type LoginCompleteInput struct {
|
|
33
|
+
Code string
|
|
34
|
+
CodeVerifier string
|
|
35
|
+
Nonce string
|
|
36
|
+
}
|
|
37
|
+
|
|
38
|
+
// ExternalIdentity is the normalized identity returned by every provider
|
|
39
|
+
// adapter. Provider must be the registry name and Subject must be the
|
|
40
|
+
// provider's stable subject identifier, never an email address.
|
|
41
|
+
type ExternalIdentity struct {
|
|
42
|
+
Provider string
|
|
43
|
+
Subject string
|
|
44
|
+
Email string
|
|
45
|
+
EmailVerified bool
|
|
46
|
+
Name string
|
|
47
|
+
AvatarURL string
|
|
48
|
+
}
|
|
49
|
+
|
|
50
|
+
// ProviderRegistry is an immutable-by-convention lookup boundary. It is
|
|
51
|
+
// assembled in the composition root from providers whose configuration is
|
|
52
|
+
// complete; an absent provider is therefore an expected unavailable state,
|
|
53
|
+
// not an invitation to construct an empty redirect URL.
|
|
54
|
+
type ProviderRegistry struct {
|
|
55
|
+
providers map[string]LoginProvider
|
|
56
|
+
}
|
|
57
|
+
|
|
58
|
+
func NewProviderRegistry(providers ...LoginProvider) ProviderRegistry {
|
|
59
|
+
byName := make(map[string]LoginProvider, len(providers))
|
|
60
|
+
for _, provider := range providers {
|
|
61
|
+
if provider == nil || provider.Name() == "" {
|
|
62
|
+
continue
|
|
63
|
+
}
|
|
64
|
+
byName[provider.Name()] = provider
|
|
65
|
+
}
|
|
66
|
+
return ProviderRegistry{providers: byName}
|
|
67
|
+
}
|
|
68
|
+
|
|
69
|
+
func (r ProviderRegistry) Lookup(name string) (LoginProvider, bool) {
|
|
70
|
+
provider, ok := r.providers[name]
|
|
71
|
+
return provider, ok
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
type OAuthErrorCode string
|
|
75
|
+
|
|
76
|
+
const (
|
|
77
|
+
OAuthDenied OAuthErrorCode = "oauth_denied"
|
|
78
|
+
OAuthStateInvalid OAuthErrorCode = "oauth_state_invalid"
|
|
79
|
+
OAuthProviderUnavailable OAuthErrorCode = "oauth_provider_unavailable"
|
|
80
|
+
OAuthFailed OAuthErrorCode = "oauth_failed"
|
|
81
|
+
)
|
|
82
|
+
|
|
83
|
+
// OAuthError is the public, controlled error contract for the browser login
|
|
84
|
+
// start/exchange boundary. Cause is retained only for server-side logs and
|
|
85
|
+
// unwrapping; handlers never serialize Error() to the browser.
|
|
86
|
+
type OAuthError struct {
|
|
87
|
+
Code OAuthErrorCode
|
|
88
|
+
Cause error
|
|
89
|
+
}
|
|
90
|
+
|
|
91
|
+
func (e *OAuthError) Error() string {
|
|
92
|
+
if e.Cause == nil {
|
|
93
|
+
return string(e.Code)
|
|
94
|
+
}
|
|
95
|
+
return fmt.Sprintf("%s: %v", e.Code, e.Cause)
|
|
96
|
+
}
|
|
97
|
+
|
|
98
|
+
func (e *OAuthError) Unwrap() error { return e.Cause }
|
|
99
|
+
|
|
100
|
+
func NewOAuthError(code OAuthErrorCode, cause error) *OAuthError {
|
|
101
|
+
return &OAuthError{Code: code, Cause: cause}
|
|
102
|
+
}
|
|
103
|
+
|
|
104
|
+
type ProviderError struct {
|
|
105
|
+
Unavailable bool
|
|
106
|
+
Cause error
|
|
107
|
+
}
|
|
108
|
+
|
|
109
|
+
func (e *ProviderError) Error() string {
|
|
110
|
+
if e.Cause == nil {
|
|
111
|
+
if e.Unavailable {
|
|
112
|
+
return "oauth provider unavailable"
|
|
113
|
+
}
|
|
114
|
+
return "oauth provider failed"
|
|
115
|
+
}
|
|
116
|
+
return e.Cause.Error()
|
|
117
|
+
}
|
|
118
|
+
|
|
119
|
+
func (e *ProviderError) Unwrap() error { return e.Cause }
|
|
120
|
+
|
|
121
|
+
func NewProviderUnavailable(cause error) error {
|
|
122
|
+
return &ProviderError{Unavailable: true, Cause: cause}
|
|
123
|
+
}
|
|
124
|
+
|
|
125
|
+
func NewProviderFailure(cause error) error {
|
|
126
|
+
return &ProviderError{Cause: cause}
|
|
127
|
+
}
|
|
128
|
+
|
|
129
|
+
func IsProviderUnavailable(err error) bool {
|
|
130
|
+
var providerErr *ProviderError
|
|
131
|
+
return errors.As(err, &providerErr) && providerErr.Unavailable
|
|
132
|
+
}
|
|
@@ -0,0 +1,285 @@
|
|
|
1
|
+
package application
|
|
2
|
+
|
|
3
|
+
import (
|
|
4
|
+
"context"
|
|
5
|
+
"errors"
|
|
6
|
+
"net/url"
|
|
7
|
+
"strings"
|
|
8
|
+
"testing"
|
|
9
|
+
"time"
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
type fakeLoginProvider struct {
|
|
13
|
+
name string
|
|
14
|
+
beginFn func(LoginStartInput) (Authorization, error)
|
|
15
|
+
completeFn func(LoginCompleteInput) (ExternalIdentity, error)
|
|
16
|
+
beginCnt int
|
|
17
|
+
completeCnt int
|
|
18
|
+
}
|
|
19
|
+
|
|
20
|
+
func (p *fakeLoginProvider) Name() string { return p.name }
|
|
21
|
+
|
|
22
|
+
func (p *fakeLoginProvider) Begin(_ context.Context, in LoginStartInput) (Authorization, error) {
|
|
23
|
+
p.beginCnt++
|
|
24
|
+
return p.beginFn(in)
|
|
25
|
+
}
|
|
26
|
+
|
|
27
|
+
func (p *fakeLoginProvider) Complete(_ context.Context, in LoginCompleteInput) (ExternalIdentity, error) {
|
|
28
|
+
p.completeCnt++
|
|
29
|
+
return p.completeFn(in)
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
func newTestServiceWithProviders(repo repository, tokens testTokenStore, providers ...LoginProvider) *Service {
|
|
33
|
+
return NewService(Dependencies{
|
|
34
|
+
Repository: repo,
|
|
35
|
+
Passwords: fakePasswordHasher{},
|
|
36
|
+
RefreshTokens: tokens,
|
|
37
|
+
OAuthTransactions: tokens,
|
|
38
|
+
RecoveryTokens: tokens,
|
|
39
|
+
MFA: newFakeMFAStore(),
|
|
40
|
+
Mailer: fakeMailer{},
|
|
41
|
+
Providers: NewProviderRegistry(providers...),
|
|
42
|
+
}, AuthConfig{
|
|
43
|
+
JWTSecret: "test-secret",
|
|
44
|
+
JWTAccessTTL: time.Minute,
|
|
45
|
+
JWTRefreshTTL: time.Hour,
|
|
46
|
+
})
|
|
47
|
+
}
|
|
48
|
+
|
|
49
|
+
func oauthErrorCode(t *testing.T, err error) OAuthErrorCode {
|
|
50
|
+
t.Helper()
|
|
51
|
+
var oauthErr *OAuthError
|
|
52
|
+
if !errors.As(err, &oauthErr) {
|
|
53
|
+
t.Fatalf("expected *OAuthError, got %T: %v", err, err)
|
|
54
|
+
}
|
|
55
|
+
return oauthErr.Code
|
|
56
|
+
}
|
|
57
|
+
|
|
58
|
+
const testOAuthVerifier = "client-code-verifier-0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZ"
|
|
59
|
+
|
|
60
|
+
func validLoginStart() LoginStartInput {
|
|
61
|
+
return LoginStartInput{
|
|
62
|
+
State: "client state",
|
|
63
|
+
CodeChallenge: pkceChallenge(testOAuthVerifier),
|
|
64
|
+
CodeChallengeMethod: "S256",
|
|
65
|
+
}
|
|
66
|
+
}
|
|
67
|
+
|
|
68
|
+
func validLoginExchange() LoginExchangeInput {
|
|
69
|
+
return LoginExchangeInput{Code: "code", State: "client state", CodeVerifier: testOAuthVerifier}
|
|
70
|
+
}
|
|
71
|
+
|
|
72
|
+
func validFakeProvider(t *testing.T) (*fakeLoginProvider, *Service) {
|
|
73
|
+
t.Helper()
|
|
74
|
+
provider := &fakeLoginProvider{name: "fake"}
|
|
75
|
+
provider.beginFn = func(in LoginStartInput) (Authorization, error) {
|
|
76
|
+
expected := validLoginStart()
|
|
77
|
+
if in.State != expected.State || in.CodeChallenge != expected.CodeChallenge || in.CodeChallengeMethod != expected.CodeChallengeMethod || strings.TrimSpace(in.Nonce) == "" {
|
|
78
|
+
t.Fatalf("unexpected begin input: %+v", in)
|
|
79
|
+
}
|
|
80
|
+
return Authorization{URL: "https://provider.example.test/authorize?" + url.Values{
|
|
81
|
+
"state": {in.State},
|
|
82
|
+
"code_challenge": {in.CodeChallenge},
|
|
83
|
+
"code_challenge_method": {in.CodeChallengeMethod},
|
|
84
|
+
"nonce": {in.Nonce},
|
|
85
|
+
}.Encode()}, nil
|
|
86
|
+
}
|
|
87
|
+
provider.completeFn = func(in LoginCompleteInput) (ExternalIdentity, error) {
|
|
88
|
+
if in.Code != "code" || in.CodeVerifier != testOAuthVerifier || strings.TrimSpace(in.Nonce) == "" {
|
|
89
|
+
t.Fatalf("unexpected complete input: %+v", in)
|
|
90
|
+
}
|
|
91
|
+
return ExternalIdentity{
|
|
92
|
+
Provider: "fake",
|
|
93
|
+
Subject: "subject-1",
|
|
94
|
+
Email: "user@example.com",
|
|
95
|
+
EmailVerified: true,
|
|
96
|
+
Name: "User",
|
|
97
|
+
}, nil
|
|
98
|
+
}
|
|
99
|
+
return provider, newTestServiceWithProviders(&fakeRepo{}, newFakeTokenStore(), provider)
|
|
100
|
+
}
|
|
101
|
+
|
|
102
|
+
func TestService_BeginLoginUsesRegistryAndPreservesClientStateAndS256PKCE(t *testing.T) {
|
|
103
|
+
provider, svc := validFakeProvider(t)
|
|
104
|
+
start, err := svc.BeginLogin(context.Background(), "fake", validLoginStart())
|
|
105
|
+
if err != nil {
|
|
106
|
+
t.Fatalf("begin login: %v", err)
|
|
107
|
+
}
|
|
108
|
+
if provider.beginCnt != 1 || start == nil || start.URL == "" {
|
|
109
|
+
t.Fatalf("unexpected begin result: provider=%+v start=%+v", provider, start)
|
|
110
|
+
}
|
|
111
|
+
parsed, err := url.Parse(start.URL)
|
|
112
|
+
if err != nil {
|
|
113
|
+
t.Fatalf("parse authorization URL: %v", err)
|
|
114
|
+
}
|
|
115
|
+
if parsed.Query().Get("state") != "client state" || parsed.Query().Get("code_challenge") != pkceChallenge(testOAuthVerifier) || parsed.Query().Get("code_challenge_method") != "S256" || parsed.Query().Get("nonce") == "" {
|
|
116
|
+
t.Fatalf("authorization URL did not preserve client state/PKCE: %s", start.URL)
|
|
117
|
+
}
|
|
118
|
+
}
|
|
119
|
+
|
|
120
|
+
func TestService_BeginLoginRejectsMissingStateOrNonS256PKCE(t *testing.T) {
|
|
121
|
+
provider, svc := validFakeProvider(t)
|
|
122
|
+
tests := []struct {
|
|
123
|
+
name string
|
|
124
|
+
in LoginStartInput
|
|
125
|
+
}{
|
|
126
|
+
{name: "missing state", in: LoginStartInput{CodeChallenge: "challenge", CodeChallengeMethod: "S256"}},
|
|
127
|
+
{name: "missing challenge", in: LoginStartInput{State: "state", CodeChallengeMethod: "S256"}},
|
|
128
|
+
{name: "plain challenge method", in: LoginStartInput{State: "state", CodeChallenge: "challenge", CodeChallengeMethod: "plain"}},
|
|
129
|
+
{name: "short challenge", in: LoginStartInput{State: "state", CodeChallenge: "short", CodeChallengeMethod: "S256"}},
|
|
130
|
+
{name: "invalid challenge character", in: LoginStartInput{State: "state", CodeChallenge: strings.Repeat("!", 43), CodeChallengeMethod: "S256"}},
|
|
131
|
+
{name: "lowercase challenge method", in: LoginStartInput{State: "state", CodeChallenge: strings.Repeat("A", 43), CodeChallengeMethod: "s256"}},
|
|
132
|
+
}
|
|
133
|
+
for _, tt := range tests {
|
|
134
|
+
t.Run(tt.name, func(t *testing.T) {
|
|
135
|
+
_, err := svc.BeginLogin(context.Background(), "fake", tt.in)
|
|
136
|
+
if got := oauthErrorCode(t, err); got != OAuthStateInvalid {
|
|
137
|
+
t.Fatalf("oauth code = %q, want %q", got, OAuthStateInvalid)
|
|
138
|
+
}
|
|
139
|
+
if provider.beginCnt != 0 {
|
|
140
|
+
t.Fatal("provider must not receive an invalid state/PKCE request")
|
|
141
|
+
}
|
|
142
|
+
})
|
|
143
|
+
}
|
|
144
|
+
}
|
|
145
|
+
|
|
146
|
+
func TestService_ExchangeLoginReturnsSessionForNormalizedProviderIdentity(t *testing.T) {
|
|
147
|
+
provider, svc := validFakeProvider(t)
|
|
148
|
+
if _, err := svc.BeginLogin(context.Background(), "fake", validLoginStart()); err != nil {
|
|
149
|
+
t.Fatalf("begin login: %v", err)
|
|
150
|
+
}
|
|
151
|
+
auth, err := svc.ExchangeLogin(context.Background(), "fake", validLoginExchange())
|
|
152
|
+
if err != nil {
|
|
153
|
+
t.Fatalf("exchange login: %v", err)
|
|
154
|
+
}
|
|
155
|
+
if auth.AccessToken == "" || auth.RefreshToken == "" || provider.completeCnt != 1 {
|
|
156
|
+
t.Fatalf("expected local session after provider completion: auth=%+v complete calls=%d", auth, provider.completeCnt)
|
|
157
|
+
}
|
|
158
|
+
}
|
|
159
|
+
|
|
160
|
+
func TestService_ExchangeLoginRejectsIdentityFromAnotherProvider(t *testing.T) {
|
|
161
|
+
provider, svc := validFakeProvider(t)
|
|
162
|
+
if _, err := svc.BeginLogin(context.Background(), "fake", validLoginStart()); err != nil {
|
|
163
|
+
t.Fatalf("begin login: %v", err)
|
|
164
|
+
}
|
|
165
|
+
provider.completeFn = func(LoginCompleteInput) (ExternalIdentity, error) {
|
|
166
|
+
return ExternalIdentity{
|
|
167
|
+
Provider: "another-provider",
|
|
168
|
+
Subject: "subject-1",
|
|
169
|
+
Email: "user@example.com",
|
|
170
|
+
EmailVerified: true,
|
|
171
|
+
}, nil
|
|
172
|
+
}
|
|
173
|
+
_, err := svc.ExchangeLogin(context.Background(), "fake", validLoginExchange())
|
|
174
|
+
if got := oauthErrorCode(t, err); got != OAuthFailed {
|
|
175
|
+
t.Fatalf("oauth code = %q, want %q", got, OAuthFailed)
|
|
176
|
+
}
|
|
177
|
+
}
|
|
178
|
+
|
|
179
|
+
func TestService_ExchangeLoginMapsProviderErrorsToControlledCodes(t *testing.T) {
|
|
180
|
+
tests := []struct {
|
|
181
|
+
name string
|
|
182
|
+
providerError error
|
|
183
|
+
wantCode OAuthErrorCode
|
|
184
|
+
}{
|
|
185
|
+
{name: "provider exchange failure", providerError: NewProviderFailure(errors.New("exchange failed")), wantCode: OAuthFailed},
|
|
186
|
+
{name: "provider unavailable", providerError: NewProviderUnavailable(errors.New("upstream timeout")), wantCode: OAuthProviderUnavailable},
|
|
187
|
+
}
|
|
188
|
+
|
|
189
|
+
for _, tt := range tests {
|
|
190
|
+
t.Run(tt.name, func(t *testing.T) {
|
|
191
|
+
provider, svc := validFakeProvider(t)
|
|
192
|
+
if _, err := svc.BeginLogin(context.Background(), "fake", validLoginStart()); err != nil {
|
|
193
|
+
t.Fatalf("begin login: %v", err)
|
|
194
|
+
}
|
|
195
|
+
provider.completeFn = func(LoginCompleteInput) (ExternalIdentity, error) {
|
|
196
|
+
return ExternalIdentity{}, tt.providerError
|
|
197
|
+
}
|
|
198
|
+
_, err := svc.ExchangeLogin(context.Background(), "fake", validLoginExchange())
|
|
199
|
+
if got := oauthErrorCode(t, err); got != tt.wantCode {
|
|
200
|
+
t.Fatalf("oauth code = %q, want %q", got, tt.wantCode)
|
|
201
|
+
}
|
|
202
|
+
})
|
|
203
|
+
}
|
|
204
|
+
}
|
|
205
|
+
|
|
206
|
+
func TestService_ExchangeLoginRejectsMissingStatePKCEAndCode(t *testing.T) {
|
|
207
|
+
_, svc := validFakeProvider(t)
|
|
208
|
+
tests := []struct {
|
|
209
|
+
name string
|
|
210
|
+
in LoginExchangeInput
|
|
211
|
+
want OAuthErrorCode
|
|
212
|
+
}{
|
|
213
|
+
{name: "missing state", in: LoginExchangeInput{Code: "code", CodeVerifier: testOAuthVerifier}, want: OAuthStateInvalid},
|
|
214
|
+
{name: "missing verifier", in: LoginExchangeInput{Code: "code", State: "state"}, want: OAuthStateInvalid},
|
|
215
|
+
{name: "missing code", in: LoginExchangeInput{State: "state", CodeVerifier: testOAuthVerifier}, want: OAuthFailed},
|
|
216
|
+
}
|
|
217
|
+
for _, tt := range tests {
|
|
218
|
+
t.Run(tt.name, func(t *testing.T) {
|
|
219
|
+
_, err := svc.ExchangeLogin(context.Background(), "fake", tt.in)
|
|
220
|
+
if got := oauthErrorCode(t, err); got != tt.want {
|
|
221
|
+
t.Fatalf("oauth code = %q, want %q", got, tt.want)
|
|
222
|
+
}
|
|
223
|
+
})
|
|
224
|
+
}
|
|
225
|
+
}
|
|
226
|
+
|
|
227
|
+
func TestService_ExchangeLoginConsumesStateAndBindsThePKCEVerifier(t *testing.T) {
|
|
228
|
+
provider, svc := validFakeProvider(t)
|
|
229
|
+
if _, err := svc.BeginLogin(context.Background(), "fake", validLoginStart()); err != nil {
|
|
230
|
+
t.Fatalf("begin login: %v", err)
|
|
231
|
+
}
|
|
232
|
+
wrong := validLoginExchange()
|
|
233
|
+
wrong.CodeVerifier = testOAuthVerifier + "x"
|
|
234
|
+
if _, err := svc.ExchangeLogin(context.Background(), "fake", wrong); oauthErrorCode(t, err) != OAuthStateInvalid {
|
|
235
|
+
t.Fatalf("wrong verifier should be rejected as state invalid: %v", err)
|
|
236
|
+
}
|
|
237
|
+
if _, err := svc.ExchangeLogin(context.Background(), "fake", validLoginExchange()); oauthErrorCode(t, err) != OAuthStateInvalid {
|
|
238
|
+
t.Fatalf("a rejected exchange must consume the one-time transaction: %v", err)
|
|
239
|
+
}
|
|
240
|
+
if provider.completeCnt != 0 {
|
|
241
|
+
t.Fatal("provider must not receive a wrong PKCE verifier")
|
|
242
|
+
}
|
|
243
|
+
}
|
|
244
|
+
|
|
245
|
+
func TestService_UnconfiguredProviderDoesNotBuildAURL(t *testing.T) {
|
|
246
|
+
svc := newTestServiceWithProviders(&fakeRepo{}, newFakeTokenStore())
|
|
247
|
+
start, err := svc.BeginLogin(context.Background(), "google", validLoginStart())
|
|
248
|
+
if start != nil {
|
|
249
|
+
t.Fatalf("unconfigured provider returned a login start: %+v", start)
|
|
250
|
+
}
|
|
251
|
+
if got := oauthErrorCode(t, err); got != OAuthProviderUnavailable {
|
|
252
|
+
t.Fatalf("oauth code = %q, want %q", got, OAuthProviderUnavailable)
|
|
253
|
+
}
|
|
254
|
+
_, err = svc.ExchangeLogin(context.Background(), "google", validLoginExchange())
|
|
255
|
+
if got := oauthErrorCode(t, err); got != OAuthProviderUnavailable {
|
|
256
|
+
t.Fatalf("exchange oauth code = %q, want %q", got, OAuthProviderUnavailable)
|
|
257
|
+
}
|
|
258
|
+
}
|
|
259
|
+
|
|
260
|
+
func TestProviderRegistry_UsesProviderNameAsTheOnlyLookupKey(t *testing.T) {
|
|
261
|
+
provider := &fakeLoginProvider{name: "fake"}
|
|
262
|
+
registry := NewProviderRegistry(provider)
|
|
263
|
+
if got, ok := registry.Lookup("fake"); !ok || got != provider {
|
|
264
|
+
t.Fatalf("registry lookup = (%v, %v), want fake provider", got, ok)
|
|
265
|
+
}
|
|
266
|
+
if _, ok := registry.Lookup("google"); ok {
|
|
267
|
+
t.Fatal("registry must not synthesize an unconfigured provider")
|
|
268
|
+
}
|
|
269
|
+
}
|
|
270
|
+
|
|
271
|
+
func TestService_BeginLoginRejectsEmptyProviderAuthorizationURL(t *testing.T) {
|
|
272
|
+
provider := &fakeLoginProvider{name: "fake", beginFn: func(LoginStartInput) (Authorization, error) {
|
|
273
|
+
return Authorization{}, nil
|
|
274
|
+
}}
|
|
275
|
+
svc := newTestServiceWithProviders(&fakeRepo{}, newFakeTokenStore(), provider)
|
|
276
|
+
_, err := svc.BeginLogin(context.Background(), "fake", validLoginStart())
|
|
277
|
+
if got := oauthErrorCode(t, err); got != OAuthProviderUnavailable {
|
|
278
|
+
t.Fatalf("oauth code = %q, want %q", got, OAuthProviderUnavailable)
|
|
279
|
+
}
|
|
280
|
+
if strings.Contains(err.Error(), "Location") {
|
|
281
|
+
t.Fatalf("empty authorization must not become a redirect: %v", err)
|
|
282
|
+
}
|
|
283
|
+
}
|
|
284
|
+
|
|
285
|
+
var _ LoginProvider = (*fakeLoginProvider)(nil)
|
|
@@ -0,0 +1,82 @@
|
|
|
1
|
+
package application
|
|
2
|
+
|
|
3
|
+
import (
|
|
4
|
+
"context"
|
|
5
|
+
"errors"
|
|
6
|
+
"fmt"
|
|
7
|
+
|
|
8
|
+
"{{goModule}}/internal/app/user/domain"
|
|
9
|
+
"{{goModule}}/internal/app/user/ports"
|
|
10
|
+
|
|
11
|
+
"github.com/google/uuid"
|
|
12
|
+
)
|
|
13
|
+
|
|
14
|
+
type Recovery struct {
|
|
15
|
+
repo ports.UserRepository
|
|
16
|
+
tokens ports.RecoveryTokenStore
|
|
17
|
+
hasher PasswordHasher
|
|
18
|
+
}
|
|
19
|
+
|
|
20
|
+
func NewRecovery(repo ports.UserRepository, tokens ports.RecoveryTokenStore, hasher PasswordHasher) *Recovery {
|
|
21
|
+
return &Recovery{repo: repo, tokens: tokens, hasher: hasher}
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
// ResetPassword consumes the token and updates the local identity in one
|
|
25
|
+
// unit of work. The token adapter owns the transaction boundary.
|
|
26
|
+
func (r *Recovery) ResetPassword(ctx context.Context, tokenHash, newPassword string) (uuid.UUID, error) {
|
|
27
|
+
hash, err := r.hasher.Hash(newPassword)
|
|
28
|
+
if err != nil {
|
|
29
|
+
return uuid.Nil, fmt.Errorf("hash password: %w", err)
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
var userID uuid.UUID
|
|
33
|
+
err = r.tokens.WithTransaction(ctx, func(txctx context.Context) error {
|
|
34
|
+
id, ok, err := r.tokens.ConsumePasswordResetToken(txctx, tokenHash)
|
|
35
|
+
if err != nil {
|
|
36
|
+
return fmt.Errorf("consume password reset token: %w", err)
|
|
37
|
+
}
|
|
38
|
+
if !ok {
|
|
39
|
+
return domain.ErrInvalidToken
|
|
40
|
+
}
|
|
41
|
+
|
|
42
|
+
identity, err := r.repo.FindIdentity(txctx, id, domain.ProviderLocal)
|
|
43
|
+
if err != nil {
|
|
44
|
+
return fmt.Errorf("find local identity: %w", err)
|
|
45
|
+
}
|
|
46
|
+
if identity.PasswordHash == nil {
|
|
47
|
+
return errors.New("local identity has no password")
|
|
48
|
+
}
|
|
49
|
+
identity.PasswordHash = &hash
|
|
50
|
+
if err := r.repo.UpdateIdentity(txctx, identity); err != nil {
|
|
51
|
+
return fmt.Errorf("update password: %w", err)
|
|
52
|
+
}
|
|
53
|
+
userID = id
|
|
54
|
+
return nil
|
|
55
|
+
})
|
|
56
|
+
return userID, err
|
|
57
|
+
}
|
|
58
|
+
|
|
59
|
+
func (r *Recovery) VerifyEmail(ctx context.Context, tokenHash string) (uuid.UUID, error) {
|
|
60
|
+
var userID uuid.UUID
|
|
61
|
+
err := r.tokens.WithTransaction(ctx, func(txctx context.Context) error {
|
|
62
|
+
id, ok, err := r.tokens.ConsumeEmailVerifyToken(txctx, tokenHash)
|
|
63
|
+
if err != nil {
|
|
64
|
+
return fmt.Errorf("consume email verification token: %w", err)
|
|
65
|
+
}
|
|
66
|
+
if !ok {
|
|
67
|
+
return domain.ErrInvalidToken
|
|
68
|
+
}
|
|
69
|
+
|
|
70
|
+
user, err := r.repo.FindByID(txctx, id)
|
|
71
|
+
if err != nil {
|
|
72
|
+
return fmt.Errorf("find user for email verification: %w", err)
|
|
73
|
+
}
|
|
74
|
+
user.EmailVerified = true
|
|
75
|
+
if err := r.repo.UpdateUser(txctx, user); err != nil {
|
|
76
|
+
return fmt.Errorf("mark email verified: %w", err)
|
|
77
|
+
}
|
|
78
|
+
userID = id
|
|
79
|
+
return nil
|
|
80
|
+
})
|
|
81
|
+
return userID, err
|
|
82
|
+
}
|