Reject merged conditional branches during flow validation

This commit is contained in:
Lucio Lelii 2026-03-12 12:17:41 +01:00
parent 022d916558
commit e7ce657367
2 changed files with 151 additions and 0 deletions

View File

@ -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<ValidFlowStructure
for (Connection connection : connections) {
validateConnection(connection, blocksById);
}
validateConditionalBranches(blocks, connections, blocksById);
}
private void validateBlock(Block<?> block) {
@ -123,6 +130,59 @@ public class FlowDataValidator implements ConstraintValidator<ValidFlowStructure
return descriptors.stream().anyMatch(descriptor -> name.equals(descriptor.getName()));
}
private void validateConditionalBranches(List<Block<?>> blocks, List<Connection> connections,
Map<String, Block<?>> blocksById) {
Map<String, List<Connection>> 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<String> trueBranch = collectReachableTargets(block.getId(), ConditionalBlockFactory.TRUE_OUTPUT, outgoingBySource);
Set<String> falseBranch = collectReachableTargets(block.getId(), ConditionalBlockFactory.FALSE_OUTPUT, outgoingBySource);
Set<String> 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<String> collectReachableTargets(String conditionalId, String sourceOutput,
Map<String, List<Connection>> outgoingBySource) {
Set<String> visited = new LinkedHashSet<>();
List<Connection> 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<String, List<Connection>> outgoingBySource, Set<String> 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()

View File

@ -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<LLMBlockType> source = blocksController.create(LLMBlockConfiguration.builder()
.name("Source")
.llmDescriptor(llmDescriptor)
.prompt("Classify candidate from ${{candidate}}")
.build());
Block<ConditionalBlockType> conditional = blocksController.create(ConditionalBlockConfiguration.builder()
.name("Decision")
.condition("${{response}} == 'yes'")
.useLlm(false)
.outputTemplate("${{response}}")
.build());
Block<LLMBlockType> trueBranch = blocksController.create(LLMBlockConfiguration.builder()
.name("True branch")
.llmDescriptor(llmDescriptor)
.prompt("Summarize positive outcome from ${{response}}")
.build());
Block<LLMBlockType> falseBranch = blocksController.create(LLMBlockConfiguration.builder()
.name("False branch")
.llmDescriptor(llmDescriptor)
.prompt("Summarize negative outcome from ${{response}}")
.build());
Block<LLMBlockType> 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"));
}
}