Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -6,17 +6,27 @@
import com.devkor.ifive.nadab.domain.askchat.core.dto.AskChatAnswerReferenceDocument;
import com.devkor.ifive.nadab.domain.askchat.core.properties.AskChatAnswerProperties;
import com.devkor.ifive.nadab.global.core.prompt.askchat.AskChatAnswerPromptLoader;
import com.devkor.ifive.nadab.global.core.response.ErrorCode;
import com.devkor.ifive.nadab.global.exception.ai.AiServiceException;
import lombok.RequiredArgsConstructor;
import org.springframework.stereotype.Component;

import java.util.List;
import java.util.regex.Matcher;
import java.util.regex.Pattern;

@Component
@RequiredArgsConstructor
public class AskChatAnswerPromptComposer implements AskChatAnswerPromptAugmenter {

private static final String NO_REFERENCE_DOCUMENT = "검색된 사용자 기록이 없습니다.";
private static final String NO_RECENT_MESSAGE = "최근 대화가 없습니다.";
private static final Pattern PROMPT_VERSION_SECTION_PATTERN = Pattern.compile(
"(?m)^\\[프롬프트 버전]\\R\\{promptVersion}(?:\\R){1,2}"
);
private static final Pattern TEMPLATE_VARIABLE_PATTERN = Pattern.compile(
"\\{([A-Za-z][A-Za-z0-9_]*)}"
);

private final AskChatAnswerProperties properties;
private final AskChatAnswerPromptLoader promptLoader;
Expand All @@ -34,12 +44,33 @@ private String systemPrompt() {
}

private String userPrompt(AskChatAnswerPromptContext context) {
return promptLoader.loadUserPrompt()
.replace("{promptVersion}", String.valueOf(properties.getPromptVersion()))
.replace("{question}", context.question())
.replace("{recentMessages}", formatRecentMessages(context.recentMessages()))
.replace("{referenceDocuments}", formatReferenceDocuments(context.referenceDocuments()))
.replace("{followUpQuestionCount}", String.valueOf(properties.getFollowUpQuestionCount()));
String template = PROMPT_VERSION_SECTION_PATTERN
.matcher(promptLoader.loadUserPrompt())
.replaceFirst("");

return renderTemplate(template, context);
}

private String renderTemplate(String template, AskChatAnswerPromptContext context) {
String question = escapePromptData(context.question());
String recentMessages = formatRecentMessages(context.recentMessages());
String referenceDocuments = formatReferenceDocuments(context.referenceDocuments());
String followUpQuestionCount = String.valueOf(properties.getFollowUpQuestionCount());

Matcher matcher = TEMPLATE_VARIABLE_PATTERN.matcher(template);
StringBuilder builder = new StringBuilder();
while (matcher.find()) {
String replacement = switch (matcher.group(1)) {
case "question" -> question;
case "recentMessages" -> recentMessages;
case "referenceDocuments" -> referenceDocuments;
case "followUpQuestionCount" -> followUpQuestionCount;
default -> throw new AiServiceException(ErrorCode.PROMPT_ASK_CHAT_VARIABLE_UNSUPPORTED);
};
matcher.appendReplacement(builder, Matcher.quoteReplacement(replacement));
}
matcher.appendTail(builder);
return builder.toString();
}

private String formatRecentMessages(List<AskChatAnswerConversationMessage> messages) {
Expand All @@ -54,7 +85,7 @@ private String formatRecentMessages(List<AskChatAnswerConversationMessage> messa
.append(". ")
.append(message.role())
.append(": ")
.append(message.content())
.append(escapePromptData(message.content()))
.append(System.lineSeparator());
}
return builder.toString().trim();
Expand All @@ -68,19 +99,25 @@ private String formatReferenceDocuments(List<AskChatAnswerReferenceDocument> doc
StringBuilder builder = new StringBuilder();
for (int i = 0; i < documents.size(); i++) {
AskChatAnswerReferenceDocument document = documents.get(i);
builder.append(i + 1)
.append(". documentId=")
.append(document.documentId())
.append(", sourceType=")
.append(document.sourceType())
.append(", interestCode=")
.append(document.interestCode())
.append(", distance=")
.append(document.distance())
builder.append("[기록 ")
.append(i + 1)
.append("]")
.append(System.lineSeparator())
.append(escapePromptData(document.content()))
.append(System.lineSeparator())
.append(document.content())
.append(System.lineSeparator());
}
return builder.toString().trim();
}

private String escapePromptData(String value) {
if (value == null) {
return "";
}

return value
.replace("&", "&amp;")
.replace("<", "&lt;")
.replace(">", "&gt;");
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,9 @@ public class AskChatAnswerProperties {
@NotBlank
private String model = "gpt-5.6-luna";

@NotBlank
private String reasoningEffort = "low";

@DecimalMin("0.0")
@DecimalMax("2.0")
private double temperature = 1.0;
Expand All @@ -37,7 +40,4 @@ public class AskChatAnswerProperties {

@Min(0)
private int followUpQuestionCount = 2;

@Min(1)
private int promptVersion = 2;
}
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,8 @@
@RequiredArgsConstructor
public class AskChatAnswerLlmClient {

private static final int MAX_FOLLOW_UP_QUESTION_LENGTH = 30;

/*
* Keep prompt augmentation behind this boundary while evidence documents are stored in
* ask_chat_message_references. A future Spring AI Advisor implementation can replace this
Expand Down Expand Up @@ -75,6 +77,7 @@ public AskChatAnswerGenerationResult generate(AskChatAnswerPromptContext context
private OpenAiChatOptions options() {
return OpenAiChatOptions.builder()
.model(properties.getModel())
.reasoningEffort(properties.getReasoningEffort())
.temperature(properties.getTemperature())
.maxCompletionTokens(properties.getMaxTokens())
.build();
Expand All @@ -93,17 +96,33 @@ private void validateAnswer(AskChatGeneratedAnswer answer) {
throw new AiResponseParseException(ErrorCode.AI_RESPONSE_FORMAT_INVALID);
}

if (containsUnsupportedScript(answer.answer())) {
throw new AiResponseParseException(ErrorCode.AI_RESPONSE_UNSUPPORTED_SCRIPT);
}

if (answer.followUpQuestions().size() > properties.getFollowUpQuestionCount()) {
throw new AiResponseParseException(ErrorCode.AI_RESPONSE_FORMAT_INVALID);
}

for (String followUpQuestion : answer.followUpQuestions()) {
if (isBlank(followUpQuestion)) {
if (isBlank(followUpQuestion)
|| followUpQuestion.codePointCount(0, followUpQuestion.length()) > MAX_FOLLOW_UP_QUESTION_LENGTH) {
throw new AiResponseParseException(ErrorCode.AI_RESPONSE_FORMAT_INVALID);
}

if (containsUnsupportedScript(followUpQuestion)) {
throw new AiResponseParseException(ErrorCode.AI_RESPONSE_UNSUPPORTED_SCRIPT);
}
}
}

private boolean containsUnsupportedScript(String value) {
return value.codePoints().anyMatch(codePoint -> switch (Character.UnicodeScript.of(codePoint)) {
case COMMON, INHERITED, HANGUL, LATIN -> false;
default -> true;
});
}

private List<Long> referenceDocumentIds(AskChatAnswerPromptContext context) {
return context.referenceDocuments().stream()
.map(AskChatAnswerReferenceDocument::documentId)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -5,11 +5,13 @@
import com.devkor.ifive.nadab.domain.dailyreport.api.dto.response.CreateAnswerImageUploadUrlResponse;
import com.devkor.ifive.nadab.domain.dailyreport.api.dto.response.CreateDailyReportResponse;
import com.devkor.ifive.nadab.domain.dailyreport.api.dto.response.ImageStatusResponse;
import com.devkor.ifive.nadab.domain.dailyreport.application.helper.DailyReportModelSelector;
import com.devkor.ifive.nadab.domain.dailyreport.core.dto.ConfirmDailyAndRewardDto;
import com.devkor.ifive.nadab.domain.dailyreport.core.dto.PrepareDailyResultDto;
import com.devkor.ifive.nadab.domain.dailyreport.core.dto.AiDailyReportResultDto;
import com.devkor.ifive.nadab.domain.dailyreport.core.entity.AnswerEntry;
import com.devkor.ifive.nadab.domain.dailyreport.core.entity.ImageStatus;
import com.devkor.ifive.nadab.domain.dailyreport.core.properties.DailyReportLlmProperties.ModelCandidate;
import com.devkor.ifive.nadab.domain.dailyreport.infra.DailyReportLlmClient;
import com.devkor.ifive.nadab.domain.question.core.entity.DailyQuestion;
import com.devkor.ifive.nadab.domain.question.core.entity.UserDailyQuestion;
Expand Down Expand Up @@ -48,6 +50,7 @@ public class DailyReportService {
private final DailyReportTxService dailyReportTxService;
private final ProfileImageService profileImageService;

private final DailyReportModelSelector dailyReportModelSelector;
private final DailyReportLlmClient dailyReportLlmClient;
private final ReportGenerationLogRecorder reportGenerationLogRecorder;

Expand Down Expand Up @@ -88,19 +91,20 @@ public CreateDailyReportResponse generateDailyReport(Long userId, DailyReportReq
PrepareDailyResultDto prep = dailyReportTxService.prepareDaily(user, question, request.answer(), isDayPassed, request.objectKey());

AnswerEntry answerEntry = prep.entry();
ModelCandidate modelCandidate = dailyReportModelSelector.select();
Long generationLogId = reportGenerationLogRecorder.start(
userId,
ReportGenerationType.DAILY,
prep.reportId(),
ReportGenerationStep.DAILY_GENERATE,
LlmProvider.OPENAI,
dailyReportLlmClient.model()
modelCandidate.getModel()
);

AiDailyReportResultDto dto;
try {
LlmGenerationResult<AiDailyReportResultDto> generationResult =
dailyReportLlmClient.generate(question.getQuestionText(), answerEntry);
dailyReportLlmClient.generate(question.getQuestionText(), answerEntry, modelCandidate);
dto = generationResult.content();
LlmTokenUsage tokenUsage = generationResult.tokenUsage();
reportGenerationLogRecorder.recordTokenUsage(
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
package com.devkor.ifive.nadab.domain.dailyreport.application.helper;

import com.devkor.ifive.nadab.domain.dailyreport.core.properties.DailyReportLlmProperties;
import com.devkor.ifive.nadab.domain.dailyreport.core.properties.DailyReportLlmProperties.ModelCandidate;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.stereotype.Component;

import java.util.concurrent.ThreadLocalRandom;
import java.util.function.IntUnaryOperator;

@Component
public class DailyReportModelSelector {

private static final int TOTAL_WEIGHT = 100;

private final DailyReportLlmProperties properties;
private final IntUnaryOperator randomValueGenerator;

@Autowired
public DailyReportModelSelector(DailyReportLlmProperties properties) {
this(properties, bound -> ThreadLocalRandom.current().nextInt(bound));
}

DailyReportModelSelector(
DailyReportLlmProperties properties,
IntUnaryOperator randomValueGenerator
) {
this.properties = properties;
this.randomValueGenerator = randomValueGenerator;
}

public ModelCandidate select() {
int randomValue = randomValueGenerator.applyAsInt(TOTAL_WEIGHT);
int cumulativeWeight = 0;

for (ModelCandidate candidate : properties.getCandidates()) {
cumulativeWeight += candidate.getWeight();
if (randomValue < cumulativeWeight) {
return candidate;
}
}

throw new IllegalStateException("Failed to select a DailyReport LLM model candidate");
}
}
Original file line number Diff line number Diff line change
@@ -1,37 +1,86 @@
package com.devkor.ifive.nadab.domain.dailyreport.core.properties;

import jakarta.validation.Valid;
import jakarta.validation.constraints.AssertTrue;
import jakarta.validation.constraints.DecimalMax;
import jakarta.validation.constraints.DecimalMin;
import jakarta.validation.constraints.Max;
import jakarta.validation.constraints.Min;
import jakarta.validation.constraints.NotBlank;
import jakarta.validation.constraints.NotEmpty;
import jakarta.validation.constraints.NotNull;
import lombok.Getter;
import lombok.Setter;
import org.springframework.boot.context.properties.ConfigurationProperties;
import org.springframework.stereotype.Component;
import org.springframework.validation.annotation.Validated;

import java.util.ArrayList;
import java.util.HashSet;
import java.util.List;
import java.util.Objects;

@Component
@Getter
@Setter
@Validated
@ConfigurationProperties(prefix = "daily-report.llm")
public class DailyReportLlmProperties {

@NotBlank
private String model = "gpt-4o-mini";
@Valid
@NotEmpty
private List<@NotNull ModelCandidate> candidates = new ArrayList<>();

@AssertTrue(message = "daily-report.llm.candidates weights must total 100")
public boolean isCandidateWeightTotalValid() {
if (candidates == null || candidates.isEmpty()) {
return true;
}

return candidates.stream()
.filter(Objects::nonNull)
.mapToInt(ModelCandidate::getWeight)
.sum() == 100;
}

@AssertTrue(message = "daily-report.llm.candidates models must be unique")
public boolean isCandidateModelUnique() {
if (candidates == null || candidates.isEmpty()) {
return true;
}

List<String> models = candidates.stream()
.filter(Objects::nonNull)
.map(ModelCandidate::getModel)
.filter(Objects::nonNull)
.toList();

@DecimalMin("0.0")
@DecimalMax("2.0")
private double temperature = 0.3;
return new HashSet<>(models).size() == models.size();
}

@Getter
@Setter
public static class ModelCandidate {

@NotBlank
private String model;

@Min(1)
private int maxOutputTokens = 512;
@Min(1)
@Max(100)
private int weight;

@NotNull
private TokenLimitParameter tokenLimitParameter = TokenLimitParameter.MAX_TOKENS;
@DecimalMin("0.0")
@DecimalMax("2.0")
private double temperature;

private String reasoningEffort;
@Min(1)
private int maxOutputTokens;

@NotNull
private TokenLimitParameter tokenLimitParameter;

private String reasoningEffort;
}

public enum TokenLimitParameter {
MAX_TOKENS,
Expand Down
Loading
Loading