diff --git a/server/skillhub-app/src/main/java/com/iflytek/skillhub/service/SkillLifecycleAppService.java b/server/skillhub-app/src/main/java/com/iflytek/skillhub/service/SkillLifecycleAppService.java index 5e38644f..7468ef4f 100644 --- a/server/skillhub-app/src/main/java/com/iflytek/skillhub/service/SkillLifecycleAppService.java +++ b/server/skillhub-app/src/main/java/com/iflytek/skillhub/service/SkillLifecycleAppService.java @@ -121,7 +121,7 @@ public class SkillLifecycleAppService { Map userNamespaceRoles, AuditRequestContext auditContext) { Skill skill = findSkill(namespace, slug, userId); - SkillVersion skillVersion = findVersion(skill.getId(), version); + SkillVersion skillVersion = findVersionForUpdate(skill.getId(), version); SkillVersion yanked = skillGovernanceService.yankVersion( skill, skillVersion, diff --git a/server/skillhub-app/src/test/java/com/iflytek/skillhub/controller/portal/SkillLifecycleControllerTest.java b/server/skillhub-app/src/test/java/com/iflytek/skillhub/controller/portal/SkillLifecycleControllerTest.java index 44215f5d..d1412ae0 100644 --- a/server/skillhub-app/src/test/java/com/iflytek/skillhub/controller/portal/SkillLifecycleControllerTest.java +++ b/server/skillhub-app/src/test/java/com/iflytek/skillhub/controller/portal/SkillLifecycleControllerTest.java @@ -177,7 +177,7 @@ class SkillLifecycleControllerTest { given(namespaceRepository.findBySlug("global")).willReturn(java.util.Optional.of(namespace)); given(skillSlugResolutionService.resolve(1L, "demo-skill", "usr_1", SkillSlugResolutionService.Preference.CURRENT_USER)) .willReturn(skill); - given(skillVersionRepository.findBySkillIdAndVersion(1L, "1.2.3")).willReturn(java.util.Optional.of(version)); + given(skillVersionRepository.findBySkillIdForUpdate(1L)).willReturn(java.util.List.of(version)); given(skillGovernanceService.yankVersion( eq(skill), eq(version), eq("usr_1"), anyMap(), nullable(String.class), nullable(String.class), eq("broken"))) .willReturn(yanked); @@ -216,7 +216,7 @@ class SkillLifecycleControllerTest { given(namespaceRepository.findBySlug("global")).willReturn(java.util.Optional.of(namespace)); given(skillSlugResolutionService.resolve(1L, "demo-skill", "usr_1", SkillSlugResolutionService.Preference.CURRENT_USER)) .willReturn(skill); - given(skillVersionRepository.findBySkillIdAndVersion(1L, "1.2.3")).willReturn(java.util.Optional.of(version)); + given(skillVersionRepository.findBySkillIdForUpdate(1L)).willReturn(java.util.List.of(version)); given(skillGovernanceService.yankVersion( eq(skill), eq(version), eq("usr_1"), anyMap(), nullable(String.class), nullable(String.class), isNull())) .willReturn(yanked); diff --git a/server/skillhub-app/src/test/java/com/iflytek/skillhub/service/SkillLifecycleAppServiceTest.java b/server/skillhub-app/src/test/java/com/iflytek/skillhub/service/SkillLifecycleAppServiceTest.java index b007621f..eff554ba 100644 --- a/server/skillhub-app/src/test/java/com/iflytek/skillhub/service/SkillLifecycleAppServiceTest.java +++ b/server/skillhub-app/src/test/java/com/iflytek/skillhub/service/SkillLifecycleAppServiceTest.java @@ -124,4 +124,57 @@ class SkillLifecycleAppServiceTest { "global" ); } + + @Test + void yankVersion_locksAllSkillVersionsBeforeDelegatingLifecycleMutation() { + Namespace namespace = new Namespace("global", "Global", "owner-1"); + ReflectionTestUtils.setField(namespace, "id", 7L); + Skill skill = new Skill(7L, "demo-skill", "owner-1", SkillVisibility.PUBLIC); + ReflectionTestUtils.setField(skill, "id", 11L); + SkillVersion version = new SkillVersion(11L, "1.0.0", "owner-1"); + ReflectionTestUtils.setField(version, "id", 13L); + version.setStatus(SkillVersionStatus.PUBLISHED); + + when(namespaceRepository.findBySlug("global")).thenReturn(Optional.of(namespace)); + when(skillSlugResolutionService.resolve( + 7L, + "demo-skill", + "owner-1", + SkillSlugResolutionService.Preference.CURRENT_USER + )).thenReturn(skill); + when(skillVersionRepository.findBySkillIdForUpdate(11L)) + .thenReturn(java.util.List.of(version)); + when(skillGovernanceService.yankVersion( + eq(skill), + eq(version), + eq("owner-1"), + anyMap(), + nullable(String.class), + nullable(String.class), + eq("broken") + )).thenReturn(version); + + var response = service.yankVersion( + "global", + "demo-skill", + "1.0.0", + new AdminSkillActionRequest("broken"), + "owner-1", + Map.of(7L, NamespaceRole.OWNER), + new AuditRequestContext("127.0.0.1", "JUnit") + ); + + assertThat(response.versionId()).isEqualTo(13L); + assertThat(response.action()).isEqualTo("YANK"); + verify(skillVersionRepository).findBySkillIdForUpdate(11L); + verify(skillGovernanceService).yankVersion( + skill, + version, + "owner-1", + Map.of(7L, NamespaceRole.OWNER), + "127.0.0.1", + "JUnit", + "broken" + ); + } } diff --git a/server/skillhub-app/src/test/java/com/iflytek/skillhub/service/SkillYankLockingTest.java b/server/skillhub-app/src/test/java/com/iflytek/skillhub/service/SkillYankLockingTest.java new file mode 100644 index 00000000..5bbac28e --- /dev/null +++ b/server/skillhub-app/src/test/java/com/iflytek/skillhub/service/SkillYankLockingTest.java @@ -0,0 +1,203 @@ +package com.iflytek.skillhub.service; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.doAnswer; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import com.iflytek.skillhub.domain.audit.AuditLogService; +import com.iflytek.skillhub.domain.event.SkillVersionYankedEvent; +import com.iflytek.skillhub.domain.review.ReviewTaskRepository; +import com.iflytek.skillhub.domain.security.SecurityScanService; +import com.iflytek.skillhub.domain.shared.exception.DomainBadRequestException; +import com.iflytek.skillhub.domain.namespace.Namespace; +import com.iflytek.skillhub.domain.skill.Skill; +import com.iflytek.skillhub.domain.skill.SkillFileRepository; +import com.iflytek.skillhub.domain.skill.SkillRepository; +import com.iflytek.skillhub.domain.skill.SkillVersion; +import com.iflytek.skillhub.domain.skill.SkillVersionRepository; +import com.iflytek.skillhub.domain.skill.SkillVersionStatus; +import com.iflytek.skillhub.domain.skill.SkillVisibility; +import com.iflytek.skillhub.domain.skill.service.SkillGovernanceService; +import com.iflytek.skillhub.domain.skill.service.SkillStorageDeletionCompensationService; +import com.iflytek.skillhub.domain.user.UserAccount; +import com.iflytek.skillhub.infra.jpa.SkillVersionJpaRepository; +import com.iflytek.skillhub.storage.ObjectStorageService; +import jakarta.persistence.EntityManager; +import jakarta.persistence.PersistenceContext; +import java.time.Clock; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; +import java.util.UUID; +import org.junit.jupiter.api.Test; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.autoconfigure.jdbc.AutoConfigureTestDatabase; +import org.springframework.boot.test.autoconfigure.orm.jpa.DataJpaTest; +import org.springframework.context.ApplicationEventPublisher; +import org.springframework.test.context.ActiveProfiles; +import org.springframework.test.context.DynamicPropertyRegistry; +import org.springframework.test.context.DynamicPropertySource; +import org.springframework.transaction.PlatformTransactionManager; +import org.springframework.transaction.annotation.Propagation; +import org.springframework.transaction.annotation.Transactional; +import org.springframework.transaction.support.TransactionTemplate; +import org.testcontainers.containers.PostgreSQLContainer; +import org.testcontainers.junit.jupiter.Container; +import org.testcontainers.junit.jupiter.Testcontainers; + +@DataJpaTest +@AutoConfigureTestDatabase(replace = AutoConfigureTestDatabase.Replace.NONE) +@ActiveProfiles("test") +@Testcontainers +class SkillYankLockingTest { + + @Container + private static final PostgreSQLContainer POSTGRES = + new PostgreSQLContainer<>("postgres:16-alpine"); + + @DynamicPropertySource + static void configurePostgres(DynamicPropertyRegistry registry) { + registry.add("spring.datasource.url", POSTGRES::getJdbcUrl); + registry.add("spring.datasource.username", POSTGRES::getUsername); + registry.add("spring.datasource.password", POSTGRES::getPassword); + registry.add("spring.datasource.driver-class-name", () -> "org.postgresql.Driver"); + registry.add("spring.jpa.database-platform", () -> "org.hibernate.dialect.PostgreSQLDialect"); + } + + @Autowired + private SkillVersionJpaRepository skillVersionRepository; + + @Autowired + private PlatformTransactionManager transactionManager; + + @PersistenceContext + private EntityManager entityManager; + + @Test + @Transactional(propagation = Propagation.NOT_SUPPORTED) + void secondYankLockSeesTheCommittedYankedStatus() throws Exception { + Fixture fixture = persistPublishedVersion(); + CountDownLatch firstLocked = new CountDownLatch(1); + CountDownLatch secondAboutToLock = new CountDownLatch(1); + CountDownLatch releaseFirst = new CountDownLatch(1); + TransactionTemplate transactions = new TransactionTemplate(transactionManager); + + try (var executor = Executors.newVirtualThreadPerTaskExecutor()) { + var first = executor.submit(() -> transactions.executeWithoutResult(status -> { + SkillVersion version = skillVersionRepository.findBySkillIdForUpdate(fixture.skillId()).getFirst(); + version.setStatus(SkillVersionStatus.YANKED); + firstLocked.countDown(); + await(releaseFirst); + })); + + assertThat(firstLocked.await(10, TimeUnit.SECONDS)).isTrue(); + var second = executor.submit(() -> transactions.execute(status -> { + assertThat(skillVersionRepository.findStatusByIdAndSkillId( + fixture.versionId(), fixture.skillId())).contains(SkillVersionStatus.PUBLISHED); + secondAboutToLock.countDown(); + return skillVersionRepository.findBySkillIdForUpdate(fixture.skillId()).getFirst().getStatus(); + })); + + assertThat(secondAboutToLock.await(10, TimeUnit.SECONDS)).isTrue(); + releaseFirst.countDown(); + first.get(10, TimeUnit.SECONDS); + assertThat(second.get(10, TimeUnit.SECONDS)).isEqualTo(SkillVersionStatus.YANKED); + } finally { + releaseFirst.countDown(); + } + } + + @Test + @Transactional(propagation = Propagation.NOT_SUPPORTED) + void concurrentAdminYanksRecordOneAuditAndOneEvent() throws Exception { + Fixture fixture = persistPublishedVersion(); + SkillRepository skillRepository = mock(SkillRepository.class); + AuditLogService auditLogService = mock(AuditLogService.class); + ApplicationEventPublisher eventPublisher = mock(ApplicationEventPublisher.class); + when(skillRepository.findById(fixture.skillId())).thenReturn(java.util.Optional.empty()); + SkillGovernanceService service = new SkillGovernanceService( + skillRepository, + skillVersionRepository, + mock(SkillFileRepository.class), + mock(ReviewTaskRepository.class), + mock(ObjectStorageService.class), + auditLogService, + eventPublisher, + mock(SecurityScanService.class), + mock(SkillStorageDeletionCompensationService.class), + Clock.systemUTC()); + CountDownLatch firstAudited = new CountDownLatch(1); + CountDownLatch secondReady = new CountDownLatch(1); + CountDownLatch releaseFirst = new CountDownLatch(1); + doAnswer(invocation -> { + firstAudited.countDown(); + await(releaseFirst); + return null; + }).when(auditLogService).record(any(), any(), any(), any(), any(), any(), any(), any()); + TransactionTemplate transactions = new TransactionTemplate(transactionManager); + + try (var executor = Executors.newVirtualThreadPerTaskExecutor()) { + var first = executor.submit(() -> transactions.executeWithoutResult(status -> + service.yankVersion(fixture.versionId(), "admin-one", "127.0.0.1", "JUnit", "broken"))); + assertThat(firstAudited.await(10, TimeUnit.SECONDS)).isTrue(); + var second = executor.submit(() -> { + secondReady.countDown(); + return transactions.execute(status -> + service.yankVersion(fixture.versionId(), "admin-two", "127.0.0.1", "JUnit", "broken")); + }); + assertThat(secondReady.await(10, TimeUnit.SECONDS)).isTrue(); + releaseFirst.countDown(); + first.get(10, TimeUnit.SECONDS); + assertThatThrownBy(() -> second.get(10, TimeUnit.SECONDS)) + .hasCauseInstanceOf(DomainBadRequestException.class); + } finally { + releaseFirst.countDown(); + } + + SkillVersionStatus finalStatus = new TransactionTemplate(transactionManager).execute(status -> + ((SkillVersionRepository) skillVersionRepository).findById(fixture.versionId()) + .orElseThrow().getStatus()); + assertThat(finalStatus).isEqualTo(SkillVersionStatus.YANKED); + verify(auditLogService, times(1)).record(any(), any(), any(), any(), any(), any(), any(), any()); + verify(eventPublisher, times(1)).publishEvent(any(SkillVersionYankedEvent.class)); + } + + private Fixture persistPublishedVersion() { + TransactionTemplate transaction = new TransactionTemplate(transactionManager); + return transaction.execute(status -> { + String suffix = UUID.randomUUID().toString().substring(0, 8); + UserAccount user = new UserAccount("yank-lock-user-" + suffix, "Yank Lock User", null, null); + entityManager.persist(user); + Namespace namespace = new Namespace("yank-lock-" + suffix, "Yank Lock", user.getId()); + entityManager.persist(namespace); + entityManager.flush(); + Skill skill = new Skill(namespace.getId(), "yank-lock", user.getId(), SkillVisibility.PUBLIC); + entityManager.persist(skill); + entityManager.flush(); + SkillVersion version = new SkillVersion(skill.getId(), "1.0.0", user.getId()); + version.setStatus(SkillVersionStatus.PUBLISHED); + entityManager.persist(version); + entityManager.flush(); + return new Fixture(skill.getId(), version.getId()); + }); + } + + private void await(CountDownLatch latch) { + try { + if (!latch.await(10, TimeUnit.SECONDS)) { + throw new IllegalStateException("Timed out waiting for concurrent yank test"); + } + } catch (InterruptedException error) { + Thread.currentThread().interrupt(); + throw new IllegalStateException("Concurrent yank test interrupted", error); + } + } + + private record Fixture(Long skillId, Long versionId) { + } +} diff --git a/server/skillhub-domain/src/main/java/com/iflytek/skillhub/domain/skill/service/SkillGovernanceService.java b/server/skillhub-domain/src/main/java/com/iflytek/skillhub/domain/skill/service/SkillGovernanceService.java index 5737e70e..72b32eab 100644 --- a/server/skillhub-domain/src/main/java/com/iflytek/skillhub/domain/skill/service/SkillGovernanceService.java +++ b/server/skillhub-domain/src/main/java/com/iflytek/skillhub/domain/skill/service/SkillGovernanceService.java @@ -266,7 +266,7 @@ public class SkillGovernanceService { */ @Transactional public SkillVersion yankVersion(Long versionId, String actorUserId, String clientIp, String userAgent, String reason) { - SkillVersion version = skillVersionRepository.findById(versionId) + SkillVersion version = skillVersionRepository.findByIdForUpdate(versionId) .orElseThrow(() -> new DomainNotFoundException("error.skill.version.notFound", versionId)); return yankVersionInternal(version, actorUserId, clientIp, userAgent, reason); } diff --git a/server/skillhub-domain/src/test/java/com/iflytek/skillhub/domain/skill/service/SkillGovernanceServiceTest.java b/server/skillhub-domain/src/test/java/com/iflytek/skillhub/domain/skill/service/SkillGovernanceServiceTest.java index dd9e9a75..2a149024 100644 --- a/server/skillhub-domain/src/test/java/com/iflytek/skillhub/domain/skill/service/SkillGovernanceServiceTest.java +++ b/server/skillhub-domain/src/test/java/com/iflytek/skillhub/domain/skill/service/SkillGovernanceServiceTest.java @@ -150,7 +150,7 @@ class SkillGovernanceServiceTest { SkillVersion version = new SkillVersion(2L, "1.0.0", "owner"); setField(version, "id", 22L); version.setStatus(SkillVersionStatus.PUBLISHED); - given(skillVersionRepository.findById(22L)).willReturn(Optional.of(version)); + given(skillVersionRepository.findByIdForUpdate(22L)).willReturn(Optional.of(version)); given(skillVersionRepository.save(version)).willReturn(version); given(skillRepository.findById(2L)).willReturn(Optional.empty()); @@ -159,6 +159,7 @@ class SkillGovernanceServiceTest { assertThat(result.getStatus()).isEqualTo(SkillVersionStatus.YANKED); assertThat(result.getYankedBy()).isEqualTo("admin"); assertThat(result.getYankedAt()).isEqualTo(Instant.now(CLOCK)); + verify(skillVersionRepository).findByIdForUpdate(22L); verify(auditLogService).record("admin", "YANK_SKILL_VERSION", "SKILL_VERSION", 22L, null, "127.0.0.1", "JUnit", "{\"reason\":\"broken\"}"); } @@ -212,6 +213,7 @@ class SkillGovernanceServiceTest { assertThat(version.getStatus()).isEqualTo(SkillVersionStatus.PUBLISHED); verify(skillVersionRepository, never()).save(any()); verify(auditLogService, never()).record(any(), any(), any(), any(), any(), any(), any(), any()); + verify(eventPublisher, never()).publishEvent(any()); } @Test @@ -226,6 +228,8 @@ class SkillGovernanceServiceTest { () -> service.yankVersion(skill, version, "owner", Map.of(), "127.0.0.1", "JUnit", null)); verify(skillVersionRepository, never()).save(any()); + verify(auditLogService, never()).record(any(), any(), any(), any(), any(), any(), any(), any()); + verify(eventPublisher, never()).publishEvent(any()); } @Test @@ -262,7 +266,7 @@ class SkillGovernanceServiceTest { setField(skill, "id", 2L); skill.setLatestVersionId(22L); - given(skillVersionRepository.findById(22L)).willReturn(Optional.of(yanked)); + given(skillVersionRepository.findByIdForUpdate(22L)).willReturn(Optional.of(yanked)); given(skillVersionRepository.save(yanked)).willReturn(yanked); given(skillRepository.findById(2L)).willReturn(Optional.of(skill)); given(skillVersionRepository.findBySkillIdAndStatus(2L, SkillVersionStatus.PUBLISHED)).willReturn(java.util.List.of(fallback));