diff --git a/src/main/kotlin/conversation/ConversationEngine.kt b/src/main/kotlin/conversation/ConversationEngine.kt index 1646e4c..e720cbe 100644 --- a/src/main/kotlin/conversation/ConversationEngine.kt +++ b/src/main/kotlin/conversation/ConversationEngine.kt @@ -5,6 +5,7 @@ import com.aallam.openai.api.chat.ChatCompletionRequest import com.aallam.openai.api.chat.ChatMessage import com.aallam.openai.api.chat.ChatRole import com.aallam.openai.api.chat.ToolCall +import com.aallam.openai.api.chat.ToolChoice import com.aallam.openai.api.core.Usage import com.aallam.openai.api.model.ModelId import io.ktor.util.collections.ConcurrentSet @@ -284,6 +285,7 @@ internal object ConversationEngine { temperature = endpoint.temperature, messages = history, tools = availableTools, + toolChoice = ToolChoice.Required, ) JChatGPT.logger.info("API Requesting... Model=${endpoint.model} [${endpoint.label}]") return endpoint.service.chatCompletions(request, onCacheUsage) diff --git a/src/test/kotlin/llm/ModelServiceTest.kt b/src/test/kotlin/llm/ModelServiceTest.kt index cc8ebdd..f218a66 100644 --- a/src/test/kotlin/llm/ModelServiceTest.kt +++ b/src/test/kotlin/llm/ModelServiceTest.kt @@ -2,6 +2,7 @@ package top.jie65535.mirai.llm import com.aallam.openai.api.chat.ChatCompletionRequest import com.aallam.openai.api.chat.ChatMessage +import com.aallam.openai.api.chat.ToolChoice import com.aallam.openai.api.model.ModelId import com.sun.net.httpserver.HttpServer import io.ktor.utils.io.ByteReadChannel @@ -96,6 +97,23 @@ class ModelServiceTest { ) } + @Test + fun serializesRequiredToolChoice() { + val body = service.buildRequestBody( + request = ChatCompletionRequest( + model = ModelId("test-model"), + messages = listOf(ChatMessage.User("hello")), + toolChoice = ToolChoice.Required, + ), + stream = true, + ) + + assertEquals( + "required", + Json.parseToJsonElement(body).jsonObject.getValue("tool_choice").jsonPrimitive.content, + ) + } + @Test fun acceptsStandaloneUsageAndEmptyHeartbeatEvents() { assertNull(service.decodeStreamChunk("{}"))