diff --git a/src/main/java/it/cnr/isti/workflow/manager/flows/validation/FlowDataValidator.java b/src/main/java/it/cnr/isti/workflow/manager/flows/validation/FlowDataValidator.java index fe9f8b1..3c50536 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/flows/validation/FlowDataValidator.java +++ b/src/main/java/it/cnr/isti/workflow/manager/flows/validation/FlowDataValidator.java @@ -1,8 +1,11 @@ package it.cnr.isti.workflow.manager.flows.validation; import java.util.HashMap; +import java.util.LinkedHashSet; import java.util.List; +import java.util.Map; import java.util.Objects; +import java.util.Set; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.stereotype.Component; @@ -10,6 +13,8 @@ import org.springframework.stereotype.Component; import it.cnr.isti.workflow.manager.blocks.Block; import it.cnr.isti.workflow.manager.blocks.configurations.BlockConfiguration; import it.cnr.isti.workflow.manager.blocks.factories.BlockFactory; +import it.cnr.isti.workflow.manager.blocks.factories.ConditionalBlockFactory; +import it.cnr.isti.workflow.manager.blocks.types.ConditionalBlockType; import it.cnr.isti.workflow.manager.flows.model.Connection; import it.cnr.isti.workflow.manager.flows.model.FlowData; import it.cnr.isti.workflow.manager.ios.IODescriptor; @@ -53,6 +58,8 @@ public class FlowDataValidator implements ConstraintValidator block) { @@ -123,6 +130,59 @@ public class FlowDataValidator implements ConstraintValidator name.equals(descriptor.getName())); } + private void validateConditionalBranches(List> blocks, List connections, + Map> blocksById) { + Map> outgoingBySource = new HashMap<>(); + for (Connection connection : connections) { + outgoingBySource.computeIfAbsent(connection.getSourceId(), ignored -> new java.util.ArrayList<>()) + .add(connection); + } + + for (Block block : blocks) { + if (!ConditionalBlockType.TYPE.equals(block.getType().getName())) { + continue; + } + + Set trueBranch = collectReachableTargets(block.getId(), ConditionalBlockFactory.TRUE_OUTPUT, outgoingBySource); + Set falseBranch = collectReachableTargets(block.getId(), ConditionalBlockFactory.FALSE_OUTPUT, outgoingBySource); + + Set mergedNodes = new LinkedHashSet<>(trueBranch); + mergedNodes.retainAll(falseBranch); + if (!mergedNodes.isEmpty()) { + String mergedNodeNames = mergedNodes.stream() + .map(nodeId -> { + Block target = blocksById.get(nodeId); + return target == null ? nodeId : target.getName() + " (" + nodeId + ")"; + }) + .reduce((left, right) -> left + ", " + right) + .orElse("unknown"); + throw validationError(error("block", block.getId(), "outputs", + "Conditional true/false branches must not merge. Shared downstream nodes: " + mergedNodeNames)); + } + } + } + + private Set collectReachableTargets(String conditionalId, String sourceOutput, + Map> outgoingBySource) { + Set visited = new LinkedHashSet<>(); + List initialConnections = outgoingBySource.getOrDefault(conditionalId, List.of()).stream() + .filter(connection -> sourceOutput.equals(connection.getSourceName())) + .toList(); + for (Connection connection : initialConnections) { + walkTargets(connection.getTargetId(), outgoingBySource, visited); + } + return visited; + } + + private void walkTargets(String nodeId, Map> outgoingBySource, Set visited) { + if (nodeId == null || !visited.add(nodeId)) { + return; + } + for (Connection connection : outgoingBySource.getOrDefault(nodeId, List.of())) { + walkTargets(connection.getTargetId(), outgoingBySource, visited); + } + } + @SuppressWarnings({ "rawtypes", "unchecked" }) private Block recreateBlock(BlockConfiguration configuration) { BlockFactory factory = blockFactories.stream() diff --git a/src/test/java/it/cnr/isti/workflow/manager/controllers/FlowControllerTest.java b/src/test/java/it/cnr/isti/workflow/manager/controllers/FlowControllerTest.java index 98605e6..9c08113 100644 --- a/src/test/java/it/cnr/isti/workflow/manager/controllers/FlowControllerTest.java +++ b/src/test/java/it/cnr/isti/workflow/manager/controllers/FlowControllerTest.java @@ -26,9 +26,13 @@ import it.cnr.isti.workflow.manager.app.ObjectMapperHolder; import it.cnr.isti.workflow.manager.auth.config.JwtUtil; import it.cnr.isti.workflow.manager.auth.repo.LoginEntity; import it.cnr.isti.workflow.manager.blocks.Block; +import it.cnr.isti.workflow.manager.blocks.configurations.ConditionalBlockConfiguration; import it.cnr.isti.workflow.manager.blocks.configurations.LLMBlockConfiguration; +import it.cnr.isti.workflow.manager.blocks.factories.ConditionalBlockFactory; +import it.cnr.isti.workflow.manager.blocks.types.ConditionalBlockType; import it.cnr.isti.workflow.manager.blocks.types.LLMBlockType; import it.cnr.isti.workflow.manager.flows.FlowTestCreator; +import it.cnr.isti.workflow.manager.flows.model.Connection; import it.cnr.isti.workflow.manager.flows.model.Flow; import it.cnr.isti.workflow.manager.flows.model.FlowCreateRequest; import it.cnr.isti.workflow.manager.flows.model.FlowData; @@ -381,4 +385,91 @@ public class FlowControllerTest { assertNotNull(createResponse.getBody()); assertEquals(FlowViewStatus.DRAFT, createResponse.getBody().status()); } + + @Test + public void createFlowRejectsConditionalBranchMerge() { + LLMDescriptor llmDescriptor = LLMDescriptor.builder() + .provider("testProvider") + .model("testModel") + .build(); + + Block source = blocksController.create(LLMBlockConfiguration.builder() + .name("Source") + .llmDescriptor(llmDescriptor) + .prompt("Classify candidate from ${{candidate}}") + .build()); + + Block conditional = blocksController.create(ConditionalBlockConfiguration.builder() + .name("Decision") + .condition("${{response}} == 'yes'") + .useLlm(false) + .outputTemplate("${{response}}") + .build()); + + Block trueBranch = blocksController.create(LLMBlockConfiguration.builder() + .name("True branch") + .llmDescriptor(llmDescriptor) + .prompt("Summarize positive outcome from ${{response}}") + .build()); + + Block falseBranch = blocksController.create(LLMBlockConfiguration.builder() + .name("False branch") + .llmDescriptor(llmDescriptor) + .prompt("Summarize negative outcome from ${{response}}") + .build()); + + Block merged = blocksController.create(LLMBlockConfiguration.builder() + .name("Merged result") + .llmDescriptor(llmDescriptor) + .prompt("Combine true branch: ${{fromTrue}} and false branch: ${{fromFalse}}") + .build()); + + FlowCreateRequest request = new FlowCreateRequest( + "Invalid conditional merge", + "Conditional branches must not merge", + FlowData.builder() + .block(source) + .block(conditional) + .block(trueBranch) + .block(falseBranch) + .block(merged) + .connection(Connection.builder() + .sourceId(source.getId()) + .sourceName("response") + .targetId(conditional.getId()) + .targetName("response") + .build()) + .connection(Connection.builder() + .sourceId(conditional.getId()) + .sourceName(ConditionalBlockFactory.TRUE_OUTPUT) + .targetId(trueBranch.getId()) + .targetName("response") + .build()) + .connection(Connection.builder() + .sourceId(conditional.getId()) + .sourceName(ConditionalBlockFactory.FALSE_OUTPUT) + .targetId(falseBranch.getId()) + .targetName("response") + .build()) + .connection(Connection.builder() + .sourceId(trueBranch.getId()) + .sourceName("response") + .targetId(merged.getId()) + .targetName("fromTrue") + .build()) + .connection(Connection.builder() + .sourceId(falseBranch.getId()) + .sourceName("response") + .targetId(merged.getId()) + .targetName("fromFalse") + .build()) + .build()); + + ResponseStatusException exception = assertThrows( + ResponseStatusException.class, + () -> flowController.createFlow(request, new LoginEntity("testuser", "testpassword"))); + + assertEquals(HttpStatus.BAD_REQUEST, exception.getStatusCode()); + assertTrue(exception.getReason().contains("Conditional true/false branches must not merge")); + } }