Implement session-based assistant orchestration and harden flow assembly
This commit is contained in:
parent
7e328ea276
commit
b2a5cd2673
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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<String, SessionState> sessions = new ConcurrentHashMap<>();
|
||||
private final ConcurrentHashMap<String, CallState> 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<AssistantMessageView> messages = session.toView().messages();
|
||||
int fromIndex = Math.max(0, messages.size() - 6);
|
||||
List<AssistantMessageView> 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<ValidationError> lastValidationErrors = List.of();
|
||||
private final List<AssistantMessageView> 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);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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<AssistantPromptFieldDescriptor> configurationFields,
|
||||
List<String> inputs,
|
||||
List<String> outputs) {
|
||||
}
|
||||
|
||||
@Autowired
|
||||
private Map<String, BlockType> blockTypes;
|
||||
|
||||
|
|
@ -42,6 +64,12 @@ public class BlockCatalogService {
|
|||
.toList();
|
||||
}
|
||||
|
||||
public List<AssistantPromptBlockDescriptor> getPromptCatalog() {
|
||||
return getCatalog().stream()
|
||||
.map(this::toPromptDescriptor)
|
||||
.toList();
|
||||
}
|
||||
|
||||
@SuppressWarnings("unchecked")
|
||||
private AssistantBlockDescriptor toDescriptor(BlockType blockType) {
|
||||
Class<? extends BlockConfiguration<?>> 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<AssistantPromptFieldDescriptor> extractConfigurationFields(JsonNode schema) {
|
||||
if (schema == null || !schema.has("properties")) {
|
||||
return List.of();
|
||||
}
|
||||
|
||||
Set<String> 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<String> extractIoNames(Block<?> exampleBlock, boolean inputs) {
|
||||
if (exampleBlock == null) {
|
||||
return List.of();
|
||||
}
|
||||
List<IODescriptor> descriptors = inputs ? exampleBlock.getInputs() : exampleBlock.getOutputs();
|
||||
if (descriptors == null) {
|
||||
return List.of();
|
||||
}
|
||||
return descriptors.stream().map(IODescriptor::getName).toList();
|
||||
}
|
||||
|
||||
private Set<String> extractRequiredFields(JsonNode requiredNode) {
|
||||
if (!(requiredNode instanceof ArrayNode requiredArray)) {
|
||||
return Set.of();
|
||||
}
|
||||
java.util.LinkedHashSet<String> 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<Map.Entry<String, JsonNode>> iterable(java.util.Iterator<Map.Entry<String, JsonNode>> fields) {
|
||||
java.util.ArrayList<Map.Entry<String, JsonNode>> entries = new java.util.ArrayList<>();
|
||||
fields.forEachRemaining(entries::add);
|
||||
return entries;
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -11,81 +11,41 @@ import it.cnr.isti.workflow.manager.flows.validation.ValidationError;
|
|||
@Service
|
||||
public class FlowAssistantPromptService {
|
||||
|
||||
public String buildDraftPrompt(String userPrompt, List<BlockCatalogService.AssistantBlockDescriptor> catalog) {
|
||||
public enum OperationMode {
|
||||
DRAFT,
|
||||
REFINE,
|
||||
FIX
|
||||
}
|
||||
|
||||
public String buildPlanPrompt(OperationMode mode, String userPrompt, FlowCreateRequest currentFlow,
|
||||
List<ValidationError> errors, List<BlockCatalogService.AssistantPromptBlockDescriptor> 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<BlockCatalogService.AssistantBlockDescriptor> 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<ValidationError> errors,
|
||||
List<BlockCatalogService.AssistantBlockDescriptor> 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<ValidationError> 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<ValidationError> 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<ValidationError> errors) {
|
||||
if (errors == null || errors.isEmpty()) {
|
||||
return "(none)";
|
||||
}
|
||||
return toJson(errors);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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<AssistantBlockPlan> 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<AssistantConnectionDraft> connections, String rationale) {
|
||||
}
|
||||
|
||||
private record ConfiguredBlockSummary(String blockId, String blockType, String name, String purpose, List<String> inputs,
|
||||
List<String> outputs) {
|
||||
}
|
||||
|
||||
private record AssembledFlow(FlowCreateRequest flow, String rationale) {
|
||||
}
|
||||
|
||||
@Autowired
|
||||
|
|
@ -44,67 +87,380 @@ public class FlowAssistantService {
|
|||
@Autowired
|
||||
private FlowAssistantPromptService promptService;
|
||||
|
||||
@Autowired
|
||||
private List<BlockFactory<?, ?>> 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<ValidationError> 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<ValidationError> 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<ValidationError> errorContext = initialErrors == null ? List.of() : initialErrors;
|
||||
AssembledFlow assembled = null;
|
||||
List<ValidationError> errors = List.of();
|
||||
|
||||
ParsedAssistantFlow generated = parseAssistantFlow(invokeProvider(provider, model, initialPrompt));
|
||||
FlowCreateRequest currentFlow = generated.flow();
|
||||
List<ValidationError> 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<ValidationError> errors, ProgressListener progressListener) {
|
||||
List<BlockCatalogService.AssistantPromptBlockDescriptor> catalog = blockCatalogService.getPromptCatalog();
|
||||
Map<String, BlockCatalogService.AssistantPromptBlockDescriptor> 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<Block<?>> assembledBlocks = new ArrayList<>();
|
||||
Map<String, Block<?>> blocksByPlanId = new LinkedHashMap<>();
|
||||
Map<String, Block<?>> blocksByAlias = new LinkedHashMap<>();
|
||||
List<ConfiguredBlockSummary> configuredBlocks = new ArrayList<>();
|
||||
List<String> 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<Connection> 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<String, Block<?>> blocksByPlanId,
|
||||
Map<String, Block<?>> 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<String, Block<?>> blocksByPlanId,
|
||||
Map<String, Block<?>> 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<Block<?>> blocks, boolean output) {
|
||||
String normalizedIo = normalizeBlockReference(ioName);
|
||||
if (normalizedIo == null) {
|
||||
return null;
|
||||
}
|
||||
Block<?> match = null;
|
||||
for (Block<?> block : blocks) {
|
||||
List<IODescriptor> 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<String, Block<?>> 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<AssistantConnectionDraft> 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<String> 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<ValidationError> validate(FlowCreateRequest flow) {
|
||||
Set<ConstraintViolation<FlowCreateRequest>> 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<String> 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) {
|
||||
|
|
|
|||
|
|
@ -0,0 +1,4 @@
|
|||
package it.cnr.isti.workflow.manager.assistant.model;
|
||||
|
||||
public record AssistantCallAcceptedResponse(String sessionId, String callId) {
|
||||
}
|
||||
|
|
@ -0,0 +1,8 @@
|
|||
package it.cnr.isti.workflow.manager.assistant.model;
|
||||
|
||||
public enum AssistantCallStatus {
|
||||
QUEUED,
|
||||
RUNNING,
|
||||
COMPLETED,
|
||||
FAILED
|
||||
}
|
||||
|
|
@ -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) {
|
||||
}
|
||||
|
|
@ -0,0 +1,7 @@
|
|||
package it.cnr.isti.workflow.manager.assistant.model;
|
||||
|
||||
public record AssistantConfigView(
|
||||
String provider,
|
||||
String defaultModel,
|
||||
String availableModelsRetrieverUrl) {
|
||||
}
|
||||
|
|
@ -0,0 +1,8 @@
|
|||
package it.cnr.isti.workflow.manager.assistant.model;
|
||||
|
||||
public enum AssistantIntent {
|
||||
DRAFT,
|
||||
REFINE,
|
||||
FIX,
|
||||
EXPLAIN
|
||||
}
|
||||
|
|
@ -0,0 +1,6 @@
|
|||
package it.cnr.isti.workflow.manager.assistant.model;
|
||||
|
||||
public enum AssistantMessageRole {
|
||||
USER,
|
||||
ASSISTANT
|
||||
}
|
||||
|
|
@ -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) {
|
||||
}
|
||||
|
|
@ -0,0 +1,6 @@
|
|||
package it.cnr.isti.workflow.manager.assistant.model;
|
||||
|
||||
import jakarta.validation.constraints.NotBlank;
|
||||
|
||||
public record AssistantSessionCreateRequest(@NotBlank String model) {
|
||||
}
|
||||
|
|
@ -0,0 +1,6 @@
|
|||
package it.cnr.isti.workflow.manager.assistant.model;
|
||||
|
||||
import jakarta.validation.constraints.NotBlank;
|
||||
|
||||
public record AssistantSessionMessageRequest(@NotBlank String message) {
|
||||
}
|
||||
|
|
@ -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<ValidationError> lastValidationErrors,
|
||||
List<AssistantMessageView> messages) {
|
||||
}
|
||||
|
|
@ -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<Map<String, String>> 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<ValidationError> toValidationErrors(ObjectError error) {
|
||||
List<ValidationError> 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<String, String> toMap(ValidationError error) {
|
||||
if (error == null) {
|
||||
return Map.of("message", "Request failed");
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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<String> 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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
app.import.enabled=true
|
||||
|
|
|
|||
Loading…
Reference in New Issue