diff --git a/database/migrations/20260812_01_oauth_callback_status.sql b/database/migrations/20260812_01_oauth_callback_status.sql new file mode 100644 index 0000000..9a2cd2f --- /dev/null +++ b/database/migrations/20260812_01_oauth_callback_status.sql @@ -0,0 +1,18 @@ +ALTER TABLE oauth_authorization_request + ADD COLUMN status VARCHAR(20) NULL AFTER used_at, + ADD COLUMN failure_stage VARCHAR(30) NULL AFTER status, + ADD COLUMN failure_type VARCHAR(30) NULL AFTER failure_stage, + ADD COLUMN processing_started_at DATETIME(6) NULL AFTER failure_type, + ADD COLUMN completed_at DATETIME(6) NULL AFTER processing_started_at; + +UPDATE oauth_authorization_request +SET status = CASE WHEN used_at IS NULL THEN 'ISSUED' ELSE 'LEGACY' END, + processing_started_at = used_at, + completed_at = used_at +WHERE status IS NULL; + +ALTER TABLE oauth_authorization_request + MODIFY COLUMN status VARCHAR(20) NOT NULL; + +CREATE INDEX idx_oauth_authorization_request_status_processing + ON oauth_authorization_request (status, processing_started_at); diff --git a/src/main/java/com/example/todayEng/domain/user/entity/OAuthAuthorizationRequest.java b/src/main/java/com/example/todayEng/domain/user/entity/OAuthAuthorizationRequest.java index a0b2b0c..1658e8f 100644 --- a/src/main/java/com/example/todayEng/domain/user/entity/OAuthAuthorizationRequest.java +++ b/src/main/java/com/example/todayEng/domain/user/entity/OAuthAuthorizationRequest.java @@ -1,6 +1,9 @@ package com.example.todayEng.domain.user.entity; import com.example.todayEng.domain.user.entity.enums.ExternalServiceProvider; +import com.example.todayEng.domain.user.entity.enums.OAuthAuthorizationRequestStatus; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureStage; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureType; import com.example.todayEng.global.common.BaseTimeEntity; import jakarta.persistence.*; import lombok.AccessLevel; @@ -23,6 +26,10 @@ @Index( name = "idx_oauth_authorization_request_expires_at", columnList = "expires_at" + ), + @Index( + name = "idx_oauth_authorization_request_status_processing", + columnList = "status, processing_started_at" ) } ) @@ -62,6 +69,24 @@ public class OAuthAuthorizationRequest extends BaseTimeEntity { @Column(name = "used_at") private LocalDateTime usedAt; + @Enumerated(EnumType.STRING) + @Column(name = "status", nullable = false, length = 20) + private OAuthAuthorizationRequestStatus status; + + @Enumerated(EnumType.STRING) + @Column(name = "failure_stage", length = 30) + private OAuthCallbackFailureStage failureStage; + + @Enumerated(EnumType.STRING) + @Column(name = "failure_type", length = 30) + private OAuthCallbackFailureType failureType; + + @Column(name = "processing_started_at") + private LocalDateTime processingStartedAt; + + @Column(name = "completed_at") + private LocalDateTime completedAt; + private OAuthAuthorizationRequest( User user, ExternalServiceProvider provider, @@ -72,6 +97,7 @@ private OAuthAuthorizationRequest( this.provider = provider; this.stateHash = stateHash; this.expiresAt = expiresAt; + this.status = OAuthAuthorizationRequestStatus.ISSUED; } public static OAuthAuthorizationRequest create( @@ -93,6 +119,14 @@ public boolean isExpired(LocalDateTime now) { } public boolean isUsed() { - return usedAt != null; + return status != OAuthAuthorizationRequestStatus.ISSUED; + } + + public void succeed(LocalDateTime completedAt) { + if (status != OAuthAuthorizationRequestStatus.PROCESSING) { + throw new IllegalStateException("OAuth authorization request is not processing"); + } + this.status = OAuthAuthorizationRequestStatus.SUCCEEDED; + this.completedAt = completedAt; } -} \ No newline at end of file +} diff --git a/src/main/java/com/example/todayEng/domain/user/entity/enums/OAuthAuthorizationRequestStatus.java b/src/main/java/com/example/todayEng/domain/user/entity/enums/OAuthAuthorizationRequestStatus.java new file mode 100644 index 0000000..33273b6 --- /dev/null +++ b/src/main/java/com/example/todayEng/domain/user/entity/enums/OAuthAuthorizationRequestStatus.java @@ -0,0 +1,9 @@ +package com.example.todayEng.domain.user.entity.enums; + +public enum OAuthAuthorizationRequestStatus { + LEGACY, + ISSUED, + PROCESSING, + SUCCEEDED, + FAILED +} diff --git a/src/main/java/com/example/todayEng/domain/user/entity/enums/OAuthCallbackFailureStage.java b/src/main/java/com/example/todayEng/domain/user/entity/enums/OAuthCallbackFailureStage.java new file mode 100644 index 0000000..00f4f24 --- /dev/null +++ b/src/main/java/com/example/todayEng/domain/user/entity/enums/OAuthCallbackFailureStage.java @@ -0,0 +1,9 @@ +package com.example.todayEng.domain.user.entity.enums; + +public enum OAuthCallbackFailureStage { + CALLBACK_VALIDATION, + TOKEN_EXCHANGE, + USER_INFO, + ACCOUNT_SAVE, + STALE_PROCESSING +} diff --git a/src/main/java/com/example/todayEng/domain/user/entity/enums/OAuthCallbackFailureType.java b/src/main/java/com/example/todayEng/domain/user/entity/enums/OAuthCallbackFailureType.java new file mode 100644 index 0000000..5448b3e --- /dev/null +++ b/src/main/java/com/example/todayEng/domain/user/entity/enums/OAuthCallbackFailureType.java @@ -0,0 +1,10 @@ +package com.example.todayEng.domain.user.entity.enums; + +public enum OAuthCallbackFailureType { + AUTHORIZATION_DENIED, + AUTHORIZATION_CODE_MISSING, + EXTERNAL_API, + DATABASE, + CONFLICT, + INTERNAL +} diff --git a/src/main/java/com/example/todayEng/domain/user/repository/OAuthAuthorizationRequestRepository.java b/src/main/java/com/example/todayEng/domain/user/repository/OAuthAuthorizationRequestRepository.java index fd0b76b..95138d9 100644 --- a/src/main/java/com/example/todayEng/domain/user/repository/OAuthAuthorizationRequestRepository.java +++ b/src/main/java/com/example/todayEng/domain/user/repository/OAuthAuthorizationRequestRepository.java @@ -2,6 +2,9 @@ import com.example.todayEng.domain.user.entity.OAuthAuthorizationRequest; import com.example.todayEng.domain.user.entity.enums.ExternalServiceProvider; +import com.example.todayEng.domain.user.entity.enums.OAuthAuthorizationRequestStatus; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureStage; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureType; import jakarta.persistence.LockModeType; import org.springframework.data.jpa.repository.JpaRepository; import org.springframework.data.jpa.repository.Lock; @@ -27,16 +30,54 @@ Optional findByStateHashAndProvider( ) @Query(""" UPDATE OAuthAuthorizationRequest request - SET request.usedAt = :usedAt + SET request.usedAt = :usedAt, request.processingStartedAt = :usedAt, + request.status = :processing WHERE request.id = :requestId - AND request.usedAt IS NULL + AND request.status = :issued AND request.expiresAt > :usedAt """) - int consumeIfAvailable( + int startProcessingIfIssued( @Param("requestId") Long requestId, - @Param("usedAt") LocalDateTime usedAt + @Param("usedAt") LocalDateTime usedAt, + @Param("issued") OAuthAuthorizationRequestStatus issued, + @Param("processing") OAuthAuthorizationRequestStatus processing ); + @Lock(LockModeType.PESSIMISTIC_WRITE) + @Query("SELECT request FROM OAuthAuthorizationRequest request WHERE request.id = :requestId") + Optional findByIdForUpdate( + @Param("requestId") Long requestId); + + @Modifying(clearAutomatically = true, flushAutomatically = true) + @Query(""" + UPDATE OAuthAuthorizationRequest request + SET request.status = :failed, request.failureStage = :stage, + request.failureType = :failureType, request.completedAt = :completedAt + WHERE request.id = :requestId AND request.status = :processing + """) + int markFailed(@Param("requestId") Long requestId, + @Param("processing") OAuthAuthorizationRequestStatus processing, + @Param("failed") OAuthAuthorizationRequestStatus failed, + @Param("stage") OAuthCallbackFailureStage stage, + @Param("failureType") OAuthCallbackFailureType failureType, + @Param("completedAt") LocalDateTime completedAt); + + @Modifying(clearAutomatically = true, flushAutomatically = true) + @Query(""" + UPDATE OAuthAuthorizationRequest request + SET request.status = :failed, request.failureStage = :stage, + request.failureType = :failureType, request.completedAt = :now + WHERE request.status = :processing + AND request.processingStartedAt <= :staleBefore + """) + int failStaleProcessing( + @Param("processing") OAuthAuthorizationRequestStatus processing, + @Param("failed") OAuthAuthorizationRequestStatus failed, + @Param("stage") OAuthCallbackFailureStage stage, + @Param("failureType") OAuthCallbackFailureType failureType, + @Param("staleBefore") LocalDateTime staleBefore, + @Param("now") LocalDateTime now); + @Modifying( clearAutomatically = true, flushAutomatically = true diff --git a/src/main/java/com/example/todayEng/domain/user/service/ExternalAccountOAuthService.java b/src/main/java/com/example/todayEng/domain/user/service/ExternalAccountOAuthService.java index 102439c..dd798d2 100644 --- a/src/main/java/com/example/todayEng/domain/user/service/ExternalAccountOAuthService.java +++ b/src/main/java/com/example/todayEng/domain/user/service/ExternalAccountOAuthService.java @@ -6,10 +6,14 @@ import com.example.todayEng.domain.user.dto.oauth.OAuthTokenResponse; import com.example.todayEng.domain.user.dto.response.OAuthAuthorizationResponse; import com.example.todayEng.domain.user.entity.enums.ExternalServiceProvider; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureStage; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureType; +import com.example.todayEng.domain.user.service.OAuthAuthorizationRequestService.ProcessingClaim; import com.example.todayEng.global.error.ErrorCode; import com.example.todayEng.global.error.exception.BaseException; import lombok.RequiredArgsConstructor; import org.springframework.stereotype.Service; +import org.springframework.dao.DataAccessException; @Service @RequiredArgsConstructor @@ -21,8 +25,7 @@ public class ExternalAccountOAuthService { private final OAuthProviderClientRegistry oauthProviderClientRegistry; - private final ExternalAccountConnectionService - externalAccountConnectionService; + private final OAuthCallbackCompletionService oauthCallbackCompletionService; public OAuthAuthorizationResponse createAuthorizationUrl( Long userId, @@ -50,35 +53,53 @@ public void connectExternalAccount( String state, String error ) { - Long userId = + ProcessingClaim claim = oauthAuthorizationRequestService - .validateAndConsume( + .startProcessing( state, provider ); + OAuthCallbackFailureStage stage = OAuthCallbackFailureStage.CALLBACK_VALIDATION; + try { + validateOAuthCallback(code, error); + OAuthProviderClient providerClient = oauthProviderClientRegistry.getClient(provider); + + stage = OAuthCallbackFailureStage.TOKEN_EXCHANGE; + OAuthTokenResponse tokenResponse = providerClient.exchangeToken(code); + + stage = OAuthCallbackFailureStage.USER_INFO; + ExternalUserInfo externalUserInfo = + providerClient.fetchUserInfo(tokenResponse.accessToken()); + + stage = OAuthCallbackFailureStage.ACCOUNT_SAVE; + oauthCallbackCompletionService.saveAccountAndSucceed( + claim.requestId(), claim.userId(), provider, + tokenResponse, externalUserInfo); + } catch (RuntimeException exception) { + oauthAuthorizationRequestService.fail( + claim.requestId(), stage, classifyFailure(exception)); + throw exception; + } + } - validateOAuthCallback( - code, - error - ); - - OAuthProviderClient providerClient = - oauthProviderClientRegistry.getClient(provider); - - OAuthTokenResponse tokenResponse = - providerClient.exchangeToken(code); - - ExternalUserInfo externalUserInfo = - providerClient.fetchUserInfo( - tokenResponse.accessToken() - ); - - externalAccountConnectionService.saveOrUpdate( - userId, - provider, - tokenResponse, - externalUserInfo - ); + private OAuthCallbackFailureType classifyFailure(RuntimeException exception) { + if (exception instanceof BaseException baseException) { + return switch (baseException.getErrorCode()) { + case OAUTH_AUTHORIZATION_DENIED -> + OAuthCallbackFailureType.AUTHORIZATION_DENIED; + case OAUTH_AUTHORIZATION_CODE_MISSING -> + OAuthCallbackFailureType.AUTHORIZATION_CODE_MISSING; + case EXTERNAL_ACCOUNT_ALREADY_LINKED -> + OAuthCallbackFailureType.CONFLICT; + case OAUTH_TOKEN_EXCHANGE_FAILED, OAUTH_USER_INFO_FAILED, EXTERNAL_API_ERROR -> + OAuthCallbackFailureType.EXTERNAL_API; + default -> OAuthCallbackFailureType.INTERNAL; + }; + } + if (exception instanceof DataAccessException) { + return OAuthCallbackFailureType.DATABASE; + } + return OAuthCallbackFailureType.INTERNAL; } private void validateOAuthCallback( diff --git a/src/main/java/com/example/todayEng/domain/user/service/OAuthAuthorizationRequestRecoveryScheduler.java b/src/main/java/com/example/todayEng/domain/user/service/OAuthAuthorizationRequestRecoveryScheduler.java new file mode 100644 index 0000000..91cded1 --- /dev/null +++ b/src/main/java/com/example/todayEng/domain/user/service/OAuthAuthorizationRequestRecoveryScheduler.java @@ -0,0 +1,40 @@ +package com.example.todayEng.domain.user.service; + +import com.example.todayEng.domain.user.entity.enums.OAuthAuthorizationRequestStatus; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureStage; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureType; +import com.example.todayEng.domain.user.repository.OAuthAuthorizationRequestRepository; +import java.time.Duration; +import java.time.LocalDateTime; +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; +import org.springframework.scheduling.annotation.Scheduled; +import org.springframework.stereotype.Component; +import org.springframework.transaction.annotation.Transactional; + +@Slf4j +@Component +@RequiredArgsConstructor +public class OAuthAuthorizationRequestRecoveryScheduler { + + private static final Duration PROCESSING_TIMEOUT = Duration.ofMinutes(15); + + private final OAuthAuthorizationRequestRepository repository; + + @Scheduled(fixedDelayString = "${oauth.callback.recovery-interval-millis:60000}") + @Transactional + public void failStaleProcessingRequests() { + LocalDateTime now = LocalDateTime.now(); + int updated = repository.failStaleProcessing( + OAuthAuthorizationRequestStatus.PROCESSING, + OAuthAuthorizationRequestStatus.FAILED, + OAuthCallbackFailureStage.STALE_PROCESSING, + OAuthCallbackFailureType.INTERNAL, + now.minus(PROCESSING_TIMEOUT), + now + ); + if (updated > 0) { + log.warn("Stale OAuth callback requests marked failed: count={}", updated); + } + } +} diff --git a/src/main/java/com/example/todayEng/domain/user/service/OAuthAuthorizationRequestService.java b/src/main/java/com/example/todayEng/domain/user/service/OAuthAuthorizationRequestService.java index e99e507..7f935b2 100644 --- a/src/main/java/com/example/todayEng/domain/user/service/OAuthAuthorizationRequestService.java +++ b/src/main/java/com/example/todayEng/domain/user/service/OAuthAuthorizationRequestService.java @@ -3,6 +3,9 @@ import com.example.todayEng.domain.user.entity.OAuthAuthorizationRequest; import com.example.todayEng.domain.user.entity.User; import com.example.todayEng.domain.user.entity.enums.ExternalServiceProvider; +import com.example.todayEng.domain.user.entity.enums.OAuthAuthorizationRequestStatus; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureStage; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureType; import com.example.todayEng.domain.user.repository.OAuthAuthorizationRequestRepository; import com.example.todayEng.domain.user.repository.UserRepository; import com.example.todayEng.global.error.ErrorCode; @@ -64,7 +67,7 @@ public String issue( } @Transactional(propagation = Propagation.REQUIRES_NEW) - public Long validateAndConsume( + public ProcessingClaim startProcessing( String rawState, ExternalServiceProvider provider ) { @@ -88,9 +91,11 @@ public Long validateAndConsume( int updatedCount = oauthAuthorizationRequestRepository - .consumeIfAvailable( + .startProcessingIfIssued( authorizationRequest.getId(), - now + now, + OAuthAuthorizationRequestStatus.ISSUED, + OAuthAuthorizationRequestStatus.PROCESSING ); if (updatedCount == 0) { @@ -101,7 +106,23 @@ public Long validateAndConsume( ); } - return authorizationRequest.getUser().getId(); + return new ProcessingClaim( + authorizationRequest.getId(), + authorizationRequest.getUser().getId() + ); + } + + @Transactional(propagation = Propagation.REQUIRES_NEW) + public void fail(Long requestId, OAuthCallbackFailureStage stage, + OAuthCallbackFailureType failureType) { + oauthAuthorizationRequestRepository.markFailed( + requestId, + OAuthAuthorizationRequestStatus.PROCESSING, + OAuthAuthorizationRequestStatus.FAILED, + stage, + failureType, + LocalDateTime.now() + ); } private void validateState( @@ -114,7 +135,7 @@ private void validateState( ); } - if (authorizationRequest.isUsed()) { + if (authorizationRequest.getStatus() != OAuthAuthorizationRequestStatus.ISSUED) { throw new BaseException( ErrorCode.OAUTH_STATE_ALREADY_USED ); @@ -142,7 +163,7 @@ private void throwStateConsumptionFailure( ); } - if (currentRequest.isUsed()) { + if (currentRequest.getStatus() != OAuthAuthorizationRequestStatus.ISSUED) { throw new BaseException( ErrorCode.OAUTH_STATE_ALREADY_USED ); @@ -153,6 +174,9 @@ private void throwStateConsumptionFailure( ); } + public record ProcessingClaim(Long requestId, Long userId) { + } + private String generateRawState() { byte[] randomBytes = new byte[STATE_BYTE_LENGTH]; secureRandom.nextBytes(randomBytes); @@ -186,4 +210,4 @@ private String hashState(String rawState) { ); } } -} \ No newline at end of file +} diff --git a/src/main/java/com/example/todayEng/domain/user/service/OAuthCallbackCompletionService.java b/src/main/java/com/example/todayEng/domain/user/service/OAuthCallbackCompletionService.java new file mode 100644 index 0000000..a3644a9 --- /dev/null +++ b/src/main/java/com/example/todayEng/domain/user/service/OAuthCallbackCompletionService.java @@ -0,0 +1,41 @@ +package com.example.todayEng.domain.user.service; + +import com.example.todayEng.domain.user.dto.oauth.ExternalUserInfo; +import com.example.todayEng.domain.user.dto.oauth.OAuthTokenResponse; +import com.example.todayEng.domain.user.entity.OAuthAuthorizationRequest; +import com.example.todayEng.domain.user.entity.enums.ExternalServiceProvider; +import com.example.todayEng.domain.user.repository.OAuthAuthorizationRequestRepository; +import com.example.todayEng.global.error.ErrorCode; +import com.example.todayEng.global.error.exception.BaseException; +import java.time.LocalDateTime; +import lombok.RequiredArgsConstructor; +import org.springframework.stereotype.Service; +import org.springframework.transaction.annotation.Transactional; + +@Service +@RequiredArgsConstructor +public class OAuthCallbackCompletionService { + + private final OAuthAuthorizationRequestRepository authorizationRequestRepository; + private final ExternalAccountConnectionService externalAccountConnectionService; + + @Transactional + public void saveAccountAndSucceed( + Long requestId, + Long userId, + ExternalServiceProvider provider, + OAuthTokenResponse tokenResponse, + ExternalUserInfo externalUserInfo + ) { + OAuthAuthorizationRequest request = authorizationRequestRepository + .findByIdForUpdate(requestId) + .orElseThrow(() -> new BaseException(ErrorCode.OAUTH_STATE_INVALID)); + try { + request.succeed(LocalDateTime.now()); + } catch (IllegalStateException exception) { + throw new BaseException(ErrorCode.OAUTH_STATE_CONSUME_FAILED); + } + externalAccountConnectionService.saveOrUpdate( + userId, provider, tokenResponse, externalUserInfo); + } +} diff --git a/src/test/java/com/example/todayEng/domain/user/repository/OAuthAuthorizationRequestRepositoryTest.java b/src/test/java/com/example/todayEng/domain/user/repository/OAuthAuthorizationRequestRepositoryTest.java new file mode 100644 index 0000000..83cdac9 --- /dev/null +++ b/src/test/java/com/example/todayEng/domain/user/repository/OAuthAuthorizationRequestRepositoryTest.java @@ -0,0 +1,205 @@ +package com.example.todayEng.domain.user.repository; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.example.todayEng.domain.user.entity.OAuthAuthorizationRequest; +import com.example.todayEng.domain.user.entity.User; +import com.example.todayEng.domain.user.entity.enums.ExternalServiceProvider; +import com.example.todayEng.domain.user.entity.enums.OAuthAuthorizationRequestStatus; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureStage; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureType; +import java.time.LocalDateTime; +import java.time.temporal.ChronoUnit; +import java.util.List; +import java.util.UUID; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.Future; +import java.util.concurrent.TimeUnit; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.autoconfigure.orm.jpa.DataJpaTest; +import org.springframework.transaction.PlatformTransactionManager; +import org.springframework.transaction.annotation.Propagation; +import org.springframework.transaction.annotation.Transactional; +import org.springframework.transaction.support.TransactionTemplate; + +@DataJpaTest +class OAuthAuthorizationRequestRepositoryTest { + + @Autowired OAuthAuthorizationRequestRepository repository; + @Autowired UserRepository userRepository; + @Autowired PlatformTransactionManager transactionManager; + + @Test + void onlyOneCallbackCanAtomicallyStartProcessing() { + OAuthAuthorizationRequest request = issueRequest(); + LocalDateTime now = LocalDateTime.now().truncatedTo(ChronoUnit.MICROS); + + int first = repository.startProcessingIfIssued( + request.getId(), now, + OAuthAuthorizationRequestStatus.ISSUED, + OAuthAuthorizationRequestStatus.PROCESSING); + int duplicate = repository.startProcessingIfIssued( + request.getId(), now, + OAuthAuthorizationRequestStatus.ISSUED, + OAuthAuthorizationRequestStatus.PROCESSING); + + assertThat(first).isEqualTo(1); + assertThat(duplicate).isZero(); + assertThat(repository.findById(request.getId()).orElseThrow().getStatus()) + .isEqualTo(OAuthAuthorizationRequestStatus.PROCESSING); + } + + @Test + void recordsFailureClassificationWithoutCallbackSecrets() { + OAuthAuthorizationRequest request = issueRequest(); + LocalDateTime now = LocalDateTime.now(); + repository.startProcessingIfIssued( + request.getId(), now, + OAuthAuthorizationRequestStatus.ISSUED, + OAuthAuthorizationRequestStatus.PROCESSING); + + int updated = repository.markFailed( + request.getId(), + OAuthAuthorizationRequestStatus.PROCESSING, + OAuthAuthorizationRequestStatus.FAILED, + OAuthCallbackFailureStage.TOKEN_EXCHANGE, + OAuthCallbackFailureType.EXTERNAL_API, + now.plusSeconds(1)); + + OAuthAuthorizationRequest failed = repository.findById(request.getId()).orElseThrow(); + assertThat(updated).isEqualTo(1); + assertThat(failed.getStatus()).isEqualTo(OAuthAuthorizationRequestStatus.FAILED); + assertThat(failed.getFailureStage()).isEqualTo(OAuthCallbackFailureStage.TOKEN_EXCHANGE); + assertThat(failed.getFailureType()).isEqualTo(OAuthCallbackFailureType.EXTERNAL_API); + } + + @Test + void staleProcessingRequestIsFailedInsteadOfRetried() { + OAuthAuthorizationRequest request = issueRequest(); + LocalDateTime now = LocalDateTime.now().truncatedTo(ChronoUnit.MICROS); + repository.startProcessingIfIssued( + request.getId(), now.minusMinutes(15), + OAuthAuthorizationRequestStatus.ISSUED, + OAuthAuthorizationRequestStatus.PROCESSING); + + int updated = repository.failStaleProcessing( + OAuthAuthorizationRequestStatus.PROCESSING, + OAuthAuthorizationRequestStatus.FAILED, + OAuthCallbackFailureStage.STALE_PROCESSING, + OAuthCallbackFailureType.INTERNAL, + now.minusMinutes(15), + now); + + OAuthAuthorizationRequest failed = repository.findById(request.getId()).orElseThrow(); + assertThat(updated).isEqualTo(1); + assertThat(failed.getStatus()).isEqualTo(OAuthAuthorizationRequestStatus.FAILED); + assertThat(failed.getFailureStage()) + .isEqualTo(OAuthCallbackFailureStage.STALE_PROCESSING); + } + + @Test + @Transactional(propagation = Propagation.NOT_SUPPORTED) + void concurrentCallbacksAllowExactlyOneProcessingClaim() throws Exception { + TransactionTemplate transaction = new TransactionTemplate(transactionManager); + Long requestId = transaction.execute(status -> issueRequest().getId()); + CountDownLatch ready = new CountDownLatch(2); + CountDownLatch start = new CountDownLatch(1); + ExecutorService executor = Executors.newFixedThreadPool(2); + try { + List> results = List.of( + executor.submit(() -> claimConcurrently(requestId, ready, start)), + executor.submit(() -> claimConcurrently(requestId, ready, start))); + ready.await(); + start.countDown(); + + int successfulClaims = 0; + for (Future result : results) { + successfulClaims += result.get(); + } + + assertThat(successfulClaims).isEqualTo(1); + OAuthAuthorizationRequestStatus finalStatus = transaction.execute(status -> + repository.findById(requestId).orElseThrow().getStatus()); + assertThat(finalStatus).isEqualTo(OAuthAuthorizationRequestStatus.PROCESSING); + } finally { + executor.shutdownNow(); + } + } + + @Test + @Transactional(propagation = Propagation.NOT_SUPPORTED) + void completionLockPreventsSchedulerFromChangingSucceededRequestToFailed() throws Exception { + TransactionTemplate transaction = new TransactionTemplate(transactionManager); + Long requestId = transaction.execute(status -> { + OAuthAuthorizationRequest request = issueRequest(); + repository.startProcessingIfIssued( + request.getId(), LocalDateTime.now().minusMinutes(15), + OAuthAuthorizationRequestStatus.ISSUED, + OAuthAuthorizationRequestStatus.PROCESSING); + return request.getId(); + }); + CountDownLatch locked = new CountDownLatch(1); + CountDownLatch allowCommit = new CountDownLatch(1); + ExecutorService executor = Executors.newFixedThreadPool(2); + try { + Future completion = executor.submit(() -> { + transaction.executeWithoutResult(status -> { + OAuthAuthorizationRequest request = + repository.findByIdForUpdate(requestId).orElseThrow(); + request.succeed(LocalDateTime.now()); + locked.countDown(); + try { + allowCommit.await(); + } catch (InterruptedException exception) { + Thread.currentThread().interrupt(); + throw new IllegalStateException(exception); + } + }); + return null; + }); + assertThat(locked.await(5, TimeUnit.SECONDS)).isTrue(); + + Future recovery = executor.submit(() -> + transaction.execute(status -> repository.failStaleProcessing( + OAuthAuthorizationRequestStatus.PROCESSING, + OAuthAuthorizationRequestStatus.FAILED, + OAuthCallbackFailureStage.STALE_PROCESSING, + OAuthCallbackFailureType.INTERNAL, + LocalDateTime.now().minusMinutes(15), + LocalDateTime.now()))); + + allowCommit.countDown(); + completion.get(); + assertThat(recovery.get()).isZero(); + + OAuthAuthorizationRequestStatus finalStatus = transaction.execute(status -> + repository.findById(requestId).orElseThrow().getStatus()); + assertThat(finalStatus).isEqualTo(OAuthAuthorizationRequestStatus.SUCCEEDED); + } finally { + executor.shutdownNow(); + } + } + + private int claimConcurrently(Long requestId, CountDownLatch ready, CountDownLatch start) + throws Exception { + ready.countDown(); + start.await(); + return new TransactionTemplate(transactionManager).execute(status -> + repository.startProcessingIfIssued( + requestId, LocalDateTime.now(), + OAuthAuthorizationRequestStatus.ISSUED, + OAuthAuthorizationRequestStatus.PROCESSING)); + } + + private OAuthAuthorizationRequest issueRequest() { + User user = userRepository.save(User.create()); + return repository.saveAndFlush(OAuthAuthorizationRequest.create( + user, + ExternalServiceProvider.SPOTIFY, + UUID.randomUUID().toString().replace("-", "").repeat(2), + LocalDateTime.now().plusMinutes(10))); + } +} diff --git a/src/test/java/com/example/todayEng/domain/user/service/ExternalAccountOAuthServiceTest.java b/src/test/java/com/example/todayEng/domain/user/service/ExternalAccountOAuthServiceTest.java index d9bd9e6..da8bbe9 100644 --- a/src/test/java/com/example/todayEng/domain/user/service/ExternalAccountOAuthServiceTest.java +++ b/src/test/java/com/example/todayEng/domain/user/service/ExternalAccountOAuthServiceTest.java @@ -4,8 +4,15 @@ import static org.mockito.BDDMockito.given; import static org.mockito.Mockito.never; import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.doThrow; import com.example.todayEng.domain.user.client.OAuthProviderClientRegistry; +import com.example.todayEng.domain.user.client.OAuthProviderClient; +import com.example.todayEng.domain.user.dto.oauth.ExternalUserInfo; +import com.example.todayEng.domain.user.dto.oauth.OAuthTokenResponse; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureStage; +import com.example.todayEng.domain.user.entity.enums.OAuthCallbackFailureType; +import com.example.todayEng.domain.user.service.OAuthAuthorizationRequestService.ProcessingClaim; import com.example.todayEng.domain.user.entity.enums.ExternalServiceProvider; import com.example.todayEng.global.error.ErrorCode; import com.example.todayEng.global.error.exception.BaseException; @@ -14,6 +21,7 @@ import org.mockito.InjectMocks; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.dao.DataIntegrityViolationException; @ExtendWith(MockitoExtension.class) class ExternalAccountOAuthServiceTest { @@ -27,8 +35,10 @@ class ExternalAccountOAuthServiceTest { oauthProviderClientRegistry; @Mock - private ExternalAccountConnectionService - externalAccountConnectionService; + private OAuthCallbackCompletionService oauthCallbackCompletionService; + + @Mock + private OAuthProviderClient providerClient; @InjectMocks private ExternalAccountOAuthService service; @@ -54,4 +64,86 @@ void createAuthorizationUrl_unsupportedProvider_doesNotIssueState() { verify(oauthAuthorizationRequestService, never()) .issue(1L, provider); } + + @Test + void connectExternalAccount_marksRequestSucceededAfterAccountSave() { + ExternalServiceProvider provider = ExternalServiceProvider.SPOTIFY; + OAuthTokenResponse token = new OAuthTokenResponse("access", "refresh", 3600L); + ExternalUserInfo userInfo = new ExternalUserInfo("provider-id", "user@example.com"); + given(oauthAuthorizationRequestService.startProcessing("state", provider)) + .willReturn(new ProcessingClaim(10L, 1L)); + given(oauthProviderClientRegistry.getClient(provider)).willReturn(providerClient); + given(providerClient.exchangeToken("code")).willReturn(token); + given(providerClient.fetchUserInfo("access")).willReturn(userInfo); + + service.connectExternalAccount(provider, "code", "state", null); + + verify(oauthCallbackCompletionService) + .saveAccountAndSucceed(10L, 1L, provider, token, userInfo); + } + + @Test + void connectExternalAccount_recordsTokenExchangeFailure() { + ExternalServiceProvider provider = ExternalServiceProvider.SPOTIFY; + given(oauthAuthorizationRequestService.startProcessing("state", provider)) + .willReturn(new ProcessingClaim(10L, 1L)); + given(oauthProviderClientRegistry.getClient(provider)).willReturn(providerClient); + given(providerClient.exchangeToken("code")) + .willThrow(new BaseException(ErrorCode.OAUTH_TOKEN_EXCHANGE_FAILED)); + + assertThatThrownBy(() -> + service.connectExternalAccount(provider, "code", "state", null)) + .isInstanceOf(BaseException.class); + + verify(oauthAuthorizationRequestService).fail( + 10L, OAuthCallbackFailureStage.TOKEN_EXCHANGE, + OAuthCallbackFailureType.EXTERNAL_API); + verify(oauthCallbackCompletionService, never()) + .saveAccountAndSucceed( + org.mockito.ArgumentMatchers.anyLong(), + org.mockito.ArgumentMatchers.anyLong(), + org.mockito.ArgumentMatchers.any(), + org.mockito.ArgumentMatchers.any(), + org.mockito.ArgumentMatchers.any()); + } + + @Test + void connectExternalAccount_recordsDatabaseFailureWithoutSensitiveValues() { + ExternalServiceProvider provider = ExternalServiceProvider.SPOTIFY; + OAuthTokenResponse token = + new OAuthTokenResponse("sensitive-access", "sensitive-refresh", 3600L); + ExternalUserInfo userInfo = new ExternalUserInfo("provider-id", "user@example.com"); + given(oauthAuthorizationRequestService.startProcessing("state", provider)) + .willReturn(new ProcessingClaim(10L, 1L)); + given(oauthProviderClientRegistry.getClient(provider)).willReturn(providerClient); + given(providerClient.exchangeToken("sensitive-code")).willReturn(token); + given(providerClient.fetchUserInfo("sensitive-access")).willReturn(userInfo); + doThrow(new DataIntegrityViolationException("duplicate")) + .when(oauthCallbackCompletionService) + .saveAccountAndSucceed(10L, 1L, provider, token, userInfo); + + assertThatThrownBy(() -> service.connectExternalAccount( + provider, "sensitive-code", "state", null)) + .isInstanceOf(DataIntegrityViolationException.class); + + verify(oauthAuthorizationRequestService).fail( + 10L, OAuthCallbackFailureStage.ACCOUNT_SAVE, + OAuthCallbackFailureType.DATABASE); + } + + @Test + void deniedCallbackIsClaimedThenMarkedFailed() { + ExternalServiceProvider provider = ExternalServiceProvider.SPOTIFY; + given(oauthAuthorizationRequestService.startProcessing("state", provider)) + .willReturn(new ProcessingClaim(10L, 1L)); + + assertThatThrownBy(() -> + service.connectExternalAccount(provider, null, "state", "access_denied")) + .isInstanceOf(BaseException.class); + + verify(oauthAuthorizationRequestService).fail( + 10L, OAuthCallbackFailureStage.CALLBACK_VALIDATION, + OAuthCallbackFailureType.AUTHORIZATION_DENIED); + verify(oauthProviderClientRegistry, never()).getClient(provider); + } } diff --git a/src/test/java/com/example/todayEng/domain/user/service/OAuthCallbackCompletionServiceTest.java b/src/test/java/com/example/todayEng/domain/user/service/OAuthCallbackCompletionServiceTest.java new file mode 100644 index 0000000..2318b1c --- /dev/null +++ b/src/test/java/com/example/todayEng/domain/user/service/OAuthCallbackCompletionServiceTest.java @@ -0,0 +1,46 @@ +package com.example.todayEng.domain.user.service; + +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.BDDMockito.given; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; + +import com.example.todayEng.domain.user.dto.oauth.ExternalUserInfo; +import com.example.todayEng.domain.user.dto.oauth.OAuthTokenResponse; +import com.example.todayEng.domain.user.entity.OAuthAuthorizationRequest; +import com.example.todayEng.domain.user.entity.User; +import com.example.todayEng.domain.user.entity.enums.ExternalServiceProvider; +import com.example.todayEng.domain.user.repository.OAuthAuthorizationRequestRepository; +import com.example.todayEng.global.error.exception.BaseException; +import java.time.LocalDateTime; +import java.util.Optional; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +@ExtendWith(MockitoExtension.class) +class OAuthCallbackCompletionServiceTest { + + @Mock OAuthAuthorizationRequestRepository repository; + @Mock ExternalAccountConnectionService connectionService; + + @Test + void doesNotSaveAccountWhenRequestIsNoLongerProcessing() { + OAuthAuthorizationRequest request = OAuthAuthorizationRequest.create( + User.create(), ExternalServiceProvider.SPOTIFY, "a".repeat(64), + LocalDateTime.now().plusMinutes(10)); + given(repository.findByIdForUpdate(10L)).willReturn(Optional.of(request)); + OAuthCallbackCompletionService service = + new OAuthCallbackCompletionService(repository, connectionService); + OAuthTokenResponse token = new OAuthTokenResponse("access", "refresh", 3600L); + ExternalUserInfo userInfo = new ExternalUserInfo("id", "email"); + + assertThatThrownBy(() -> service.saveAccountAndSucceed( + 10L, 1L, ExternalServiceProvider.SPOTIFY, token, userInfo)) + .isInstanceOf(BaseException.class); + + verify(connectionService, never()).saveOrUpdate( + 1L, ExternalServiceProvider.SPOTIFY, token, userInfo); + } +}