diff --git a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/factory/ai/stage/oneclick/GenerateStageHandler.java b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/factory/ai/stage/oneclick/GenerateStageHandler.java index d7ba394..bc2b6f5 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/factory/ai/stage/oneclick/GenerateStageHandler.java +++ b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/factory/ai/stage/oneclick/GenerateStageHandler.java @@ -1,5 +1,9 @@ package com.ruoyi.generator.factory.ai.stage.oneclick; +import java.util.HashSet; +import java.util.List; +import java.util.Map; +import java.util.Set; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.stereotype.Component; import com.ruoyi.common.exception.ServiceException; @@ -78,11 +82,28 @@ public class GenerateStageHandler implements OneClickGenerationStageHandler } result.setGenerationRunId(generated.getGenerationRunId()); stagePipeline.linkGenerationRun(context.getTask(), stageCode(), generated.getGenerationRunId()); - for (String type : previewService.getSupportedTemplateTypes(userId, projectId)) + List templateTypes = previewService.getSupportedTemplateTypes(userId, projectId); + if (templateTypes == null || templateTypes.isEmpty()) { - previewService.getStructure(userId, projectId, type); + throw new ServiceException("代码模板没有可生成的项目类型"); + } + Set renderedTypes = new HashSet(); + for (String type : templateTypes) + { + if (type == null || type.trim().length() == 0 || !renderedTypes.add(type)) + { + throw new ServiceException("代码模板包含无效或重复的项目类型"); + } + List> structure = previewService.getStructure(userId, projectId, type); + if (structure == null || structure.isEmpty()) + { + throw new ServiceException("生成结果缺少项目结构: " + type); + } + } + if (previewService.markPreviewReady(userId, projectId) != 1) + { + throw new ServiceException("无法标记项目源码为可预览状态"); } - previewService.markPreviewReady(userId, projectId); result.setDownloadReady(true); checkpointService.save(context.getTask(), stageCode(), result, result.getSpecVersionId(), result.getGenerationRunId()); diff --git a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/factory/ai/stage/oneclick/RunPreviewStageHandler.java b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/factory/ai/stage/oneclick/RunPreviewStageHandler.java index a82076f..62389d6 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/factory/ai/stage/oneclick/RunPreviewStageHandler.java +++ b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/factory/ai/stage/oneclick/RunPreviewStageHandler.java @@ -47,8 +47,9 @@ public class RunPreviewStageHandler implements OneClickGenerationStageHandler context.getResult().getSpecVersionId(), context.getResult().getGenerationRunId()); if (stageCode().equals(context.getResult().getFailedStage())) { - stagePipeline.fail(context.getTask(), stageCode(), - new ServiceException(context.getResult().getErrorMessage())); + ServiceException failure = new ServiceException(context.getResult().getErrorMessage()); + stagePipeline.fail(context.getTask(), stageCode(), failure); + throw failure; } } diff --git a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/mapper/front/AiGenerationTaskMapper.java b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/mapper/front/AiGenerationTaskMapper.java index 806892e..64159e0 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/mapper/front/AiGenerationTaskMapper.java +++ b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/mapper/front/AiGenerationTaskMapper.java @@ -24,6 +24,11 @@ public interface AiGenerationTaskMapper public List selectClaimableTasks(@Param("limit") Integer limit); public int insertAiGenerationTask(AiGenerationTask aiGenerationTask); public int updateAiGenerationTask(AiGenerationTask aiGenerationTask); + public int finishClaimedTask(@Param("task") AiGenerationTask task, + @Param("lockedBy") String lockedBy); + public int cancelPendingTask(@Param("userId") Long userId, @Param("projectId") Long projectId, + @Param("taskId") Long taskId); + public int retryFailedTask(@Param("task") AiGenerationTask task); public int claimTask(@Param("taskId") Long taskId, @Param("lockedBy") String lockedBy); public int renewTaskLock(@Param("taskId") Long taskId, @Param("lockedBy") String lockedBy); public int clearTaskLock(@Param("taskId") Long taskId); diff --git a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/mapper/front/AiQuotaBucketMapper.java b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/mapper/front/AiQuotaBucketMapper.java index 40fad4a..64ec396 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/mapper/front/AiQuotaBucketMapper.java +++ b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/mapper/front/AiQuotaBucketMapper.java @@ -6,6 +6,8 @@ import com.ruoyi.generator.domain.front.AiQuotaBucket; public interface AiQuotaBucketMapper { public AiQuotaBucket selectQuotaBucket(@Param("userId") Long userId, @Param("periodType") String periodType, @Param("periodKey") String periodKey); + public AiQuotaBucket selectQuotaBucketForUpdate(@Param("userId") Long userId, + @Param("periodType") String periodType, @Param("periodKey") String periodKey); public int insertQuotaBucket(AiQuotaBucket aiQuotaBucket); public int updateQuotaBucket(AiQuotaBucket aiQuotaBucket); } diff --git a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/mapper/front/FrontProjectMapper.java b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/mapper/front/FrontProjectMapper.java index e9926f5..a0093a0 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/mapper/front/FrontProjectMapper.java +++ b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/mapper/front/FrontProjectMapper.java @@ -14,5 +14,6 @@ public interface FrontProjectMapper public int updateFrontProjectIfRevision(FrontProject frontProject); public int deleteFrontProjectById(Long projectId); public FrontProject selectFrontProjectByUserAndId(@Param("userId") Long userId, @Param("projectId") Long projectId); + public FrontProject lockFrontProjectByUserAndId(@Param("userId") Long userId, @Param("projectId") Long projectId); public int deleteFrontProjectByUserAndId(@Param("userId") Long userId, @Param("projectId") Long projectId); } diff --git a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/GenProjectServiceImpl.java b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/GenProjectServiceImpl.java index b1eaea6..827950f 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/GenProjectServiceImpl.java +++ b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/GenProjectServiceImpl.java @@ -33,6 +33,7 @@ import com.ruoyi.generator.util.TableOperationResolver; import com.ruoyi.generator.util.TypedBusinessActionCompiler; import com.ruoyi.generator.util.VelocityInitializer; import com.ruoyi.generator.util.VelocityUtils; +import com.ruoyi.generator.util.ZipEntryPathValidator; import org.apache.velocity.VelocityContext; import org.apache.velocity.app.Velocity; import org.springframework.beans.factory.annotation.Autowired; @@ -95,6 +96,7 @@ public class GenProjectServiceImpl implements IGenProjectService { private static final String ADMIN_INDEX_TEMPLATE_FILE = "admin-index.vue.vm"; private static final String BUNDLED_ADMIN_INDEX_TEMPLATE = "qing/admin-index.vue.vm"; private static final String FRONTEND_PACKAGE_TEMPLATE_FILE = "package.json.vm"; + private static final String BUNDLED_FRONTEND_PACKAGE_TEMPLATE = "qing/vue-package.json.vm"; private static final String ADMIN_MAIN_TEMPLATE_FILE = "admin-main.js.vm"; private static final String FRONTEND_MAIN_TEMPLATE_FILE = "frontend-main.js.vm"; private static final String RESOURCE_URL_PROTOCOL_CHECK = @@ -262,7 +264,7 @@ public class GenProjectServiceImpl implements IGenProjectService { Map metadata = parseArtifactMetadata(artifact.getVariableMetadata()); String root = generatedRootName(project, type); - String rootEntry = root + "/"; + String rootEntry = ZipEntryPathValidator.requireRelative(root + "/", "Generated project"); if (zipEntries.add(rootEntry)) { zip.putNextEntry(new ZipEntry(rootEntry)); zip.closeEntry(); @@ -276,7 +278,8 @@ public class GenProjectServiceImpl implements IGenProjectService { input.closeEntry(); continue; } - String relativePath = replaceArtifactPathVariables(sourcePath, metadata, project, type); + String relativePath = ZipEntryPathValidator.requireRelative( + replaceArtifactPathVariables(sourcePath, metadata, project, type), "Template skeleton"); String outputPath = root + "/" + relativePath; if (entry.isDirectory()) { String folderPath = outputPath.endsWith("/") ? outputPath : outputPath + "/"; @@ -303,14 +306,7 @@ public class GenProjectServiceImpl implements IGenProjectService { } private String normalizeSkeletonEntry(String entryName) { - String normalized = StringUtils.defaultString(entryName).replace('\\', '/'); - while (normalized.startsWith("/")) { - normalized = normalized.substring(1); - } - if (normalized.equals("..") || normalized.startsWith("../") || normalized.contains("/../")) { - throw new ServiceException("Template skeleton contains an invalid path: " + entryName); - } - return normalized; + return ZipEntryPathValidator.requireRelative(entryName, "Template skeleton"); } private byte[] readZipEntry(ZipInputStream input) throws IOException { @@ -399,11 +395,18 @@ public class GenProjectServiceImpl implements IGenProjectService { private void processStructureWithPath(List> structure, ZipOutputStream zip, GenProject project, String currentPath, String type, Set zipEntries) throws IOException { for (Map node : structure) { + if (node == null) { + throw new ServiceException("Generated project structure contains an empty node"); + } String nodeType = (String) node.get("type"); String name = (String) node.get("name"); + if (!("folder".equals(nodeType) || "file".equals(nodeType)) || StringUtils.isBlank(name)) { + throw new ServiceException("Generated project structure contains an invalid node"); + } String category = (String) node.get("category"); Object tableId = node.get("tableId"); - String fullPath = currentPath + name; + String fullPath = ZipEntryPathValidator.requireRelative(currentPath + name, + "Generated project structure"); if ("folder".equals(nodeType)) { String folderPath = fullPath + "/"; @@ -422,6 +425,7 @@ public class GenProjectServiceImpl implements IGenProjectService { if (outputPath == null) { continue; } + outputPath = ZipEntryPathValidator.requireRelative(outputPath, "Generated project structure"); String content = generateFileContent(project, category, (Long) tableId, type); if (content != null && zipEntries.add(outputPath)) { zip.putNextEntry(new ZipEntry(outputPath)); @@ -1079,6 +1083,12 @@ public class GenProjectServiceImpl implements IGenProjectService { private String resolveTemplateContent(TemplateFile templateFile) { String content = upgradeRichTextEditorSupport(templateFile, StringUtils.defaultString(templateFile.getFileContent())); + if (shouldUseBundledFrontendPackageTemplate(templateFile)) { + String bundledContent = readClasspathTemplate(BUNDLED_FRONTEND_PACKAGE_TEMPLATE); + if (StringUtils.isNotEmpty(bundledContent)) { + return bundledContent; + } + } if (shouldUseBundledBackendControllerTemplate(templateFile, content)) { String bundledContent = readClasspathTemplate(BUNDLED_BACKEND_CONTROLLER_TEMPLATE); if (StringUtils.isNotEmpty(bundledContent)) { @@ -1148,6 +1158,13 @@ public class GenProjectServiceImpl implements IGenProjectService { return content; } + private boolean shouldUseBundledFrontendPackageTemplate(TemplateFile templateFile) { + return templateFile != null + && (RUNNABLE_ADMIN_FRONTEND_TEMPLATE_ID.equals(templateFile.getTemplateId()) + || RUNNABLE_FRONTEND_TEMPLATE_ID.equals(templateFile.getTemplateId())) + && matchesTemplateFile(templateFile, FRONTEND_PACKAGE_TEMPLATE_FILE); + } + private String upgradeRichTextEditorSupport(TemplateFile templateFile, String content) { if (templateFile == null || (!RUNNABLE_ADMIN_FRONTEND_TEMPLATE_ID.equals(templateFile.getTemplateId()) && !RUNNABLE_FRONTEND_TEMPLATE_ID.equals(templateFile.getTemplateId()))) { diff --git a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/AiGenerationTaskServiceImpl.java b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/AiGenerationTaskServiceImpl.java index 877f92c..91e34db 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/AiGenerationTaskServiceImpl.java +++ b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/AiGenerationTaskServiceImpl.java @@ -43,6 +43,10 @@ public class AiGenerationTaskServiceImpl implements IAiGenerationTaskService { private static final int MAX_ATTEMPTS = 3; private static final int ONE_CLICK_HISTORY_LIMIT = 20; + private static final int MAX_TASK_REQUEST_BYTES = 1024 * 1024; + private static final int MAX_PROJECT_NAME_LENGTH = 100; + private static final int MAX_PROJECT_DESCRIPTION_LENGTH = 20000; + private static final int MAX_OPTION_CODE_LENGTH = 100; private final OneClickTaskProgressComposer oneClickTaskProgressComposer = new OneClickTaskProgressComposer(); @@ -73,10 +77,16 @@ public class AiGenerationTaskServiceImpl implements IAiGenerationTaskService @Transactional public AiGenerationTaskStatusResponse createTask(Long userId, Long projectId, AiGenerationTaskCreateRequest request) { - FrontProject project = assertOwnedProject(userId, projectId); - validateRequest(request); + validateGenerateType(request); + FrontProject project = lockOwnedProject(userId, projectId); + normalizeOneClickRequest(project, request); + validateRequestContent(request); releaseExpiredRunningTasks(); String requestPayload = JSON.toJSONString(request); + if (requestPayload.getBytes(StandardCharsets.UTF_8).length > MAX_TASK_REQUEST_BYTES) + { + throw new ServiceException("生成请求内容过大"); + } AiInvocationPlan invocation = invocationPlan(request.getGenerateType()); AiStageManifest stageManifest = stageManifest(request.getGenerateType()); String stageManifestJson = stageManifest == null ? null : aiStageManifestService.canonical(stageManifest); @@ -135,6 +145,13 @@ public class AiGenerationTaskServiceImpl implements IAiGenerationTaskService public AiGenerationTaskStatusResponse retryTask(Long userId, Long projectId, Long taskId) { AiGenerationTask task = selectOwnedTask(userId, projectId, taskId); + if (frontProjectMapper.lockFrontProjectByUserAndId(userId, projectId) == null) + { + throw new ServiceException("项目不存在或无权访问"); + } + // Re-read after acquiring the project lock so concurrent retries cannot + // reserve quota and dispatch the same task more than once. + task = selectOwnedTask(userId, projectId, taskId); if (!"FAILED".equals(task.getStatus()) && !"RETRY_WAITING".equals(task.getStatus())) { throw new ServiceException("当前任务状态不允许重试"); @@ -158,8 +175,12 @@ public class AiGenerationTaskServiceImpl implements IAiGenerationTaskService task.setErrorMessage(""); task.setNextRetryTime(new Date()); clearTaskLock(task); - aiGenerationTaskMapper.updateAiGenerationTask(task); - aiGenerationTaskMapper.clearTaskLock(taskId); + task.setStartedAt(null); + task.setFinishedAt(null); + if (aiGenerationTaskMapper.retryFailedTask(task) == 0) + { + throw new ServiceException("当前任务状态不允许重试"); + } dispatchAfterCommit(taskId); return toResponse(task); } @@ -173,11 +194,15 @@ public class AiGenerationTaskServiceImpl implements IAiGenerationTaskService { throw new ServiceException("当前任务状态不允许取消"); } + if (aiGenerationTaskMapper.cancelPendingTask(userId, projectId, taskId) == 0) + { + throw new ServiceException("当前任务状态不允许取消"); + } task.setStatus("CANCELED"); task.setProgress(0); task.setCurrentStep("已取消"); task.setFinishedAt(new Date()); - aiGenerationTaskMapper.updateAiGenerationTask(task); + clearTaskLock(task); aiQuotaService.releaseRunning(userId); return toResponse(task); } @@ -279,7 +304,7 @@ public class AiGenerationTaskServiceImpl implements IAiGenerationTaskService return task; } - private void validateRequest(AiGenerationTaskCreateRequest request) + private void validateGenerateType(AiGenerationTaskCreateRequest request) { if (request == null || StringUtils.isEmpty(request.getGenerateType())) { @@ -298,6 +323,59 @@ public class AiGenerationTaskServiceImpl implements IAiGenerationTaskService } } + private void normalizeOneClickRequest(FrontProject project, AiGenerationTaskCreateRequest request) + { + if (project == null || request == null || !"one_click_project".equals(request.getGenerateType())) + { + return; + } + request.setProjectName(project.getProjectName()); + request.setProjectDesc(project.getProjectDesc()); + request.setIndustryTemplate(project.getIndustryTemplate()); + request.setCodeTemplate(project.getCodeTemplate()); + request.setStylePreset(project.getStylePreset()); + } + + private void validateRequestContent(AiGenerationTaskCreateRequest request) + { + if (!"one_click_project".equals(request.getGenerateType())) + { + return; + } + requireText(request.getProjectName(), "项目名称不能为空", MAX_PROJECT_NAME_LENGTH, "项目名称过长"); + requireText(request.getStylePreset(), "视觉方案不能为空", MAX_OPTION_CODE_LENGTH, "视觉方案编码过长"); + requireText(request.getCodeTemplate(), "代码模板不能为空", MAX_OPTION_CODE_LENGTH, "代码模板编码过长"); + if (StringUtils.defaultString(request.getProjectDesc()).length() > MAX_PROJECT_DESCRIPTION_LENGTH) + { + throw new ServiceException("项目描述过长"); + } + } + + private void requireText(String value, String emptyMessage, int maxLength, String longMessage) + { + if (StringUtils.isBlank(value)) + { + throw new ServiceException(emptyMessage); + } + if (value.trim().length() > maxLength) + { + throw new ServiceException(longMessage); + } + } + + private FrontProject lockOwnedProject(Long userId, Long projectId) + { + FrontProject project = assertOwnedProject(userId, projectId); + // Serialize task creation/retry per project without relying on an in-process lock. + // The database row lock also works when multiple application instances are running. + FrontProject lockedProject = frontProjectMapper.lockFrontProjectByUserAndId(userId, projectId); + if (lockedProject == null) + { + throw new ServiceException("项目不存在或无权访问"); + } + return lockedProject; + } + private Long createGenerationRecord(FrontProject project, Long userId, String generateType, String requestPayload, AiInvocationPlan invocation) { diff --git a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/AiGenerationTaskWorker.java b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/AiGenerationTaskWorker.java index ca7590e..1f0205c 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/AiGenerationTaskWorker.java +++ b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/AiGenerationTaskWorker.java @@ -106,11 +106,16 @@ public class AiGenerationTaskWorker task.setProgress(initialProgress(task.getGenerateType())); task.setCurrentStep(initialCurrentStep(task.getGenerateType())); task.setStartedAt(new Date()); - aiGenerationTaskMapper.updateAiGenerationTask(task); + if (aiGenerationTaskMapper.updateAiGenerationTask(task) == 0) + { + log.warn("AI generation task lease was lost before execution, taskId={}", taskId); + return; + } renewTaskLock(task); AiUsageScope usageScope = taskUsageScope(task); ScheduledFuture heartbeat = startHeartbeat(task); + boolean finalized = false; try { Object result = executeWithUsageScope(task, usageScope); @@ -124,50 +129,74 @@ public class AiGenerationTaskWorker applyUsage(task, usageScope); int totalTokens = safe(task.getTotalTokens()); int costCents = safe(task.getCostCents()); + if (!finishClaimedTask(task, lockedBy)) + { + return; + } + finalized = true; updateGenerationRecord(task, "1"); - clearTaskLock(task); - aiGenerationTaskMapper.updateAiGenerationTask(task); - clearTaskLockInStore(task); insertEstimatedLedgerIfNeeded(task, usageScope, "1"); aiQuotaService.settleCost(task.getUserId(), totalTokens, costCents); } catch (RuntimeException e) { - String message = sanitize(e.getMessage()); - task.setErrorMessage(message); - task.setProgress(0); - applyMeasuredUsage(task, usageScope); - preserveOneClickFailureResult(task, e); - if (retryPolicy.canRetry(message, attempts, safeMax(task.getMaxAttempts()))) + if (finalized) { - task.setStatus("RETRY_WAITING"); - task.setCurrentStep("等待重试"); - task.setNextRetryTime(retryPolicy.nextRetryTime(attempts, new Date())); + log.error("AI generation task post-finalization processing failed, taskId={}", taskId, e); } else { - task.setStatus("FAILED"); - task.setCurrentStep("生成失败"); - task.setFinishedAt(new Date()); - updateGenerationRecord(task, "0"); - } - clearTaskLock(task); - aiGenerationTaskMapper.updateAiGenerationTask(task); - clearTaskLockInStore(task); - insertEstimatedLedgerIfNeeded(task, usageScope, "0"); - if (safe(task.getTotalTokens()) > 0 || safe(task.getCostCents()) > 0) - { - aiQuotaService.settleCost(task.getUserId(), safe(task.getTotalTokens()), safe(task.getCostCents())); + String message = sanitize(e.getMessage()); + task.setErrorMessage(message); + task.setProgress(0); + applyMeasuredUsage(task, usageScope); + preserveOneClickFailureResult(task, e); + if (retryPolicy.canRetry(message, attempts, safeMax(task.getMaxAttempts()))) + { + task.setStatus("RETRY_WAITING"); + task.setCurrentStep("等待重试"); + task.setNextRetryTime(retryPolicy.nextRetryTime(attempts, new Date())); + } + else + { + task.setStatus("FAILED"); + task.setCurrentStep("生成失败"); + task.setFinishedAt(new Date()); + } + if (!finishClaimedTask(task, lockedBy)) + { + return; + } + finalized = true; + if ("FAILED".equals(task.getStatus())) + { + updateGenerationRecord(task, "0"); + } + insertEstimatedLedgerIfNeeded(task, usageScope, "0"); + if (safe(task.getTotalTokens()) > 0 || safe(task.getCostCents()) > 0) + { + aiQuotaService.settleCost(task.getUserId(), safe(task.getTotalTokens()), safe(task.getCostCents())); + } } } catch (Throwable e) { - handleFailure(task, usageScope, attempts, e); + if (finalized) + { + log.error("AI generation task post-finalization processing failed, taskId={}", taskId, e); + } + else + { + finalized = handleFailure(task, usageScope, attempts, lockedBy, e); + } } finally { cancelHeartbeat(heartbeat); - aiQuotaService.releaseRunning(task.getUserId()); + if (finalized) + { + aiQuotaService.releaseRunning(task.getUserId()); + } } } @@ -188,7 +217,8 @@ public class AiGenerationTaskWorker return count; } - private void handleFailure(AiGenerationTask task, AiUsageScope usageScope, int attempts, Throwable error) + private boolean handleFailure(AiGenerationTask task, AiUsageScope usageScope, int attempts, + String lockedBy, Throwable error) { String message = sanitize(error); task.setErrorMessage(message); @@ -210,11 +240,15 @@ public class AiGenerationTaskWorker task.setStatus("FAILED"); task.setCurrentStep("Generation failed"); task.setFinishedAt(new Date()); + } + if (!finishClaimedTask(task, lockedBy)) + { + return false; + } + if ("FAILED".equals(task.getStatus())) + { updateGenerationRecord(task, "0"); } - clearTaskLock(task); - aiGenerationTaskMapper.updateAiGenerationTask(task); - clearTaskLockInStore(task); insertEstimatedLedgerIfNeeded(task, usageScope, "0"); if (safe(task.getTotalTokens()) > 0 || safe(task.getCostCents()) > 0) { @@ -224,6 +258,7 @@ public class AiGenerationTaskWorker { log.error("AI generation task failed with non-runtime error, taskId={}", task.getTaskId(), error); } + return true; } private void releaseExpiredRunningTasks() @@ -566,23 +601,21 @@ public class AiGenerationTaskWorker generationRecordService.updateTaskResult(task, success); } - private void clearTaskLock(AiGenerationTask task) + private boolean finishClaimedTask(AiGenerationTask task, String lockedBy) { - if (task == null) + if (task == null || task.getTaskId() == null || StringUtils.isBlank(lockedBy)) { - return; + return false; + } + if (aiGenerationTaskMapper.finishClaimedTask(task, lockedBy) == 0) + { + log.warn("Ignored stale AI generation task completion, taskId={}, lockedBy={}", + task.getTaskId(), lockedBy); + return false; } task.setLockedBy(null); task.setLockedUntil(null); - } - - private void clearTaskLockInStore(AiGenerationTask task) - { - if (task == null || task.getTaskId() == null) - { - return; - } - aiGenerationTaskMapper.clearTaskLock(task.getTaskId()); + return true; } private GenerateAppBlueprintRequest toAppBlueprintRequest(AiGenerationTaskCreateRequest source) diff --git a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/AiQuotaService.java b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/AiQuotaService.java index 586bb1e..b0da193 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/AiQuotaService.java +++ b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/AiQuotaService.java @@ -103,7 +103,7 @@ public class AiQuotaService private AiQuotaBucket ensureBucket(Long userId, String periodType, String periodKey) { - AiQuotaBucket bucket = aiQuotaBucketMapper.selectQuotaBucket(userId, periodType, periodKey); + AiQuotaBucket bucket = aiQuotaBucketMapper.selectQuotaBucketForUpdate(userId, periodType, periodKey); if (bucket != null) { return bucket; @@ -121,7 +121,8 @@ public class AiQuotaService bucket.setRunningLimit("DAY".equals(periodType) ? RUNNING_LIMIT : 0); bucket.setRunningCount(0); aiQuotaBucketMapper.insertQuotaBucket(bucket); - return bucket; + AiQuotaBucket locked = aiQuotaBucketMapper.selectQuotaBucketForUpdate(userId, periodType, periodKey); + return locked == null ? bucket : locked; } private void reconcileRunningCount(Long userId, AiQuotaBucket day) diff --git a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/FrontProjectPreviewServiceImpl.java b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/FrontProjectPreviewServiceImpl.java index 9b5b251..e83ddb9 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/FrontProjectPreviewServiceImpl.java +++ b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/FrontProjectPreviewServiceImpl.java @@ -17,6 +17,7 @@ import com.ruoyi.generator.service.IGenProjectService; import com.ruoyi.generator.service.ITemplateBundleService; import com.ruoyi.common.exception.ServiceException; import com.ruoyi.common.utils.StringUtils; +import com.ruoyi.generator.util.ZipEntryPathValidator; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.stereotype.Service; @@ -376,15 +377,6 @@ public class FrontProjectPreviewServiceImpl implements IFrontProjectPreviewServi { return null; } - String normalized = entryName.replace('\\', '/'); - while (normalized.startsWith("/")) - { - normalized = normalized.substring(1); - } - if (normalized.startsWith("../") || normalized.contains("/../")) - { - throw new ServiceException("项目源码压缩包路径不合法"); - } - return normalized; + return ZipEntryPathValidator.requireRelative(entryName, "项目源码压缩包"); } } diff --git a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/FrontProjectServiceImpl.java b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/FrontProjectServiceImpl.java index f066cb0..9fffeb9 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/FrontProjectServiceImpl.java +++ b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/service/front/FrontProjectServiceImpl.java @@ -53,6 +53,7 @@ public class FrontProjectServiceImpl implements IFrontProjectService { private static final Pattern SAFE_PROJECT_FILE_NAME = Pattern.compile("^[a-z][a-z0-9-]*[a-z0-9]$|^[a-z]$"); private static final Pattern DB_NAME_PATTERN = Pattern.compile("^[a-z][a-z0-9_]{1,63}$"); + private static final Pattern GENERATOR_IDENTIFIER_PATTERN = Pattern.compile("^[A-Za-z][A-Za-z0-9_]{0,63}$"); private static final Pattern BIGINT_TYPE_PATTERN = Pattern.compile("^bigint(\\(20\\))?$"); private static final Pattern INT_TYPE_PATTERN = Pattern.compile("^int(\\(11\\))?$"); private static final Pattern VARCHAR_TYPE_PATTERN = Pattern.compile("^varchar\\((\\d{1,5})\\)$"); @@ -592,6 +593,15 @@ public class FrontProjectServiceImpl implements IFrontProjectService { throw new ServiceException("Table name duplicated: " + table.getTableName()); } + if (StringUtils.isNotBlank(table.getModuleName()) + && !GENERATOR_IDENTIFIER_PATTERN.matcher(table.getModuleName()).matches()) + { + throw new ServiceException("模块名只能包含字母、数字和下划线,且必须以字母开头"); + } + if (!GENERATOR_IDENTIFIER_PATTERN.matcher(table.getBusinessName()).matches()) + { + throw new ServiceException("业务名只能包含字母、数字和下划线,且必须以字母开头"); + } if (table.getColumns() == null || table.getColumns().isEmpty()) { throw new ServiceException("每张表至少需要一个字段"); @@ -660,7 +670,8 @@ public class FrontProjectServiceImpl implements IFrontProjectService } table.setTableName(trimToNull(table.getTableName())); table.setClassName(StringUtils.convertToCamelCase(table.getTableName())); - table.setBusinessName(StringUtils.defaultIfEmpty(table.getBusinessName(), table.getTableName())); + table.setModuleName(trimToNull(table.getModuleName())); + table.setBusinessName(StringUtils.defaultIfEmpty(trimToNull(table.getBusinessName()), table.getTableName())); table.setFunctionName(StringUtils.defaultIfEmpty(table.getFunctionName(), StringUtils.defaultIfEmpty(table.getTableComment(), table.getTableName()))); } diff --git a/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/util/ZipEntryPathValidator.java b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/util/ZipEntryPathValidator.java new file mode 100644 index 0000000..4061211 --- /dev/null +++ b/RuoYi-Vue/ruoyi-generator/src/main/java/com/ruoyi/generator/util/ZipEntryPathValidator.java @@ -0,0 +1,88 @@ +package com.ruoyi.generator.util; + +import java.util.Locale; +import com.ruoyi.common.exception.ServiceException; +import com.ruoyi.common.utils.StringUtils; + +/** Validates generated ZIP entry paths for portable, relative extraction. */ +public final class ZipEntryPathValidator +{ + private static final int MAX_ENTRY_PATH_LENGTH = 1024; + private static final String WINDOWS_RESERVED_CHARACTERS = "<>:\"|?*"; + + private ZipEntryPathValidator() + { + } + + public static String requireRelative(String entryName, String sourceLabel) + { + String normalized = StringUtils.defaultString(entryName).replace('\\', '/'); + if (normalized.length() == 0 || normalized.length() > MAX_ENTRY_PATH_LENGTH + || normalized.startsWith("/") || hasControlCharacter(normalized)) + { + throw invalid(sourceLabel, entryName); + } + boolean directory = normalized.endsWith("/"); + String path = directory ? normalized.substring(0, normalized.length() - 1) : normalized; + if (path.length() == 0) + { + throw invalid(sourceLabel, entryName); + } + String[] segments = path.split("/", -1); + for (String segment : segments) + { + if (segment.length() == 0 || ".".equals(segment) || "..".equals(segment) + || containsWindowsReservedCharacter(segment) || segment.endsWith(" ") + || segment.endsWith(".") || isWindowsDeviceName(segment)) + { + throw invalid(sourceLabel, entryName); + } + } + return directory ? path + "/" : path; + } + + private static boolean hasControlCharacter(String value) + { + for (int i = 0; i < value.length(); i++) + { + if (Character.isISOControl(value.charAt(i))) + { + return true; + } + } + return false; + } + + private static boolean isWindowsDeviceName(String segment) + { + String upper = segment.toUpperCase(Locale.ROOT); + int extension = upper.indexOf('.'); + String baseName = extension < 0 ? upper : upper.substring(0, extension); + if ("CON".equals(baseName) || "PRN".equals(baseName) || "AUX".equals(baseName) + || "NUL".equals(baseName)) + { + return true; + } + return baseName.matches("COM[1-9]") || baseName.matches("LPT[1-9]"); + } + + private static boolean containsWindowsReservedCharacter(String segment) + { + for (int i = 0; i < segment.length(); i++) + { + if (WINDOWS_RESERVED_CHARACTERS.indexOf(segment.charAt(i)) >= 0) + { + return true; + } + } + return false; + } + + private static ServiceException invalid(String sourceLabel, String entryName) + { + String value = StringUtils.defaultString(entryName); + String display = value.length() <= 160 ? value : value.substring(0, 160) + "..."; + return new ServiceException(StringUtils.defaultIfEmpty(sourceLabel, "ZIP entry") + + " contains an invalid path: " + display); + } +} diff --git a/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/factory/AiTaskStageMapper.xml b/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/factory/AiTaskStageMapper.xml index a8390e4..ab7db31 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/factory/AiTaskStageMapper.xml +++ b/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/factory/AiTaskStageMapper.xml @@ -53,6 +53,7 @@ set status = 'SUCCEEDED', finished_at = #{finishedAt}, duration_millis = #{durationMillis}, error_message = '', update_time = now() where task_id = #{taskId} and stage_code = #{stageCode} and attempt_no = #{attemptNo} + and status = 'RUNNING' @@ -60,6 +61,7 @@ set status = 'FAILED', finished_at = #{finishedAt}, duration_millis = #{durationMillis}, error_message = #{errorMessage}, update_time = now() where task_id = #{taskId} and stage_code = #{stageCode} and attempt_no = #{attemptNo} + and status = 'RUNNING' diff --git a/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/front/AiGenerationTaskMapper.xml b/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/front/AiGenerationTaskMapper.xml index 6ee9f19..a272392 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/front/AiGenerationTaskMapper.xml +++ b/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/front/AiGenerationTaskMapper.xml @@ -275,6 +275,73 @@ PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" update_time = sysdate(), where task_id = #{taskId} + + and status = 'RUNNING' + and locked_by = #{lockedBy} + + + + + update front_ai_generation_task + + generation_id = #{task.generationId}, + status = #{task.status}, + result_payload = #{task.resultPayload}, + error_code = #{task.errorCode}, + error_message = #{task.errorMessage}, + attempts = #{task.attempts}, + max_attempts = #{task.maxAttempts}, + next_retry_time = #{task.nextRetryTime}, + progress = #{task.progress}, + current_step = #{task.currentStep}, + input_tokens = #{task.inputTokens}, + output_tokens = #{task.outputTokens}, + total_tokens = #{task.totalTokens}, + cost_cents = #{task.costCents}, + started_at = #{task.startedAt}, + finished_at = #{task.finishedAt}, + locked_by = null, + locked_until = null, + update_time = sysdate(), + + where task_id = #{task.taskId} + and status = 'RUNNING' + and locked_by = #{lockedBy} + + + + update front_ai_generation_task + set status = 'CANCELED', + progress = 0, + current_step = '已取消', + locked_by = null, + locked_until = null, + finished_at = sysdate(), + update_time = sysdate() + where task_id = #{taskId} + and user_id = #{userId} + and project_id = #{projectId} + and status in ('QUEUED', 'RETRY_WAITING') + + + + update front_ai_generation_task + set status = 'QUEUED', + progress = 0, + current_step = #{task.currentStep}, + result_payload = #{task.resultPayload}, + error_code = '', + error_message = '', + next_retry_time = #{task.nextRetryTime}, + started_at = null, + finished_at = null, + locked_by = null, + locked_until = null, + update_time = sysdate() + where task_id = #{task.taskId} + and user_id = #{task.userId} + and project_id = #{task.projectId} + and status in ('FAILED', 'RETRY_WAITING') diff --git a/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/front/AiQuotaBucketMapper.xml b/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/front/AiQuotaBucketMapper.xml index c7156e2..5a2ed44 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/front/AiQuotaBucketMapper.xml +++ b/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/front/AiQuotaBucketMapper.xml @@ -28,6 +28,14 @@ PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" where user_id = #{userId} and period_type = #{periodType} and period_key = #{periodKey} + + insert into front_ai_quota_bucket (user_id, period_type, period_key, task_limit, task_used, token_limit, token_used, @@ -35,6 +43,7 @@ PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" values (#{userId}, #{periodType}, #{periodKey}, #{taskLimit}, #{taskUsed}, #{tokenLimit}, #{tokenUsed}, #{costLimitCents}, #{costUsedCents}, #{runningLimit}, #{runningCount}, sysdate(), sysdate()) + on duplicate key update bucket_id = last_insert_id(bucket_id) diff --git a/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/front/FrontProjectMapper.xml b/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/front/FrontProjectMapper.xml index 27dec67..50521b8 100644 --- a/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/front/FrontProjectMapper.xml +++ b/RuoYi-Vue/ruoyi-generator/src/main/resources/mapper/front/FrontProjectMapper.xml @@ -47,6 +47,12 @@ PUBLIC "-//mybatis.org//DTD Mapper 3.0//EN" where user_id = #{userId} and project_id = #{projectId} + + = 0 && end > start); + return xml.substring(start, end); + } +} diff --git a/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/AiGenerationTaskServiceImplTest.java b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/AiGenerationTaskServiceImplTest.java index edf9af2..4142bee 100644 --- a/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/AiGenerationTaskServiceImplTest.java +++ b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/AiGenerationTaskServiceImplTest.java @@ -94,6 +94,9 @@ public class AiGenerationTaskServiceImplTest setField("deepSeekProperties", deepSeekProperties); when(aiInvocationPlanner.plan(any(String.class))).thenAnswer(invocation -> invocationPlan(invocation.getArgument(0))); + when(frontProjectMapper.lockFrontProjectByUserAndId(anyLong(), anyLong())).thenAnswer(invocation -> + Long.valueOf(20L).equals(invocation.getArgument(1)) ? project20() : project()); + when(aiGenerationTaskMapper.retryFailedTask(any(AiGenerationTask.class))).thenReturn(1); } @Test @@ -368,8 +371,10 @@ public class AiGenerationTaskServiceImplTest AiGenerationTaskCreateRequest request = new AiGenerationTaskCreateRequest(); request.setGenerateType("one_click_project"); - request.setProjectName("客户关系管理系统"); - request.setProjectDesc(""); + request.setProjectName("客户端伪造名称"); + request.setProjectDesc("客户端伪造描述"); + request.setCodeTemplate("client-template"); + request.setStylePreset("client-style"); request.setExtraRequirements(""); AiGenerationTaskStatusResponse response = service.createTask(10L, 20L, request); @@ -383,6 +388,11 @@ public class AiGenerationTaskServiceImplTest assertEquals(Long.valueOf(10L), task.getUserId()); assertEquals(Long.valueOf(20L), task.getProjectId()); assertEquals("one_click_project", task.getGenerateType()); + assertTrue(task.getRequestPayload().contains("客户关系管理系统")); + assertTrue(task.getRequestPayload().contains("管理客户资料")); + assertTrue(task.getRequestPayload().contains("qing")); + assertTrue(task.getRequestPayload().contains("dark-tech")); + assertFalse(task.getRequestPayload().contains("客户端伪造")); assertNotNull(task.getStageManifestJson()); assertEquals(64, task.getStageManifestHash().length()); assertTrue(task.getStageManifestJson().contains("front.flow_config")); @@ -418,6 +428,50 @@ public class AiGenerationTaskServiceImplTest assertEquals("QUEUED", response.getStatus()); } + @Test + public void createOneClickTaskRejectsBlankPersistedTemplateBeforeQuotaReservation() + { + FrontProject persisted = project(); + persisted.setCodeTemplate(" "); + when(frontProjectMapper.selectFrontProjectByUserAndId(7L, 10L)).thenReturn(persisted); + when(frontProjectMapper.lockFrontProjectByUserAndId(7L, 10L)).thenReturn(persisted); + AiGenerationTaskCreateRequest request = request("one_click_project"); + + try + { + service.createTask(7L, 10L, request); + } + catch (ServiceException error) + { + assertTrue(error.getMessage().contains("代码模板不能为空")); + verify(aiQuotaService, never()).reserve(anyLong()); + verify(aiGenerationTaskMapper, never()).insertAiGenerationTask(any(AiGenerationTask.class)); + return; + } + throw new AssertionError("Expected blank persisted template to be rejected"); + } + + @Test + public void createTaskRejectsOversizedSerializedPayloadBeforeQuotaReservation() + { + when(frontProjectMapper.selectFrontProjectByUserAndId(7L, 10L)).thenReturn(project()); + AiGenerationTaskCreateRequest request = request("code_analysis"); + request.setPreviousMarkdown(repeat('x', 1024 * 1024 + 1)); + + try + { + service.createTask(7L, 10L, request); + } + catch (ServiceException error) + { + assertTrue(error.getMessage().contains("生成请求内容过大")); + verify(aiQuotaService, never()).reserve(anyLong()); + verify(aiGenerationTaskMapper, never()).insertAiGenerationTask(any(AiGenerationTask.class)); + return; + } + throw new AssertionError("Expected oversized task payload to be rejected"); + } + @Test public void createTaskDefersDispatchUntilTransactionCommit() { @@ -550,7 +604,7 @@ public class AiGenerationTaskServiceImplTest AiGenerationTaskStatusResponse response = service.retryTask(7L, 10L, 99L); ArgumentCaptor taskCaptor = ArgumentCaptor.forClass(AiGenerationTask.class); - verify(aiGenerationTaskMapper).updateAiGenerationTask(taskCaptor.capture()); + verify(aiGenerationTaskMapper).retryFailedTask(taskCaptor.capture()); AiGenerationTask retriedTask = taskCaptor.getValue(); assertEquals("QUEUED", retriedTask.getStatus()); assertEquals(Integer.valueOf(0), retriedTask.getProgress()); @@ -588,12 +642,69 @@ public class AiGenerationTaskServiceImplTest { assertTrue(error.getMessage().contains("已有同类型任务正在生成")); verify(aiQuotaService, never()).reserve(anyLong()); - verify(aiGenerationTaskMapper, never()).updateAiGenerationTask(any(AiGenerationTask.class)); + verify(aiGenerationTaskMapper, never()).retryFailedTask(any(AiGenerationTask.class)); return; } throw new AssertionError("Expected retry conflict to be rejected"); } + @Test + public void retryTaskRechecksStatusAfterProjectLockBeforeReservingQuota() + { + AiGenerationTask initiallyFailed = new AiGenerationTask(); + initiallyFailed.setTaskId(99L); + initiallyFailed.setProjectId(10L); + initiallyFailed.setUserId(7L); + initiallyFailed.setGenerateType("one_click_project"); + initiallyFailed.setStatus("FAILED"); + AiGenerationTask concurrentlyQueued = new AiGenerationTask(); + concurrentlyQueued.setTaskId(99L); + concurrentlyQueued.setProjectId(10L); + concurrentlyQueued.setUserId(7L); + concurrentlyQueued.setGenerateType("one_click_project"); + concurrentlyQueued.setStatus("QUEUED"); + when(aiGenerationTaskMapper.selectTaskForUser(7L, 10L, 99L)) + .thenReturn(initiallyFailed, concurrentlyQueued); + + try + { + service.retryTask(7L, 10L, 99L); + } + catch (ServiceException error) + { + assertTrue(error.getMessage().contains("当前任务状态不允许重试")); + verify(aiQuotaService, never()).reserve(anyLong()); + verify(aiGenerationTaskMapper, never()).retryFailedTask(any(AiGenerationTask.class)); + return; + } + throw new AssertionError("Expected concurrent retry to be rejected"); + } + + @Test + public void retryTaskRejectsWorkerClaimThatWinsAtomicTransition() + { + AiGenerationTask task = new AiGenerationTask(); + task.setTaskId(99L); + task.setProjectId(10L); + task.setUserId(7L); + task.setGenerateType("database"); + task.setStatus("RETRY_WAITING"); + when(aiGenerationTaskMapper.selectTaskForUser(7L, 10L, 99L)).thenReturn(task); + when(aiGenerationTaskMapper.retryFailedTask(any(AiGenerationTask.class))).thenReturn(0); + + try + { + service.retryTask(7L, 10L, 99L); + } + catch (ServiceException error) + { + assertTrue(error.getMessage().contains("当前任务状态不允许重试")); + verify(aiQuotaService).reserve(7L); + return; + } + throw new AssertionError("Expected worker claim race to be rejected"); + } + @Test public void retryTaskClearsStaleWorkerLockBeforeQueueing() { @@ -614,7 +725,7 @@ public class AiGenerationTaskServiceImplTest service.retryTask(7L, 10L, 99L); ArgumentCaptor taskCaptor = ArgumentCaptor.forClass(AiGenerationTask.class); - verify(aiGenerationTaskMapper).updateAiGenerationTask(taskCaptor.capture()); + verify(aiGenerationTaskMapper).retryFailedTask(taskCaptor.capture()); AiGenerationTask retriedTask = taskCaptor.getValue(); assertEquals("QUEUED", retriedTask.getStatus()); assertNull(retriedTask.getLockedBy()); @@ -639,13 +750,59 @@ public class AiGenerationTaskServiceImplTest AiGenerationTaskStatusResponse response = service.retryTask(7L, 10L, 99L); ArgumentCaptor taskCaptor = ArgumentCaptor.forClass(AiGenerationTask.class); - verify(aiGenerationTaskMapper).updateAiGenerationTask(taskCaptor.capture()); + verify(aiGenerationTaskMapper).retryFailedTask(taskCaptor.capture()); AiGenerationTask retriedTask = taskCaptor.getValue(); assertEquals("QUEUED", retriedTask.getStatus()); assertEquals("", retriedTask.getResultPayload()); assertEquals("", response.getResultPayload()); } + @Test + public void cancelTaskUsesAtomicPendingTransitionBeforeReleasingQuota() + { + AiGenerationTask task = new AiGenerationTask(); + task.setTaskId(99L); + task.setProjectId(10L); + task.setUserId(7L); + task.setGenerateType("database"); + task.setStatus("QUEUED"); + task.setLockedBy("stale-worker"); + when(aiGenerationTaskMapper.selectTaskForUser(7L, 10L, 99L)).thenReturn(task); + when(aiGenerationTaskMapper.cancelPendingTask(7L, 10L, 99L)).thenReturn(1); + + AiGenerationTaskStatusResponse response = service.cancelTask(7L, 10L, 99L); + + assertEquals("CANCELED", response.getStatus()); + assertNull(task.getLockedBy()); + verify(aiQuotaService).releaseRunning(7L); + verify(aiGenerationTaskMapper, never()).updateAiGenerationTask(any(AiGenerationTask.class)); + } + + @Test + public void cancelTaskDoesNotReleaseQuotaWhenConcurrentTransitionWins() + { + AiGenerationTask task = new AiGenerationTask(); + task.setTaskId(99L); + task.setProjectId(10L); + task.setUserId(7L); + task.setGenerateType("database"); + task.setStatus("QUEUED"); + when(aiGenerationTaskMapper.selectTaskForUser(7L, 10L, 99L)).thenReturn(task); + when(aiGenerationTaskMapper.cancelPendingTask(7L, 10L, 99L)).thenReturn(0); + + try + { + service.cancelTask(7L, 10L, 99L); + } + catch (ServiceException error) + { + assertTrue(error.getMessage().contains("当前任务状态不允许取消")); + verify(aiQuotaService, never()).releaseRunning(anyLong()); + return; + } + throw new AssertionError("Expected concurrent cancellation conflict"); + } + private AiGenerationTaskCreateRequest request(String generateType) { AiGenerationTaskCreateRequest request = new AiGenerationTaskCreateRequest(); @@ -658,12 +815,22 @@ public class AiGenerationTaskServiceImplTest return request; } + private String repeat(char value, int count) + { + char[] result = new char[count]; + Arrays.fill(result, value); + return new String(result); + } + private FrontProject project() { FrontProject project = new FrontProject(); project.setProjectId(10L); project.setUserId(7L); project.setProjectName("客户中心"); + project.setProjectDesc("管理客户资料"); + project.setCodeTemplate("qing"); + project.setStylePreset("dark-tech"); return project; } @@ -693,6 +860,9 @@ public class AiGenerationTaskServiceImplTest project.setProjectId(20L); project.setUserId(10L); project.setProjectName("客户关系管理系统"); + project.setProjectDesc("管理客户资料"); + project.setCodeTemplate("qing"); + project.setStylePreset("dark-tech"); return project; } } diff --git a/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/AiGenerationTaskWorkerTest.java b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/AiGenerationTaskWorkerTest.java index 12af0b9..d9f728d 100644 --- a/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/AiGenerationTaskWorkerTest.java +++ b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/AiGenerationTaskWorkerTest.java @@ -3,11 +3,12 @@ package com.ruoyi.generator.service.front; import static org.junit.Assert.assertEquals; import static org.junit.Assert.assertTrue; import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.anyInt; import static org.mockito.ArgumentMatchers.anyString; import static org.mockito.ArgumentMatchers.eq; -import static org.mockito.Mockito.atLeast; import static org.mockito.Mockito.doThrow; import static org.mockito.Mockito.inOrder; +import static org.mockito.Mockito.never; import static org.mockito.Mockito.verify; import static org.mockito.Mockito.when; @@ -52,6 +53,8 @@ public class AiGenerationTaskWorkerTest setField("aiQuotaService", aiQuotaService); setField("retryPolicy", retryPolicy); setField("costCalculator", costCalculator); + when(taskMapper.updateAiGenerationTask(any(AiGenerationTask.class))).thenReturn(1); + when(taskMapper.finishClaimedTask(any(AiGenerationTask.class), anyString())).thenReturn(1); } @Test @@ -66,13 +69,57 @@ public class AiGenerationTaskWorkerTest worker.processTask(99L); ArgumentCaptor taskCaptor = ArgumentCaptor.forClass(AiGenerationTask.class); - verify(taskMapper, atLeast(2)).updateAiGenerationTask(taskCaptor.capture()); - AiGenerationTask failed = taskCaptor.getAllValues().get(taskCaptor.getAllValues().size() - 1); + verify(taskMapper).finishClaimedTask(taskCaptor.capture(), anyString()); + AiGenerationTask failed = taskCaptor.getValue(); assertEquals("FAILED", failed.getStatus()); assertEquals("Generation failed", failed.getCurrentStep()); assertTrue(failed.getErrorMessage().contains("AssertionError")); verify(taskMapper).renewTaskLock(eq(99L), anyString()); - verify(taskMapper).clearTaskLock(99L); + verify(aiQuotaService).releaseRunning(7L); + } + + @Test + public void staleWorkerCannotFinalizeOrReleaseQuotaAfterLeaseIsLost() + { + AiGenerationTask task = task(); + when(taskMapper.selectAiGenerationTaskById(99L)).thenReturn(task); + when(taskMapper.claimTask(eq(99L), anyString())).thenReturn(1); + when(taskMapper.finishClaimedTask(any(AiGenerationTask.class), anyString())).thenReturn(0); + + worker.processTask(99L); + + verify(taskMapper).finishClaimedTask(any(AiGenerationTask.class), anyString()); + verify(aiQuotaService, never()).settleCost(any(Long.class), anyInt(), anyInt()); + verify(aiQuotaService, never()).releaseRunning(7L); + } + + @Test + public void postFinalizationFailureStillReleasesRunningQuota() + { + AiGenerationTask task = task(); + when(taskMapper.selectAiGenerationTaskById(99L)).thenReturn(task); + when(taskMapper.claimTask(eq(99L), anyString())).thenReturn(1); + doThrow(new AssertionError("settlement failed")).when(aiQuotaService) + .settleCost(eq(7L), anyInt(), anyInt()); + + worker.processTask(99L); + + verify(taskMapper).finishClaimedTask(any(AiGenerationTask.class), anyString()); + verify(aiQuotaService).releaseRunning(7L); + } + + @Test + public void postFinalizationRuntimeFailureDoesNotRefinalizeTask() + { + AiGenerationTask task = task(); + when(taskMapper.selectAiGenerationTaskById(99L)).thenReturn(task); + when(taskMapper.claimTask(eq(99L), anyString())).thenReturn(1); + doThrow(new IllegalStateException("settlement failed")).when(aiQuotaService) + .settleCost(eq(7L), anyInt(), anyInt()); + + worker.processTask(99L); + + verify(taskMapper).finishClaimedTask(any(AiGenerationTask.class), anyString()); verify(aiQuotaService).releaseRunning(7L); } diff --git a/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/AiQuotaServiceTest.java b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/AiQuotaServiceTest.java index a945d65..87c1bdb 100644 --- a/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/AiQuotaServiceTest.java +++ b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/AiQuotaServiceTest.java @@ -39,9 +39,9 @@ public class AiQuotaServiceTest AiQuotaBucket day = bucket("DAY", "20260703"); day.setRunningCount(1); AiQuotaBucket month = bucket("MONTH", "202607"); - when(aiQuotaBucketMapper.selectQuotaBucket(org.mockito.ArgumentMatchers.eq(7L), + when(aiQuotaBucketMapper.selectQuotaBucketForUpdate(org.mockito.ArgumentMatchers.eq(7L), org.mockito.ArgumentMatchers.eq("DAY"), org.mockito.ArgumentMatchers.anyString())).thenReturn(day); - when(aiQuotaBucketMapper.selectQuotaBucket(org.mockito.ArgumentMatchers.eq(7L), + when(aiQuotaBucketMapper.selectQuotaBucketForUpdate(org.mockito.ArgumentMatchers.eq(7L), org.mockito.ArgumentMatchers.eq("MONTH"), org.mockito.ArgumentMatchers.anyString())).thenReturn(month); when(aiGenerationTaskMapper.countActiveTasksForUser(7L)).thenReturn(0); @@ -52,6 +52,8 @@ public class AiQuotaServiceTest InOrder inOrder = inOrder(aiGenerationTaskMapper, aiQuotaBucketMapper); inOrder.verify(aiGenerationTaskMapper).releaseExpiredRunningTasks(); inOrder.verify(aiGenerationTaskMapper).countActiveTasksForUser(7L); + verify(aiQuotaBucketMapper).selectQuotaBucketForUpdate(org.mockito.ArgumentMatchers.eq(7L), + org.mockito.ArgumentMatchers.eq("DAY"), org.mockito.ArgumentMatchers.anyString()); verify(aiQuotaBucketMapper).updateQuotaBucket(day); } diff --git a/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/FrontProjectPreviewServiceImplTest.java b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/FrontProjectPreviewServiceImplTest.java index e47050d..a678b5e 100644 --- a/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/FrontProjectPreviewServiceImplTest.java +++ b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/FrontProjectPreviewServiceImplTest.java @@ -287,6 +287,27 @@ public class FrontProjectPreviewServiceImplTest verify(genProjectService, never()).downloadStructure(any(GenProject.class), eq("sql")); } + @Test + public void downloadAllRejectsUnsafeEntriesFromGeneratedArchives() throws Exception + { + FrontProject project = frontProjectWithOneTable(); + when(frontProjectService.getProject(100L, 200L)).thenReturn(project); + when(genProjectService.downloadStructure(any(GenProject.class), eq("backend"))) + .thenReturn(zipWithEntry("../outside.txt", "unsafe")); + + ServiceException error = expectServiceException(new ThrowingRunnable() + { + @Override + public void run() + { + previewService.downloadAll(100L, 200L); + } + }); + + assertTrue(error.getMessage().contains("invalid path")); + verify(genProjectService, never()).downloadStructure(any(GenProject.class), eq("frontend")); + } + @Test public void rejectsFrontendStructureWhenProjectDisablesFrontend() { diff --git a/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/FrontProjectServiceImplTest.java b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/FrontProjectServiceImplTest.java index 5ef6993..f35ef13 100644 --- a/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/FrontProjectServiceImplTest.java +++ b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/service/front/FrontProjectServiceImplTest.java @@ -238,6 +238,38 @@ public class FrontProjectServiceImplTest verify(frontProjectTableMapper, never()).deleteTablesByProjectId(any(Long.class)); } + @Test + public void saveDatabaseRejectsGeneratorPathTraversalMetadata() + { + when(frontProjectMapper.selectFrontProjectByUserAndId(7L, 10L)).thenReturn(project()); + final DatabaseTableDesign unsafeModule = table("sys_user", pkColumn()); + unsafeModule.setModuleName("../outside"); + + ServiceException moduleError = expectServiceException(new ThrowingRunnable() + { + @Override + public void run() + { + service.saveDatabase(7L, 10L, database(unsafeModule)); + } + }); + assertEquals("模块名只能包含字母、数字和下划线,且必须以字母开头", moduleError.getMessage()); + + final DatabaseTableDesign unsafeBusiness = table("sys_user", pkColumn()); + unsafeBusiness.setBusinessName("C:/outside"); + ServiceException businessError = expectServiceException(new ThrowingRunnable() + { + @Override + public void run() + { + service.saveDatabase(7L, 10L, database(unsafeBusiness)); + } + }); + assertEquals("业务名只能包含字母、数字和下划线,且必须以字母开头", businessError.getMessage()); + verify(frontProjectColumnMapper, never()).deleteColumnsByProjectId(any(Long.class)); + verify(frontProjectTableMapper, never()).deleteTablesByProjectId(any(Long.class)); + } + @Test public void saveDatabaseAddsDefaultAutoIncrementPrimaryKeyWhenMissing() { diff --git a/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/util/QingTemplateSupportTest.java b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/util/QingTemplateSupportTest.java index 93b74ef..84ab173 100644 --- a/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/util/QingTemplateSupportTest.java +++ b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/util/QingTemplateSupportTest.java @@ -68,12 +68,15 @@ public class QingTemplateSupportTest } @Test - public void qingPackageIncludesEcharts() + public void qingPackageIncludesRuntimeAndLintTooling() { VelocityInitializer.initVelocity(); String content = render("qing/vue-package.json.vm", VelocityUtils.prepareContextProject(project())); assertTrue(content.contains("\"echarts\": \"5.4.0\"")); + assertTrue(content.contains("\"lint\": \"vue-cli-service lint --no-fix\"")); + assertTrue(content.contains("\"@vue/cli-plugin-eslint\": \"^4.5.19\"")); + assertTrue(content.contains("\"eslint-plugin-vue\": \"^6.2.2\"")); } @Test diff --git a/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/util/ZipEntryPathValidatorTest.java b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/util/ZipEntryPathValidatorTest.java new file mode 100644 index 0000000..10ea8c3 --- /dev/null +++ b/RuoYi-Vue/ruoyi-generator/src/test/java/com/ruoyi/generator/util/ZipEntryPathValidatorTest.java @@ -0,0 +1,45 @@ +package com.ruoyi.generator.util; + +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertTrue; +import org.junit.Test; +import com.ruoyi.common.exception.ServiceException; + +public class ZipEntryPathValidatorTest +{ + @Test + public void acceptsPortableRelativeProjectPaths() + { + assertEquals("demo/src/main/App.java", + ZipEntryPathValidator.requireRelative("demo\\src/main/App.java", "test")); + assertEquals("demo/src/", ZipEntryPathValidator.requireRelative("demo/src/", "test")); + } + + @Test + public void rejectsTraversalAbsoluteDriveAndNonPortablePaths() + { + assertInvalid("../outside.txt"); + assertInvalid("demo/../outside.txt"); + assertInvalid("/absolute/path.txt"); + assertInvalid("C:\\outside.txt"); + assertInvalid("demo//file.txt"); + assertInvalid("demo/file?.txt"); + assertInvalid("demo/CON.txt"); + assertInvalid("demo/name. "); + assertInvalid("demo/name."); + } + + private void assertInvalid(String value) + { + try + { + ZipEntryPathValidator.requireRelative(value, "test archive"); + } + catch (ServiceException error) + { + assertTrue(error.getMessage().contains("invalid path")); + return; + } + throw new AssertionError("Expected invalid path: " + value); + } +}