From b2a5cd2673c0c2edb9ef76aa457ed5417809482f Mon Sep 17 00:00:00 2001 From: Lucio Lelii Date: Wed, 11 Mar 2026 12:45:22 +0100 Subject: [PATCH] Implement session-based assistant orchestration and harden flow assembly --- .../manager/HumainFlowApplication.java | 6 + .../manager/app/config/WebConfig.java | 9 + .../AssistantConversationService.java | 294 ++++++++++++ .../assistant/BlockCatalogService.java | 113 +++++ .../assistant/FlowAssistantPromptService.java | 218 ++++++--- .../assistant/FlowAssistantService.java | 448 ++++++++++++++++-- .../model/AssistantCallAcceptedResponse.java | 4 + .../assistant/model/AssistantCallStatus.java | 8 + .../assistant/model/AssistantCallView.java | 17 + .../assistant/model/AssistantConfigView.java | 7 + .../assistant/model/AssistantIntent.java | 8 + .../assistant/model/AssistantMessageRole.java | 6 + .../assistant/model/AssistantMessageView.java | 11 + .../model/AssistantSessionCreateRequest.java | 6 + .../model/AssistantSessionMessageRequest.java | 6 + .../assistant/model/AssistantSessionView.java | 19 + .../controllers/ApiExceptionHandler.java | 40 +- .../controllers/AssistantController.java | 64 +++ .../ollama/InternalOllamaLLMProvider.java | 10 +- src/main/resources/application.properties | 2 +- .../controllers/AssistantControllerTest.java | 199 +++++++- src/test/resources/test.properties | 3 +- 22 files changed, 1360 insertions(+), 138 deletions(-) create mode 100644 src/main/java/it/cnr/isti/workflow/manager/assistant/AssistantConversationService.java create mode 100644 src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantCallAcceptedResponse.java create mode 100644 src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantCallStatus.java create mode 100644 src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantCallView.java create mode 100644 src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantConfigView.java create mode 100644 src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantIntent.java create mode 100644 src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantMessageRole.java create mode 100644 src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantMessageView.java create mode 100644 src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantSessionCreateRequest.java create mode 100644 src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantSessionMessageRequest.java create mode 100644 src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantSessionView.java diff --git a/src/main/java/it/cnr/isti/workflow/manager/HumainFlowApplication.java b/src/main/java/it/cnr/isti/workflow/manager/HumainFlowApplication.java index 7290153..fde3b3c 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/HumainFlowApplication.java +++ b/src/main/java/it/cnr/isti/workflow/manager/HumainFlowApplication.java @@ -1,5 +1,7 @@ package it.cnr.isti.workflow.manager; +import java.util.Locale; + import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.CommandLineRunner; import org.springframework.boot.SpringApplication; @@ -13,6 +15,10 @@ import it.cnr.isti.workflow.manager.flows.FlowImportComponent; @SpringBootApplication public class HumainFlowApplication { + static { + Locale.setDefault(Locale.ENGLISH); + } + @Autowired UserImportComponent userImportComponent; diff --git a/src/main/java/it/cnr/isti/workflow/manager/app/config/WebConfig.java b/src/main/java/it/cnr/isti/workflow/manager/app/config/WebConfig.java index de555b0..1f69699 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/app/config/WebConfig.java +++ b/src/main/java/it/cnr/isti/workflow/manager/app/config/WebConfig.java @@ -1,10 +1,14 @@ package it.cnr.isti.workflow.manager.app.config; +import java.util.Locale; + import org.springframework.beans.factory.annotation.Value; import org.springframework.context.annotation.Configuration; import org.springframework.lang.NonNull; +import org.springframework.web.servlet.LocaleResolver; import org.springframework.web.servlet.config.annotation.CorsRegistry; import org.springframework.web.servlet.config.annotation.WebMvcConfigurer; +import org.springframework.web.servlet.i18n.FixedLocaleResolver; import jakarta.validation.Validator; @@ -32,4 +36,9 @@ public class WebConfig implements WebMvcConfigurer { return new LocalValidatorFactoryBean(); } + @Bean + public LocaleResolver localeResolver() { + return new FixedLocaleResolver(Locale.ENGLISH); + } + } diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/AssistantConversationService.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/AssistantConversationService.java new file mode 100644 index 0000000..10862c6 --- /dev/null +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/AssistantConversationService.java @@ -0,0 +1,294 @@ +package it.cnr.isti.workflow.manager.assistant; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.List; +import java.util.Locale; +import java.util.Objects; +import java.util.UUID; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; + +import org.springframework.stereotype.Service; +import org.springframework.web.server.ResponseStatusException; +import org.springframework.http.HttpStatus; + +import it.cnr.isti.workflow.manager.assistant.model.AssistantCallAcceptedResponse; +import it.cnr.isti.workflow.manager.assistant.model.AssistantCallStatus; +import it.cnr.isti.workflow.manager.assistant.model.AssistantCallView; +import it.cnr.isti.workflow.manager.assistant.model.AssistantExplainRequest; +import it.cnr.isti.workflow.manager.assistant.model.AssistantExplainResponse; +import it.cnr.isti.workflow.manager.assistant.model.AssistantFixRequest; +import it.cnr.isti.workflow.manager.assistant.model.AssistantFlowResponse; +import it.cnr.isti.workflow.manager.assistant.model.AssistantGenerationRequest; +import it.cnr.isti.workflow.manager.assistant.model.AssistantIntent; +import it.cnr.isti.workflow.manager.assistant.model.AssistantMessageRole; +import it.cnr.isti.workflow.manager.assistant.model.AssistantMessageView; +import it.cnr.isti.workflow.manager.assistant.model.AssistantRefineRequest; +import it.cnr.isti.workflow.manager.assistant.model.AssistantSessionCreateRequest; +import it.cnr.isti.workflow.manager.assistant.model.AssistantSessionMessageRequest; +import it.cnr.isti.workflow.manager.assistant.model.AssistantSessionView; +import it.cnr.isti.workflow.manager.flows.model.FlowCreateRequest; +import it.cnr.isti.workflow.manager.flows.validation.ValidationError; +import jakarta.annotation.PreDestroy; + +@Service +public class AssistantConversationService { + + private static final int DEFAULT_MAX_REPAIR_ATTEMPTS = 1; + + private final ConcurrentHashMap sessions = new ConcurrentHashMap<>(); + private final ConcurrentHashMap calls = new ConcurrentHashMap<>(); + private final ExecutorService executor = Executors.newVirtualThreadPerTaskExecutor(); + + private final FlowAssistantService flowAssistantService; + + public AssistantConversationService(FlowAssistantService flowAssistantService) { + this.flowAssistantService = flowAssistantService; + } + + public AssistantSessionView createSession(String owner, AssistantSessionCreateRequest request) { + Instant now = Instant.now(); + SessionState session = new SessionState( + UUID.randomUUID().toString(), + owner, + request.model(), + now, + now); + sessions.put(session.id, session); + return session.toView(); + } + + public AssistantSessionView getSession(String sessionId, String owner) { + return requireSession(sessionId, owner).toView(); + } + + public AssistantCallView getCall(String callId, String owner) { + CallState call = calls.get(callId); + if (call == null) { + throw new ResponseStatusException(HttpStatus.NOT_FOUND, "Assistant call not found"); + } + SessionState session = requireSession(call.sessionId, owner); + if (!Objects.equals(session.id, call.sessionId)) { + throw new ResponseStatusException(HttpStatus.FORBIDDEN, "Assistant call does not belong to the user"); + } + return call.toView(); + } + + public AssistantCallAcceptedResponse submitMessage(String sessionId, String owner, AssistantSessionMessageRequest request) { + SessionState session = requireSession(sessionId, owner); + session.appendMessage(AssistantMessageRole.USER, request.message(), null); + + Instant now = Instant.now(); + CallState call = new CallState(UUID.randomUUID().toString(), sessionId, now, now); + calls.put(call.id, call); + session.lastCallId = call.id; + session.touch(); + + executor.submit(() -> processCall(session, call, request.message())); + return new AssistantCallAcceptedResponse(sessionId, call.id); + } + + @PreDestroy + void shutdown() { + executor.shutdownNow(); + } + + private void processCall(SessionState session, CallState call, String message) { + try { + call.status = AssistantCallStatus.RUNNING; + call.updatePhase("routing", "Routing the request"); + AssistantIntent intent = inferIntent(session, message); + call.intent = intent; + + String contextualPrompt = buildContextualPrompt(session, message); + + switch (intent) { + case DRAFT -> { + AssistantFlowResponse result = flowAssistantService.draft( + new AssistantGenerationRequest(contextualPrompt, session.model, DEFAULT_MAX_REPAIR_ATTEMPTS), + call::updatePhase); + call.flowResult = result; + session.currentFlow = result.flow(); + session.lastValidationErrors = result.validationErrors(); + session.appendMessage(AssistantMessageRole.ASSISTANT, buildFlowMessage(intent, result), call.id); + } + case REFINE -> { + AssistantFlowResponse result = flowAssistantService.refine( + new AssistantRefineRequest(contextualPrompt, session.currentFlow, session.model, + DEFAULT_MAX_REPAIR_ATTEMPTS), + call::updatePhase); + call.flowResult = result; + session.currentFlow = result.flow(); + session.lastValidationErrors = result.validationErrors(); + session.appendMessage(AssistantMessageRole.ASSISTANT, buildFlowMessage(intent, result), call.id); + } + case FIX -> { + AssistantFlowResponse result = flowAssistantService.fix( + new AssistantFixRequest(contextualPrompt, session.currentFlow, session.lastValidationErrors, + session.model, DEFAULT_MAX_REPAIR_ATTEMPTS), + call::updatePhase); + call.flowResult = result; + session.currentFlow = result.flow(); + session.lastValidationErrors = result.validationErrors(); + session.appendMessage(AssistantMessageRole.ASSISTANT, buildFlowMessage(intent, result), call.id); + } + case EXPLAIN -> { + AssistantExplainResponse result = flowAssistantService.explain( + new AssistantExplainRequest(session.currentFlow, contextualPrompt, session.model)); + call.explainResult = result; + session.appendMessage(AssistantMessageRole.ASSISTANT, result.explanation(), call.id); + } + } + + call.status = AssistantCallStatus.COMPLETED; + call.updatePhase("completed", "Assistant request completed"); + session.touch(); + } catch (Exception e) { + call.status = AssistantCallStatus.FAILED; + call.errorMessage = e.getMessage(); + call.updatePhase("failed", "Assistant request failed"); + session.appendMessage(AssistantMessageRole.ASSISTANT, + "The assistant request failed: " + (e.getMessage() == null ? "unknown error" : e.getMessage()), + call.id); + session.touch(); + } + } + + private AssistantIntent inferIntent(SessionState session, String message) { + String normalized = message.toLowerCase(Locale.ROOT); + if (normalized.contains("explain") || normalized.contains("what does") || normalized.contains("why")) { + return AssistantIntent.EXPLAIN; + } + if (session.currentFlow == null) { + return AssistantIntent.DRAFT; + } + boolean flowInvalid = session.lastValidationErrors != null && !session.lastValidationErrors.isEmpty(); + if (flowInvalid || normalized.contains("fix") || normalized.contains("problem") || normalized.contains("error")) { + return AssistantIntent.FIX; + } + return AssistantIntent.REFINE; + } + + private String buildContextualPrompt(SessionState session, String message) { + List messages = session.toView().messages(); + int fromIndex = Math.max(0, messages.size() - 6); + List recent = messages.subList(fromIndex, messages.size()); + StringBuilder builder = new StringBuilder(); + builder.append("Conversation context:\n"); + for (AssistantMessageView item : recent) { + builder.append(item.role().name()).append(": ").append(item.content()).append("\n"); + } + builder.append("\nCurrent user request:\n").append(message); + return builder.toString(); + } + + private String buildFlowMessage(AssistantIntent intent, AssistantFlowResponse result) { + String rationale = result.assistantRationale() == null || result.assistantRationale().isBlank() + ? "" + : result.assistantRationale().trim(); + String validity = result.valid() ? "The flow is valid." : "The flow still has validation issues."; + return switch (intent) { + case DRAFT -> (rationale + " " + validity).trim(); + case REFINE -> (rationale + " " + validity).trim(); + case FIX -> (rationale + " " + validity).trim(); + case EXPLAIN -> rationale; + }; + } + + private SessionState requireSession(String sessionId, String owner) { + SessionState session = sessions.get(sessionId); + if (session == null) { + throw new ResponseStatusException(HttpStatus.NOT_FOUND, "Assistant session not found"); + } + if (!Objects.equals(session.owner, owner)) { + throw new ResponseStatusException(HttpStatus.FORBIDDEN, "Assistant session does not belong to the user"); + } + return session; + } + + private static final class SessionState { + private final String id; + private final String owner; + private final String model; + private final Instant createdAt; + private volatile Instant updatedAt; + private volatile String lastCallId; + private volatile FlowCreateRequest currentFlow; + private volatile List lastValidationErrors = List.of(); + private final List messages = java.util.Collections.synchronizedList(new ArrayList<>()); + + private SessionState(String id, String owner, String model, Instant createdAt, Instant updatedAt) { + this.id = id; + this.owner = owner; + this.model = model; + this.createdAt = createdAt; + this.updatedAt = updatedAt; + } + + private void appendMessage(AssistantMessageRole role, String content, String callId) { + messages.add(new AssistantMessageView(UUID.randomUUID().toString(), role, content, Instant.now(), callId)); + touch(); + } + + private void touch() { + updatedAt = Instant.now(); + } + + private AssistantSessionView toView() { + return new AssistantSessionView( + id, + owner, + model, + createdAt, + updatedAt, + lastCallId, + currentFlow, + lastValidationErrors == null ? List.of() : List.copyOf(lastValidationErrors), + List.copyOf(messages)); + } + } + + private static final class CallState { + private final String id; + private final String sessionId; + private final Instant createdAt; + private volatile Instant updatedAt; + private volatile AssistantCallStatus status = AssistantCallStatus.QUEUED; + private volatile String phase = "queued"; + private volatile String progressMessage = "Assistant request queued"; + private volatile AssistantIntent intent; + private volatile AssistantFlowResponse flowResult; + private volatile AssistantExplainResponse explainResult; + private volatile String errorMessage; + + private CallState(String id, String sessionId, Instant createdAt, Instant updatedAt) { + this.id = id; + this.sessionId = sessionId; + this.createdAt = createdAt; + this.updatedAt = updatedAt; + } + + private void updatePhase(String phase, String progressMessage) { + this.phase = phase; + this.progressMessage = progressMessage; + this.updatedAt = Instant.now(); + } + + private AssistantCallView toView() { + return new AssistantCallView( + id, + sessionId, + status, + phase, + progressMessage, + createdAt, + updatedAt, + intent, + flowResult, + explainResult, + errorMessage); + } + } +} diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/BlockCatalogService.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/BlockCatalogService.java index 2b46338..cccbc91 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/assistant/BlockCatalogService.java +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/BlockCatalogService.java @@ -1,18 +1,22 @@ package it.cnr.isti.workflow.manager.assistant; +import java.util.Comparator; import java.util.List; import java.util.Map; +import java.util.Set; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.stereotype.Service; import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.node.ArrayNode; import it.cnr.isti.workflow.manager.blocks.Block; import it.cnr.isti.workflow.manager.blocks.configurations.BlockConfiguration; import it.cnr.isti.workflow.manager.blocks.configurations.JsonSchemaProducer; import it.cnr.isti.workflow.manager.blocks.factories.BlockFactory; import it.cnr.isti.workflow.manager.blocks.types.BlockType; +import it.cnr.isti.workflow.manager.ios.IODescriptor; @Service public class BlockCatalogService { @@ -26,6 +30,24 @@ public class BlockCatalogService { Block exampleBlock) { } + public record AssistantPromptFieldDescriptor( + String name, + String type, + boolean required, + String placeholder, + boolean structural) { + } + + public record AssistantPromptBlockDescriptor( + String type, + String description, + boolean userInteractive, + String configurationType, + List configurationFields, + List inputs, + List outputs) { + } + @Autowired private Map blockTypes; @@ -42,6 +64,12 @@ public class BlockCatalogService { .toList(); } + public List getPromptCatalog() { + return getCatalog().stream() + .map(this::toPromptDescriptor) + .toList(); + } + @SuppressWarnings("unchecked") private AssistantBlockDescriptor toDescriptor(BlockType blockType) { Class> configurationClass = blockType.getBlockConfigurationClass(); @@ -63,4 +91,89 @@ public class BlockCatalogService { schema, exampleBlock); } + + private AssistantPromptBlockDescriptor toPromptDescriptor(AssistantBlockDescriptor descriptor) { + return new AssistantPromptBlockDescriptor( + descriptor.type(), + descriptor.description(), + descriptor.userInteractive(), + extractConfigurationType(descriptor), + extractConfigurationFields(descriptor.schema()), + extractIoNames(descriptor.exampleBlock(), true), + extractIoNames(descriptor.exampleBlock(), false)); + } + + private String extractConfigurationType(AssistantBlockDescriptor descriptor) { + if (descriptor.schema() == null) { + return null; + } + JsonNode typeNode = descriptor.schema().path("properties").path("type").path("enum"); + if (typeNode instanceof ArrayNode enumValues && !enumValues.isEmpty()) { + return enumValues.get(0).asText(); + } + return descriptor.configurationClass(); + } + + private List extractConfigurationFields(JsonNode schema) { + if (schema == null || !schema.has("properties")) { + return List.of(); + } + + Set requiredFields = extractRequiredFields(schema.path("required")); + return iterable(schema.path("properties").fields()).stream() + .filter(entry -> !"type".equals(entry.getKey())) + .map(entry -> new AssistantPromptFieldDescriptor( + entry.getKey(), + extractFieldType(entry.getValue()), + requiredFields.contains(entry.getKey()), + entry.getValue().path("x-ui-placeholder").asText(null), + entry.getValue().path("x-ui-structural").asBoolean(false))) + .sorted(Comparator + .comparing(AssistantPromptFieldDescriptor::required).reversed() + .thenComparing(AssistantPromptFieldDescriptor::name, String.CASE_INSENSITIVE_ORDER)) + .toList(); + } + + private List extractIoNames(Block exampleBlock, boolean inputs) { + if (exampleBlock == null) { + return List.of(); + } + List descriptors = inputs ? exampleBlock.getInputs() : exampleBlock.getOutputs(); + if (descriptors == null) { + return List.of(); + } + return descriptors.stream().map(IODescriptor::getName).toList(); + } + + private Set extractRequiredFields(JsonNode requiredNode) { + if (!(requiredNode instanceof ArrayNode requiredArray)) { + return Set.of(); + } + java.util.LinkedHashSet result = new java.util.LinkedHashSet<>(); + for (JsonNode node : requiredArray) { + result.add(node.asText()); + } + return result; + } + + private String extractFieldType(JsonNode node) { + if (node == null || node.isMissingNode()) { + return "unknown"; + } + if (node.has("type")) { + return node.get("type").asText(); + } + if (node.has("$ref")) { + String ref = node.get("$ref").asText(); + int separator = ref.lastIndexOf('/'); + return separator >= 0 ? ref.substring(separator + 1) : ref; + } + return "object"; + } + + private List> iterable(java.util.Iterator> fields) { + java.util.ArrayList> entries = new java.util.ArrayList<>(); + fields.forEachRemaining(entries::add); + return entries; + } } diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/FlowAssistantPromptService.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/FlowAssistantPromptService.java index abd72fe..412ffab 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/assistant/FlowAssistantPromptService.java +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/FlowAssistantPromptService.java @@ -11,81 +11,41 @@ import it.cnr.isti.workflow.manager.flows.validation.ValidationError; @Service public class FlowAssistantPromptService { - public String buildDraftPrompt(String userPrompt, List catalog) { + public enum OperationMode { + DRAFT, + REFINE, + FIX + } + + public String buildPlanPrompt(OperationMode mode, String userPrompt, FlowCreateRequest currentFlow, + List errors, List catalog) { return """ - TASK: DRAFT - You are a workflow planner. Return only JSON. - Produce a JSON object with this exact shape: + TASK: PLAN + MODE: %s + You are planning a workflow using the available block types. + Return only JSON with this exact shape: { "rationale": "short explanation", - "flow": { + "plan": { "name": "...", "description": "...", - "flow": { - "blocks": [], - "connections": [] - } + "blocks": [ + { + "blockId": "b1", + "blockType": "LLMBlock", + "purpose": "..." + } + ] } } Rules: - - Use only block types present in the catalog. - - Do not invent fields or connection names. - - Each block must include specificConfiguration. - - Prefer placeholders like ${{variable}} when an input must be derived from previous blocks. - - Return valid JSON with no markdown fences. - - Available block catalog: - %s - - User request: - %s - """.formatted(toJson(catalog), userPrompt); - } - - public String buildRefinePrompt(FlowCreateRequest currentFlow, String userPrompt, - List catalog) { - return """ - TASK: REFINE - You are updating an existing workflow. Return only JSON. - Produce the same wrapper object used for draft: - { - "rationale": "short explanation", - "flow": { ... FlowCreateRequest ... } - } - - Rules: - - Preserve valid parts of the current flow when possible. - - Use only block types from the catalog. - - Keep ids stable when you can. New blocks may have new ids. - - Return valid JSON with no markdown fences. - - Available block catalog: - %s - - Current flow: - %s - - User refinement request: - %s - """.formatted(toJson(catalog), toJson(currentFlow), userPrompt); - } - - public String buildFixPrompt(FlowCreateRequest currentFlow, List errors, - List catalog, String userPrompt) { - return """ - TASK: FIX - You are repairing a workflow that failed backend validation. Return only JSON. - Produce the same wrapper object used for draft: - { - "rationale": "short explanation", - "flow": { ... FlowCreateRequest ... } - } - - Rules: - - Fix only what is necessary to resolve validation errors. - - Keep ids stable when possible. - Use only block types from the catalog. + - Keep the plan minimal. + - One block per logical action. + - blockId must be stable symbolic ids like b1, b2, b3. + - Do not return block configuration yet. + - Do not return connections yet. - Return valid JSON with no markdown fences. Available block catalog: @@ -97,9 +57,119 @@ public class FlowAssistantPromptService { Validation errors: %s - Additional user context: + User request: %s - """.formatted(toJson(catalog), toJson(currentFlow), toJson(errors), + """.formatted( + mode.name(), + toJson(catalog), + summarizeFlow(currentFlow), + summarizeErrors(errors), + userPrompt == null || userPrompt.isBlank() ? "(none)" : userPrompt); + } + + public String buildBlockConfigurationPrompt(OperationMode mode, String userPrompt, + BlockCatalogService.AssistantPromptBlockDescriptor descriptor, Object flowPlan, Object blockPlan, + FlowCreateRequest currentFlow, List errors, String selectedModel) { + return """ + TASK: BLOCK_CONFIG + MODE: %s + You are configuring one workflow block. + Return only JSON with this exact shape: + { + "rationale": "short explanation", + "block": { + "blockId": "...", + "name": "...", + "config": { } + } + } + + Rules: + - Configure only the requested block. + - The config object must contain only task-specific fields. + - Do not include system-managed fields like provider, model, llmDescriptor, simulateWith, ids, inputs, outputs. + - Use placeholders like ${{variable}} when needed. + - Return valid JSON with no markdown fences. + + Selected internal model: + %s + + Flow plan: + %s + + Current block to configure: + %s + + Current flow: + %s + + Validation errors: + %s + + Block descriptor: + %s + + User request: + %s + """.formatted( + mode.name(), + selectedModel, + toJson(flowPlan), + toJson(blockPlan), + summarizeFlow(currentFlow), + summarizeErrors(errors), + toJson(descriptor), + userPrompt == null || userPrompt.isBlank() ? "(none)" : userPrompt); + } + + public String buildConnectionsPrompt(OperationMode mode, String userPrompt, Object flowPlan, Object configuredBlocks, + FlowCreateRequest currentFlow, List errors) { + return """ + TASK: CONNECTIONS + MODE: %s + You are connecting already configured workflow blocks. + Return only JSON with this exact shape: + { + "rationale": "short explanation", + "connections": [ + { + "fromBlockId": "b1", + "fromOutput": "response", + "toBlockId": "b2", + "toInput": "input" + } + ] + } + + Rules: + - Use only block ids from the flow plan. + - Copy block ids exactly as provided in the flow plan, for example b1, b2, b3. + - Never use block names or purposes in fromBlockId/toBlockId. + - Use only output/input names that exist in the configured blocks. + - Allowed block ids are listed in the flow plan and configured blocks. Use only those exact values. + - Return the minimal set of connections required by the user request. + - Return valid JSON with no markdown fences. + + Flow plan: + %s + + Configured blocks: + %s + + Current flow: + %s + + Validation errors: + %s + + User request: + %s + """.formatted( + mode.name(), + toJson(flowPlan), + toJson(configuredBlocks), + summarizeFlow(currentFlow), + summarizeErrors(errors), userPrompt == null || userPrompt.isBlank() ? "(none)" : userPrompt); } @@ -125,4 +195,18 @@ public class FlowAssistantPromptService { throw new IllegalStateException("Unable to serialize assistant prompt payload", e); } } + + private String summarizeFlow(FlowCreateRequest flow) { + if (flow == null || flow.flow() == null) { + return "(current flow unavailable or structurally invalid)"; + } + return toJson(flow); + } + + private String summarizeErrors(List errors) { + if (errors == null || errors.isEmpty()) { + return "(none)"; + } + return toJson(errors); + } } diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/FlowAssistantService.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/FlowAssistantService.java index 4aac070..e9fb9f2 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/assistant/FlowAssistantService.java +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/FlowAssistantService.java @@ -1,9 +1,13 @@ package it.cnr.isti.workflow.manager.assistant; import java.util.ArrayList; +import java.util.Collection; +import java.util.LinkedHashMap; +import java.util.LinkedHashSet; import java.util.List; import java.util.Map; import java.util.Objects; +import java.util.Locale; import java.util.Set; import org.springframework.beans.factory.annotation.Autowired; @@ -12,17 +16,25 @@ import org.springframework.stereotype.Service; import org.springframework.web.server.ResponseStatusException; import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.node.ObjectNode; import it.cnr.isti.workflow.manager.app.ObjectMapperHolder; +import it.cnr.isti.workflow.manager.assistant.FlowAssistantPromptService.OperationMode; import it.cnr.isti.workflow.manager.assistant.model.AssistantExplainRequest; import it.cnr.isti.workflow.manager.assistant.model.AssistantExplainResponse; import it.cnr.isti.workflow.manager.assistant.model.AssistantFixRequest; import it.cnr.isti.workflow.manager.assistant.model.AssistantFlowResponse; import it.cnr.isti.workflow.manager.assistant.model.AssistantGenerationRequest; import it.cnr.isti.workflow.manager.assistant.model.AssistantRefineRequest; +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.flows.model.Connection; import it.cnr.isti.workflow.manager.flows.model.FlowCreateRequest; +import it.cnr.isti.workflow.manager.flows.model.FlowData; import it.cnr.isti.workflow.manager.flows.validation.ValidationError; import it.cnr.isti.workflow.manager.flows.validation.ValidationErrorCodec; +import it.cnr.isti.workflow.manager.ios.IODescriptor; import it.cnr.isti.workflow.manager.llms.providers.LLMProvider; import jakarta.validation.ConstraintViolation; import jakarta.validation.Validator; @@ -31,8 +43,39 @@ import jakarta.validation.Validator; public class FlowAssistantService { private static final String INTERNAL_PROVIDER_NAME = "InternalOllama"; + @FunctionalInterface + public interface ProgressListener { + void onProgress(String phase, String message); + } + private static final ProgressListener NOOP_PROGRESS = (phase, message) -> { + }; - private record ParsedAssistantFlow(FlowCreateRequest flow, String rationale) { + private record AssistantFlowPlan(String name, String description, List blocks) { + } + + private record AssistantBlockPlan(String blockId, String blockType, String purpose) { + } + + private record AssistantConfiguredBlockDraft(String blockId, String name, JsonNode config) { + } + + private record AssistantConnectionDraft(String fromBlockId, String fromOutput, String toBlockId, String toInput) { + } + + private record ParsedPlan(AssistantFlowPlan plan, String rationale) { + } + + private record ParsedBlockDraft(AssistantConfiguredBlockDraft block, String rationale) { + } + + private record ParsedConnections(List connections, String rationale) { + } + + private record ConfiguredBlockSummary(String blockId, String blockType, String name, String purpose, List inputs, + List outputs) { + } + + private record AssembledFlow(FlowCreateRequest flow, String rationale) { } @Autowired @@ -44,67 +87,380 @@ public class FlowAssistantService { @Autowired private FlowAssistantPromptService promptService; + @Autowired + private List> blockFactories; + @Autowired private Validator validator; public AssistantFlowResponse draft(AssistantGenerationRequest request) { - LLMProvider provider = resolveInternalProvider(); - String prompt = promptService.buildDraftPrompt(request.userPrompt(), blockCatalogService.getCatalog()); - return generateFlow(provider, request.model(), prompt, - request.maxRepairAttempts(), request.userPrompt()); + return draft(request, NOOP_PROGRESS); + } + + public AssistantFlowResponse draft(AssistantGenerationRequest request, ProgressListener progressListener) { + return generateFlow(OperationMode.DRAFT, request.userPrompt(), null, List.of(), request.model(), + request.maxRepairAttempts(), progressListener); } public AssistantFlowResponse refine(AssistantRefineRequest request) { - LLMProvider provider = resolveInternalProvider(); - String prompt = promptService.buildRefinePrompt(request.flow(), request.userPrompt(), blockCatalogService.getCatalog()); - return generateFlow(provider, request.model(), prompt, - request.maxRepairAttempts(), request.userPrompt()); + return refine(request, NOOP_PROGRESS); + } + + public AssistantFlowResponse refine(AssistantRefineRequest request, ProgressListener progressListener) { + return generateFlow(OperationMode.REFINE, request.userPrompt(), request.flow(), List.of(), request.model(), + request.maxRepairAttempts(), progressListener); } public AssistantFlowResponse fix(AssistantFixRequest request) { - LLMProvider provider = resolveInternalProvider(); + return fix(request, NOOP_PROGRESS); + } + + public AssistantFlowResponse fix(AssistantFixRequest request, ProgressListener progressListener) { List initialErrors = request.validationErrors() == null || request.validationErrors().isEmpty() ? validate(request.flow()) : request.validationErrors(); - String prompt = promptService.buildFixPrompt(request.flow(), initialErrors, blockCatalogService.getCatalog(), - request.userPrompt()); - return generateFlow(provider, request.model(), prompt, - request.maxRepairAttempts(), request.userPrompt()); + return generateFlow(OperationMode.FIX, request.userPrompt(), request.flow(), initialErrors, request.model(), + request.maxRepairAttempts(), progressListener); } public AssistantExplainResponse explain(AssistantExplainRequest request) { - LLMProvider provider = resolveInternalProvider(); - String prompt = promptService.buildExplainPrompt(request.flow(), request.userPrompt()); - return new AssistantExplainResponse(invokeProvider(provider, request.model(), prompt)); + return explain(request, NOOP_PROGRESS); } - private AssistantFlowResponse generateFlow(LLMProvider provider, String model, - String initialPrompt, Integer maxRepairAttempts, String userPrompt) { + public AssistantExplainResponse explain(AssistantExplainRequest request, ProgressListener progressListener) { + LLMProvider provider = resolveInternalProvider(); + progressListener.onProgress("explaining", "Explaining the current flow"); + String prompt = promptService.buildExplainPrompt(request.flow(), request.userPrompt()); + AssistantExplainResponse response = new AssistantExplainResponse(invokeProvider(provider, request.model(), prompt)); + progressListener.onProgress("completed", "Flow explanation ready"); + return response; + } + + private AssistantFlowResponse generateFlow(OperationMode initialMode, String userPrompt, FlowCreateRequest currentFlow, + List initialErrors, String model, Integer maxRepairAttempts, ProgressListener progressListener) { + LLMProvider provider = resolveInternalProvider(); int allowedRepairs = maxRepairAttempts == null ? 1 : maxRepairAttempts; int repairs = 0; + OperationMode mode = initialMode; + FlowCreateRequest flowContext = currentFlow; + List errorContext = initialErrors == null ? List.of() : initialErrors; + AssembledFlow assembled = null; + List errors = List.of(); - ParsedAssistantFlow generated = parseAssistantFlow(invokeProvider(provider, model, initialPrompt)); - FlowCreateRequest currentFlow = generated.flow(); - List errors = validate(currentFlow); - - while (!errors.isEmpty() && repairs < allowedRepairs) { + while (true) { + assembled = assembleFlow(provider, model, mode, userPrompt, flowContext, errorContext, progressListener); + progressListener.onProgress("validating", "Validating the assembled flow"); + errors = validate(assembled.flow()); + if (errors.isEmpty() || repairs >= allowedRepairs) { + break; + } repairs++; - String repairPrompt = promptService.buildFixPrompt(currentFlow, errors, blockCatalogService.getCatalog(), - userPrompt); - generated = parseAssistantFlow(invokeProvider(provider, model, repairPrompt)); - currentFlow = generated.flow(); - errors = validate(currentFlow); + mode = OperationMode.FIX; + flowContext = assembled.flow(); + errorContext = errors; + progressListener.onProgress("fixing", "Validation failed, retrying with fix mode"); } + progressListener.onProgress("completed", "Assistant flow generation completed"); return new AssistantFlowResponse( - currentFlow, + assembled.flow(), errors.isEmpty(), errors, List.of(), - generated.rationale(), + assembled.rationale(), repairs); } + private AssembledFlow assembleFlow(LLMProvider provider, String model, OperationMode mode, String userPrompt, + FlowCreateRequest currentFlow, List errors, ProgressListener progressListener) { + List catalog = blockCatalogService.getPromptCatalog(); + Map catalogByType = new LinkedHashMap<>(); + for (BlockCatalogService.AssistantPromptBlockDescriptor descriptor : catalog) { + catalogByType.put(descriptor.type(), descriptor); + } + + progressListener.onProgress("planning", "Planning workflow blocks"); + ParsedPlan parsedPlan = parsePlan(invokeProvider(provider, model, + promptService.buildPlanPrompt(mode, userPrompt, currentFlow, errors, catalog))); + validatePlan(parsedPlan.plan()); + + List> assembledBlocks = new ArrayList<>(); + Map> blocksByPlanId = new LinkedHashMap<>(); + Map> blocksByAlias = new LinkedHashMap<>(); + List configuredBlocks = new ArrayList<>(); + List rationaleParts = new ArrayList<>(); + appendRationale(rationaleParts, parsedPlan.rationale()); + + for (AssistantBlockPlan blockPlan : parsedPlan.plan().blocks()) { + progressListener.onProgress("configuring_blocks", "Configuring block " + blockPlan.blockId()); + BlockCatalogService.AssistantPromptBlockDescriptor descriptor = catalogByType.get(blockPlan.blockType()); + if (descriptor == null) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, + "Assistant selected an unknown block type: " + blockPlan.blockType()); + } + + ParsedBlockDraft parsedBlock = parseBlockDraft(invokeProvider(provider, model, + promptService.buildBlockConfigurationPrompt(mode, userPrompt, descriptor, parsedPlan.plan(), blockPlan, + currentFlow, errors, model))); + appendRationale(rationaleParts, parsedBlock.rationale()); + + AssistantConfiguredBlockDraft normalizedDraft = normalizeBlockDraft(blockPlan, parsedBlock.block()); + Block assembledBlock = buildBlock(descriptor, blockPlan, normalizedDraft, model); + assembledBlocks.add(assembledBlock); + blocksByPlanId.put(blockPlan.blockId(), assembledBlock); + registerBlockAlias(blocksByAlias, blockPlan.blockId(), assembledBlock); + registerBlockAlias(blocksByAlias, assembledBlock.getName(), assembledBlock); + registerBlockAlias(blocksByAlias, blockPlan.purpose(), assembledBlock); + configuredBlocks.add(new ConfiguredBlockSummary( + blockPlan.blockId(), + blockPlan.blockType(), + assembledBlock.getName(), + blockPlan.purpose(), + assembledBlock.getInputs() == null ? List.of() + : assembledBlock.getInputs().stream().map(io -> io.getName()).toList(), + assembledBlock.getOutputs() == null ? List.of() + : assembledBlock.getOutputs().stream().map(io -> io.getName()).toList())); + } + + progressListener.onProgress("connecting_blocks", "Connecting configured blocks"); + ParsedConnections parsedConnections = parseConnections(invokeProvider(provider, model, + promptService.buildConnectionsPrompt(mode, userPrompt, parsedPlan.plan(), configuredBlocks, currentFlow, + errors))); + appendRationale(rationaleParts, parsedConnections.rationale()); + + List connections = parsedConnections.connections().stream() + .map(connection -> toConnection(connection, blocksByPlanId, blocksByAlias)) + .toList(); + + FlowCreateRequest flow = new FlowCreateRequest( + defaultIfBlank(parsedPlan.plan().name(), "Assistant flow"), + parsedPlan.plan().description(), + FlowData.builder() + .blocks(assembledBlocks) + .connections(connections) + .build()); + + return new AssembledFlow(flow, String.join(" ", rationaleParts).trim()); + } + + private AssistantConfiguredBlockDraft normalizeBlockDraft(AssistantBlockPlan blockPlan, + AssistantConfiguredBlockDraft draft) { + if (draft == null) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, + "Assistant returned an empty block configuration payload"); + } + if (draft.config() == null || draft.config().isMissingNode() || draft.config().isNull()) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, + "Assistant returned a block configuration without config"); + } + if (draft.blockId() == null || draft.blockId().isBlank() || !Objects.equals(blockPlan.blockId(), draft.blockId())) { + return new AssistantConfiguredBlockDraft(blockPlan.blockId(), draft.name(), draft.config()); + } + return draft; + } + + private Block buildBlock(BlockCatalogService.AssistantPromptBlockDescriptor descriptor, AssistantBlockPlan blockPlan, + AssistantConfiguredBlockDraft draft, String model) { + if (!(draft.config() instanceof ObjectNode configNode)) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, + "Assistant returned a non-object config for block " + blockPlan.blockId()); + } + + ObjectNode normalizedConfig = configNode.deepCopy(); + normalizedConfig.put("type", descriptor.configurationType()); + normalizedConfig.put("name", defaultIfBlank(draft.name(), defaultIfBlank(blockPlan.purpose(), blockPlan.blockType()))); + injectSystemManagedFields(normalizedConfig, descriptor, model); + + try { + BlockConfiguration configuration = ObjectMapperHolder.mapper.treeToValue(normalizedConfig, + BlockConfiguration.class); + return createBlock(configuration); + } catch (Exception e) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, + "Assistant returned an invalid block configuration for " + blockPlan.blockType() + ": " + + e.getMessage()); + } + } + + private void injectSystemManagedFields(ObjectNode config, BlockCatalogService.AssistantPromptBlockDescriptor descriptor, + String model) { + switch (descriptor.configurationType()) { + case "LLMBlockConfiguration" -> config.set("llmDescriptor", llmDescriptorNode(model)); + case "HumanInteractiveBlockConfiguration" -> config.set("simulateWith", llmDescriptorNode(model)); + case "ConditionalBlockConfiguration" -> { + boolean useLlm = inferConditionalUseLlm(config); + config.put("useLlm", useLlm); + if (useLlm) { + config.set("llmDescriptor", llmDescriptorNode(model)); + } else { + config.remove("llmDescriptor"); + } + } + default -> { + } + } + } + + private boolean inferConditionalUseLlm(ObjectNode config) { + if (config.has("useLlm")) { + return config.get("useLlm").asBoolean(false); + } + if (config.hasNonNull("prompt")) { + return true; + } + return false; + } + + @SuppressWarnings({ "rawtypes", "unchecked" }) + private Block createBlock(BlockConfiguration configuration) { + BlockFactory factory = blockFactories.stream() + .filter(candidate -> candidate.getBlockType().equals(configuration.getBlockType())) + .findFirst() + .orElseThrow(() -> new IllegalArgumentException( + "Block factory not found for type: " + configuration.getBlockType().getSimpleName())); + return (Block) factory.create(configuration); + } + + private Connection toConnection(AssistantConnectionDraft connection, Map> blocksByPlanId, + Map> blocksByAlias) { + Block source = resolveConnectionBlock(connection.fromBlockId(), blocksByPlanId, blocksByAlias); + Block target = resolveConnectionBlock(connection.toBlockId(), blocksByPlanId, blocksByAlias); + if (source == null) { + source = inferBlockByIo(connection.fromOutput(), blocksByPlanId.values(), true); + } + if (target == null) { + target = inferBlockByIo(connection.toInput(), blocksByPlanId.values(), false); + } + if (source == null || target == null) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, + "Assistant returned a connection with unknown block ids"); + } + return Connection.builder() + .sourceId(source.getId()) + .sourceName(connection.fromOutput()) + .targetId(target.getId()) + .targetName(connection.toInput()) + .build(); + } + + private Block resolveConnectionBlock(String rawReference, Map> blocksByPlanId, + Map> blocksByAlias) { + if (rawReference == null || rawReference.isBlank()) { + return null; + } + Block direct = blocksByPlanId.get(rawReference); + if (direct != null) { + return direct; + } + return blocksByAlias.get(normalizeBlockReference(rawReference)); + } + + private Block inferBlockByIo(String ioName, Collection> blocks, boolean output) { + String normalizedIo = normalizeBlockReference(ioName); + if (normalizedIo == null) { + return null; + } + Block match = null; + for (Block block : blocks) { + List ioDescriptors = output ? block.getOutputs() : block.getInputs(); + if (ioDescriptors == null) { + continue; + } + boolean hasMatch = ioDescriptors.stream() + .anyMatch(io -> normalizedIo.equals(normalizeBlockReference(io.getName()))); + if (!hasMatch) { + continue; + } + if (match != null) { + return null; + } + match = block; + } + return match; + } + + private void registerBlockAlias(Map> blocksByAlias, String reference, Block block) { + String normalized = normalizeBlockReference(reference); + if (normalized != null) { + blocksByAlias.putIfAbsent(normalized, block); + } + } + + private String normalizeBlockReference(String reference) { + if (reference == null) { + return null; + } + String normalized = reference.trim().toLowerCase(Locale.ROOT); + return normalized.isBlank() ? null : normalized; + } + + private ParsedPlan parsePlan(String rawResponse) { + try { + JsonNode root = ObjectMapperHolder.mapper.readTree(extractJsonObject(rawResponse)); + JsonNode planNode = root.has("plan") ? root.get("plan") : root; + AssistantFlowPlan plan = ObjectMapperHolder.mapper.treeToValue(planNode, AssistantFlowPlan.class); + return new ParsedPlan(plan, root.path("rationale").asText("")); + } catch (Exception e) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, + "Assistant returned an invalid plan payload: " + e.getMessage()); + } + } + + private ParsedBlockDraft parseBlockDraft(String rawResponse) { + try { + JsonNode root = ObjectMapperHolder.mapper.readTree(extractJsonObject(rawResponse)); + JsonNode blockNode = root.has("block") ? root.get("block") : root; + AssistantConfiguredBlockDraft block = new AssistantConfiguredBlockDraft( + blockNode.path("blockId").asText(null), + blockNode.path("name").asText(null), + blockNode.path("config")); + if (block.blockId() == null || block.config().isMissingNode()) { + throw new IllegalArgumentException("Missing blockId or config"); + } + return new ParsedBlockDraft(block, root.path("rationale").asText("")); + } catch (Exception e) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, + "Assistant returned an invalid block configuration payload: " + e.getMessage()); + } + } + + private ParsedConnections parseConnections(String rawResponse) { + try { + JsonNode root = ObjectMapperHolder.mapper.readTree(extractJsonObject(rawResponse)); + JsonNode connectionsNode = root.has("connections") ? root.get("connections") : root.path("connections"); + List connections = new ArrayList<>(); + if (connectionsNode.isArray()) { + for (JsonNode node : connectionsNode) { + connections.add(ObjectMapperHolder.mapper.treeToValue(node, AssistantConnectionDraft.class)); + } + } + return new ParsedConnections(connections, root.path("rationale").asText("")); + } catch (Exception e) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, + "Assistant returned an invalid connections payload: " + e.getMessage()); + } + } + + private void validatePlan(AssistantFlowPlan plan) { + if (plan == null) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, "Assistant returned an empty plan"); + } + if (plan.blocks() == null || plan.blocks().isEmpty()) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, "Assistant returned a plan with no blocks"); + } + Set ids = new LinkedHashSet<>(); + for (AssistantBlockPlan block : plan.blocks()) { + if (block.blockId() == null || block.blockId().isBlank() || !ids.add(block.blockId())) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, + "Assistant returned invalid or duplicate block ids in the plan"); + } + if (block.blockType() == null || block.blockType().isBlank()) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, + "Assistant returned a block without blockType"); + } + } + } + private List validate(FlowCreateRequest flow) { Set> violations = validator.validate(flow); if (violations.isEmpty()) { @@ -133,19 +489,6 @@ public class FlowAssistantService { return new ValidationError("flow", null, violation.getPropertyPath().toString(), violation.getMessage()); } - private ParsedAssistantFlow parseAssistantFlow(String rawResponse) { - try { - JsonNode root = ObjectMapperHolder.mapper.readTree(extractJsonObject(rawResponse)); - JsonNode flowNode = root.has("flow") ? root.get("flow") : root; - FlowCreateRequest flow = ObjectMapperHolder.mapper.treeToValue(flowNode, FlowCreateRequest.class); - String rationale = root.has("rationale") ? root.get("rationale").asText("") : ""; - return new ParsedAssistantFlow(flow, rationale); - } catch (Exception e) { - throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, - "Assistant returned an invalid JSON payload: " + e.getMessage()); - } - } - private String extractJsonObject(String rawResponse) { if (rawResponse == null || rawResponse.isBlank()) { throw new IllegalArgumentException("Empty assistant response"); @@ -160,6 +503,23 @@ public class FlowAssistantService { return trimmed.substring(start, end + 1); } + private ObjectNode llmDescriptorNode(String model) { + ObjectNode llmDescriptor = ObjectMapperHolder.mapper.createObjectNode(); + llmDescriptor.put("provider", INTERNAL_PROVIDER_NAME); + llmDescriptor.put("model", model); + return llmDescriptor; + } + + private void appendRationale(List target, String rationale) { + if (rationale != null && !rationale.isBlank()) { + target.add(rationale.trim()); + } + } + + private String defaultIfBlank(String value, String fallback) { + return value == null || value.isBlank() ? fallback : value; + } + private LLMProvider resolveInternalProvider() { LLMProvider provider = llmProviders.get("internalOllamaLLMProvider"); if (provider != null) { diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantCallAcceptedResponse.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantCallAcceptedResponse.java new file mode 100644 index 0000000..7dc72cb --- /dev/null +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantCallAcceptedResponse.java @@ -0,0 +1,4 @@ +package it.cnr.isti.workflow.manager.assistant.model; + +public record AssistantCallAcceptedResponse(String sessionId, String callId) { +} diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantCallStatus.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantCallStatus.java new file mode 100644 index 0000000..2b44878 --- /dev/null +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantCallStatus.java @@ -0,0 +1,8 @@ +package it.cnr.isti.workflow.manager.assistant.model; + +public enum AssistantCallStatus { + QUEUED, + RUNNING, + COMPLETED, + FAILED +} diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantCallView.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantCallView.java new file mode 100644 index 0000000..36fc47c --- /dev/null +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantCallView.java @@ -0,0 +1,17 @@ +package it.cnr.isti.workflow.manager.assistant.model; + +import java.time.Instant; + +public record AssistantCallView( + String id, + String sessionId, + AssistantCallStatus status, + String phase, + String progressMessage, + Instant createdAt, + Instant updatedAt, + AssistantIntent intent, + AssistantFlowResponse flowResult, + AssistantExplainResponse explainResult, + String errorMessage) { +} diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantConfigView.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantConfigView.java new file mode 100644 index 0000000..8f1c25c --- /dev/null +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantConfigView.java @@ -0,0 +1,7 @@ +package it.cnr.isti.workflow.manager.assistant.model; + +public record AssistantConfigView( + String provider, + String defaultModel, + String availableModelsRetrieverUrl) { +} diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantIntent.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantIntent.java new file mode 100644 index 0000000..92ae858 --- /dev/null +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantIntent.java @@ -0,0 +1,8 @@ +package it.cnr.isti.workflow.manager.assistant.model; + +public enum AssistantIntent { + DRAFT, + REFINE, + FIX, + EXPLAIN +} diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantMessageRole.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantMessageRole.java new file mode 100644 index 0000000..b4b4cdf --- /dev/null +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantMessageRole.java @@ -0,0 +1,6 @@ +package it.cnr.isti.workflow.manager.assistant.model; + +public enum AssistantMessageRole { + USER, + ASSISTANT +} diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantMessageView.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantMessageView.java new file mode 100644 index 0000000..bad7224 --- /dev/null +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantMessageView.java @@ -0,0 +1,11 @@ +package it.cnr.isti.workflow.manager.assistant.model; + +import java.time.Instant; + +public record AssistantMessageView( + String id, + AssistantMessageRole role, + String content, + Instant createdAt, + String callId) { +} diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantSessionCreateRequest.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantSessionCreateRequest.java new file mode 100644 index 0000000..12ae50e --- /dev/null +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantSessionCreateRequest.java @@ -0,0 +1,6 @@ +package it.cnr.isti.workflow.manager.assistant.model; + +import jakarta.validation.constraints.NotBlank; + +public record AssistantSessionCreateRequest(@NotBlank String model) { +} diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantSessionMessageRequest.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantSessionMessageRequest.java new file mode 100644 index 0000000..2f7a181 --- /dev/null +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantSessionMessageRequest.java @@ -0,0 +1,6 @@ +package it.cnr.isti.workflow.manager.assistant.model; + +import jakarta.validation.constraints.NotBlank; + +public record AssistantSessionMessageRequest(@NotBlank String message) { +} diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantSessionView.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantSessionView.java new file mode 100644 index 0000000..ba3734a --- /dev/null +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantSessionView.java @@ -0,0 +1,19 @@ +package it.cnr.isti.workflow.manager.assistant.model; + +import java.time.Instant; +import java.util.List; + +import it.cnr.isti.workflow.manager.flows.model.FlowCreateRequest; +import it.cnr.isti.workflow.manager.flows.validation.ValidationError; + +public record AssistantSessionView( + String id, + String owner, + String model, + Instant createdAt, + Instant updatedAt, + String lastCallId, + FlowCreateRequest currentFlow, + List lastValidationErrors, + List messages) { +} diff --git a/src/main/java/it/cnr/isti/workflow/manager/controllers/ApiExceptionHandler.java b/src/main/java/it/cnr/isti/workflow/manager/controllers/ApiExceptionHandler.java index 8ffa9b9..30a85ae 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/controllers/ApiExceptionHandler.java +++ b/src/main/java/it/cnr/isti/workflow/manager/controllers/ApiExceptionHandler.java @@ -7,6 +7,8 @@ import java.util.stream.Collectors; import org.slf4j.Logger; import org.springframework.http.HttpStatus; import org.springframework.http.ProblemDetail; +import org.springframework.validation.FieldError; +import org.springframework.validation.ObjectError; import org.springframework.web.bind.MethodArgumentNotValidException; import org.springframework.web.bind.annotation.ExceptionHandler; import org.springframework.web.bind.annotation.RestControllerAdvice; @@ -23,7 +25,7 @@ public class ApiExceptionHandler { @ExceptionHandler(MethodArgumentNotValidException.class) public ProblemDetail handleMethodArgumentNotValid(MethodArgumentNotValidException e) { List> errors = e.getBindingResult().getAllErrors().stream() - .flatMap(error -> ValidationErrorCodec.decode(error.getDefaultMessage()).stream()) + .flatMap(error -> toValidationErrors(error).stream()) .map(this::toMap) .toList(); String detail = errors.stream() @@ -50,6 +52,42 @@ public class ApiExceptionHandler { return problem; } + private List toValidationErrors(ObjectError error) { + List decoded = ValidationErrorCodec.decode(error.getDefaultMessage()); + if (decoded.isEmpty()) { + return List.of(toFallbackError(error)); + } + + return decoded.stream() + .map(decodedError -> mergeWithFallback(decodedError, error)) + .toList(); + } + + private ValidationError mergeWithFallback(ValidationError decoded, ObjectError source) { + ValidationError fallback = toFallbackError(source); + if (decoded == null) { + return fallback; + } + return new ValidationError( + isBlank(decoded.entity()) ? fallback.entity() : decoded.entity(), + isBlank(decoded.id()) ? fallback.id() : decoded.id(), + isBlank(decoded.field()) ? fallback.field() : decoded.field(), + isBlank(decoded.message()) ? fallback.message() : decoded.message()); + } + + private ValidationError toFallbackError(ObjectError error) { + String entity = error.getObjectName(); + String field = error instanceof FieldError fieldError ? fieldError.getField() : ""; + String message = error.getDefaultMessage() == null || error.getDefaultMessage().isBlank() + ? "Validation failed" + : error.getDefaultMessage(); + return new ValidationError(entity, null, field, message); + } + + private boolean isBlank(String value) { + return value == null || value.isBlank(); + } + private Map toMap(ValidationError error) { if (error == null) { return Map.of("message", "Request failed"); diff --git a/src/main/java/it/cnr/isti/workflow/manager/controllers/AssistantController.java b/src/main/java/it/cnr/isti/workflow/manager/controllers/AssistantController.java index c02c20b..34f3029 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/controllers/AssistantController.java +++ b/src/main/java/it/cnr/isti/workflow/manager/controllers/AssistantController.java @@ -3,27 +3,48 @@ package it.cnr.isti.workflow.manager.controllers; import org.eclipse.microprofile.openapi.annotations.Operation; import org.eclipse.microprofile.openapi.annotations.security.SecurityRequirement; import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.beans.factory.annotation.Value; +import org.springframework.security.core.annotation.AuthenticationPrincipal; +import org.springframework.web.bind.annotation.GetMapping; +import org.springframework.web.bind.annotation.PathVariable; import org.springframework.web.bind.annotation.PostMapping; import org.springframework.web.bind.annotation.RequestBody; import org.springframework.web.bind.annotation.RequestMapping; import org.springframework.web.bind.annotation.RestController; +import it.cnr.isti.workflow.manager.assistant.AssistantConversationService; import it.cnr.isti.workflow.manager.assistant.FlowAssistantService; import it.cnr.isti.workflow.manager.assistant.model.AssistantExplainRequest; import it.cnr.isti.workflow.manager.assistant.model.AssistantExplainResponse; +import it.cnr.isti.workflow.manager.assistant.model.AssistantCallAcceptedResponse; +import it.cnr.isti.workflow.manager.assistant.model.AssistantCallView; +import it.cnr.isti.workflow.manager.assistant.model.AssistantConfigView; import it.cnr.isti.workflow.manager.assistant.model.AssistantFixRequest; import it.cnr.isti.workflow.manager.assistant.model.AssistantFlowResponse; import it.cnr.isti.workflow.manager.assistant.model.AssistantGenerationRequest; import it.cnr.isti.workflow.manager.assistant.model.AssistantRefineRequest; +import it.cnr.isti.workflow.manager.assistant.model.AssistantSessionCreateRequest; +import it.cnr.isti.workflow.manager.assistant.model.AssistantSessionMessageRequest; +import it.cnr.isti.workflow.manager.assistant.model.AssistantSessionView; +import it.cnr.isti.workflow.manager.auth.repo.LoginEntity; import jakarta.validation.Valid; @RestController @RequestMapping("/assistant") public class AssistantController { + private static final String INTERNAL_PROVIDER_NAME = "InternalOllama"; + private static final String MODELS_RETRIEVER_URL = "/retriever/LLM/models?provider=InternalOllama"; + @Autowired private FlowAssistantService flowAssistantService; + @Autowired + private AssistantConversationService assistantConversationService; + + @Value("${app.assistant.default-model}") + private String defaultAssistantModel; + @PostMapping("/flows/draft") @SecurityRequirement(name = "bearerAuth") @Operation(summary = "Generate a flow draft", description = "Creates a flow draft from a natural language request.") @@ -51,4 +72,47 @@ public class AssistantController { public AssistantExplainResponse explain(@RequestBody @Valid AssistantExplainRequest request) { return flowAssistantService.explain(request); } + + @GetMapping("/config") + @SecurityRequirement(name = "bearerAuth") + @Operation(summary = "Get assistant config", description = "Returns assistant configuration for the GUI, including the default model and retriever URL for dynamic model selection.") + public AssistantConfigView getConfig() { + return new AssistantConfigView( + INTERNAL_PROVIDER_NAME, + defaultAssistantModel, + MODELS_RETRIEVER_URL); + } + + @PostMapping("/sessions") + @SecurityRequirement(name = "bearerAuth") + @Operation(summary = "Create assistant session", description = "Creates a chat assistant session bound to the selected internal model.") + public AssistantSessionView createSession(@RequestBody @Valid AssistantSessionCreateRequest request, + @AuthenticationPrincipal LoginEntity userDetails) { + return assistantConversationService.createSession(userDetails.getUsername(), request); + } + + @GetMapping("/sessions/{sessionId}") + @SecurityRequirement(name = "bearerAuth") + @Operation(summary = "Get assistant session", description = "Returns the current state of an assistant chat session.") + public AssistantSessionView getSession(@PathVariable String sessionId, + @AuthenticationPrincipal LoginEntity userDetails) { + return assistantConversationService.getSession(sessionId, userDetails.getUsername()); + } + + @PostMapping("/sessions/{sessionId}/messages") + @SecurityRequirement(name = "bearerAuth") + @Operation(summary = "Submit assistant message", description = "Submits a user message to the assistant session and returns a call id for polling.") + public AssistantCallAcceptedResponse submitMessage(@PathVariable String sessionId, + @RequestBody @Valid AssistantSessionMessageRequest request, + @AuthenticationPrincipal LoginEntity userDetails) { + return assistantConversationService.submitMessage(sessionId, userDetails.getUsername(), request); + } + + @GetMapping("/calls/{callId}") + @SecurityRequirement(name = "bearerAuth") + @Operation(summary = "Get assistant call", description = "Returns the execution status and result of an assistant call.") + public AssistantCallView getCall(@PathVariable String callId, + @AuthenticationPrincipal LoginEntity userDetails) { + return assistantConversationService.getCall(callId, userDetails.getUsername()); + } } diff --git a/src/main/java/it/cnr/isti/workflow/manager/llms/providers/ollama/InternalOllamaLLMProvider.java b/src/main/java/it/cnr/isti/workflow/manager/llms/providers/ollama/InternalOllamaLLMProvider.java index 7eb1c67..36c0ae9 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/llms/providers/ollama/InternalOllamaLLMProvider.java +++ b/src/main/java/it/cnr/isti/workflow/manager/llms/providers/ollama/InternalOllamaLLMProvider.java @@ -51,14 +51,6 @@ public class InternalOllamaLLMProvider implements LLMProvider { "model", model, "prompt", prompt, "stream", false); - String requestBody = null; - try { - requestBody = mapper.writeValueAsString(bodyMap); - } catch (Exception e) { - log.error("Error serializing request body: {}", e.getMessage(), e); - throw new RuntimeException("Error serializing request body", e); - } - log.debug("ollama body request: {} ", requestBody); // Implement the logic to call the Ollama API and return the response WebClient webClient = webClientBuilder.baseUrl(this.ollamaURL).build(); @@ -68,7 +60,7 @@ public class InternalOllamaLLMProvider implements LLMProvider { .build()) .header("Authorization", "Bearer " + ollamaKey) .contentType(MediaType.APPLICATION_JSON) - .bodyValue(requestBody) + .bodyValue(bodyMap) .retrieve() .onStatus( status -> status.is5xxServerError(), diff --git a/src/main/resources/application.properties b/src/main/resources/application.properties index 2b603c1..3cd4d68 100644 --- a/src/main/resources/application.properties +++ b/src/main/resources/application.properties @@ -31,9 +31,9 @@ app.db.init.enabled=true app.security.key=${WFEDITOR_SECRET_KEY:088c65fd2a5ca418a79cd10df5dff15c0a79781c0da4fd43c1c14e4e2d7af1ff} app.ollama.internal.key=${OLLAMA_INTERNAL_KEY:ollama} app.ollama.internal.url=${OLLAMA_INTERNAL_URL:https://ollama.internal/api} +app.assistant.default-model=${ASSISTANT_DEFAULT_MODEL:gemma3:12b} cors.allowed-origins=${CORS_ALLOWED_ORIGINS:http://localhost:4200} app.import.path=${IMPORT_PATH:/workflow-editor-init} app.import.enabled=true logging.level.it.cnr.isti.workflow.manager=DEBUG logging.level.root=ERROR - diff --git a/src/test/java/it/cnr/isti/workflow/manager/controllers/AssistantControllerTest.java b/src/test/java/it/cnr/isti/workflow/manager/controllers/AssistantControllerTest.java index 9d6d696..ec31ac3 100644 --- a/src/test/java/it/cnr/isti/workflow/manager/controllers/AssistantControllerTest.java +++ b/src/test/java/it/cnr/isti/workflow/manager/controllers/AssistantControllerTest.java @@ -7,16 +7,28 @@ import static org.junit.jupiter.api.Assertions.assertTrue; import org.junit.jupiter.api.Test; import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.autoconfigure.web.servlet.AutoConfigureMockMvc; import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.http.MediaType; import org.springframework.test.context.TestPropertySource; +import org.springframework.test.web.servlet.MockMvc; import it.cnr.isti.workflow.manager.app.ObjectMapperHolder; +import it.cnr.isti.workflow.manager.assistant.model.AssistantCallAcceptedResponse; +import it.cnr.isti.workflow.manager.assistant.model.AssistantCallStatus; +import it.cnr.isti.workflow.manager.assistant.model.AssistantCallView; +import it.cnr.isti.workflow.manager.assistant.model.AssistantConfigView; import it.cnr.isti.workflow.manager.assistant.model.AssistantExplainRequest; import it.cnr.isti.workflow.manager.assistant.model.AssistantExplainResponse; import it.cnr.isti.workflow.manager.assistant.model.AssistantFixRequest; import it.cnr.isti.workflow.manager.assistant.model.AssistantFlowResponse; import it.cnr.isti.workflow.manager.assistant.model.AssistantGenerationRequest; import it.cnr.isti.workflow.manager.assistant.model.AssistantRefineRequest; +import it.cnr.isti.workflow.manager.assistant.model.AssistantSessionCreateRequest; +import it.cnr.isti.workflow.manager.assistant.model.AssistantSessionMessageRequest; +import it.cnr.isti.workflow.manager.assistant.model.AssistantSessionView; +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.HumanInteractiveBlockConfiguration; import it.cnr.isti.workflow.manager.blocks.configurations.LLMBlockConfiguration; @@ -30,9 +42,15 @@ import it.cnr.isti.workflow.manager.ios.IOType; import it.cnr.isti.workflow.manager.llms.LLMDescriptor; import it.cnr.isti.workflow.manager.llms.providers.ollama.InternalOllamaLLMProvider; import org.mockito.Mockito; +import org.mockito.stubbing.Answer; import org.springframework.test.context.bean.override.mockito.MockitoBean; +import static org.springframework.test.web.servlet.request.MockMvcRequestBuilders.post; +import static org.springframework.test.web.servlet.result.MockMvcResultMatchers.jsonPath; +import static org.springframework.test.web.servlet.result.MockMvcResultMatchers.status; + @SpringBootTest +@AutoConfigureMockMvc @TestPropertySource(locations = "classpath:test.properties") public class AssistantControllerTest { @@ -41,14 +59,18 @@ public class AssistantControllerTest { @Autowired private AssistantController assistantController; + @Autowired + private MockMvc mockMvc; + + @Autowired + private JwtUtil jwtUtil; + @MockitoBean private InternalOllamaLLMProvider internalOllamaLLMProvider; @Test public void draftGeneratesValidFlow() { - Mockito.when(internalOllamaLLMProvider.generate(Mockito.eq(MODEL), Mockito.contains("TASK: DRAFT"))) - .thenReturn(TestAssistantResponses.wrap(TestAssistantResponses.singleBlockFlow(), - "Created a simple LLM-first draft.")); + mockAssistantResponses(); AssistantFlowResponse response = assistantController.draft( new AssistantGenerationRequest( @@ -66,9 +88,7 @@ public class AssistantControllerTest { @Test public void refineReturnsExpandedFlow() { - Mockito.when(internalOllamaLLMProvider.generate(Mockito.eq(MODEL), Mockito.contains("TASK: REFINE"))) - .thenReturn(TestAssistantResponses.wrap(TestAssistantResponses.refinedFlow(), - "Added a human review step after the LLM classification.")); + mockAssistantResponses(); FlowCreateRequest initialFlow = TestAssistantResponses.singleBlockFlow(); AssistantFlowResponse response = assistantController.refine( @@ -85,9 +105,7 @@ public class AssistantControllerTest { @Test public void fixRepairsInvalidFlow() { - Mockito.when(internalOllamaLLMProvider.generate(Mockito.eq(MODEL), Mockito.contains("TASK: FIX"))) - .thenReturn(TestAssistantResponses.wrap(TestAssistantResponses.singleBlockFlow(), - "Repaired the invalid output definition.")); + mockAssistantResponses(); AssistantFlowResponse response = assistantController.fix( new AssistantFixRequest( @@ -104,8 +122,7 @@ public class AssistantControllerTest { @Test public void explainReturnsNarrativeText() { - Mockito.when(internalOllamaLLMProvider.generate(Mockito.eq(MODEL), Mockito.contains("TASK: EXPLAIN"))) - .thenReturn("This workflow classifies the input and sends it to human review when needed."); + mockAssistantResponses(); AssistantExplainResponse response = assistantController.explain( new AssistantExplainRequest( @@ -117,6 +134,160 @@ public class AssistantControllerTest { assertTrue(response.explanation().contains("human review")); } + @Test + public void configReturnsDefaultModelAndRetrieverUrl() { + AssistantConfigView config = assistantController.getConfig(); + + assertNotNull(config); + assertEquals("InternalOllama", config.provider()); + assertEquals(MODEL, config.defaultModel()); + assertEquals("/retriever/LLM/models?provider=InternalOllama", config.availableModelsRetrieverUrl()); + } + + @Test + public void blankSessionMessageReturnsFieldNameInValidationErrors() throws Exception { + mockMvc.perform(post("/assistant/sessions/test-session/messages") + .contentType(MediaType.APPLICATION_JSON) + .header("Authorization", "Bearer " + jwtUtil.generateToken("testuser")) + .content(""" + { + "message": " " + } + """)) + .andExpect(status().isBadRequest()) + .andExpect(jsonPath("$.errors[0].field").value("message")); + } + + @Test + public void sessionMessageFlowCompletesAndStoresConversation() throws Exception { + mockAssistantResponses(); + LoginEntity user = new LoginEntity("testuser", "testpassword"); + + AssistantSessionView session = assistantController.createSession( + new AssistantSessionCreateRequest(MODEL), + user); + + assertNotNull(session); + assertEquals(MODEL, session.model()); + + AssistantCallAcceptedResponse accepted = assistantController.submitMessage( + session.id(), + new AssistantSessionMessageRequest("Create a flow that classifies incoming tickets"), + user); + + assertNotNull(accepted.callId()); + + AssistantCallView call = waitForCallCompletion(accepted.callId(), user); + assertEquals(AssistantCallStatus.COMPLETED, call.status()); + assertNotNull(call.flowResult()); + assertTrue(call.flowResult().valid()); + + AssistantSessionView updatedSession = assistantController.getSession(session.id(), user); + assertNotNull(updatedSession.currentFlow()); + assertTrue(updatedSession.messages().size() >= 2); + assertEquals(accepted.callId(), updatedSession.lastCallId()); + } + + private void mockAssistantResponses() { + Answer answer = invocation -> { + String prompt = invocation.getArgument(1, String.class); + if (prompt.contains("TASK: EXPLAIN")) { + return "This workflow classifies the input and sends it to human review when needed."; + } + if (prompt.contains("TASK: PLAN") && prompt.contains("MODE: DRAFT")) { + return TestAssistantResponses.wrap(java.util.Map.of( + "rationale", "Planned a minimal draft flow.", + "plan", java.util.Map.of( + "name", "Ticket classification", + "description", "Classify incoming tickets.", + "blocks", java.util.List.of( + java.util.Map.of( + "blockId", "b1", + "blockType", "LLMBlock", + "purpose", "Classify incoming ticket"))))); + } + if (prompt.contains("TASK: PLAN") && prompt.contains("MODE: REFINE")) { + return TestAssistantResponses.wrap(java.util.Map.of( + "rationale", "Planned an extra human review block.", + "plan", java.util.Map.of( + "name", "Ticket classification with review", + "description", "Classify incoming tickets and review urgent ones.", + "blocks", java.util.List.of( + java.util.Map.of( + "blockId", "b1", + "blockType", "LLMBlock", + "purpose", "Classify incoming ticket"), + java.util.Map.of( + "blockId", "b2", + "blockType", "HumanInteractionBlock", + "purpose", "Review urgent tickets"))))); + } + if (prompt.contains("TASK: PLAN") && prompt.contains("MODE: FIX")) { + return TestAssistantResponses.wrap(java.util.Map.of( + "rationale", "Planned the repaired flow.", + "plan", java.util.Map.of( + "name", "Ticket classification", + "description", "Classify incoming tickets.", + "blocks", java.util.List.of( + java.util.Map.of( + "blockId", "b1", + "blockType", "LLMBlock", + "purpose", "Classify incoming ticket"))))); + } + if (prompt.contains("TASK: BLOCK_CONFIG") + && prompt.contains("Current block to configure:\n{\n \"blockId\" : \"b2\"")) { + return TestAssistantResponses.wrap(java.util.Map.of( + "rationale", "Configured the human review block.", + "block", java.util.Map.of( + "blockId", "review-node", + "name", "Human review", + "config", java.util.Map.of( + "actionDescription", "Review high-risk tickets")))); + } + if (prompt.contains("TASK: BLOCK_CONFIG") + && prompt.contains("Current block to configure:\n{\n \"blockId\" : \"b1\"")) { + return TestAssistantResponses.wrap(java.util.Map.of( + "rationale", "Configured the LLM block.", + "block", java.util.Map.of( + "blockId", "b1", + "name", "Ticket classifier", + "config", java.util.Map.of( + "prompt", "Classify the ticket: ${{ticket}}")))); + } + if (prompt.contains("TASK: CONNECTIONS") && prompt.contains("\"blockId\" : \"b2\"")) { + return TestAssistantResponses.wrap(java.util.Map.of( + "rationale", "Connected classification to review.", + "connections", java.util.List.of( + java.util.Map.of( + "fromBlockId", "classifier-node", + "fromOutput", "response", + "toBlockId", "human-review-node", + "toInput", "input")))); + } + if (prompt.contains("TASK: CONNECTIONS")) { + return TestAssistantResponses.wrap(java.util.Map.of( + "rationale", "No extra connections needed.", + "connections", java.util.List.of())); + } + throw new IllegalStateException("Unexpected assistant prompt:\n" + prompt); + }; + + Mockito.when(internalOllamaLLMProvider.generate(Mockito.eq(MODEL), Mockito.anyString())) + .thenAnswer(answer); + } + + private AssistantCallView waitForCallCompletion(String callId, LoginEntity user) throws Exception { + long deadline = System.currentTimeMillis() + 5_000; + while (System.currentTimeMillis() < deadline) { + AssistantCallView call = assistantController.getCall(callId, user); + if (call.status() == AssistantCallStatus.COMPLETED || call.status() == AssistantCallStatus.FAILED) { + return call; + } + Thread.sleep(25); + } + throw new AssertionError("Assistant call did not complete in time"); + } + static class TestAssistantResponses { private static final String PROVIDER = "InternalOllama"; @@ -223,11 +394,9 @@ public class AssistantControllerTest { .build()); } - static String wrap(FlowCreateRequest flow, String rationale) { + static String wrap(Object payload) { try { - return ObjectMapperHolder.mapper.writeValueAsString(java.util.Map.of( - "rationale", rationale, - "flow", flow)); + return ObjectMapperHolder.mapper.writeValueAsString(payload); } catch (Exception e) { throw new IllegalStateException("Unable to serialize fake assistant response", e); } diff --git a/src/test/resources/test.properties b/src/test/resources/test.properties index 15423ac..eff760c 100644 --- a/src/test/resources/test.properties +++ b/src/test/resources/test.properties @@ -7,6 +7,7 @@ spring.jpa.database-platform=org.hibernate.dialect.MySQLDialect spring.jpa.hibernate.ddl-auto=create-drop spring.jpa.properties.jakarta.persistence.validation.mode=none app.db.init.enabled=true +app.assistant.default-model=assistant-test-model app.import.path=src/test/resources/workflow-editor-init -app.import.enabled=true \ No newline at end of file +app.import.enabled=true