diff --git a/backend-java/src/test/java/com/nanri/aiimage/modules/task/contract/RollbackSemanticsContractTest.java b/backend-java/src/test/java/com/nanri/aiimage/modules/task/contract/RollbackSemanticsContractTest.java new file mode 100644 index 00000000..c6f6cf27 --- /dev/null +++ b/backend-java/src/test/java/com/nanri/aiimage/modules/task/contract/RollbackSemanticsContractTest.java @@ -0,0 +1,324 @@ +package com.nanri.aiimage.modules.task.contract; + +import cn.hutool.crypto.digest.DigestUtil; +import com.baomidou.mybatisplus.core.MybatisConfiguration; +import com.baomidou.mybatisplus.core.metadata.TableInfoHelper; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.nanri.aiimage.common.service.DistributedJobLockService; +import com.nanri.aiimage.config.InstanceMetadata; +import com.nanri.aiimage.config.SimilarAsinProperties; +import com.nanri.aiimage.config.StorageProperties; +import com.nanri.aiimage.modules.appearancepatent.client.AppearancePatentLlmClient; +import com.nanri.aiimage.modules.appearancepatent.service.AppearancePatentTaskCacheService; +import com.nanri.aiimage.modules.appearancepatent.service.AppearancePatentTaskService; +import com.nanri.aiimage.modules.appearancepatent.model.dto.AppearancePatentSubmitResultRequest; +import com.nanri.aiimage.modules.file.service.LocalFileStorageService; +import com.nanri.aiimage.modules.file.service.oss.OssStorageService; +import com.nanri.aiimage.modules.similarasin.mapper.SimilarAsinFilterConditionMapper; +import com.nanri.aiimage.modules.similarasin.model.dto.SimilarAsinSubmitResultRequest; +import com.nanri.aiimage.modules.similarasin.util.SimilarAsinImageEmbedder; +import com.nanri.aiimage.modules.similarasin.service.SimilarAsinImagePrefetchService; +import com.nanri.aiimage.modules.similarasin.service.SimilarAsinLlmService; +import com.nanri.aiimage.modules.similarasin.service.SimilarAsinTaskCacheService; +import com.nanri.aiimage.modules.similarasin.service.SimilarAsinTaskService; +import com.nanri.aiimage.modules.task.mapper.FileResultMapper; +import com.nanri.aiimage.modules.task.mapper.FileTaskMapper; +import com.nanri.aiimage.modules.task.mapper.TaskChunkMapper; +import com.nanri.aiimage.modules.task.mapper.TaskScopeStateMapper; +import com.nanri.aiimage.modules.task.model.entity.FileTaskEntity; +import com.nanri.aiimage.modules.task.model.entity.TaskScopeStateEntity; +import com.nanri.aiimage.modules.task.service.TaskDistributedLockService; +import com.nanri.aiimage.modules.task.service.TaskFileJobService; +import com.nanri.aiimage.modules.task.service.TaskProgressLightAssembler; +import com.nanri.aiimage.modules.task.service.TaskProgressSnapshotService; +import com.nanri.aiimage.modules.task.service.TransientPayloadStorageService; +import org.apache.ibatis.builder.MapperBuilderAssistant; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.Spy; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.test.util.ReflectionTestUtils; +import org.springframework.transaction.PlatformTransactionManager; +import org.springframework.transaction.TransactionDefinition; +import org.springframework.transaction.TransactionStatus; + +import java.time.Duration; +import java.util.List; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyLong; +import static org.mockito.ArgumentMatchers.anyString; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.doAnswer; +import static org.mockito.Mockito.lenient; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +/** + * task-133:回滚语义硬约束契约(spec 07 §2)。 + * 落库异常 → 完整回滚:任务状态/结果不被改写、payload 保留、异常抛出; + * 纯计算(prepare)无副作用且幂等;重试安全。similarasin + appearancepatent + * 双模块统一守门。 + */ +@ExtendWith(MockitoExtension.class) +class RollbackSemanticsContractTest { + + private static final Long TASK_ID = 9999L; + private static final String PARSED_POINTER = "rustfs:task-parsed/similar-asin/9999/payload.json"; + private static final String CHUNK_POINTER = "rustfs:task-chunk/similar-asin/9999/chunk.json"; + private static final String STORED_CHUNK_POINTER = "\"" + CHUNK_POINTER + "\""; + + @Mock private LocalFileStorageService localFileStorageService; + @Mock private OssStorageService ossStorageService; + @Mock private StorageProperties storageProperties; + @Mock private FileTaskMapper fileTaskMapper; + @Mock private FileResultMapper fileResultMapper; + @Mock private TaskScopeStateMapper taskScopeStateMapper; + @Mock private TaskChunkMapper taskChunkMapper; + @Mock private SimilarAsinFilterConditionMapper filterConditionMapper; + @Spy private ObjectMapper objectMapper = new ObjectMapper(); + @Mock private SimilarAsinLlmService similarAsinLlmService; + @Mock private SimilarAsinTaskCacheService taskCacheService; + @Mock private SimilarAsinProperties properties; + @Mock private TaskFileJobService taskFileJobService; + @Mock private TaskDistributedLockService taskDistributedLockService; + @Mock private TaskProgressSnapshotService taskProgressSnapshotService; + @Mock private TransientPayloadStorageService transientPayloadStorageService; + @Mock private PlatformTransactionManager transactionManager; + @Mock private DistributedJobLockService distributedJobLockService; + @Mock private InstanceMetadata instanceMetadata; + @Mock private SimilarAsinImageEmbedder imageEmbedder; + @Mock private SimilarAsinImagePrefetchService imagePrefetchService; + @Mock private TransactionStatus transactionStatus; + + @InjectMocks private SimilarAsinTaskService service; + + private final AtomicBoolean transactionActive = new AtomicBoolean(); + private final AtomicReference storedPayloadJson = new AtomicReference<>(); + + @BeforeAll + static void initializeMybatisMetadata() { + MapperBuilderAssistant assistant = new MapperBuilderAssistant(new MybatisConfiguration(), ""); + TableInfoHelper.initTableInfo(assistant, FileTaskEntity.class); + TableInfoHelper.initTableInfo(assistant, TaskScopeStateEntity.class); + } + + @BeforeEach + void setUp() { + lenient().when(instanceMetadata.getInstanceId()).thenReturn("instance-a"); + lenient().when(taskDistributedLockService.acquire( + eq(SimilarAsinTaskService.MODULE_TYPE), anyLong(), any(Duration.class), eq(10_000L))) + .thenReturn(mock(TaskDistributedLockService.LockHandle.class)); + lenient().when(transactionManager.getTransaction(any(TransactionDefinition.class))) + .thenAnswer(invocation -> { + transactionActive.set(true); + return transactionStatus; + }); + lenient().doAnswer(invocation -> { + transactionActive.set(false); + return null; + }).when(transactionManager).commit(transactionStatus); + lenient().doAnswer(invocation -> { + transactionActive.set(false); + return null; + }).when(transactionManager).rollback(transactionStatus); + lenient().when(taskChunkMapper.selectList(any())).thenReturn(List.of()); + lenient().when(fileResultMapper.selectList(any())).thenReturn(List.of()); + lenient().when(taskChunkMapper.selectCount(any())).thenReturn(1L); + lenient().when(taskChunkMapper.selectOne(any())).thenReturn(null); + lenient().when(transientPayloadStorageService.storeChunkPayloadVersioned( + eq(SimilarAsinTaskService.MODULE_TYPE), eq(TASK_ID), anyString(), any(), anyString())) + .thenAnswer(invocation -> { + storedPayloadJson.set(invocation.getArgument(4)); + return STORED_CHUNK_POINTER; + }); + lenient().when(transientPayloadStorageService.wasLastStoreLocalFallback()).thenReturn(false); + lenient().doAnswer(invocation -> { + com.nanri.aiimage.modules.task.model.entity.TaskChunkEntity chunk = invocation.getArgument(0); + chunk.setId(301L); + return 1; + }).when(taskChunkMapper).insert(any(com.nanri.aiimage.modules.task.model.entity.TaskChunkEntity.class)); + lenient().doAnswer(invocation -> { + TaskScopeStateEntity scope = invocation.getArgument(0); + scope.setId(401L); + return 1; + }).when(taskScopeStateMapper).insert(any(TaskScopeStateEntity.class)); + lenient().when(fileTaskMapper.updateById(any(FileTaskEntity.class))).thenReturn(1); + } + + @AfterEach + void shutdownExecutors() { + service.shutdownAssembleExecutor(); + } + + @Test + void rollbackIsFullAndExceptionPropagates() { + FileTaskEntity task = runningTask(); + when(fileTaskMapper.selectById(TASK_ID)).thenReturn(task); + org.mockito.Mockito.doThrow(new IllegalStateException("scope upsert failed")) + .when(taskScopeStateMapper).insert(any(TaskScopeStateEntity.class)); + + assertThrows(IllegalStateException.class, () -> service.submitResult(TASK_ID, request())); + + verify(transactionManager).rollback(transactionStatus); + verify(transactionManager, never()).commit(transactionStatus); + } + + @Test + void rollbackLeavesTaskRowUntouched() { + FileTaskEntity task = runningTask(); + when(fileTaskMapper.selectById(TASK_ID)).thenReturn(task); + org.mockito.Mockito.doThrow(new IllegalStateException("scope upsert failed")) + .when(taskScopeStateMapper).insert(any(TaskScopeStateEntity.class)); + + assertThrows(IllegalStateException.class, () -> service.submitResult(TASK_ID, request())); + + verify(fileTaskMapper, never()).updateById(any(FileTaskEntity.class)); + } + + @Test + void rollbackKeepsPayloadNotDeleted() { + FileTaskEntity task = runningTask(); + when(fileTaskMapper.selectById(TASK_ID)).thenReturn(task); + org.mockito.Mockito.doThrow(new IllegalStateException("scope upsert failed")) + .when(taskScopeStateMapper).insert(any(TaskScopeStateEntity.class)); + + assertThrows(IllegalStateException.class, () -> service.submitResult(TASK_ID, request())); + + verify(transientPayloadStorageService, never()).deletePayloadIfPresent(anyString()); + } + + @Test + void retryAfterRollbackSucceeds() { + FileTaskEntity task = runningTask(); + when(fileTaskMapper.selectById(TASK_ID)).thenReturn(task); + AtomicInteger failures = new AtomicInteger(); + lenient().doAnswer(invocation -> { + if (failures.getAndIncrement() == 0) { + throw new IllegalStateException("first attempt failed"); + } + TaskScopeStateEntity scope = invocation.getArgument(0); + scope.setId(401L); + return 1; + }).when(taskScopeStateMapper).insert(any(TaskScopeStateEntity.class)); + + assertThrows(IllegalStateException.class, () -> service.submitResult(TASK_ID, request())); + service.submitResult(TASK_ID, request()); + + verify(transactionManager, times(1)).rollback(transactionStatus); + verify(transactionManager).commit(transactionStatus); + verify(taskCacheService).touchTaskHeartbeat(TASK_ID); + } + + @Test + void computePhaseIsSideEffectFreeAndIdempotent() throws Exception { + FileTaskEntity task = runningTask(); + when(fileTaskMapper.selectById(TASK_ID)).thenReturn(task); + + Object first = ReflectionTestUtils.invokeMethod(service, "prepareSubmittedChunk", TASK_ID, request()); + Object second = ReflectionTestUtils.invokeMethod(service, "prepareSubmittedChunk", TASK_ID, request()); + + assertEquals(DigestUtil.sha256Hex(storedPayloadJson.get()), + ReflectionTestUtils.getField(first, "payloadHash")); + assertEquals(ReflectionTestUtils.getField(first, "payloadHash"), + ReflectionTestUtils.getField(second, "payloadHash"), "prepare 幂等:重复调用哈希一致"); + verify(taskChunkMapper, never()).insert(any(com.nanri.aiimage.modules.task.model.entity.TaskChunkEntity.class)); + verify(taskScopeStateMapper, never()).insert(any(TaskScopeStateEntity.class)); + } + + @Test + void partialWriteDoesNotContinueAfterFailure() { + FileTaskEntity task = runningTask(); + when(fileTaskMapper.selectById(TASK_ID)).thenReturn(task); + org.mockito.Mockito.doThrow(new IllegalStateException("scope upsert failed")) + .when(taskScopeStateMapper).insert(any(TaskScopeStateEntity.class)); + + assertThrows(IllegalStateException.class, () -> service.submitResult(TASK_ID, request())); + + // chunk insert 已发生(部分写入),但后续 updateById/心跳/schedule 全部不执行 + verify(taskChunkMapper).insert(any(com.nanri.aiimage.modules.task.model.entity.TaskChunkEntity.class)); + verify(fileTaskMapper, never()).updateById(any(FileTaskEntity.class)); + verify(taskCacheService, never()).touchTaskHeartbeat(TASK_ID); + verify(taskFileJobService, never()).enqueueAssembleResult(anyLong(), anyString(), anyLong(), anyString()); + } + + @Test + void appearancePatentRollbackSkipsSecondTransaction() { + AppearancePatentTaskService apService = appearancePatentService(); + FileTaskEntity task = new FileTaskEntity(); + task.setId(TASK_ID); + task.setModuleType("APPEARANCE_PATENT"); + task.setStatus("RUNNING"); + task.setUserId(7L); + task.setResultJson("{\"parsedPayloadRef\":\"" + PARSED_POINTER + "\",\"ownerInstanceId\":\"instance-a\"}"); + when(fileTaskMapper.selectById(TASK_ID)).thenReturn(task); + lenient().when(transientPayloadStorageService.isSharedWriteEnabled()).thenReturn(true); + when(transientPayloadStorageService.storeChunkPayload( + eq("APPEARANCE_PATENT"), eq(TASK_ID), anyString(), any(), anyString())) + .thenReturn(STORED_CHUNK_POINTER); + lenient().when(taskScopeStateMapper.selectOne(any())).thenReturn(null); + when(fileTaskMapper.updateById(any(FileTaskEntity.class))).thenThrow( + new IllegalStateException("persist tx failed")); + + assertThrows(IllegalStateException.class, () -> apService.submitResult(TASK_ID, new AppearancePatentSubmitResultRequest())); + + verify(transactionManager, times(1)).getTransaction(any(TransactionDefinition.class)); + verify(transactionManager).rollback(transactionStatus); + verify(transactionManager, never()).commit(transactionStatus); + } + + @Test + void rollbackConsistencyAcrossModules() { + // 双模块回滚观测一致:rollback 调用、无 commit、异常传播(同类断言已覆盖) + FileTaskEntity task = runningTask(); + when(fileTaskMapper.selectById(TASK_ID)).thenReturn(task); + org.mockito.Mockito.doThrow(new IllegalStateException("boom")) + .when(taskScopeStateMapper).insert(any(TaskScopeStateEntity.class)); + + assertThrows(IllegalStateException.class, () -> service.submitResult(TASK_ID, request())); + + verify(transactionManager, never()).commit(transactionStatus); + } + + private AppearancePatentTaskService appearancePatentService() { + return new AppearancePatentTaskService( + localFileStorageService, ossStorageService, storageProperties, fileTaskMapper, fileResultMapper, + taskScopeStateMapper, taskChunkMapper, objectMapper, mock(AppearancePatentLlmClient.class), + mock(AppearancePatentTaskCacheService.class), mock(com.nanri.aiimage.config.AppearancePatentProperties.class), + taskFileJobService, taskProgressSnapshotService, transientPayloadStorageService, transactionManager, + distributedJobLockService, taskDistributedLockService, instanceMetadata, + mock(TaskProgressLightAssembler.class)); + } + + private SimilarAsinSubmitResultRequest request() { + SimilarAsinSubmitResultRequest request = new SimilarAsinSubmitResultRequest(); + request.setSubmissionId("similar-asin-9999"); + request.setChunkIndex(0); + request.setChunkTotal(1); + request.setDone(false); + return request; + } + + private FileTaskEntity runningTask() { + FileTaskEntity task = new FileTaskEntity(); + task.setId(TASK_ID); + task.setModuleType(SimilarAsinTaskService.MODULE_TYPE); + task.setStatus("RUNNING"); + task.setUserId(7L); + task.setResultJson("{\"parsedPayloadRef\":\"" + PARSED_POINTER + "\",\"ownerInstanceId\":\"instance-a\"}"); + return task; + } +}