mirror of
https://github.com/jie65535/JChatGPT.git
synced 2026-09-15 02:56:10 +08:00
conversation: require tool calls
This commit is contained in:
@@ -5,6 +5,7 @@ import com.aallam.openai.api.chat.ChatCompletionRequest
|
|||||||
import com.aallam.openai.api.chat.ChatMessage
|
import com.aallam.openai.api.chat.ChatMessage
|
||||||
import com.aallam.openai.api.chat.ChatRole
|
import com.aallam.openai.api.chat.ChatRole
|
||||||
import com.aallam.openai.api.chat.ToolCall
|
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.core.Usage
|
||||||
import com.aallam.openai.api.model.ModelId
|
import com.aallam.openai.api.model.ModelId
|
||||||
import io.ktor.util.collections.ConcurrentSet
|
import io.ktor.util.collections.ConcurrentSet
|
||||||
@@ -284,6 +285,7 @@ internal object ConversationEngine {
|
|||||||
temperature = endpoint.temperature,
|
temperature = endpoint.temperature,
|
||||||
messages = history,
|
messages = history,
|
||||||
tools = availableTools,
|
tools = availableTools,
|
||||||
|
toolChoice = ToolChoice.Required,
|
||||||
)
|
)
|
||||||
JChatGPT.logger.info("API Requesting... Model=${endpoint.model} [${endpoint.label}]")
|
JChatGPT.logger.info("API Requesting... Model=${endpoint.model} [${endpoint.label}]")
|
||||||
return endpoint.service.chatCompletions(request, onCacheUsage)
|
return endpoint.service.chatCompletions(request, onCacheUsage)
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package top.jie65535.mirai.llm
|
|||||||
|
|
||||||
import com.aallam.openai.api.chat.ChatCompletionRequest
|
import com.aallam.openai.api.chat.ChatCompletionRequest
|
||||||
import com.aallam.openai.api.chat.ChatMessage
|
import com.aallam.openai.api.chat.ChatMessage
|
||||||
|
import com.aallam.openai.api.chat.ToolChoice
|
||||||
import com.aallam.openai.api.model.ModelId
|
import com.aallam.openai.api.model.ModelId
|
||||||
import com.sun.net.httpserver.HttpServer
|
import com.sun.net.httpserver.HttpServer
|
||||||
import io.ktor.utils.io.ByteReadChannel
|
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
|
@Test
|
||||||
fun acceptsStandaloneUsageAndEmptyHeartbeatEvents() {
|
fun acceptsStandaloneUsageAndEmptyHeartbeatEvents() {
|
||||||
assertNull(service.decodeStreamChunk("{}"))
|
assertNull(service.decodeStreamChunk("{}"))
|
||||||
|
|||||||
Reference in New Issue
Block a user