Add MCP agent chat block and safe execution cancellation
This commit is contained in:
parent
9f558cd4a4
commit
068fa695ba
|
|
@ -189,7 +189,7 @@ public class BlockCatalogService {
|
|||
}
|
||||
|
||||
private AssistantInteractionContractDescriptor resolveInteractionContract(BlockType blockType) {
|
||||
if ("ChatInteraction".equals(blockType.getName())) {
|
||||
if ("ChatInteraction".equals(blockType.getName()) || "MCPAgentChat".equals(blockType.getName())) {
|
||||
return new AssistantInteractionContractDescriptor(
|
||||
"chat-session",
|
||||
"message",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,72 @@
|
|||
package it.cnr.isti.workflow.manager.blocks.configurations;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
import com.fasterxml.jackson.annotation.JsonProperty;
|
||||
import com.fasterxml.jackson.databind.JsonNode;
|
||||
|
||||
import it.cnr.isti.workflow.manager.blocks.types.MCPAgentChatBlockType;
|
||||
import it.cnr.isti.workflow.manager.configurations.annotations.ConfigurableAsInput;
|
||||
import it.cnr.isti.workflow.manager.configurations.annotations.DynamicSchema;
|
||||
import it.cnr.isti.workflow.manager.configurations.annotations.FieldRetriever;
|
||||
import it.cnr.isti.workflow.manager.configurations.annotations.Structural;
|
||||
import jakarta.validation.Valid;
|
||||
import jakarta.validation.constraints.NotBlank;
|
||||
import jakarta.validation.constraints.NotNull;
|
||||
import lombok.Builder;
|
||||
import lombok.EqualsAndHashCode;
|
||||
import lombok.Getter;
|
||||
import lombok.NoArgsConstructor;
|
||||
import lombok.NonNull;
|
||||
|
||||
@NoArgsConstructor(access = lombok.AccessLevel.PROTECTED)
|
||||
@Getter
|
||||
@EqualsAndHashCode(callSuper = true)
|
||||
public class MCPAgentChatBlockConfiguration extends BlockConfiguration<MCPAgentChatBlockType> {
|
||||
|
||||
@Structural
|
||||
@ConfigurableAsInput
|
||||
@FieldRetriever(name = "LLM", url = "/retriever/LLM/models?provider=InternalOllama")
|
||||
@JsonProperty(required = false)
|
||||
private String model;
|
||||
|
||||
@JsonProperty(required = false)
|
||||
@Valid
|
||||
private List<ChatInteractionInput> inputs = List.of();
|
||||
|
||||
@JsonProperty(required = false)
|
||||
@Valid
|
||||
private List<MCPServerBinding> mcpServers = List.of();
|
||||
|
||||
@Builder
|
||||
public MCPAgentChatBlockConfiguration(@NonNull String name, String model, List<ChatInteractionInput> inputs,
|
||||
List<MCPServerBinding> mcpServers) {
|
||||
super(name);
|
||||
this.model = model;
|
||||
this.inputs = inputs == null ? List.of() : List.copyOf(inputs);
|
||||
this.mcpServers = mcpServers == null ? List.of() : List.copyOf(mcpServers);
|
||||
}
|
||||
|
||||
@Override
|
||||
public Class<MCPAgentChatBlockType> getBlockType() {
|
||||
return MCPAgentChatBlockType.class;
|
||||
}
|
||||
|
||||
public static MCPAgentChatBlockConfiguration empty() {
|
||||
MCPAgentChatBlockConfiguration configuration = new MCPAgentChatBlockConfiguration();
|
||||
configuration.name = MCPAgentChatBlockType.TYPE;
|
||||
configuration.inputs = List.of();
|
||||
configuration.mcpServers = List.of();
|
||||
return configuration;
|
||||
}
|
||||
|
||||
@Builder
|
||||
public record MCPServerBinding(
|
||||
@NotBlank
|
||||
@FieldRetriever(name = "MCPServers", url = "/retriever/MCPServers/servers")
|
||||
String serverName,
|
||||
@NotNull
|
||||
@DynamicSchema(url = "/retriever/MCPServers/definitions/schema", dependsOn = { "serverName" })
|
||||
JsonNode configuration) {
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,121 @@
|
|||
package it.cnr.isti.workflow.manager.blocks.factories;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Objects;
|
||||
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.stereotype.Component;
|
||||
|
||||
import it.cnr.isti.workflow.manager.blocks.Block;
|
||||
import it.cnr.isti.workflow.manager.blocks.IOCapability;
|
||||
import it.cnr.isti.workflow.manager.blocks.IOCapabilityType;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.ChatInteractionInput;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.MCPAgentChatBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.MCPAgentChatBlockType;
|
||||
import it.cnr.isti.workflow.manager.ios.IODescriptor;
|
||||
import it.cnr.isti.workflow.manager.ios.IOType;
|
||||
|
||||
@Component
|
||||
public class MCPAgentChatBlockFactory
|
||||
implements BlockFactory<MCPAgentChatBlockType, MCPAgentChatBlockConfiguration> {
|
||||
|
||||
public static final String INTERACTION_FIELD = "message";
|
||||
public static final String FINAL_RESPONSE_FIELD = "response";
|
||||
public static final String RESPONSE_OUTPUT = "response";
|
||||
public static final String HISTORY_OUTPUT = "history";
|
||||
public static final String MODEL_INPUT = "model";
|
||||
|
||||
private static final List<IOCapability> RESPONSE_CAPABILITIES = List.of(
|
||||
new IOCapability(IOCapabilityType.TEXT, false));
|
||||
private static final List<IOCapability> HISTORY_CAPABILITIES = List.of(
|
||||
new IOCapability(IOCapabilityType.TEXT, true));
|
||||
|
||||
@Autowired
|
||||
private MCPAgentChatBlockType blockType;
|
||||
|
||||
@Override
|
||||
public Block<MCPAgentChatBlockType> create(MCPAgentChatBlockConfiguration configuration) {
|
||||
validateConfiguration(configuration);
|
||||
return Block.<MCPAgentChatBlockType>builder()
|
||||
.inputs(resolveInputs(configuration))
|
||||
.output(IODescriptor.output(RESPONSE_OUTPUT, IOType.TEXT, false, RESPONSE_CAPABILITIES))
|
||||
.output(IODescriptor.output(HISTORY_OUTPUT, IOType.TEXT, true, HISTORY_CAPABILITIES))
|
||||
.specificConfiguration(configuration)
|
||||
.type(blockType)
|
||||
.build();
|
||||
}
|
||||
|
||||
@Override
|
||||
public Block<MCPAgentChatBlockType> createEmpty() {
|
||||
return create(MCPAgentChatBlockConfiguration.empty());
|
||||
}
|
||||
|
||||
@Override
|
||||
public Class<MCPAgentChatBlockType> getBlockType() {
|
||||
return MCPAgentChatBlockType.class;
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<IOCapability> supportedInputCapabilities() {
|
||||
return List.of(
|
||||
new IOCapability(IOCapabilityType.TEXT, false),
|
||||
new IOCapability(IOCapabilityType.TEXT, true));
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<IOCapability> supportedOutputCapabilities() {
|
||||
return List.of(
|
||||
new IOCapability(IOCapabilityType.TEXT, false),
|
||||
new IOCapability(IOCapabilityType.TEXT, true));
|
||||
}
|
||||
|
||||
private List<IODescriptor> resolveInputs(MCPAgentChatBlockConfiguration configuration) {
|
||||
java.util.ArrayList<IODescriptor> inputs = new java.util.ArrayList<>(configurableInputDescriptors(configuration));
|
||||
if (configuration.getInputs() != null) {
|
||||
inputs.addAll(configuration.getInputs().stream().map(this::toDescriptor).toList());
|
||||
}
|
||||
return inputs.stream().distinct().toList();
|
||||
}
|
||||
|
||||
private IODescriptor toDescriptor(ChatInteractionInput input) {
|
||||
IOType type = input.ioType() == null ? IOType.TEXT : input.ioType();
|
||||
return IODescriptor.input(input.name(), type, input.multiple(),
|
||||
List.of(new IOCapability(toCapabilityType(type), input.multiple())));
|
||||
}
|
||||
|
||||
private void validateConfiguration(MCPAgentChatBlockConfiguration configuration) {
|
||||
if (configuration == null || configuration.getInputs() == null || configuration.getInputs().isEmpty()) {
|
||||
return;
|
||||
}
|
||||
long distinctNames = configuration.getInputs().stream()
|
||||
.map(ChatInteractionInput::name)
|
||||
.filter(Objects::nonNull)
|
||||
.distinct()
|
||||
.count();
|
||||
long names = configuration.getInputs().stream()
|
||||
.map(ChatInteractionInput::name)
|
||||
.filter(Objects::nonNull)
|
||||
.count();
|
||||
if (distinctNames != names) {
|
||||
throw new IllegalArgumentException("inputs must have unique names");
|
||||
}
|
||||
boolean unsupportedType = configuration.getInputs().stream()
|
||||
.anyMatch(input -> input.ioType() != null && input.ioType() != IOType.TEXT);
|
||||
if (unsupportedType) {
|
||||
throw new IllegalArgumentException("MCPAgentChat inputs support only TEXT or TEXT[]");
|
||||
}
|
||||
}
|
||||
|
||||
private IOCapabilityType toCapabilityType(IOType type) {
|
||||
if (type == IOType.TEXT) {
|
||||
return IOCapabilityType.TEXT;
|
||||
}
|
||||
if (type == IOType.FILE || type == IOType.CSV) {
|
||||
return IOCapabilityType.FILE;
|
||||
}
|
||||
if (type == IOType.ANY) {
|
||||
return IOCapabilityType.ANY;
|
||||
}
|
||||
throw new IllegalArgumentException("Unsupported IOType: " + type);
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,37 @@
|
|||
package it.cnr.isti.workflow.manager.blocks.types;
|
||||
|
||||
import org.springframework.stereotype.Component;
|
||||
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.BlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.MCPAgentChatBlockConfiguration;
|
||||
|
||||
@Component(MCPAgentChatBlockType.TYPE)
|
||||
public class MCPAgentChatBlockType implements BlockType {
|
||||
|
||||
public static final String TYPE = "MCPAgentChat";
|
||||
|
||||
@Override
|
||||
public String getName() {
|
||||
return TYPE;
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getDescription() {
|
||||
return "A human-interactive MCP agent chat block with a persistent MCP session";
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean validate() {
|
||||
return true;
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean isUserInteractive() {
|
||||
return true;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Class<? extends BlockConfiguration<?>> getBlockConfigurationClass() {
|
||||
return MCPAgentChatBlockConfiguration.class;
|
||||
}
|
||||
}
|
||||
|
|
@ -135,7 +135,7 @@ public class BlocksController {
|
|||
}
|
||||
|
||||
private InteractionContractDescriptor resolveInteractionContract(BlockType blockType) {
|
||||
if ("ChatInteraction".equals(blockType.getName())) {
|
||||
if ("ChatInteraction".equals(blockType.getName()) || "MCPAgentChat".equals(blockType.getName())) {
|
||||
return new InteractionContractDescriptor(
|
||||
"chat-session",
|
||||
"message",
|
||||
|
|
|
|||
|
|
@ -175,6 +175,12 @@ public class ExecutionsController {
|
|||
return executionService.startExecution(id);
|
||||
}
|
||||
|
||||
@PutMapping(path = "{id}/cancel")
|
||||
@Operation(summary = "Cancels an execution", description = "Cancels the execution, clears runtime state and marks it as CANCELLED.")
|
||||
public ExecutionObject cancel(@PathVariable String id) {
|
||||
return executionService.cancelExecution(id);
|
||||
}
|
||||
|
||||
|
||||
/**
|
||||
* Prepares the input for an execution by associating a given input string
|
||||
|
|
|
|||
|
|
@ -52,8 +52,10 @@ public class ExecutionContext implements ExecutionListener {
|
|||
protected void setStatus(ExecutionStatus status) {
|
||||
this.status = status;
|
||||
if (status == ExecutionStatus.RUNNING) {
|
||||
this.startTime = System.currentTimeMillis();
|
||||
} else if (status == ExecutionStatus.SUCCESS || status == ExecutionStatus.ERROR) {
|
||||
if (this.startTime == null) {
|
||||
this.startTime = System.currentTimeMillis();
|
||||
}
|
||||
} else if (status == ExecutionStatus.SUCCESS || status == ExecutionStatus.ERROR || status == ExecutionStatus.CANCELLED) {
|
||||
this.endTime = System.currentTimeMillis();
|
||||
}
|
||||
}
|
||||
|
|
@ -84,6 +86,9 @@ public class ExecutionContext implements ExecutionListener {
|
|||
|
||||
@Override
|
||||
public void completed(String id, Map<String, Object> result) {
|
||||
if (this.status == ExecutionStatus.CANCELLED) {
|
||||
return;
|
||||
}
|
||||
logger.info("Step " + id + " completed with result: " + result);
|
||||
clearPartialResults(id);
|
||||
Step<?> completedStep = this.steps.get(id);
|
||||
|
|
@ -103,6 +108,9 @@ public class ExecutionContext implements ExecutionListener {
|
|||
|
||||
@Override
|
||||
public void skipped(String id) {
|
||||
if (this.status == ExecutionStatus.CANCELLED) {
|
||||
return;
|
||||
}
|
||||
logger.info("Step " + id + " skipped");
|
||||
Step<?> skippedStep = this.steps.get(id);
|
||||
skippedStep.getOutputs().forEach(output -> output.markUnavailable());
|
||||
|
|
@ -111,6 +119,9 @@ public class ExecutionContext implements ExecutionListener {
|
|||
|
||||
@Override
|
||||
public void failed(String id, String error) {
|
||||
if (this.status == ExecutionStatus.CANCELLED) {
|
||||
return;
|
||||
}
|
||||
this.addError(id, error);
|
||||
this.setStatus(ExecutionStatus.ERROR);
|
||||
logger.severe("Step " + id + " failed with error: " + error);
|
||||
|
|
@ -118,11 +129,17 @@ public class ExecutionContext implements ExecutionListener {
|
|||
|
||||
@Override
|
||||
public synchronized void started(String id) {
|
||||
if (this.status == ExecutionStatus.CANCELLED) {
|
||||
return;
|
||||
}
|
||||
logger.info("Step " + id + " started");
|
||||
}
|
||||
|
||||
@Override
|
||||
public synchronized void paused(String id) {
|
||||
if (this.status == ExecutionStatus.CANCELLED) {
|
||||
return;
|
||||
}
|
||||
logger.info("Step " + id + " paused");
|
||||
synchronized (this.waitingSteps) {
|
||||
if (!this.waitingSteps.contains(id)) {
|
||||
|
|
@ -136,6 +153,9 @@ public class ExecutionContext implements ExecutionListener {
|
|||
|
||||
@Override
|
||||
public synchronized void resumed(String id) {
|
||||
if (this.status == ExecutionStatus.CANCELLED) {
|
||||
return;
|
||||
}
|
||||
logger.info("Step " + id + " resumed");
|
||||
synchronized (this.waitingSteps) {
|
||||
this.waitingSteps.remove(id);
|
||||
|
|
@ -147,6 +167,9 @@ public class ExecutionContext implements ExecutionListener {
|
|||
|
||||
@Override
|
||||
public synchronized void partialUpdated(String id, Map<String, Object> partialResult) {
|
||||
if (this.status == ExecutionStatus.CANCELLED) {
|
||||
return;
|
||||
}
|
||||
clearPartialResults(id);
|
||||
partialResult.forEach((key, value) -> addPartialResult(id, key, value));
|
||||
}
|
||||
|
|
@ -191,7 +214,22 @@ public class ExecutionContext implements ExecutionListener {
|
|||
this.steps.values().forEach(step -> step.start(executorService));
|
||||
}
|
||||
|
||||
protected synchronized void cancel() {
|
||||
this.inputs.clear();
|
||||
this.authorizations.clear();
|
||||
this.result.clear();
|
||||
this.partialResult.clear();
|
||||
this.errors.clear();
|
||||
this.warnings.clear();
|
||||
this.waitingSteps.clear();
|
||||
this.steps.values().forEach(Step::cancel);
|
||||
this.setStatus(ExecutionStatus.CANCELLED);
|
||||
}
|
||||
|
||||
private void updateTerminalStatus() {
|
||||
if (this.status == ExecutionStatus.CANCELLED) {
|
||||
return;
|
||||
}
|
||||
boolean allTerminal = this.steps.values().stream()
|
||||
.allMatch(step -> step.getStatus() == StepStatus.COMPLETED || step.getStatus() == StepStatus.SKIPPED);
|
||||
if (allTerminal) {
|
||||
|
|
|
|||
|
|
@ -117,6 +117,13 @@ public class ExecutionObject {
|
|||
+ " is not in READY status (CURRENT STATUS is " + this.getContext().getStatus() + ")");
|
||||
}
|
||||
|
||||
protected void cancel() {
|
||||
if (this.executorService != null) {
|
||||
this.executorService.shutdownNow();
|
||||
}
|
||||
this.context.cancel();
|
||||
}
|
||||
|
||||
public List<String> getMissingAuthorizationKeys() {
|
||||
return this.requiredAuthorizations.stream()
|
||||
.map(ExecutionAuthorizationRequirement::key)
|
||||
|
|
|
|||
|
|
@ -7,7 +7,8 @@ public enum ExecutionStatus {
|
|||
RUNNING(false, false, false),
|
||||
WAITING(false, false, true),
|
||||
SUCCESS(false, true, false),
|
||||
ERROR(false, true, false);
|
||||
ERROR(false, true, false),
|
||||
CANCELLED(false, true, false);
|
||||
|
||||
private boolean initState;
|
||||
|
||||
|
|
|
|||
|
|
@ -81,6 +81,15 @@ public class ExecutionsService {
|
|||
executions.remove(id);
|
||||
}
|
||||
|
||||
public ExecutionObject cancelExecution(String id) {
|
||||
ExecutionObject eo = getExecution(id);
|
||||
if (eo.getContext().getStatus().isFinalState()) {
|
||||
return eo;
|
||||
}
|
||||
eo.cancel();
|
||||
return eo;
|
||||
}
|
||||
|
||||
public ExecutionObject prepareInput(String executionId, String blockId, String inputName, Object input) {
|
||||
ExecutionObject eo = getExecution(executionId);
|
||||
eo.setInput(blockId, inputName, input);
|
||||
|
|
|
|||
|
|
@ -41,4 +41,11 @@ public final class NodeExecutors {
|
|||
}
|
||||
throw new IllegalStateException("No interactive executor found for node type " + node.getClass().getName());
|
||||
}
|
||||
|
||||
public static void cancel(FlowNode node, List<Input> inputs, Map<String, Object> partialResults,
|
||||
Map<String, Object> authorizations) {
|
||||
if (node instanceof Block<?> block) {
|
||||
BlockExecutors.get(block.getType()).cancel((Block) block, inputs, partialResults, authorizations);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -28,4 +28,8 @@ public interface BlockExecutor<B extends BlockType> {
|
|||
throw new UnsupportedOperationException("Interaction not supported for this block type");
|
||||
}
|
||||
|
||||
default void cancel(Block<B> block, List<Input> inputs, Map<String, Object> partialResults,
|
||||
Map<String, Object> authorizations) {
|
||||
}
|
||||
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1,180 @@
|
|||
package it.cnr.isti.workflow.manager.executions.executors.blocks;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collection;
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.stream.Collectors;
|
||||
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.stereotype.Component;
|
||||
import org.springframework.util.StringUtils;
|
||||
|
||||
import it.cnr.isti.workflow.manager.blocks.Block;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.MCPAgentBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.MCPAgentChatBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.blocks.factories.MCPAgentChatBlockFactory;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.MCPAgentChatBlockType;
|
||||
import it.cnr.isti.workflow.manager.executions.InteractionResult;
|
||||
import it.cnr.isti.workflow.manager.executions.steps.Input;
|
||||
import it.cnr.isti.workflow.manager.mcp.MCPAgentService;
|
||||
|
||||
@Component
|
||||
public class MCPAgentChatExecutor implements BlockExecutor<MCPAgentChatBlockType> {
|
||||
|
||||
private static final String SESSION_ID_STATE = "__sessionId";
|
||||
|
||||
@Autowired
|
||||
private MCPAgentService mcpAgentService;
|
||||
|
||||
@Override
|
||||
public Map<String, Object> execute(Block<MCPAgentChatBlockType> block, List<Input> inputs,
|
||||
Map<String, Object> authorizations) {
|
||||
throw new UnsupportedOperationException("MCPAgentChat blocks require user interaction.");
|
||||
}
|
||||
|
||||
@Override
|
||||
public Class<MCPAgentChatBlockType> getBlockType() {
|
||||
return MCPAgentChatBlockType.class;
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean isInteractive() {
|
||||
return true;
|
||||
}
|
||||
|
||||
@Override
|
||||
public InteractionResult interact(Block<MCPAgentChatBlockType> block, List<Input> inputs,
|
||||
Map<String, Object> interaction, Map<String, Object> partialResults, Map<String, Object> authorizations) {
|
||||
MCPAgentChatBlockConfiguration configuration =
|
||||
(MCPAgentChatBlockConfiguration) block.getSpecificConfiguration();
|
||||
|
||||
if (interaction.containsKey(MCPAgentChatBlockFactory.INTERACTION_FIELD)) {
|
||||
Object messageValue = interaction.get(MCPAgentChatBlockFactory.INTERACTION_FIELD);
|
||||
if (!(messageValue instanceof String message) || !StringUtils.hasText(message)) {
|
||||
throw new IllegalArgumentException("Missing interaction value for field: "
|
||||
+ MCPAgentChatBlockFactory.INTERACTION_FIELD);
|
||||
}
|
||||
|
||||
String resolvedMessage = resolvePlaceholders(message, inputs);
|
||||
String model = resolveConfigurableInput(MCPAgentChatBlockFactory.MODEL_INPUT, configuration.getModel(), inputs);
|
||||
if (!StringUtils.hasText(model)) {
|
||||
throw new IllegalArgumentException("MCPAgentChat requires a model either from configuration or runtime input");
|
||||
}
|
||||
|
||||
String existingSessionId = existingSessionId(partialResults);
|
||||
boolean createdSession = false;
|
||||
if (!StringUtils.hasText(existingSessionId)) {
|
||||
existingSessionId = mcpAgentService.openSession(model, mapServers(configuration.getMcpServers()));
|
||||
createdSession = true;
|
||||
}
|
||||
|
||||
try {
|
||||
// The MCP bridge keeps the conversation context inside the server-side session.
|
||||
// We send only the current user turn and keep history locally for UI/output purposes.
|
||||
String assistantResponse = mcpAgentService.querySession(existingSessionId, resolvedMessage);
|
||||
List<String> history = existingHistory(partialResults);
|
||||
List<String> updatedHistory = new ArrayList<>(history);
|
||||
updatedHistory.add(formatConversationLine("USER", resolvedMessage));
|
||||
updatedHistory.add(formatConversationLine("ASSISTANT", assistantResponse));
|
||||
return InteractionResult.partial(Map.of(
|
||||
MCPAgentChatBlockFactory.RESPONSE_OUTPUT, assistantResponse,
|
||||
MCPAgentChatBlockFactory.HISTORY_OUTPUT, List.copyOf(updatedHistory),
|
||||
SESSION_ID_STATE, existingSessionId));
|
||||
} catch (RuntimeException ex) {
|
||||
if (createdSession) {
|
||||
mcpAgentService.closeSessionQuietly(existingSessionId);
|
||||
}
|
||||
throw ex;
|
||||
}
|
||||
}
|
||||
|
||||
if (interaction.containsKey(MCPAgentChatBlockFactory.FINAL_RESPONSE_FIELD)) {
|
||||
Object responseValue = interaction.get(MCPAgentChatBlockFactory.FINAL_RESPONSE_FIELD);
|
||||
if (!(responseValue instanceof String response) || !StringUtils.hasText(response)) {
|
||||
throw new IllegalArgumentException("Missing interaction value for field: "
|
||||
+ MCPAgentChatBlockFactory.FINAL_RESPONSE_FIELD);
|
||||
}
|
||||
mcpAgentService.closeSessionQuietly(existingSessionId(partialResults));
|
||||
List<String> history = existingHistory(partialResults);
|
||||
return InteractionResult.completed(Map.of(
|
||||
MCPAgentChatBlockFactory.RESPONSE_OUTPUT, response,
|
||||
MCPAgentChatBlockFactory.HISTORY_OUTPUT, List.copyOf(history)));
|
||||
}
|
||||
|
||||
throw new IllegalArgumentException("Unsupported interaction field for MCPAgentChat");
|
||||
}
|
||||
|
||||
@Override
|
||||
public void cancel(Block<MCPAgentChatBlockType> block, List<Input> inputs, Map<String, Object> partialResults,
|
||||
Map<String, Object> authorizations) {
|
||||
mcpAgentService.closeSessionQuietly(existingSessionId(partialResults));
|
||||
}
|
||||
|
||||
@SuppressWarnings("unchecked")
|
||||
private List<String> existingHistory(Map<String, Object> partialResults) {
|
||||
Object value = partialResults.get(MCPAgentChatBlockFactory.HISTORY_OUTPUT);
|
||||
if (value instanceof List<?> list) {
|
||||
return ((List<Object>) list).stream().map(String::valueOf).toList();
|
||||
}
|
||||
return List.of();
|
||||
}
|
||||
|
||||
private String existingSessionId(Map<String, Object> partialResults) {
|
||||
Object value = partialResults.get(SESSION_ID_STATE);
|
||||
return value == null ? null : String.valueOf(value);
|
||||
}
|
||||
|
||||
private String formatConversationLine(String role, String content) {
|
||||
return "[" + role + "] " + content;
|
||||
}
|
||||
|
||||
private String resolveConfigurableInput(String inputName, String configuredValue, List<Input> inputs) {
|
||||
return inputs.stream()
|
||||
.filter(input -> inputName.equals(input.getDescriptor().getName()))
|
||||
.findFirst()
|
||||
.map(Input::getValue)
|
||||
.map(Object::toString)
|
||||
.filter(StringUtils::hasText)
|
||||
.orElse(configuredValue);
|
||||
}
|
||||
|
||||
private List<MCPAgentBlockConfiguration.MCPServerBinding> mapServers(
|
||||
List<MCPAgentChatBlockConfiguration.MCPServerBinding> servers) {
|
||||
if (servers == null || servers.isEmpty()) {
|
||||
return List.of();
|
||||
}
|
||||
return servers.stream()
|
||||
.filter(java.util.Objects::nonNull)
|
||||
.map(server -> MCPAgentBlockConfiguration.MCPServerBinding.builder()
|
||||
.serverName(server.serverName())
|
||||
.configuration(server.configuration())
|
||||
.build())
|
||||
.toList();
|
||||
}
|
||||
|
||||
private String resolvePlaceholders(String template, List<Input> inputs) {
|
||||
if (!StringUtils.hasText(template)) {
|
||||
return template;
|
||||
}
|
||||
Map<String, String> values = new LinkedHashMap<>();
|
||||
for (Input input : inputs) {
|
||||
values.put(input.getDescriptor().getName(), formatInputValue(input.getValue()));
|
||||
}
|
||||
String resolved = template;
|
||||
for (Map.Entry<String, String> entry : values.entrySet()) {
|
||||
resolved = resolved.replace("${{" + entry.getKey() + "}}", entry.getValue());
|
||||
}
|
||||
return resolved;
|
||||
}
|
||||
|
||||
private String formatInputValue(Object value) {
|
||||
if (value instanceof Collection<?> collection) {
|
||||
return collection.stream()
|
||||
.map(item -> item == null ? "null" : item.toString())
|
||||
.collect(Collectors.joining(System.lineSeparator()));
|
||||
}
|
||||
return value == null ? "null" : value.toString();
|
||||
}
|
||||
}
|
||||
|
|
@ -213,7 +213,21 @@ public class Step<N extends FlowNode> implements InputListener {
|
|||
private boolean isTerminal() {
|
||||
return this.status == StepStatus.COMPLETED
|
||||
|| this.status == StepStatus.SKIPPED
|
||||
|| this.status == StepStatus.FAILED;
|
||||
|| this.status == StepStatus.FAILED
|
||||
|| this.status == StepStatus.CANCELLED;
|
||||
}
|
||||
|
||||
public void cancel() {
|
||||
if (isTerminal()) {
|
||||
return;
|
||||
}
|
||||
try {
|
||||
NodeExecutors.cancel(this.node, this.inputs, Map.copyOf(this.partialResults), this.authorizations);
|
||||
} catch (RuntimeException ex) {
|
||||
logger.warn("Failed to cancel resources for step {}", this.id, ex);
|
||||
}
|
||||
this.partialResults.clear();
|
||||
this.status = StepStatus.CANCELLED;
|
||||
}
|
||||
|
||||
}
|
||||
|
|
|
|||
|
|
@ -8,5 +8,6 @@ public enum StepStatus {
|
|||
WAITING_FOR_INTERACTION,
|
||||
COMPLETED,
|
||||
SKIPPED,
|
||||
FAILED
|
||||
FAILED,
|
||||
CANCELLED
|
||||
}
|
||||
|
|
|
|||
|
|
@ -32,6 +32,9 @@ public class FlowExecutionValidator {
|
|||
}
|
||||
|
||||
public boolean isExecutable(FlowData flowData) {
|
||||
if (flowData == null || flowData.getNodes().isEmpty()) {
|
||||
return false;
|
||||
}
|
||||
return collectErrors(flowData).isEmpty();
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -41,15 +41,17 @@ public class MCPAgentService {
|
|||
Objects.requireNonNull(model, "model cannot be null");
|
||||
Objects.requireNonNull(prompt, "prompt cannot be null");
|
||||
|
||||
String sessionId = createSession(model, mcpServers);
|
||||
String sessionId = openSession(model, mcpServers);
|
||||
try {
|
||||
return executeQuery(sessionId, prompt);
|
||||
return querySession(sessionId, prompt);
|
||||
} finally {
|
||||
deleteSessionQuietly(sessionId);
|
||||
closeSessionQuietly(sessionId);
|
||||
}
|
||||
}
|
||||
|
||||
private String createSession(String model, List<it.cnr.isti.workflow.manager.blocks.configurations.MCPAgentBlockConfiguration.MCPServerBinding> mcpServers) {
|
||||
public String openSession(String model,
|
||||
List<it.cnr.isti.workflow.manager.blocks.configurations.MCPAgentBlockConfiguration.MCPServerBinding> mcpServers) {
|
||||
Objects.requireNonNull(model, "model cannot be null");
|
||||
Map<String, Object> request = new LinkedHashMap<>();
|
||||
request.put("llm_provider", Map.of(
|
||||
"provider", "ollama",
|
||||
|
|
@ -94,7 +96,9 @@ public class MCPAgentService {
|
|||
return response.session_id();
|
||||
}
|
||||
|
||||
private String executeQuery(String sessionId, String prompt) {
|
||||
public String querySession(String sessionId, String prompt) {
|
||||
Objects.requireNonNull(sessionId, "sessionId cannot be null");
|
||||
Objects.requireNonNull(prompt, "prompt cannot be null");
|
||||
QueryResponse response = client().post()
|
||||
.uri("/sessions/{sessionId}/query", sessionId)
|
||||
.contentType(MediaType.APPLICATION_JSON)
|
||||
|
|
@ -114,7 +118,7 @@ public class MCPAgentService {
|
|||
return response.result();
|
||||
}
|
||||
|
||||
private void deleteSessionQuietly(String sessionId) {
|
||||
public void closeSessionQuietly(String sessionId) {
|
||||
if (sessionId == null || sessionId.isBlank()) {
|
||||
return;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -12,14 +12,19 @@ import org.springframework.test.context.TestPropertySource;
|
|||
import it.cnr.isti.workflow.manager.blocks.configurations.LLMBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.HTTPServerCallBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.MCPAgentBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.MCPAgentChatBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.ChatInteractionInput;
|
||||
import it.cnr.isti.workflow.manager.app.ObjectMapperHolder;
|
||||
import it.cnr.isti.workflow.manager.blocks.factories.BlockFactory;
|
||||
import it.cnr.isti.workflow.manager.blocks.factories.HTTPServerCallBlockFactory;
|
||||
import it.cnr.isti.workflow.manager.blocks.factories.LLMBlockFactory;
|
||||
import it.cnr.isti.workflow.manager.blocks.factories.MCPAgentBlockFactory;
|
||||
import it.cnr.isti.workflow.manager.blocks.factories.MCPAgentChatBlockFactory;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.HTTPServerCallBlockType;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.LLMBlockType;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.MCPAgentBlockType;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.MCPAgentChatBlockType;
|
||||
import it.cnr.isti.workflow.manager.ios.IOType;
|
||||
import it.cnr.isti.workflow.manager.llms.LLMDescriptor;
|
||||
|
||||
@SpringBootTest
|
||||
|
|
@ -35,6 +40,9 @@ public class BlockTest {
|
|||
@Autowired
|
||||
MCPAgentBlockFactory mcpAgentBlockFactory;
|
||||
|
||||
@Autowired
|
||||
MCPAgentChatBlockFactory mcpAgentChatBlockFactory;
|
||||
|
||||
@Autowired
|
||||
HTTPServerCallBlockFactory httpServerCallBlockFactory;
|
||||
|
||||
|
|
@ -90,6 +98,29 @@ public class BlockTest {
|
|||
assertNotNull(((MCPAgentBlockConfiguration) block.getSpecificConfiguration()).getMcpServers());
|
||||
}
|
||||
|
||||
@Test
|
||||
void createMCPAgentChatBlock() {
|
||||
BlockFactory<MCPAgentChatBlockType, MCPAgentChatBlockConfiguration> factory = mcpAgentChatBlockFactory;
|
||||
|
||||
MCPAgentChatBlockConfiguration config = MCPAgentChatBlockConfiguration.builder()
|
||||
.name("agent-chat-1")
|
||||
.model("llama3.1:8b")
|
||||
.inputs(List.of(new ChatInteractionInput("candidate", IOType.TEXT, false)))
|
||||
.mcpServers(List.of(
|
||||
MCPAgentChatBlockConfiguration.MCPServerBinding.builder()
|
||||
.serverName("filesystem")
|
||||
.configuration(ObjectMapperHolder.mapper.createObjectNode().put("rootPath", "/tmp"))
|
||||
.build()))
|
||||
.build();
|
||||
|
||||
Block<MCPAgentChatBlockType> block = factory.create(config);
|
||||
assertNotNull(block);
|
||||
assertNotNull(block.getSpecificConfiguration());
|
||||
assertNotNull(block.getInputs());
|
||||
assertNotNull(block.getOutputs());
|
||||
assertNotNull(((MCPAgentChatBlockConfiguration) block.getSpecificConfiguration()).getMcpServers());
|
||||
}
|
||||
|
||||
@Test
|
||||
void createHTTPServerCallBlock() {
|
||||
BlockFactory<HTTPServerCallBlockType, HTTPServerCallBlockConfiguration> factory = httpServerCallBlockFactory;
|
||||
|
|
|
|||
|
|
@ -26,10 +26,12 @@ import it.cnr.isti.workflow.manager.blocks.types.ConditionalBlockType;
|
|||
import it.cnr.isti.workflow.manager.blocks.configurations.HTTPServerCallBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.LLMBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.MCPAgentBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.MCPAgentChatBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.controllers.BlocksController.BlockConfigurationDescriptor;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.HTTPServerCallBlockType;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.LLMBlockType;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.MCPAgentBlockType;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.MCPAgentChatBlockType;
|
||||
import it.cnr.isti.workflow.manager.llms.LLMDescriptor;
|
||||
import it.cnr.isti.workflow.manager.app.ObjectMapperHolder;
|
||||
|
||||
|
|
@ -208,6 +210,41 @@ public class BlocksControllerTest {
|
|||
assertTrue(exception.getMessage().contains("ChatInteraction inputs support only TEXT or TEXT[]"));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void createMcpAgentChatBlockExposesInteractiveOutputs() {
|
||||
MCPAgentChatBlockConfiguration config = MCPAgentChatBlockConfiguration.builder()
|
||||
.name("Agent chat")
|
||||
.model("llama3.1:8b")
|
||||
.inputs(List.of(new ChatInteractionInput("candidate", it.cnr.isti.workflow.manager.ios.IOType.TEXT)))
|
||||
.mcpServers(List.of(MCPAgentChatBlockConfiguration.MCPServerBinding.builder()
|
||||
.serverName("filesystem")
|
||||
.configuration(ObjectMapperHolder.mapper.createObjectNode().put("rootPath", "/workspace"))
|
||||
.build()))
|
||||
.build();
|
||||
|
||||
Block<MCPAgentChatBlockType> block = blocksController.create(config);
|
||||
|
||||
assertNotNull(block);
|
||||
assertEquals(MCPAgentChatBlockType.TYPE, block.getType().getName());
|
||||
assertTrue(block.getInputs().stream().anyMatch(input -> input.getName().equals("candidate")));
|
||||
assertTrue(block.getOutputs().stream().anyMatch(output -> output.getName().equals("response") && !output.isMultiple()));
|
||||
assertTrue(block.getOutputs().stream().anyMatch(output -> output.getName().equals("history") && output.isMultiple()));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void mcpAgentChatDescriptorExposesInteractionContract() {
|
||||
BlockConfigurationDescriptor descriptor = blocksController
|
||||
.getConfigurationDescriptorForType(MCPAgentChatBlockType.TYPE);
|
||||
|
||||
assertNotNull(descriptor.interactionContract());
|
||||
assertEquals("chat-session", descriptor.interactionContract().kind());
|
||||
assertEquals("message", descriptor.interactionContract().messageField());
|
||||
assertEquals("response", descriptor.interactionContract().completionField());
|
||||
assertEquals("history", descriptor.interactionContract().historyField());
|
||||
assertEquals("response", descriptor.interactionContract().responseField());
|
||||
assertTrue(descriptor.interactionContract().supportsPartialResult());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void LLMSchemaHasLongTextHints() {
|
||||
BlockConfigurationDescriptor descriptor = blocksController
|
||||
|
|
@ -342,6 +379,20 @@ public class BlocksControllerTest {
|
|||
assertEquals(ChatInteractionBlockType.TYPE, block.getName());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void getMcpAgentChatExampleForType() {
|
||||
Block<MCPAgentChatBlockType> block = blocksController.getExampleForType(MCPAgentChatBlockType.TYPE);
|
||||
|
||||
assertNotNull(block);
|
||||
assertEquals(MCPAgentChatBlockType.TYPE, block.getType().getName());
|
||||
assertEquals(MCPAgentChatBlockType.TYPE, block.getName());
|
||||
assertNotNull(block.getSpecificConfiguration());
|
||||
assertTrue(((MCPAgentChatBlockConfiguration) block.getSpecificConfiguration()).getMcpServers().isEmpty());
|
||||
assertTrue(block.getInputs().stream().anyMatch(input -> input.getName().equals("model")));
|
||||
assertTrue(block.getOutputs().stream().anyMatch(output -> output.getName().equals("response")));
|
||||
assertTrue(block.getOutputs().stream().anyMatch(output -> output.getName().equals("history") && output.isMultiple()));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void createLlmBlockWithoutConfiguredPromptUsesPromptInput() {
|
||||
Block<LLMBlockType> block = blocksController.create(LLMBlockConfiguration.builder()
|
||||
|
|
|
|||
|
|
@ -219,6 +219,43 @@ public class ExecutionControllerTest {
|
|||
org.junit.jupiter.api.Assertions.assertEquals(ExecutionStatus.READY, executionObject.getContext().getStatus());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void cancelExecutionClearsRuntimeStateAndMarksExecutionCancelled() {
|
||||
LLMDescriptor llmDescriptor = LLMDescriptor.builder()
|
||||
.provider("testProvider")
|
||||
.model("testModel")
|
||||
.build();
|
||||
|
||||
Block<HumanInteractionBlockType> reviewBlock = blocksController.create(HumanInteractiveBlockConfiguration.builder()
|
||||
.name("Recruiter review")
|
||||
.actionDescription("Validate candidate")
|
||||
.simulateWith(llmDescriptor)
|
||||
.build());
|
||||
|
||||
FlowCreateRequest request = new FlowCreateRequest(
|
||||
"Human Cancel Flow",
|
||||
"Flow to test execution cancellation",
|
||||
FlowData.builder().block(reviewBlock).build());
|
||||
|
||||
ResponseEntity<FlowView> createdFlow = flowController.createFlow(request, new LoginEntity("testuser", "testpassword"));
|
||||
ExecutionObject executionObject = executionsController.create(createdFlow.getBody().id());
|
||||
|
||||
executionsController.prepareStringInputs(executionObject.getId(), reviewBlock.getId(),
|
||||
reviewBlock.getInputs().getFirst().getName(), "Ada Lovelace");
|
||||
executionsController.start(executionObject.getId());
|
||||
waitForExecutionStatus(executionObject, ExecutionStatus.WAITING);
|
||||
|
||||
ExecutionObject cancelled = executionsController.cancel(executionObject.getId());
|
||||
|
||||
org.junit.jupiter.api.Assertions.assertEquals(ExecutionStatus.CANCELLED, cancelled.getContext().getStatus());
|
||||
org.junit.jupiter.api.Assertions.assertTrue(cancelled.getContext().getStatus().isFinalState());
|
||||
org.junit.jupiter.api.Assertions.assertTrue(cancelled.getContext().getInputs().isEmpty());
|
||||
org.junit.jupiter.api.Assertions.assertTrue(cancelled.getContext().getResult().isEmpty());
|
||||
org.junit.jupiter.api.Assertions.assertTrue(cancelled.getContext().getPartialResult().isEmpty());
|
||||
org.junit.jupiter.api.Assertions.assertTrue(cancelled.getContext().getErrors().isEmpty());
|
||||
org.junit.jupiter.api.Assertions.assertTrue(cancelled.getContext().getWarnings().isEmpty());
|
||||
}
|
||||
|
||||
private void waitForExecutionStatus(ExecutionObject executionObject, ExecutionStatus expectedStatus) {
|
||||
long deadline = System.currentTimeMillis() + 5_000;
|
||||
while (System.currentTimeMillis() < deadline) {
|
||||
|
|
|
|||
|
|
@ -471,6 +471,22 @@ public class FlowControllerTest {
|
|||
assertEquals(FlowViewStatus.DRAFT, createResponse.getBody().status());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void createEmptyFlowReturnsDraftStatus() {
|
||||
FlowCreateRequest request = new FlowCreateRequest(
|
||||
"Empty Flow",
|
||||
"Flow without nodes must remain draft",
|
||||
FlowData.builder().build());
|
||||
|
||||
ResponseEntity<FlowView> createResponse = flowController.createFlow(
|
||||
request,
|
||||
new LoginEntity("testuser", "testpassword"));
|
||||
|
||||
assertTrue(createResponse.getStatusCode().is2xxSuccessful());
|
||||
assertNotNull(createResponse.getBody());
|
||||
assertEquals(FlowViewStatus.DRAFT, createResponse.getBody().status());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void createFlowWithBlankModelReturnsDraftStatus() {
|
||||
Block<LLMBlockType> draftBlock = blocksController.create(LLMBlockConfiguration.builder()
|
||||
|
|
|
|||
|
|
@ -7,12 +7,16 @@ import static org.junit.jupiter.api.Assertions.assertTrue;
|
|||
|
||||
import java.util.List;
|
||||
import java.lang.reflect.Field;
|
||||
import java.util.concurrent.atomic.AtomicReference;
|
||||
import java.util.concurrent.atomic.AtomicInteger;
|
||||
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.mockito.Mockito;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.boot.test.context.SpringBootTest;
|
||||
import org.springframework.boot.test.context.TestConfiguration;
|
||||
import org.springframework.context.annotation.Bean;
|
||||
import org.springframework.test.context.bean.override.mockito.MockitoBean;
|
||||
import org.springframework.test.context.TestPropertySource;
|
||||
import org.springframework.util.ResourceUtils;
|
||||
import org.springframework.web.server.ResponseStatusException;
|
||||
|
|
@ -25,6 +29,7 @@ import it.cnr.isti.workflow.manager.blocks.configurations.BlockConfiguration;
|
|||
import it.cnr.isti.workflow.manager.blocks.configurations.ChatInteractionBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.ChatInteractionInput;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.HumanInteractiveBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.blocks.configurations.MCPAgentChatBlockConfiguration;
|
||||
import it.cnr.isti.workflow.manager.containers.Container;
|
||||
import it.cnr.isti.workflow.manager.containers.configurations.ContainerConfiguration;
|
||||
import it.cnr.isti.workflow.manager.containers.configurations.GenericContainerConfiguration;
|
||||
|
|
@ -45,19 +50,24 @@ import it.cnr.isti.workflow.manager.blocks.configurations.LLMBlockConfiguration;
|
|||
import it.cnr.isti.workflow.manager.blocks.factories.ChatInteractionBlockFactory;
|
||||
import it.cnr.isti.workflow.manager.blocks.factories.HTTPServerCallBlockFactory;
|
||||
import it.cnr.isti.workflow.manager.blocks.factories.LLMBlockFactory;
|
||||
import it.cnr.isti.workflow.manager.blocks.factories.MCPAgentChatBlockFactory;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.ChatInteractionBlockType;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.HTTPServerCallBlockType;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.LLMBlockType;
|
||||
import it.cnr.isti.workflow.manager.blocks.types.MCPAgentChatBlockType;
|
||||
import it.cnr.isti.workflow.manager.ios.IODescriptor;
|
||||
import it.cnr.isti.workflow.manager.ios.IOType;
|
||||
import it.cnr.isti.workflow.manager.llms.ChatMessage;
|
||||
import it.cnr.isti.workflow.manager.llms.LLMDescriptor;
|
||||
import it.cnr.isti.workflow.manager.llms.providers.LLMProvider;
|
||||
import it.cnr.isti.workflow.manager.mcp.MCPAgentService;
|
||||
|
||||
@SpringBootTest
|
||||
@TestPropertySource(locations = "classpath:test.properties")
|
||||
public class ExecutionTest {
|
||||
|
||||
private static final AtomicReference<List<ChatMessage>> lastChatMessages = new AtomicReference<>(List.of());
|
||||
|
||||
@TestConfiguration
|
||||
static class TestConfig {
|
||||
|
||||
|
|
@ -82,6 +92,7 @@ public class ExecutionTest {
|
|||
|
||||
@Override
|
||||
public String chat(String model, List<ChatMessage> messages) {
|
||||
lastChatMessages.set(List.copyOf(messages));
|
||||
return "Chat, " + model + "!";
|
||||
}
|
||||
};
|
||||
|
|
@ -100,6 +111,9 @@ public class ExecutionTest {
|
|||
@Autowired
|
||||
ChatInteractionBlockFactory chatInteractionBlockFactory;
|
||||
|
||||
@Autowired
|
||||
MCPAgentChatBlockFactory mcpAgentChatBlockFactory;
|
||||
|
||||
@Autowired
|
||||
LLMBlockType llmBlockType;
|
||||
|
||||
|
|
@ -112,6 +126,9 @@ public class ExecutionTest {
|
|||
@Autowired
|
||||
IteratorContainerFactory iteratorContainerFactory;
|
||||
|
||||
@MockitoBean
|
||||
MCPAgentService mcpAgentService;
|
||||
|
||||
LLMDescriptor llmBrick = LLMDescriptor.builder()
|
||||
.provider("testProvider")
|
||||
.model("testModel")
|
||||
|
|
@ -205,6 +222,55 @@ public class ExecutionTest {
|
|||
assertTrue(execObject.getContext().getPartialResult().isEmpty());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void chatInteractionExecutionKeepsConversationHistoryAcrossMessages() {
|
||||
Block<ChatInteractionBlockType> chatBlock = chatInteractionBlockFactory.create(ChatInteractionBlockConfiguration.builder()
|
||||
.name("Recruiter Chat")
|
||||
.llmDescriptor(llmBrick)
|
||||
.inputs(List.of(new ChatInteractionInput("cand", IOType.TEXT, false)))
|
||||
.build());
|
||||
|
||||
Flow flow = Flow.builder()
|
||||
.name("Chat flow")
|
||||
.description("Single chat block")
|
||||
.block(chatBlock)
|
||||
.build();
|
||||
|
||||
ExecutionObject execObject = executionsService.createExecution(flow);
|
||||
execObject = executionsService.prepareInput(execObject.getId(), chatBlock.getId(), "cand", "John Doe");
|
||||
execObject = executionsService.startExecution(execObject.getId());
|
||||
|
||||
while (execObject.getContext().getStatus() == ExecutionStatus.RUNNING) {
|
||||
execObject = executionsService.getExecution(execObject.getId());
|
||||
}
|
||||
|
||||
execObject.setInteractionValue(chatBlock.getId(), ChatInteractionBlockFactory.INTERACTION_FIELD, "Hello ${{cand}}");
|
||||
execObject = executionsService.getExecution(execObject.getId());
|
||||
|
||||
List<ChatMessage> firstTurn = lastChatMessages.get();
|
||||
assertEquals(1, firstTurn.size());
|
||||
assertEquals(ChatMessage.Role.USER, firstTurn.getFirst().role());
|
||||
assertEquals("Hello John Doe", firstTurn.getFirst().content());
|
||||
|
||||
execObject.setInteractionValue(chatBlock.getId(), ChatInteractionBlockFactory.INTERACTION_FIELD,
|
||||
"Can you summarize the previous answer for ${{cand}}?");
|
||||
execObject = executionsService.getExecution(execObject.getId());
|
||||
|
||||
List<ChatMessage> secondTurn = lastChatMessages.get();
|
||||
assertEquals(3, secondTurn.size());
|
||||
assertEquals(ChatMessage.Role.USER, secondTurn.get(0).role());
|
||||
assertEquals("Hello John Doe", secondTurn.get(0).content());
|
||||
assertEquals(ChatMessage.Role.ASSISTANT, secondTurn.get(1).role());
|
||||
assertEquals("Chat, testModel!", secondTurn.get(1).content());
|
||||
assertEquals(ChatMessage.Role.USER, secondTurn.get(2).role());
|
||||
assertEquals("Can you summarize the previous answer for John Doe?", secondTurn.get(2).content());
|
||||
|
||||
Object partialConversation = execObject.getContext().getPartialResult()
|
||||
.get(new FieldKey(chatBlock.getId(), ChatInteractionBlockFactory.HISTORY_OUTPUT));
|
||||
assertTrue(partialConversation instanceof List<?>);
|
||||
assertEquals(4, ((List<?>) partialConversation).size());
|
||||
}
|
||||
|
||||
@Test
|
||||
public void createExecutionRejectsChatInteractionWithoutLlmDescriptor() {
|
||||
Block<ChatInteractionBlockType> chatBlock = chatInteractionBlockFactory
|
||||
|
|
@ -223,6 +289,109 @@ public class ExecutionTest {
|
|||
assertTrue(exception.getReason().contains("\"field\":\"specificConfiguration.llmDescriptor\""));
|
||||
}
|
||||
|
||||
@Test
|
||||
public void mcpAgentChatExecutionOpensSessionAtFirstMessageAndClosesOnFinalResponse() {
|
||||
AtomicInteger queryCount = new AtomicInteger();
|
||||
Mockito.when(mcpAgentService.openSession(Mockito.eq("llama3.1:8b"), Mockito.anyList()))
|
||||
.thenReturn("session-1");
|
||||
Mockito.when(mcpAgentService.querySession(Mockito.eq("session-1"), Mockito.anyString()))
|
||||
.thenAnswer(invocation -> "MCP answer " + queryCount.incrementAndGet());
|
||||
|
||||
Block<MCPAgentChatBlockType> chatBlock = mcpAgentChatBlockFactory.create(MCPAgentChatBlockConfiguration.builder()
|
||||
.name("MCP Chat")
|
||||
.model("llama3.1:8b")
|
||||
.inputs(List.of(new ChatInteractionInput("cand", IOType.TEXT, false)))
|
||||
.mcpServers(List.of())
|
||||
.build());
|
||||
|
||||
Flow flow = Flow.builder()
|
||||
.name("MCP Chat flow")
|
||||
.description("Single MCP chat block")
|
||||
.block(chatBlock)
|
||||
.build();
|
||||
|
||||
ExecutionObject execObject = executionsService.createExecution(flow);
|
||||
execObject = executionsService.prepareInput(execObject.getId(), chatBlock.getId(), "cand", "John Doe");
|
||||
execObject = executionsService.startExecution(execObject.getId());
|
||||
while (execObject.getContext().getStatus() == ExecutionStatus.RUNNING) {
|
||||
execObject = executionsService.getExecution(execObject.getId());
|
||||
}
|
||||
|
||||
execObject.setInteractionValue(chatBlock.getId(), MCPAgentChatBlockFactory.INTERACTION_FIELD, "Hello ${{cand}}");
|
||||
execObject = executionsService.getExecution(execObject.getId());
|
||||
|
||||
Object partialConversation = execObject.getContext().getPartialResult()
|
||||
.get(new FieldKey(chatBlock.getId(), MCPAgentChatBlockFactory.HISTORY_OUTPUT));
|
||||
assertTrue(partialConversation instanceof List<?>);
|
||||
assertEquals(2, ((List<?>) partialConversation).size());
|
||||
assertTrue(((List<?>) partialConversation).contains("[USER] Hello John Doe"));
|
||||
assertTrue(((List<?>) partialConversation).contains("[ASSISTANT] MCP answer 1"));
|
||||
|
||||
execObject.setInteractionValue(chatBlock.getId(), MCPAgentChatBlockFactory.INTERACTION_FIELD,
|
||||
"Continue with ${{cand}}");
|
||||
execObject = executionsService.getExecution(execObject.getId());
|
||||
|
||||
Object updatedConversation = execObject.getContext().getPartialResult()
|
||||
.get(new FieldKey(chatBlock.getId(), MCPAgentChatBlockFactory.HISTORY_OUTPUT));
|
||||
assertTrue(updatedConversation instanceof List<?>);
|
||||
assertEquals(4, ((List<?>) updatedConversation).size());
|
||||
assertTrue(((List<?>) updatedConversation).contains("[ASSISTANT] MCP answer 2"));
|
||||
|
||||
execObject.setInteractionValue(chatBlock.getId(), MCPAgentChatBlockFactory.FINAL_RESPONSE_FIELD,
|
||||
"Candidate approved");
|
||||
execObject = executionsService.getExecution(execObject.getId());
|
||||
|
||||
assertEquals(ExecutionStatus.SUCCESS, execObject.getContext().getStatus());
|
||||
Object response = execObject.getContext().getResult()
|
||||
.get(new FieldKey(chatBlock.getId(), MCPAgentChatBlockFactory.RESPONSE_OUTPUT));
|
||||
assertEquals("Candidate approved", response);
|
||||
|
||||
Mockito.verify(mcpAgentService, Mockito.times(1)).openSession(Mockito.eq("llama3.1:8b"), Mockito.anyList());
|
||||
Mockito.verify(mcpAgentService).querySession("session-1", "Hello John Doe");
|
||||
Mockito.verify(mcpAgentService).querySession("session-1", "Continue with John Doe");
|
||||
Mockito.verify(mcpAgentService, Mockito.times(1)).closeSessionQuietly("session-1");
|
||||
}
|
||||
|
||||
@Test
|
||||
public void cancelExecutionClosesOpenMcpAgentChatSession() {
|
||||
Mockito.when(mcpAgentService.openSession(Mockito.eq("llama3.1:8b"), Mockito.anyList()))
|
||||
.thenReturn("session-cancel");
|
||||
Mockito.when(mcpAgentService.querySession("session-cancel", "Hello John Doe"))
|
||||
.thenReturn("MCP answer");
|
||||
|
||||
Block<MCPAgentChatBlockType> chatBlock = mcpAgentChatBlockFactory.create(MCPAgentChatBlockConfiguration.builder()
|
||||
.name("MCP Chat")
|
||||
.model("llama3.1:8b")
|
||||
.inputs(List.of(new ChatInteractionInput("cand", IOType.TEXT, false)))
|
||||
.mcpServers(List.of())
|
||||
.build());
|
||||
|
||||
Flow flow = Flow.builder()
|
||||
.name("MCP Chat flow")
|
||||
.description("Single MCP chat block")
|
||||
.block(chatBlock)
|
||||
.build();
|
||||
|
||||
ExecutionObject execObject = executionsService.createExecution(flow);
|
||||
execObject = executionsService.prepareInput(execObject.getId(), chatBlock.getId(), "cand", "John Doe");
|
||||
execObject = executionsService.startExecution(execObject.getId());
|
||||
while (execObject.getContext().getStatus() == ExecutionStatus.RUNNING) {
|
||||
execObject = executionsService.getExecution(execObject.getId());
|
||||
}
|
||||
|
||||
execObject.setInteractionValue(chatBlock.getId(), MCPAgentChatBlockFactory.INTERACTION_FIELD, "Hello ${{cand}}");
|
||||
execObject = executionsService.getExecution(execObject.getId());
|
||||
|
||||
assertEquals(ExecutionStatus.WAITING, execObject.getContext().getStatus());
|
||||
assertFalse(execObject.getContext().getPartialResult().isEmpty());
|
||||
|
||||
execObject = executionsService.cancelExecution(execObject.getId());
|
||||
|
||||
assertEquals(ExecutionStatus.CANCELLED, execObject.getContext().getStatus());
|
||||
assertTrue(execObject.getContext().getPartialResult().isEmpty());
|
||||
Mockito.verify(mcpAgentService, Mockito.times(1)).closeSessionQuietly("session-cancel");
|
||||
}
|
||||
|
||||
@Test
|
||||
public void createInteractiveExecutionSetInputAndStart() {
|
||||
ExecutionObject eo = createInteractiveExecutionAndSetInputInternally();
|
||||
|
|
|
|||
Loading…
Reference in New Issue