modify message type get value

This commit is contained in:
Ueibin Kim
2023-09-26 02:18:14 +09:00
committed by Mark Pollack
parent cffa790a8b
commit 5ff2c3bb92
4 changed files with 13 additions and 5 deletions

View File

@@ -73,7 +73,7 @@ public class AzureOpenAiClient implements AiClient {
List<Message> messages = prompt.getMessages();
List<ChatMessage> azureMessages = new ArrayList<>();
for (Message message : messages) {
String messageType = message.getMessageType().getValue();
String messageType = message.getMessageTypeValue();
ChatRole chatRole = ChatRole.fromString(messageType);
azureMessages.add(new ChatMessage(chatRole, message.getContent()));
}

View File

@@ -87,4 +87,9 @@ public abstract class AbstractMessage implements Message {
return this.messageType;
}
@Override
public String getMessageTypeValue() {
return this.messageType.getValue();
}
}

View File

@@ -26,4 +26,6 @@ public interface Message {
MessageType getMessageType();
String getMessageTypeValue();
}

View File

@@ -27,6 +27,7 @@ import org.springframework.ai.client.AiResponse;
import org.springframework.ai.client.Generation;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.messages.Message;
import org.springframework.ai.prompt.messages.MessageType;
import org.springframework.util.Assert;
import java.util.ArrayList;
@@ -79,7 +80,7 @@ public class OpenAiClient implements AiClient {
List<Message> messages = prompt.getMessages();
List<ChatMessage> theoMessages = new ArrayList<>();
for (Message message : messages) {
String messageType = message.getMessageType().getValue();
String messageType = message.getMessageTypeValue();
theoMessages.add(new ChatMessage(messageType, message.getContent()));
}
ChatCompletionRequest chatCompletionRequest = ChatCompletionRequest.builder()
@@ -148,14 +149,14 @@ public class OpenAiClient implements AiClient {
for (Message promptMessage : messages) {
switch (promptMessage.getMessageType()) {
case USER:
chatMessages.add(new ChatMessage("user", promptMessage.getContent()));
chatMessages.add(new ChatMessage(MessageType.USER.getValue(), promptMessage.getContent()));
break;
case ASSISTANT:
// TODO - valid?
chatMessages.add(new ChatMessage("assistant", promptMessage.getContent()));
chatMessages.add(new ChatMessage(MessageType.ASSISTANT.getValue(), promptMessage.getContent()));
break;
case SYSTEM:
chatMessages.add(new ChatMessage("system", promptMessage.getContent()));
chatMessages.add(new ChatMessage(MessageType.SYSTEM.getValue(), promptMessage.getContent()));
break;
case FUNCTION:
logger.error(