Reject merged conditional branches during flow validation
This commit is contained in:
parent
022d916558
commit
e7ce657367
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in New Issue