refs #739: govern Smilegate few-shot references
This commit is contained in:
@@ -142,7 +142,7 @@ public class SelectAiService {
|
||||
QueryPlanContext queryPlan
|
||||
) {
|
||||
EnrichedPrompt enrichedPrompt = enrichWithFewShot(
|
||||
bearerToken, selectAi, executionPrompt);
|
||||
bearerToken, selectAi, executionPrompt, queryPlan == null ? "ANY" : queryPlan.targetType());
|
||||
String generatedSql = generate(selectAi, enrichedPrompt.prompt(), "showsql");
|
||||
String normalizedSql = validateReadOnlySql(generatedSql);
|
||||
if (allowedPrefixes != null && gameScopeService != null
|
||||
@@ -287,6 +287,11 @@ public class SelectAiService {
|
||||
item.put("answer", example.answer());
|
||||
}
|
||||
item.put("embeddingModel", example.embeddingModel());
|
||||
item.put("referenceKind", example.referenceKind());
|
||||
item.put("targetType", example.targetType());
|
||||
if (example.objectRole() != null) {
|
||||
item.put("objectRole", example.objectRole());
|
||||
}
|
||||
item.put("cosineDistance", example.cosineDistance());
|
||||
}
|
||||
}
|
||||
@@ -296,7 +301,7 @@ public class SelectAiService {
|
||||
requireActiveToken(bearerToken);
|
||||
String normalizedPrompt = requiredPrompt(prompt);
|
||||
BackofficeProperties.SelectAi selectAi = requiredSelectAi();
|
||||
EnrichedPrompt enrichedPrompt = enrichWithFewShot(bearerToken, selectAi, normalizedPrompt);
|
||||
EnrichedPrompt enrichedPrompt = enrichWithFewShot(bearerToken, selectAi, normalizedPrompt, "ANY");
|
||||
String selectAiPrompt = generate(selectAi, enrichedPrompt.prompt(), "showprompt");
|
||||
|
||||
ObjectNode response = objectMapper.createObjectNode();
|
||||
@@ -370,14 +375,15 @@ public class SelectAiService {
|
||||
private EnrichedPrompt enrichWithFewShot(
|
||||
String bearerToken,
|
||||
BackofficeProperties.SelectAi selectAi,
|
||||
String prompt
|
||||
String prompt,
|
||||
String targetType
|
||||
) {
|
||||
if (!fewShotEnabled(selectAi) || qaVectorService == null) {
|
||||
return new EnrichedPrompt(prompt, "DISABLED", 0, List.of());
|
||||
}
|
||||
try {
|
||||
List<QaVectorService.VectorExample> examples = qaVectorService
|
||||
.search(bearerToken, prompt, fewShotTopK(selectAi))
|
||||
.search(bearerToken, prompt, fewShotTopK(selectAi), targetType)
|
||||
.examples();
|
||||
if (examples.isEmpty()) {
|
||||
return new EnrichedPrompt(composePolicyPrompt(prompt), "NO_MATCH", 0, List.of());
|
||||
@@ -412,12 +418,10 @@ public class SelectAiService {
|
||||
if (included >= MAX_FEW_SHOT_EXAMPLES) {
|
||||
break;
|
||||
}
|
||||
String answerSql = truncate(example.answerSql(), MAX_FEW_SHOT_SQL_CHARS);
|
||||
if (answerSql.isBlank()) {
|
||||
String candidate = referenceCandidate(included + 1, example);
|
||||
if (candidate.isBlank()) {
|
||||
continue;
|
||||
}
|
||||
String candidate = "Example " + (included + 1) + " question:\n" + example.question()
|
||||
+ "\nExample " + (included + 1) + " verified SQL:\n" + answerSql + "\n\n";
|
||||
if (enriched.length() + candidate.length() + prompt.length() > MAX_ENRICHED_PROMPT_LENGTH) {
|
||||
break;
|
||||
}
|
||||
@@ -430,6 +434,19 @@ public class SelectAiService {
|
||||
return enriched.append("Original user question:\n").append(prompt).toString();
|
||||
}
|
||||
|
||||
private static String referenceCandidate(int index, QaVectorService.VectorExample example) {
|
||||
String prefix = "Example " + index + " question:\n" + example.question()
|
||||
+ "\nExample " + index + " target type: " + example.targetType()
|
||||
+ "\nExample " + index + " reference kind: " + example.referenceKind() + "\n";
|
||||
if ("NO_TARGET".equals(example.referenceKind())
|
||||
|| "OBJECT_UNAVAILABLE".equals(example.referenceKind())) {
|
||||
String boundary = truncate(example.answer(), MAX_FEW_SHOT_SQL_CHARS);
|
||||
return boundary.isBlank() ? "" : prefix + "Boundary outcome:\n" + boundary + "\n\n";
|
||||
}
|
||||
String answerSql = truncate(example.answerSql(), MAX_FEW_SHOT_SQL_CHARS);
|
||||
return answerSql.isBlank() ? "" : prefix + "Verified SQL template:\n" + answerSql + "\n\n";
|
||||
}
|
||||
|
||||
private static String truncate(String value, int maxLength) {
|
||||
String normalized = value == null ? "" : value.trim();
|
||||
return normalized.length() <= maxLength ? normalized : normalized.substring(0, maxLength);
|
||||
|
||||
Reference in New Issue
Block a user