Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 16 additions & 15 deletions internal/api/handler_auth.go
Original file line number Diff line number Diff line change
Expand Up @@ -503,26 +503,27 @@ var (
)

// mapMFAServiceError maps a service-layer MFA error to the right
// client-facing status + message. Service errors land here as plain
// fmt.Errorf strings (the auth package doesn't export every shape as
// a sentinel because the handler is the only consumer that needs
// status-code-level discrimination). Substring matching is OK because
// the auth package owns these strings and tests pin them.
// client-facing status + message. The auth package exports typed
// sentinel errors for each user-correctable condition; we match via
// errors.Is so a renamed error message in the service never silently
// drifts the HTTP status. See issue #512.
func mapMFAServiceError(err error) error {
if err == nil {
return nil
}
msg := err.Error()
switch {
case strings.Contains(msg, "invalid password"),
strings.Contains(msg, "invalid MFA code"),
strings.Contains(msg, "MFA code or recovery code required"),
strings.Contains(msg, "no MFA enrollment in progress"),
strings.Contains(msg, "MFA enrollment expired"),
strings.Contains(msg, "MFA is not enabled"):
return NewClientError(400, msg)
case strings.Contains(msg, "authentication failed"):
return NewClientError(401, msg)
case errors.Is(err, auth.ErrMFAInvalidPassword),
errors.Is(err, auth.ErrMFAInvalidCode),
errors.Is(err, auth.ErrMFACodeRequired),
errors.Is(err, auth.ErrMFANoEnrollmentInProgress),
errors.Is(err, auth.ErrMFAEnrollmentExpired),
errors.Is(err, auth.ErrMFANotEnabled):
return NewClientError(400, err.Error())
case errors.Is(err, auth.ErrMFAAuthFailed):
// Opaque 401 to prevent user enumeration: the service returns this
// for both "user not found" and "DB lookup failed" paths so callers
// cannot distinguish the two.
return NewClientError(401, err.Error())
}
return err
}
Expand Down
80 changes: 79 additions & 1 deletion internal/api/handler_auth_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import (
"context"
"encoding/base64"
"errors"
"fmt"
"testing"

"github.com/LeanerCloud/CUDly/internal/auth"
Expand Down Expand Up @@ -1139,6 +1140,7 @@ func TestHandler_login_MFARequired_ReturnsMFASentinel(t *testing.T) {
require.True(t, ok)
assert.Equal(t, 401, ce.code)
assert.Equal(t, "mfa_required", ce.Error())
mockAuth.AssertExpectations(t)
}

func TestHandler_login_InvalidMFACode_ReturnsCodedSentinel(t *testing.T) {
Expand All @@ -1156,6 +1158,7 @@ func TestHandler_login_InvalidMFACode_ReturnsCodedSentinel(t *testing.T) {
require.True(t, ok)
assert.Equal(t, 401, ce.code)
assert.Equal(t, "invalid_mfa_code", ce.Error())
mockAuth.AssertExpectations(t)
}

func TestHandler_mfaSetup_HappyPath(t *testing.T) {
Expand All @@ -1171,15 +1174,19 @@ func TestHandler_mfaSetup_HappyPath(t *testing.T) {
require.NoError(t, err)
assert.Equal(t, "SECRET123", resp.Secret)
assert.Contains(t, resp.ProvisioningURI, "otpauth://")
mockAuth.AssertExpectations(t)
}

func TestHandler_mfaSetup_WrongPassword(t *testing.T) {
ctx := context.Background()
mockAuth := new(MockAuthService)
session := &Session{UserID: "user-1", Email: "u@x.com"}
mockAuth.On("ValidateSession", ctx, "tok").Return(session, nil)
// The real service returns a wrapped sentinel; the handler maps it to a
// 400 ClientError via errors.Is (issue #512). The mock must return the
// same sentinel so the errors.Is check in mapMFAServiceError fires.
mockAuth.On("MFASetupAPI", ctx, "user-1", "wrong").
Return("", "", errors.New("invalid password"))
Return("", "", fmt.Errorf("%w", auth.ErrMFAInvalidPassword))

handler := &Handler{auth: mockAuth}
_, err := handler.mfaSetup(ctx, authedReq("tok", `{"password":"`+b64("wrong")+`"}`))
Expand All @@ -1188,6 +1195,7 @@ func TestHandler_mfaSetup_WrongPassword(t *testing.T) {
require.True(t, ok)
assert.Equal(t, 400, ce.code)
assert.Contains(t, ce.Error(), "invalid password")
mockAuth.AssertExpectations(t)
}

func TestHandler_mfaEnable_HappyPath(t *testing.T) {
Expand All @@ -1202,6 +1210,7 @@ func TestHandler_mfaEnable_HappyPath(t *testing.T) {
resp, err := handler.mfaEnable(ctx, authedReq("tok", `{"code":"123456"}`))
require.NoError(t, err)
assert.Len(t, resp.RecoveryCodes, 2)
mockAuth.AssertExpectations(t)
}

func TestHandler_mfaEnable_NoSession(t *testing.T) {
Expand All @@ -1215,6 +1224,7 @@ func TestHandler_mfaEnable_NoSession(t *testing.T) {
ce, ok := IsClientError(err)
require.True(t, ok)
assert.Equal(t, 401, ce.code)
mockAuth.AssertExpectations(t)
}

func TestHandler_mfaDisable_HappyPath(t *testing.T) {
Expand All @@ -1227,6 +1237,7 @@ func TestHandler_mfaDisable_HappyPath(t *testing.T) {
handler := &Handler{auth: mockAuth}
_, err := handler.mfaDisable(ctx, authedReq("tok", `{"password":"`+b64("pw")+`","code":"123456"}`))
require.NoError(t, err)
mockAuth.AssertExpectations(t)
}

func TestHandler_mfaRegenerateRecoveryCodes_HappyPath(t *testing.T) {
Expand All @@ -1241,6 +1252,7 @@ func TestHandler_mfaRegenerateRecoveryCodes_HappyPath(t *testing.T) {
resp, err := handler.mfaRegenerateRecoveryCodes(ctx, authedReq("tok", `{"code":"123456"}`))
require.NoError(t, err)
assert.Len(t, resp.RecoveryCodes, 1)
mockAuth.AssertExpectations(t)
}

// ErrMFARequired_test / ErrInvalidMFACode_test return the sentinel
Expand Down Expand Up @@ -1483,3 +1495,69 @@ func TestHandler_login_ErrorEquivalence(t *testing.T) {
assert.Equal(t, ceNotFound.Error(), ceWrongPass.Error(),
"error body must be identical for unknown-user and wrong-password paths to prevent enumeration")
}

// ---------------------------------------------------------------
// mapMFAServiceError sentinel-to-HTTP-code mapping tests (issue #512).
//
// Each test verifies that a specific auth sentinel maps to the
// expected HTTP status code in mapMFAServiceError. These tests would
// fail if someone renames a sentinel value in the auth package
// without updating the switch in mapMFAServiceError.
// ---------------------------------------------------------------

func TestMapMFAServiceError_Nil(t *testing.T) {
assert.Nil(t, mapMFAServiceError(nil))
}

func TestMapMFAServiceError_NonSentinelPassesThrough(t *testing.T) {
plain := errors.New("database connection lost")
got := mapMFAServiceError(plain)
_, isClient := IsClientError(got)
assert.False(t, isClient, "non-sentinel errors must not be wrapped as ClientError")
assert.Equal(t, plain, got)
}

func testMFASentinel400(t *testing.T, sentinel error, name string) {
t.Helper()
wrapped := fmt.Errorf("some context: %w", sentinel)
got := mapMFAServiceError(wrapped)
ce, ok := IsClientError(got)
require.True(t, ok, "%s must map to a ClientError, got %T: %v", name, got, got)
assert.Equal(t, 400, ce.code, "%s must map to HTTP 400", name)
assert.Contains(t, ce.Error(), sentinel.Error())
}

func TestMapMFAServiceError_InvalidPassword_Is400(t *testing.T) {
testMFASentinel400(t, auth.ErrMFAInvalidPassword, "ErrMFAInvalidPassword")
}

func TestMapMFAServiceError_InvalidCode_Is400(t *testing.T) {
testMFASentinel400(t, auth.ErrMFAInvalidCode, "ErrMFAInvalidCode")
}

func TestMapMFAServiceError_CodeRequired_Is400(t *testing.T) {
testMFASentinel400(t, auth.ErrMFACodeRequired, "ErrMFACodeRequired")
}

func TestMapMFAServiceError_NoEnrollmentInProgress_Is400(t *testing.T) {
testMFASentinel400(t, auth.ErrMFANoEnrollmentInProgress, "ErrMFANoEnrollmentInProgress")
}

func TestMapMFAServiceError_EnrollmentExpired_Is400(t *testing.T) {
testMFASentinel400(t, auth.ErrMFAEnrollmentExpired, "ErrMFAEnrollmentExpired")
}

func TestMapMFAServiceError_NotEnabled_Is400(t *testing.T) {
testMFASentinel400(t, auth.ErrMFANotEnabled, "ErrMFANotEnabled")
}

func TestMapMFAServiceError_AuthFailed_Is401(t *testing.T) {
// ErrMFAAuthFailed must map to 401 (not 400) to prevent user enumeration:
// both "user not found" and "DB error" paths surface as opaque 401.
wrapped := fmt.Errorf("some context: %w", auth.ErrMFAAuthFailed)
got := mapMFAServiceError(wrapped)
ce, ok := IsClientError(got)
require.True(t, ok, "ErrMFAAuthFailed must map to a ClientError")
assert.Equal(t, 401, ce.code, "ErrMFAAuthFailed must map to HTTP 401")
assert.Contains(t, ce.Error(), auth.ErrMFAAuthFailed.Error())
}
25 changes: 25 additions & 0 deletions internal/auth/errors.go
Original file line number Diff line number Diff line change
Expand Up @@ -46,4 +46,29 @@ var (
// without substring-matching the human message. See issue #497.
ErrMFARequired = errors.New("mfa_required")
ErrInvalidMFACode = errors.New("invalid_mfa_code")

// MFA service-operation sentinels — returned (wrapped via fmt.Errorf
// "%w") by the MFA lifecycle methods in service_mfa.go so the API
// handler can map each error class to the right HTTP status code via
// errors.Is rather than brittle substring matching. See issue #512.
//
// ErrMFAInvalidPassword — wrong current password on setup/disable.
// ErrMFAInvalidCode — wrong TOTP or recovery code.
// ErrMFACodeRequired — MFA-enabled user supplied no code on disable.
// ErrMFANoEnrollmentInProgress — MFAEnable called before MFASetup.
// ErrMFAEnrollmentExpired — pending enrollment window elapsed.
// ErrMFANotEnabled — regenerate/disable called when MFA is off.
// ErrMFAAuthFailed — generic opaque auth failure (user not found or
// DB error; maps to 401 to prevent user enumeration).
//
// Message strings are intentionally identical to the pre-sentinel
// fmt.Errorf literals so that existing tests relying on err.Error()
// substrings continue to pass unchanged. See issue #512.
ErrMFAInvalidPassword = errors.New("invalid password")
ErrMFAInvalidCode = errors.New("invalid MFA code")
ErrMFACodeRequired = errors.New("MFA code or recovery code required")
ErrMFANoEnrollmentInProgress = errors.New("no MFA enrollment in progress")
ErrMFAEnrollmentExpired = errors.New("MFA enrollment expired")
ErrMFANotEnabled = errors.New("MFA is not enabled")
ErrMFAAuthFailed = errors.New("authentication failed")
)
28 changes: 14 additions & 14 deletions internal/auth/service_mfa.go
Original file line number Diff line number Diff line change
Expand Up @@ -262,13 +262,13 @@ func (s *Service) MFASetup(ctx context.Context, userID, password string) (*MFASe
}
user, err := s.store.GetUserByID(ctx, userID)
if err != nil {
return nil, fmt.Errorf("authentication failed")
return nil, fmt.Errorf("%w", ErrMFAAuthFailed)
}
if user == nil {
return nil, fmt.Errorf("authentication failed")
return nil, fmt.Errorf("%w", ErrMFAAuthFailed)
}
if !s.verifyPassword(password, user.PasswordHash) {
return nil, fmt.Errorf("invalid password")
return nil, fmt.Errorf("%w", ErrMFAInvalidPassword)
}

secret, err := generateMFASecret()
Expand Down Expand Up @@ -316,16 +316,16 @@ func (s *Service) generateAndHashRecoveryCodes() (plaintext, hashes []string, er
// "no enrollment in progress" rather than "expired" forever.
func (s *Service) validatePendingMFAEnrollment(ctx context.Context, user *User, code string) error {
if user.MFAPendingSecret == "" || user.MFAPendingSecretExpiresAt == nil {
return fmt.Errorf("no MFA enrollment in progress")
return fmt.Errorf("%w", ErrMFANoEnrollmentInProgress)
}
if time.Now().After(*user.MFAPendingSecretExpiresAt) {
user.MFAPendingSecret = ""
user.MFAPendingSecretExpiresAt = nil
_ = s.store.UpdateUser(ctx, user)
return fmt.Errorf("MFA enrollment expired")
return fmt.Errorf("%w", ErrMFAEnrollmentExpired)
}
if !verifyTOTP(user.MFAPendingSecret, code) {
return fmt.Errorf("invalid MFA code")
return fmt.Errorf("%w", ErrMFAInvalidCode)
}
return nil
}
Expand All @@ -350,7 +350,7 @@ func (s *Service) MFAEnable(ctx context.Context, userID, code string) ([]string,
}
user, err := s.store.GetUserByID(ctx, userID)
if err != nil || user == nil {
return nil, fmt.Errorf("authentication failed")
return nil, fmt.Errorf("%w", ErrMFAAuthFailed)
}
if err := s.validatePendingMFAEnrollment(ctx, user, code); err != nil {
return nil, err
Expand Down Expand Up @@ -410,24 +410,24 @@ func (s *Service) MFADisable(ctx context.Context, userID, password, codeOrRecove
}
user, err := s.store.GetUserByID(ctx, userID)
if err != nil || user == nil {
return fmt.Errorf("authentication failed")
return fmt.Errorf("%w", ErrMFAAuthFailed)
}
if !s.verifyPassword(password, user.PasswordHash) {
return fmt.Errorf("invalid password")
return fmt.Errorf("%w", ErrMFAInvalidPassword)
}
if !user.MFAEnabled {
return s.disableMFAAlreadyOff(ctx, user)
}
if codeOrRecovery == "" {
return fmt.Errorf("MFA code or recovery code required")
return fmt.Errorf("%w", ErrMFACodeRequired)
}

// Try TOTP first (cheap), then fall back to recovery code (bcrypt
// compare, ~constant-time per slot). Either path counts as a
// fresh proof-of-possession.
matched := verifyTOTP(user.MFASecret, codeOrRecovery) || s.consumeRecoveryCode(user, codeOrRecovery)
if !matched {
return fmt.Errorf("invalid MFA code")
return fmt.Errorf("%w", ErrMFAInvalidCode)
}

clearMFAFromUser(user)
Expand All @@ -448,13 +448,13 @@ func (s *Service) MFARegenerateRecoveryCodes(ctx context.Context, userID, code s
}
user, err := s.store.GetUserByID(ctx, userID)
if err != nil || user == nil {
return nil, fmt.Errorf("authentication failed")
return nil, fmt.Errorf("%w", ErrMFAAuthFailed)
}
if !user.MFAEnabled || user.MFASecret == "" {
return nil, fmt.Errorf("MFA is not enabled")
return nil, fmt.Errorf("%w", ErrMFANotEnabled)
}
if !verifyTOTP(user.MFASecret, code) {
return nil, fmt.Errorf("invalid MFA code")
return nil, fmt.Errorf("%w", ErrMFAInvalidCode)
}

plaintext, hashes, err := s.generateAndHashRecoveryCodes()
Expand Down
Loading
Loading