From 35db1d42f1aac99c3f62b83d1ed520bb97799ada Mon Sep 17 00:00:00 2001 From: Lucio Lelii Date: Wed, 20 May 2026 12:27:09 +0200 Subject: [PATCH] Improve assistant model routing and diagnostics --- .gitignore | 1 + .../AssistantConversationService.java | 30 +- .../assistant/FlowAssistantPromptService.java | 13 +- .../assistant/FlowAssistantService.java | 299 +++++++++++++++--- .../assistant/model/AssistantConfigView.java | 3 +- .../assistant/model/AssistantFixRequest.java | 8 +- .../model/AssistantGenerationRequest.java | 8 +- .../model/AssistantModelSelection.java | 7 + .../model/AssistantRefineRequest.java | 8 +- .../model/AssistantSessionCreateRequest.java | 9 +- .../controllers/AssistantController.java | 13 +- .../resources/application-prod.properties | 14 + src/main/resources/application.properties | 21 +- src/main/resources/logback-spring.xml | 30 ++ .../controllers/AssistantControllerTest.java | 182 +++++++++++ src/test/resources/logback-test.xml | 3 +- src/test/resources/test.properties | 3 + 17 files changed, 591 insertions(+), 61 deletions(-) create mode 100644 src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantModelSelection.java create mode 100644 src/main/resources/logback-spring.xml diff --git a/.gitignore b/.gitignore index 36fd06c..4c176e5 100644 --- a/.gitignore +++ b/.gitignore @@ -1,5 +1,6 @@ HELP.md target/ +logs/ *.tar !.mvn/wrapper/maven-wrapper.jar !**/src/main/**/target/ 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 index 9e8b758..c7aa2ef 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/assistant/AssistantConversationService.java +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/AssistantConversationService.java @@ -18,6 +18,7 @@ import java.util.concurrent.atomic.AtomicBoolean; import org.slf4j.Logger; import org.slf4j.LoggerFactory; +import org.slf4j.MDC; import org.springframework.stereotype.Service; import org.springframework.web.server.ResponseStatusException; import org.springframework.http.HttpStatus; @@ -33,6 +34,7 @@ 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.AssistantModelSelection; 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; @@ -69,6 +71,7 @@ public class AssistantConversationService { UUID.randomUUID().toString(), owner, request.model(), + request.phaseModels(), now, now); sessions.put(session.id, session); @@ -126,12 +129,17 @@ public class AssistantConversationService { } private void processCall(SessionState session, CallState call, String message) { + MDC.put("assistantSessionId", session.id); + MDC.put("assistantCallId", call.id); + MDC.put("assistantOwner", session.owner); try { ensureNotCancelled(call); call.status = AssistantCallStatus.RUNNING; call.updatePhase("routing", "Routing the request"); + MDC.put("assistantPhase", call.phase); AssistantIntent intent = inferIntent(session, message); call.intent = intent; + MDC.put("assistantIntent", intent.name()); String contextualPrompt = buildContextualPrompt(session, message); ensureNotCancelled(call); @@ -139,7 +147,8 @@ public class AssistantConversationService { switch (intent) { case DRAFT -> { AssistantFlowResponse result = flowAssistantService.draft( - new AssistantGenerationRequest(contextualPrompt, session.model, DEFAULT_MAX_REPAIR_ATTEMPTS), + new AssistantGenerationRequest(contextualPrompt, session.model, DEFAULT_MAX_REPAIR_ATTEMPTS, + session.phaseModels), (phase, progressMessage) -> updateProgress(call, phase, progressMessage)); ensureNotCancelled(call); call.flowResult = result; @@ -150,7 +159,7 @@ public class AssistantConversationService { case REFINE -> { AssistantFlowResponse result = flowAssistantService.refine( new AssistantRefineRequest(contextualPrompt, session.currentFlow, session.model, - DEFAULT_MAX_REPAIR_ATTEMPTS), + DEFAULT_MAX_REPAIR_ATTEMPTS, session.phaseModels), (phase, progressMessage) -> updateProgress(call, phase, progressMessage)); ensureNotCancelled(call); call.flowResult = result; @@ -161,7 +170,7 @@ public class AssistantConversationService { case FIX -> { AssistantFlowResponse result = flowAssistantService.fix( new AssistantFixRequest(contextualPrompt, session.currentFlow, session.lastValidationErrors, - session.model, DEFAULT_MAX_REPAIR_ATTEMPTS), + session.model, DEFAULT_MAX_REPAIR_ATTEMPTS, session.phaseModels), (phase, progressMessage) -> updateProgress(call, phase, progressMessage)); ensureNotCancelled(call); call.flowResult = result; @@ -200,6 +209,7 @@ public class AssistantConversationService { session.touch(); } finally { releaseRunning(session, call); + clearAssistantConversationMdc(); } } @@ -232,6 +242,15 @@ public class AssistantConversationService { private void updateProgress(CallState call, String phase, String progressMessage) { ensureNotCancelled(call); call.updatePhase(phase, progressMessage); + MDC.put("assistantPhase", phase); + } + + private void clearAssistantConversationMdc() { + MDC.remove("assistantSessionId"); + MDC.remove("assistantCallId"); + MDC.remove("assistantOwner"); + MDC.remove("assistantIntent"); + MDC.remove("assistantPhase"); } private void ensureNotCancelled(CallState call) { @@ -364,6 +383,7 @@ public class AssistantConversationService { private final String id; private final String owner; private final String model; + private final AssistantModelSelection phaseModels; private final Instant createdAt; private volatile Instant updatedAt; private volatile String lastCallId; @@ -373,10 +393,12 @@ public class AssistantConversationService { private final AtomicBoolean running = new AtomicBoolean(false); private final List messages = java.util.Collections.synchronizedList(new ArrayList<>()); - private SessionState(String id, String owner, String model, Instant createdAt, Instant updatedAt) { + private SessionState(String id, String owner, String model, AssistantModelSelection phaseModels, + Instant createdAt, Instant updatedAt) { this.id = id; this.owner = owner; this.model = model; + this.phaseModels = phaseModels; this.createdAt = createdAt; this.updatedAt = updatedAt; } 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 e0424d2..1013fcd 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 @@ -150,6 +150,8 @@ public class FlowAssistantPromptService { - The producer is always the first block in the chain that initialises the shared tool session. Never mark it as a consumer. - Do not include system-managed fields like provider, model, llmDescriptor, ids, inputs, outputs, skills. - Use placeholders like ${{variable}} when needed. + - If this block depends on data, context, or completion from an earlier workflow step, include a placeholder such as ${{input}}, ${{previous_response}}, or ${{state_ready}} in its prompt so the backend can create a real input for the connection. + - Do not use placeholders for literal instructions that do not depend on another block. - Return valid JSON with no markdown fences. Selected internal model: @@ -201,9 +203,9 @@ public class FlowAssistantPromptService { "rationale": "short explanation", "connections": [ { - "fromBlockId": "b1", + "fromBlockId": "EXACT_BLOCK_ID_FROM_FLOW_PLAN", "fromOutput": "response", - "toBlockId": "b2", + "toBlockId": "EXACT_BLOCK_ID_FROM_FLOW_PLAN", "toInput": "input" } ] @@ -211,10 +213,15 @@ public class FlowAssistantPromptService { Rules: - Use only block ids from the flow plan. - - Copy block ids exactly as provided in the flow plan, for example b1, b2, b3. + - Copy block ids exactly as provided in the flow plan. + - Never copy placeholder values such as EXACT_BLOCK_ID_FROM_FLOW_PLAN. - Never use block names or purposes in fromBlockId/toBlockId. - Use only output/input names that exist in the configured blocks. - Prefer canonical names such as response, output, input, and prompt when they exist. + - If there are fewer than two configured blocks, return "connections": []. + - If a target block lists no inputs, do not create a connection to that block. + - Never use an empty string for toInput or fromOutput. + - Do not connect workflow outputs to technical configuration inputs such as model. - Do not leave processing blocks disconnected when one step depends on another step's output or execution order. - For MCPAgent blocks that reuse a shared session, still connect the producer response to the consumer's upstream ordering placeholder, for example state_ready, so execution order is explicit. - In REFINE/FIX, existing connections between KEEP/UPDATE blocks are preserved by the backend; return only connections that are new or intentionally changed. 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 0c5d01b..c391f69 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 @@ -9,9 +9,11 @@ import java.util.Map; import java.util.Objects; import java.util.Locale; import java.util.Set; +import java.util.UUID; import org.slf4j.Logger; import org.slf4j.LoggerFactory; +import org.slf4j.MDC; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.http.HttpStatus; import org.springframework.stereotype.Service; @@ -32,6 +34,7 @@ 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.AssistantModelSelection; 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; @@ -51,6 +54,7 @@ import jakarta.validation.Validator; public class FlowAssistantService { private static final Logger log = LoggerFactory.getLogger(FlowAssistantService.class); + private static final Logger assistantResponseLog = LoggerFactory.getLogger("assistant.responses"); private static final String INTERNAL_PROVIDER_NAME = "InternalOllama"; private static final String SHARED_MEMORY_SESSION_NAME = "sharedMemorySession"; @@ -114,6 +118,9 @@ public class FlowAssistantService { private record AssembledFlow(FlowCreateRequest flow, String rationale) { } + private record ResolvedAssistantModels(String planningModel, String jsonModel, String repairModel) { + } + @Autowired private Map llmProviders; @@ -138,13 +145,27 @@ public class FlowAssistantService { @Value("${app.assistant.provider-retry-attempts:" + DEFAULT_PROVIDER_RETRY_ATTEMPTS + "}") private int providerRetryAttempts; + @Value("${app.assistant.default-planning-model:${app.assistant.default-model}}") + private String defaultPlanningModel; + + @Value("${app.assistant.default-json-model:${app.assistant.default-model}}") + private String defaultJsonModel; + + @Value("${app.assistant.default-repair-model:${app.assistant.default-model}}") + private String defaultRepairModel; + public AssistantFlowResponse draft(AssistantGenerationRequest request) { 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); + boolean directMdc = ensureAssistantRequestMdc(OperationMode.DRAFT.name()); + try { + return generateFlow(OperationMode.DRAFT, request.userPrompt(), null, List.of(), request.model(), + request.phaseModels(), request.maxRepairAttempts(), progressListener); + } finally { + clearAssistantRequestMdc(directMdc); + } } public AssistantFlowResponse refine(AssistantRefineRequest request) { @@ -152,8 +173,13 @@ public class FlowAssistantService { } public AssistantFlowResponse refine(AssistantRefineRequest request, ProgressListener progressListener) { - return generateFlow(OperationMode.REFINE, request.userPrompt(), request.flow(), List.of(), request.model(), - request.maxRepairAttempts(), progressListener); + boolean directMdc = ensureAssistantRequestMdc(OperationMode.REFINE.name()); + try { + return generateFlow(OperationMode.REFINE, request.userPrompt(), request.flow(), List.of(), request.model(), + request.phaseModels(), request.maxRepairAttempts(), progressListener); + } finally { + clearAssistantRequestMdc(directMdc); + } } public AssistantFlowResponse fix(AssistantFixRequest request) { @@ -164,8 +190,13 @@ public class FlowAssistantService { List initialErrors = request.validationErrors() == null || request.validationErrors().isEmpty() ? validate(request.flow()) : request.validationErrors(); - return generateFlow(OperationMode.FIX, request.userPrompt(), request.flow(), initialErrors, request.model(), - request.maxRepairAttempts(), progressListener); + boolean directMdc = ensureAssistantRequestMdc(OperationMode.FIX.name()); + try { + return generateFlow(OperationMode.FIX, request.userPrompt(), request.flow(), initialErrors, request.model(), + request.phaseModels(), request.maxRepairAttempts(), progressListener); + } finally { + clearAssistantRequestMdc(directMdc); + } } public AssistantExplainResponse explain(AssistantExplainRequest request) { @@ -173,17 +204,26 @@ public class FlowAssistantService { } 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; + boolean directMdc = ensureAssistantRequestMdc("EXPLAIN"); + try { + LLMProvider provider = resolveInternalProvider(); + progressListener.onProgress("explaining", "Explaining the current flow"); + String prompt = promptService.buildExplainPrompt(request.flow(), request.userPrompt()); + String rawResponse = invokeProvider(provider, request.model(), prompt); + logAssistantRawResponse("explain", request.model(), "text", 1, 1, rawResponse); + AssistantExplainResponse response = new AssistantExplainResponse(rawResponse); + progressListener.onProgress("completed", "Flow explanation ready"); + return response; + } finally { + clearAssistantRequestMdc(directMdc); + } } private AssistantFlowResponse generateFlow(OperationMode initialMode, String userPrompt, FlowCreateRequest currentFlow, - List initialErrors, String model, Integer maxRepairAttempts, ProgressListener progressListener) { + List initialErrors, String workflowModel, AssistantModelSelection requestedPhaseModels, + Integer maxRepairAttempts, ProgressListener progressListener) { LLMProvider provider = resolveInternalProvider(); + ResolvedAssistantModels phaseModels = resolveAssistantModels(workflowModel, requestedPhaseModels); int allowedRepairs = maxRepairAttempts == null ? 1 : maxRepairAttempts; int repairs = 0; OperationMode mode = initialMode; @@ -193,7 +233,8 @@ public class FlowAssistantService { List errors = List.of(); while (true) { - assembled = assembleFlow(provider, model, mode, userPrompt, flowContext, errorContext, progressListener); + assembled = assembleFlow(provider, workflowModel, phaseModels, mode, userPrompt, flowContext, errorContext, + progressListener); progressListener.onProgress("validating", "Validating the assembled flow"); errors = validate(assembled.flow()); if (errors.isEmpty() || repairs >= allowedRepairs) { @@ -227,8 +268,9 @@ public class FlowAssistantService { } } - private AssembledFlow assembleFlow(LLMProvider provider, String model, OperationMode mode, String userPrompt, - FlowCreateRequest currentFlow, List errors, ProgressListener progressListener) { + private AssembledFlow assembleFlow(LLMProvider provider, String workflowModel, ResolvedAssistantModels phaseModels, + OperationMode mode, String userPrompt, FlowCreateRequest currentFlow, List errors, + ProgressListener progressListener) { List catalog = blockCatalogService.getPromptCatalog(); Map catalogByType = new LinkedHashMap<>(); for (BlockCatalogService.AssistantPromptBlockDescriptor descriptor : catalog) { @@ -237,7 +279,8 @@ public class FlowAssistantService { progressListener.onProgress("planning", "Planning workflow blocks"); String planPrompt = promptService.buildPlanPrompt(mode, userPrompt, currentFlow, errors, catalog); - ParsedPlan parsedPlan = invokeStructuredAndValidate(provider, model, planPrompt, "plan", rawResponse -> { + ParsedPlan parsedPlan = invokeStructuredAndValidate(provider, planningModelFor(mode, phaseModels), + phaseModels.repairModel(), planPrompt, "plan", rawResponse -> { ParsedPlan plan = parsePlan(rawResponse); AssistantFlowPlan normalizedPlan = validateAndNormalizePlan(plan.plan(), mode, userPrompt, currentFlow, errors, catalogByType.keySet()); @@ -284,14 +327,17 @@ public class FlowAssistantService { } else { progressListener.onProgress("configuring_blocks", "Configuring block " + blockPlan.blockId()); String blockPrompt = promptService.buildBlockConfigurationPrompt(mode, userPrompt, descriptor, - parsedPlan.plan(), blockPlan, currentFlow, errors, model); - ConfiguredBlockResult configuredBlock = invokeStructuredAndValidate(provider, model, blockPrompt, + parsedPlan.plan(), blockPlan, currentFlow, errors, workflowModel); + int currentBlockIndex = blockIndex; + int blockPlanCount = blockPlans.size(); + ConfiguredBlockResult configuredBlock = invokeStructuredAndValidate(provider, + jsonModelFor(mode, phaseModels), phaseModels.repairModel(), blockPrompt, "block configuration for " + blockPlan.blockId(), rawResponse -> { ParsedBlockDraft parsedBlock = parseBlockDraft(rawResponse); AssistantConfiguredBlockDraft normalizedDraft = normalizeBlockDraft(blockPlan, parsedBlock.block()); - Block newBlock = buildBlock(descriptor, blockPlan, normalizedDraft, model, - requireSharedMemorySemantics); + Block newBlock = buildBlock(descriptor, blockPlan, normalizedDraft, workflowModel, + requireSharedMemorySemantics, currentBlockIndex, blockPlanCount); return new ConfiguredBlockResult(parsedBlock, newBlock); }); appendRationale(rationaleParts, configuredBlock.parsedBlock().rationale()); @@ -318,25 +364,28 @@ public class FlowAssistantService { } progressListener.onProgress("connecting_blocks", "Connecting configured blocks"); - String connectionsPrompt = promptService.buildConnectionsPrompt(mode, userPrompt, parsedPlan.plan(), - configuredBlocks, currentFlow, errors); - ParsedConnections parsedConnections = invokeStructuredAndValidate(provider, model, connectionsPrompt, - "connections", rawResponse -> { - ParsedConnections parsed = parseConnectionsOrInferSequential(rawResponse, assembledBlocks); - parsed = completeRequiredSequentialConnections(requireSharedMemorySemantics, parsed, - assembledBlocks, blocksByPlanId, blocksByAlias); - List candidateConnections = parsed.connections().stream() - .map(connection -> toConnection(connection, blocksByPlanId, blocksByAlias)) - .toList(); - validateSharedMemorySemantics(requireSharedMemorySemantics, blocksByPlanId.values(), - candidateConnections); - return parsed; - }); + ParsedConnections parsedConnections; + if (assembledBlocks.size() < 2) { + parsedConnections = new ParsedConnections(List.of(), "No connections needed for a single-block flow."); + } else { + String connectionsPrompt = promptService.buildConnectionsPrompt(mode, userPrompt, parsedPlan.plan(), + configuredBlocks, currentFlow, errors); + parsedConnections = invokeStructuredAndValidate(provider, jsonModelFor(mode, phaseModels), + phaseModels.repairModel(), connectionsPrompt, "connections", rawResponse -> { + ParsedConnections parsed = parseConnectionsOrInferSequential(rawResponse, assembledBlocks); + parsed = completeRequiredSequentialConnections(requireSharedMemorySemantics, parsed, + assembledBlocks, blocksByPlanId, blocksByAlias); + List candidateConnections = toValidConnections(parsed.connections(), blocksByPlanId, + blocksByAlias); + validateSharedMemorySemantics(requireSharedMemorySemantics, blocksByPlanId.values(), + candidateConnections); + return parsed; + }); + } appendRationale(rationaleParts, parsedConnections.rationale()); - List generatedConnections = parsedConnections.connections().stream() - .map(connection -> toConnection(connection, blocksByPlanId, blocksByAlias)) - .toList(); + List generatedConnections = toValidConnections(parsedConnections.connections(), blocksByPlanId, + blocksByAlias); List connections = mergeConnections( preserveCurrentConnections(currentFlow, oldBlockIdToAssembledBlock, removedExistingBlockIds), generatedConnections); @@ -370,16 +419,18 @@ public class FlowAssistantService { } } - private T invokeStructuredAndValidate(LLMProvider provider, String model, String prompt, String taskName, - StructuredResponseParser parser) { + private T invokeStructuredAndValidate(LLMProvider provider, String model, String repairModel, String prompt, + String taskName, StructuredResponseParser parser) { String currentPrompt = prompt; String rawResponse = null; ResponseStatusException lastFailure = null; int maxAttempts = maxProviderRetryAttempts(); for (int attempt = 1; attempt <= maxAttempts; attempt++) { + String attemptModel = attempt == 1 ? model : repairModel; try { - rawResponse = invokeStructuredProvider(provider, model, currentPrompt); + rawResponse = invokeStructuredProvider(provider, attemptModel, repairModel, currentPrompt, taskName, + attempt, maxAttempts); if (log.isTraceEnabled()) { log.trace( "Assistant raw response for task '{}' (attempt {}/{}):\n---BEGIN ASSISTANT RESPONSE---\n{}\n---END ASSISTANT RESPONSE---", @@ -427,10 +478,90 @@ public class FlowAssistantService { } } - throw lastFailure == null + ResponseStatusException finalFailure = lastFailure == null ? new ResponseStatusException(HttpStatus.BAD_GATEWAY, "Assistant returned an invalid " + taskName + " payload") : lastFailure; + if (HttpStatus.BAD_GATEWAY.equals(finalFailure.getStatusCode())) { + log.warn( + "Assistant returning BAD_GATEWAY for task '{}': {}. Last raw LLM response follows.\n---BEGIN ASSISTANT RESPONSE---\n{}\n---END ASSISTANT RESPONSE---", + taskName, + finalFailure.getReason(), + rawResponse == null ? "(null)" : rawResponse); + } + throw finalFailure; + } + + private void logAssistantRawResponse(String taskName, String model, String mode, int attempt, int maxAttempts, + String rawResponse) { + assistantResponseLog.info( + "task={} model={} mode={} attempt={}/{} rawLength={}\n---BEGIN ASSISTANT RESPONSE---\n{}\n---END ASSISTANT RESPONSE---", + taskName, + model, + mode, + attempt, + maxAttempts, + rawResponse == null ? 0 : rawResponse.length(), + rawResponse == null ? "(null)" : rawResponse); + } + + private boolean ensureAssistantRequestMdc(String intent) { + if (MDC.get("assistantCallId") != null) { + if (MDC.get("assistantIntent") == null) { + MDC.put("assistantIntent", intent); + } + return false; + } + MDC.put("assistantSessionId", "direct"); + MDC.put("assistantCallId", UUID.randomUUID().toString()); + MDC.put("assistantOwner", "direct"); + MDC.put("assistantIntent", intent); + return true; + } + + private void clearAssistantRequestMdc(boolean directMdc) { + if (!directMdc) { + return; + } + MDC.remove("assistantSessionId"); + MDC.remove("assistantCallId"); + MDC.remove("assistantOwner"); + MDC.remove("assistantIntent"); + MDC.remove("assistantPhase"); + } + + private ResolvedAssistantModels resolveAssistantModels(String workflowModel, + AssistantModelSelection requestedPhaseModels) { + return new ResolvedAssistantModels( + firstNonBlank( + requestedPhaseModels == null ? null : requestedPhaseModels.planningModel(), + defaultPlanningModel, + workflowModel), + firstNonBlank( + requestedPhaseModels == null ? null : requestedPhaseModels.jsonModel(), + defaultJsonModel, + workflowModel), + firstNonBlank( + requestedPhaseModels == null ? null : requestedPhaseModels.repairModel(), + defaultRepairModel, + workflowModel)); + } + + private String planningModelFor(OperationMode mode, ResolvedAssistantModels phaseModels) { + return mode == OperationMode.FIX ? phaseModels.repairModel() : phaseModels.planningModel(); + } + + private String jsonModelFor(OperationMode mode, ResolvedAssistantModels phaseModels) { + return mode == OperationMode.FIX ? phaseModels.repairModel() : phaseModels.jsonModel(); + } + + private String firstNonBlank(String... values) { + for (String value : values) { + if (value != null && !value.isBlank()) { + return value; + } + } + return null; } private int maxProviderRetryAttempts() { @@ -478,7 +609,8 @@ public class FlowAssistantService { } private Block buildBlock(BlockCatalogService.AssistantPromptBlockDescriptor descriptor, AssistantBlockPlan blockPlan, - AssistantConfiguredBlockDraft draft, String model, boolean requireSharedMemorySemantics) { + AssistantConfiguredBlockDraft draft, String model, boolean requireSharedMemorySemantics, int blockIndex, + int blockCount) { if (!(draft.config() instanceof ObjectNode configNode)) { throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, "Assistant returned a non-object config for block " + blockPlan.blockId()); @@ -490,6 +622,7 @@ public class FlowAssistantService { injectSystemManagedFields(normalizedConfig, descriptor, model); normalizeHttpServerCallAuthorization(normalizedConfig, blockPlan); normalizeMcpAgentSharedMemory(normalizedConfig, blockPlan, model, requireSharedMemorySemantics); + ensureSequentialInputPlaceholder(normalizedConfig, descriptor, blockPlan, blockIndex, blockCount); try { BlockConfiguration configuration = ObjectMapperHolder.mapper.treeToValue(normalizedConfig, @@ -531,6 +664,33 @@ public class FlowAssistantService { } } + private void ensureSequentialInputPlaceholder(ObjectNode config, + BlockCatalogService.AssistantPromptBlockDescriptor descriptor, AssistantBlockPlan blockPlan, int blockIndex, + int blockCount) { + if (blockIndex <= 0 || blockCount < 2 || !isPromptDrivenConfiguration(descriptor.configurationType())) { + return; + } + + String prompt = textOrEmpty(config.path("prompt")); + if (prompt.contains("${{")) { + return; + } + + String dependencyHint = "Use upstream workflow context from ${{input}}."; + String normalizedPrompt = prompt.isBlank() + ? dependencyHint + : prompt.stripTrailing() + "\n\n" + dependencyHint; + config.put("prompt", normalizedPrompt); + log.debug("Added default upstream input placeholder to assistant-configured block {} ({})", + blockPlan.blockId(), + blockPlan.blockType()); + } + + private boolean isPromptDrivenConfiguration(String configurationType) { + return isConfigurationType(configurationType, "LLMBlockConfiguration") + || isConfigurationType(configurationType, "MCPAgentBlockConfiguration"); + } + private void normalizeMcpAgentSharedMemory(ObjectNode config, AssistantBlockPlan blockPlan, String model, boolean requireSharedMemorySemantics) { if (!requireSharedMemorySemantics || !"MCPAgent".equals(blockPlan.blockType())) { @@ -770,6 +930,28 @@ public class FlowAssistantService { + connection.getTargetId() + "|" + connection.getTargetName(); } + private List toValidConnections(List draftedConnections, + Map> blocksByPlanId, Map> blocksByAlias) { + if (draftedConnections == null || draftedConnections.isEmpty()) { + return List.of(); + } + + List validConnections = new ArrayList<>(); + for (AssistantConnectionDraft draftedConnection : draftedConnections) { + try { + validConnections.add(toConnection(draftedConnection, blocksByPlanId, blocksByAlias)); + } catch (ResponseStatusException e) { + if (!HttpStatus.BAD_GATEWAY.equals(e.getStatusCode())) { + throw e; + } + log.warn("Skipping invalid assistant connection draft: {}. Reason: {}", + draftedConnection, + e.getReason()); + } + } + return List.copyOf(validConnections); + } + private Connection toConnection(AssistantConnectionDraft connection, Map> blocksByPlanId, Map> blocksByAlias) { Block source = resolveConnectionBlock(connection.fromBlockId(), blocksByPlanId, blocksByAlias); @@ -782,7 +964,10 @@ public class FlowAssistantService { } if (source == null || target == null) { throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, - "Assistant returned a connection with unknown block ids"); + "Assistant returned a connection with unknown block ids" + + " (fromBlockId=" + connection.fromBlockId() + + ", toBlockId=" + connection.toBlockId() + + ", allowedBlockIds=" + blocksByPlanId.keySet() + ")"); } String sourceName = resolveConnectionOutputName(source, connection.fromOutput()); String targetName = resolveConnectionInputName(target, connection.toInput()); @@ -800,6 +985,11 @@ public class FlowAssistantService { + (target.getInputs() == null ? "none" : target.getInputs().stream() .map(io -> io.getName()).toList()) + ")"); } + if (isModelInputConnection(sourceName, targetName)) { + throw new ResponseStatusException(HttpStatus.BAD_GATEWAY, + "Assistant returned a connection to technical input 'model' on block '" + target.getName() + + "'"); + } return Connection.builder() .sourceId(source.getId()) .sourceName(sourceName) @@ -808,6 +998,11 @@ public class FlowAssistantService { .build(); } + private boolean isModelInputConnection(String sourceName, String targetName) { + return "model".equals(normalizeBlockReference(targetName)) + && !"model".equals(normalizeBlockReference(sourceName)); + } + private List inferSequentialConnections(List> blocks) { if (blocks == null || blocks.size() < 2) { return List.of(); @@ -1613,10 +1808,14 @@ public class FlowAssistantService { } } - private String invokeStructuredProvider(LLMProvider provider, String model, String prompt) { + private String invokeStructuredProvider(LLMProvider provider, String model, String repairModel, String prompt, + String taskName, int attempt, int maxAttempts) { RuntimeException structuredFailure = null; try { String structuredResponse = provider.generateJson(model, prompt); + if (structuredResponse != null) { + logAssistantRawResponse(taskName, model, "json", attempt, maxAttempts, structuredResponse); + } if (structuredResponse != null && !structuredResponse.isBlank()) { return structuredResponse; } @@ -1631,11 +1830,19 @@ public class FlowAssistantService { try { String fallbackResponse = provider.generate(model, prompt); + if (fallbackResponse != null) { + logAssistantRawResponse(taskName, model, "text-fallback", attempt, maxAttempts, fallbackResponse); + } if (fallbackResponse != null && !fallbackResponse.isBlank()) { if (!looksLikeStructuredJson(fallbackResponse)) { try { - String reformatted = provider.generateJson(model, + String reformatModel = firstNonBlank(repairModel, model); + String reformatted = provider.generateJson(reformatModel, promptService.buildJsonReformatPrompt(prompt, fallbackResponse)); + if (reformatted != null) { + logAssistantRawResponse(taskName, reformatModel, "json-reformat", attempt, maxAttempts, + reformatted); + } if (reformatted != null && !reformatted.isBlank()) { log.trace("Assistant non-JSON fallback response was auto-reformatted using JSON mode"); return reformatted; 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 index 8f1c25c..aea7bf5 100644 --- 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 @@ -3,5 +3,6 @@ package it.cnr.isti.workflow.manager.assistant.model; public record AssistantConfigView( String provider, String defaultModel, - String availableModelsRetrieverUrl) { + String availableModelsRetrieverUrl, + AssistantModelSelection defaultPhaseModels) { } diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantFixRequest.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantFixRequest.java index 902dd8c..16e5836 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantFixRequest.java +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantFixRequest.java @@ -15,5 +15,11 @@ public record AssistantFixRequest( @Valid @NotNull FlowCreateRequest flow, List validationErrors, @NotBlank String model, - @Min(0) @Max(3) Integer maxRepairAttempts) { + @Min(0) @Max(3) Integer maxRepairAttempts, + @Valid AssistantModelSelection phaseModels) { + + public AssistantFixRequest(String userPrompt, FlowCreateRequest flow, + List validationErrors, String model, Integer maxRepairAttempts) { + this(userPrompt, flow, validationErrors, model, maxRepairAttempts, null); + } } diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantGenerationRequest.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantGenerationRequest.java index e607c4f..e43d427 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantGenerationRequest.java +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantGenerationRequest.java @@ -3,9 +3,15 @@ package it.cnr.isti.workflow.manager.assistant.model; import jakarta.validation.constraints.Max; import jakarta.validation.constraints.Min; import jakarta.validation.constraints.NotBlank; +import jakarta.validation.Valid; public record AssistantGenerationRequest( @NotBlank String userPrompt, @NotBlank String model, - @Min(0) @Max(3) Integer maxRepairAttempts) { + @Min(0) @Max(3) Integer maxRepairAttempts, + @Valid AssistantModelSelection phaseModels) { + + public AssistantGenerationRequest(String userPrompt, String model, Integer maxRepairAttempts) { + this(userPrompt, model, maxRepairAttempts, null); + } } diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantModelSelection.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantModelSelection.java new file mode 100644 index 0000000..6983d11 --- /dev/null +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantModelSelection.java @@ -0,0 +1,7 @@ +package it.cnr.isti.workflow.manager.assistant.model; + +public record AssistantModelSelection( + String planningModel, + String jsonModel, + String repairModel) { +} diff --git a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantRefineRequest.java b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantRefineRequest.java index d15480a..6e5eee0 100644 --- a/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantRefineRequest.java +++ b/src/main/java/it/cnr/isti/workflow/manager/assistant/model/AssistantRefineRequest.java @@ -11,5 +11,11 @@ public record AssistantRefineRequest( @NotBlank String userPrompt, @Valid @NotNull FlowCreateRequest flow, @NotBlank String model, - @Min(0) @Max(3) Integer maxRepairAttempts) { + @Min(0) @Max(3) Integer maxRepairAttempts, + @Valid AssistantModelSelection phaseModels) { + + public AssistantRefineRequest(String userPrompt, FlowCreateRequest flow, String model, + Integer maxRepairAttempts) { + this(userPrompt, flow, model, maxRepairAttempts, null); + } } 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 index 12ae50e..77fd246 100644 --- 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 @@ -1,6 +1,13 @@ package it.cnr.isti.workflow.manager.assistant.model; +import jakarta.validation.Valid; import jakarta.validation.constraints.NotBlank; -public record AssistantSessionCreateRequest(@NotBlank String model) { +public record AssistantSessionCreateRequest( + @NotBlank String model, + @Valid AssistantModelSelection phaseModels) { + + public AssistantSessionCreateRequest(String model) { + this(model, null); + } } 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 2b758ed..6cd009f 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 @@ -23,6 +23,7 @@ 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.AssistantModelSelection; 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; @@ -46,6 +47,15 @@ public class AssistantController { @Value("${app.assistant.default-model}") private String defaultAssistantModel; + @Value("${app.assistant.default-planning-model:${app.assistant.default-model}}") + private String defaultPlanningModel; + + @Value("${app.assistant.default-json-model:${app.assistant.default-model}}") + private String defaultJsonModel; + + @Value("${app.assistant.default-repair-model:${app.assistant.default-model}}") + private String defaultRepairModel; + @PostMapping("/flows/draft") @SecurityRequirement(name = "bearerAuth") @Operation(summary = "Generate a flow draft", description = "Creates a flow draft from a natural language request.") @@ -81,7 +91,8 @@ public class AssistantController { return new AssistantConfigView( INTERNAL_PROVIDER_NAME, defaultAssistantModel, - MODELS_RETRIEVER_URL); + MODELS_RETRIEVER_URL, + new AssistantModelSelection(defaultPlanningModel, defaultJsonModel, defaultRepairModel)); } @PostMapping("/sessions") diff --git a/src/main/resources/application-prod.properties b/src/main/resources/application-prod.properties index 096aa61..465fecd 100644 --- a/src/main/resources/application-prod.properties +++ b/src/main/resources/application-prod.properties @@ -31,6 +31,20 @@ app.webclient.connect-timeout-ms=5000 app.webclient.response-timeout-seconds=30 app.webclient.read-timeout-seconds=30 app.webclient.write-timeout-seconds=30 +app.llm.webclient.pool.max-connections=100 +app.llm.webclient.pool.pending-acquire-timeout-ms=30000 +app.llm.webclient.connect-timeout-ms=10000 +app.llm.webclient.response-timeout-seconds=180 +app.llm.webclient.read-timeout-seconds=180 +app.llm.webclient.write-timeout-seconds=60 +app.mcp.webclient.pool.max-connections=100 +app.mcp.webclient.pool.pending-acquire-timeout-ms=30000 +app.mcp.webclient.connect-timeout-ms=10000 +app.mcp.webclient.response-timeout-seconds=150 +app.mcp.webclient.read-timeout-seconds=150 +app.mcp.webclient.write-timeout-seconds=60 + +app.assistant.responses-log-file=${ASSISTANT_RESPONSES_LOG_FILE:logs/assistant-responses.log} logging.level.it.cnr.isti.workflow.manager=INFO logging.level.root=WARN diff --git a/src/main/resources/application.properties b/src/main/resources/application.properties index 39079a8..651b719 100644 --- a/src/main/resources/application.properties +++ b/src/main/resources/application.properties @@ -63,7 +63,26 @@ app.webclient.connect-timeout-ms=${WEBCLIENT_CONNECT_TIMEOUT_MS:5000} app.webclient.response-timeout-seconds=${WEBCLIENT_RESPONSE_TIMEOUT_SECONDS:30} app.webclient.read-timeout-seconds=${WEBCLIENT_READ_TIMEOUT_SECONDS:30} app.webclient.write-timeout-seconds=${WEBCLIENT_WRITE_TIMEOUT_SECONDS:30} -app.assistant.default-model=${ASSISTANT_DEFAULT_MODEL:qwen3:32b-q4_K_M} +app.llm.webclient.pool.max-connections=${LLM_WEBCLIENT_POOL_MAX_CONNECTIONS:100} +app.llm.webclient.pool.pending-acquire-timeout-ms=${LLM_WEBCLIENT_POOL_PENDING_ACQUIRE_TIMEOUT_MS:30000} +app.llm.webclient.connect-timeout-ms=${LLM_WEBCLIENT_CONNECT_TIMEOUT_MS:10000} +app.llm.webclient.response-timeout-seconds=${LLM_WEBCLIENT_RESPONSE_TIMEOUT_SECONDS:180} +app.llm.webclient.read-timeout-seconds=${LLM_WEBCLIENT_READ_TIMEOUT_SECONDS:180} +app.llm.webclient.write-timeout-seconds=${LLM_WEBCLIENT_WRITE_TIMEOUT_SECONDS:60} +app.mcp.webclient.pool.max-connections=${MCP_WEBCLIENT_POOL_MAX_CONNECTIONS:100} +app.mcp.webclient.pool.pending-acquire-timeout-ms=${MCP_WEBCLIENT_POOL_PENDING_ACQUIRE_TIMEOUT_MS:30000} +app.mcp.webclient.connect-timeout-ms=${MCP_WEBCLIENT_CONNECT_TIMEOUT_MS:10000} +app.mcp.webclient.response-timeout-seconds=${MCP_WEBCLIENT_RESPONSE_TIMEOUT_SECONDS:150} +app.mcp.webclient.read-timeout-seconds=${MCP_WEBCLIENT_READ_TIMEOUT_SECONDS:150} +app.mcp.webclient.write-timeout-seconds=${MCP_WEBCLIENT_WRITE_TIMEOUT_SECONDS:60} + +# Assistant LLM models configuration +app.assistant.default-model=${ASSISTANT_DEFAULT_MODEL:gpt-oss:20b} +app.assistant.default-planning-model=${ASSISTANT_PLANNING_MODEL:qwen3:14b} +app.assistant.default-json-model=${ASSISTANT_JSON_MODEL:qwen2.5-coder:14b} +app.assistant.default-repair-model=${ASSISTANT_REPAIR_MODEL:qwen2.5-coder:14b} +app.assistant.responses-log-file=${ASSISTANT_RESPONSES_LOG_FILE:logs/assistant-responses.log} + app.assistant.provider-retry-attempts=${ASSISTANT_PROVIDER_RETRY_ATTEMPTS:3} app.assistant.provider-retry-base-delay-ms=${ASSISTANT_PROVIDER_RETRY_BASE_DELAY_MS:120} app.assistant.provider-retry-max-delay-ms=${ASSISTANT_PROVIDER_RETRY_MAX_DELAY_MS:800} diff --git a/src/main/resources/logback-spring.xml b/src/main/resources/logback-spring.xml new file mode 100644 index 0000000..24dca21 --- /dev/null +++ b/src/main/resources/logback-spring.xml @@ -0,0 +1,30 @@ + + + + + + + + + ${ASSISTANT_RESPONSES_LOG_FILE} + + ${ASSISTANT_RESPONSES_LOG_FILE}.%d{yyyy-MM-dd}.%i.gz + 20MB + 14 + 500MB + + + %d{yyyy-MM-dd HH:mm:ss.SSS} level=%level sessionId=%X{assistantSessionId:-none} callId=%X{assistantCallId:-none} owner=%X{assistantOwner:-none} intent=%X{assistantIntent:-none} phase=%X{assistantPhase:-none} %msg%n + + + + + + + + + + + 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 5e740ba..f07765c 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 @@ -31,6 +31,7 @@ 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.AssistantModelSelection; 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; @@ -67,6 +68,9 @@ import static org.springframework.test.web.servlet.result.MockMvcResultMatchers. public class AssistantControllerTest { private static final String MODEL = "assistant-test-model"; + private static final String PLANNING_MODEL = "qwen3:14b"; + private static final String JSON_MODEL = "qwen2.5-coder:14b"; + private static final String REPAIR_MODEL = "qwen2.5-coder:14b-repair"; @Autowired private AssistantController assistantController; @@ -98,6 +102,47 @@ public class AssistantControllerTest { assertFalse(response.assistantRationale().isBlank()); } + @Test + public void draftUsesPhaseSpecificModels() { + mockAssistantResponsesForPhaseModels(false); + + AssistantFlowResponse response = assistantController.draft( + new AssistantGenerationRequest( + "create a flow that classifies incoming tickets", + MODEL, + 1, + new AssistantModelSelection(PLANNING_MODEL, JSON_MODEL, REPAIR_MODEL))); + + assertNotNull(response); + assertTrue(response.valid()); + LLMBlockConfiguration configuration = (LLMBlockConfiguration) response.flow().flow().getBlocks().getFirst() + .getSpecificConfiguration(); + assertEquals(MODEL, configuration.getLlmDescriptor().model()); + Mockito.verify(internalOllamaLLMProvider).generateJson(Mockito.eq(PLANNING_MODEL), + Mockito.argThat(prompt -> prompt.contains("TASK: PLAN"))); + Mockito.verify(internalOllamaLLMProvider).generateJson(Mockito.eq(JSON_MODEL), + Mockito.argThat(prompt -> prompt.contains("TASK: BLOCK_CONFIG"))); + Mockito.verify(internalOllamaLLMProvider, Mockito.never()).generateJson(Mockito.eq(JSON_MODEL), + Mockito.argThat(prompt -> prompt.contains("TASK: CONNECTIONS"))); + } + + @Test + public void draftUsesRepairModelForStructuredCorrection() { + mockAssistantResponsesForPhaseModels(true); + + AssistantFlowResponse response = assistantController.draft( + new AssistantGenerationRequest( + "create a flow that classifies incoming tickets", + MODEL, + 1, + new AssistantModelSelection(PLANNING_MODEL, JSON_MODEL, REPAIR_MODEL))); + + assertNotNull(response); + assertTrue(response.valid()); + Mockito.verify(internalOllamaLLMProvider).generateJson(Mockito.eq(REPAIR_MODEL), + Mockito.argThat(prompt -> prompt.contains("The previous response did not satisfy"))); + } + @Test public void draftRetriesAfterTransientProviderFailure() { mockAssistantResponsesWithTransientProviderFailure(); @@ -166,6 +211,25 @@ public class AssistantControllerTest { assertEquals(0, response.flow().flow().getConnections().size()); } + @Test + public void draftAddsInputPlaceholderForSequentialLlmBlocks() { + mockSequentialLlmResponsesWithoutPlaceholders(); + + AssistantFlowResponse response = assistantController.draft( + new AssistantGenerationRequest( + "create a two step workflow where each point is a node", + MODEL, + 1)); + + assertNotNull(response); + assertTrue(response.valid()); + assertEquals(2, response.flow().flow().getBlocks().size()); + assertEquals(1, response.flow().flow().getConnections().size()); + Block secondBlock = response.flow().flow().getBlocks().get(1); + assertTrue(secondBlock.getInputs().stream().anyMatch(input -> "input".equals(input.getName()))); + assertEquals("input", response.flow().flow().getConnections().getFirst().getTargetName()); + } + @Test public void draftUsesMinimalLlmPlanWhenAssistantReturnsNoBlocks() { mockAssistantResponsesWithEmptyPlan(); @@ -331,6 +395,10 @@ public class AssistantControllerTest { assertEquals("InternalOllama", config.provider()); assertEquals(MODEL, config.defaultModel()); assertEquals("/retriever/LLM/models?provider=InternalOllama", config.availableModelsRetrieverUrl()); + assertNotNull(config.defaultPhaseModels()); + assertEquals(MODEL, config.defaultPhaseModels().planningModel()); + assertEquals(MODEL, config.defaultPhaseModels().jsonModel()); + assertEquals(MODEL, config.defaultPhaseModels().repairModel()); } @Test @@ -641,6 +709,62 @@ public class AssistantControllerTest { .thenAnswer(answer); } + private void mockAssistantResponsesForPhaseModels(boolean forceStructuredRepair) { + AtomicBoolean blockConfigFailedOnce = new AtomicBoolean(false); + Answer answer = invocation -> { + String model = invocation.getArgument(0, String.class); + String prompt = invocation.getArgument(1, String.class); + if (prompt.contains("The previous response did not satisfy")) { + assertEquals(REPAIR_MODEL, model); + return TestAssistantResponses.wrap(java.util.Map.of( + "rationale", "Repaired the LLM block configuration.", + "block", java.util.Map.of( + "blockId", "b1", + "name", "Ticket classifier", + "config", java.util.Map.of( + "prompt", "Classify the ticket: ${{ticket}}")))); + } + if (prompt.contains("TASK: PLAN")) { + assertEquals(PLANNING_MODEL, model); + 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: BLOCK_CONFIG")) { + assertEquals(JSON_MODEL, model); + if (forceStructuredRepair && blockConfigFailedOnce.compareAndSet(false, true)) { + return "Configure the LLM block for ticket classification."; + } + 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")) { + assertEquals(JSON_MODEL, model); + 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.anyString(), Mockito.anyString())) + .thenAnswer(answer); + Mockito.when(internalOllamaLLMProvider.generateJson(Mockito.anyString(), Mockito.anyString())) + .thenAnswer(answer); + } + private void mockRefineKeepAndAddResponses() { Answer answer = invocation -> { String prompt = invocation.getArgument(1, String.class); @@ -963,6 +1087,64 @@ public class AssistantControllerTest { .thenAnswer(answer); } + private void mockSequentialLlmResponsesWithoutPlaceholders() { + Answer answer = invocation -> { + String prompt = invocation.getArgument(1, String.class); + if (prompt.contains("TASK: PLAN")) { + return TestAssistantResponses.wrap(java.util.Map.of( + "rationale", "Planned two sequential LLM steps.", + "plan", java.util.Map.of( + "name", "Sequential LLM flow", + "description", "Two ordered steps.", + "blocks", java.util.List.of( + java.util.Map.of( + "blockId", "b1", + "blockType", "LLMBlock", + "purpose", "Create the initial output"), + java.util.Map.of( + "blockId", "b2", + "blockType", "LLMBlock", + "purpose", "Consume the initial output"))))); + } + 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 first step without placeholders.", + "block", java.util.Map.of( + "blockId", "b1", + "name", "First step", + "config", java.util.Map.of( + "prompt", "Create the initial output.")))); + } + 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 second step without placeholders.", + "block", java.util.Map.of( + "blockId", "b2", + "name", "Second step", + "config", java.util.Map.of( + "prompt", "Consume the initial output.")))); + } + if (prompt.contains("TASK: CONNECTIONS")) { + return TestAssistantResponses.wrap(java.util.Map.of( + "rationale", "Connected first step to second step.", + "connections", java.util.List.of( + java.util.Map.of( + "fromBlockId", "b1", + "fromOutput", "response", + "toBlockId", "b2", + "toInput", "input")))); + } + throw new IllegalStateException("Unexpected assistant prompt:\n" + prompt); + }; + + Mockito.when(internalOllamaLLMProvider.generate(Mockito.eq(MODEL), Mockito.anyString())) + .thenAnswer(answer); + Mockito.when(internalOllamaLLMProvider.generateJson(Mockito.eq(MODEL), Mockito.anyString())) + .thenAnswer(answer); + } + private void mockAssistantResponsesWithEmptyPlan() { Answer answer = invocation -> { String prompt = invocation.getArgument(1, String.class); diff --git a/src/test/resources/logback-test.xml b/src/test/resources/logback-test.xml index 2fdb23a..3c46cfc 100644 --- a/src/test/resources/logback-test.xml +++ b/src/test/resources/logback-test.xml @@ -6,7 +6,8 @@ + - \ No newline at end of file + diff --git a/src/test/resources/test.properties b/src/test/resources/test.properties index a2089dc..57252b6 100644 --- a/src/test/resources/test.properties +++ b/src/test/resources/test.properties @@ -10,6 +10,9 @@ app.db.init.enabled=true app.assistant.default-model=assistant-test-model app.mcp.bridge.url=http://localhost:18080 app.turnstile.enabled=false +app.assistant.default-planning-model=assistant-test-model +app.assistant.default-json-model=assistant-test-model +app.assistant.default-repair-model=assistant-test-model app.assistant.provider-retry-attempts=3 app.assistant.provider-retry-base-delay-ms=1 app.assistant.provider-retry-max-delay-ms=2