Skip to content
Open
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 @@ -99,6 +99,11 @@ val chatParams = OpenAIChatParams(
)
```

For OpenAI-compatible providers, configure arbitrary chat template arguments with
`OpenAIChatParams().withChatTemplateKwargs(buildJsonObject { put("enable_thinking", false) })`.
The object is sent unchanged as `chat_template_kwargs`. Passing null removes the arguments.
Other additional properties and parameter values are preserved.

#### OpenAI Responses API Parameters

For the Responses API, use `OpenAIResponsesParams`:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ public final class ai/koog/prompt/executor/clients/openai/OpenAIChatParams : ai/
public static synthetic fun copy$default (Lai/koog/prompt/executor/clients/openai/OpenAIChatParams;Ljava/lang/Double;Ljava/lang/Integer;Ljava/lang/Integer;Ljava/lang/String;Lai/koog/prompt/params/LLMParams$Schema;Lai/koog/prompt/params/LLMParams$ToolChoice;Ljava/lang/String;Ljava/util/Map;Ljava/lang/Double;Ljava/lang/Double;Ljava/lang/Boolean;Ljava/lang/String;Ljava/lang/String;Lai/koog/prompt/executor/clients/openai/base/models/ServiceTier;Ljava/lang/Boolean;Lai/koog/prompt/executor/clients/openai/base/models/OpenAIAudioConfig;Ljava/lang/Boolean;Lai/koog/prompt/executor/clients/openai/base/models/ReasoningEffort;Ljava/util/List;Ljava/lang/Integer;Ljava/lang/Double;Lai/koog/prompt/executor/clients/openai/base/models/OpenAIWebSearchOptions;ILjava/lang/Object;)Lai/koog/prompt/executor/clients/openai/OpenAIChatParams;
public fun equals (Ljava/lang/Object;)Z
public final fun getAudio ()Lai/koog/prompt/executor/clients/openai/base/models/OpenAIAudioConfig;
public final fun getChatTemplateKwargs ()Lkotlinx/serialization/json/JsonObject;
public final fun getFrequencyPenalty ()Ljava/lang/Double;
public final fun getLogprobs ()Ljava/lang/Boolean;
public final fun getParallelToolCalls ()Ljava/lang/Boolean;
Expand All @@ -23,6 +24,7 @@ public final class ai/koog/prompt/executor/clients/openai/OpenAIChatParams : ai/
public final fun getWebSearchOptions ()Lai/koog/prompt/executor/clients/openai/base/models/OpenAIWebSearchOptions;
public fun hashCode ()I
public fun toString ()Ljava/lang/String;
public final fun withChatTemplateKwargs (Lkotlinx/serialization/json/JsonObject;)Lai/koog/prompt/executor/clients/openai/OpenAIChatParams;
}

public final class ai/koog/prompt/executor/clients/openai/OpenAIClientFactory {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,11 @@ import ai.koog.prompt.executor.clients.openai.models.ReasoningConfig
import ai.koog.prompt.executor.clients.openai.models.Truncation
import ai.koog.prompt.params.LLMParams
import kotlinx.serialization.json.JsonElement
import kotlinx.serialization.json.JsonObject
import org.jetbrains.annotations.ApiStatus.Experimental

private const val CHAT_TEMPLATE_KWARGS_PROPERTY: String = "chat_template_kwargs"

internal sealed interface OpenAIParams

internal fun LLMParams.toOpenAIChatParams(): OpenAIChatParams {
Expand Down Expand Up @@ -230,6 +233,31 @@ public class OpenAIChatParams(
webSearchOptions = webSearchOptions,
)

/**
* Returns a copy with provider-specific chat template arguments for Chat Completions requests.
*
* The object is sent unchanged as `chat_template_kwargs` through [additionalProperties].
* Its keys and values are defined by the OpenAI-compatible provider. Passing null removes
* previously configured arguments while preserving other additional properties.
*/
public fun withChatTemplateKwargs(chatTemplateKwargs: JsonObject?): OpenAIChatParams {
val properties = additionalProperties.orEmpty().toMutableMap().apply {
if (chatTemplateKwargs == null) {
remove(CHAT_TEMPLATE_KWARGS_PROPERTY)
} else {
put(CHAT_TEMPLATE_KWARGS_PROPERTY, chatTemplateKwargs)
}
}
return copy(additionalProperties = properties.takeIf { it.isNotEmpty() })
}

/**
* Provider-specific chat template arguments from [additionalProperties], or null when the
* property is absent or its value is not a JSON object.
*/
public val chatTemplateKwargs: JsonObject?
get() = additionalProperties?.get(CHAT_TEMPLATE_KWARGS_PROPERTY) as? JsonObject

override fun equals(other: Any?): Boolean = when {
this === other -> true
other !is OpenAIChatParams -> false
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,109 @@
package ai.koog.prompt.executor.clients.openai

import ai.koog.http.client.ktor.KtorKoogHttpClient
import ai.koog.prompt.Prompt
import ai.koog.prompt.executor.clients.openai.models.OpenAIChatCompletionRequest
import ai.koog.prompt.executor.clients.openai.models.OpenAIChatCompletionRequestSerializer
import ai.koog.test.utils.runWithBothJsonConfigurations
import io.ktor.client.HttpClient
import io.ktor.client.engine.mock.MockEngine
import io.ktor.client.engine.mock.respond
import io.ktor.http.ContentType
import io.ktor.http.HttpHeaders
import io.ktor.http.content.TextContent
import io.ktor.http.headersOf
import kotlinx.coroutines.test.runTest
import kotlinx.serialization.json.Json
import kotlinx.serialization.json.JsonNull
import kotlinx.serialization.json.JsonPrimitive
import kotlinx.serialization.json.buildJsonArray
import kotlinx.serialization.json.buildJsonObject
import kotlinx.serialization.json.jsonObject
import kotlin.test.Test
import kotlin.test.assertEquals
import kotlin.test.assertFalse
import kotlin.test.assertNull

class OpenAIChatTemplateKwargsTest {
private val kwargs = buildJsonObject {
put("enable_thinking", JsonPrimitive(false))
put(
"provider_options",
buildJsonArray {
add(JsonPrimitive(7))
add(JsonNull)
add(buildJsonObject { put("nested", JsonPrimitive("value")) })
},
)
}

@Test
fun testCopyPreservesOtherParametersAndOriginalProperties() {
val original = OpenAIChatParams(
temperature = 0.7,
additionalProperties = mapOf("other" to JsonPrimitive(true)),
)
val configured = original.withChatTemplateKwargs(kwargs)

assertNull(original.chatTemplateKwargs)
assertEquals(0.7, configured.temperature)
assertEquals(JsonPrimitive(true), configured.additionalProperties?.get("other"))
assertEquals(kwargs, configured.copy().chatTemplateKwargs)
assertEquals(original, configured.withChatTemplateKwargs(null))
assertNull(OpenAIChatParams().withChatTemplateKwargs(kwargs).withChatTemplateKwargs(null).additionalProperties)
assertEquals(buildJsonObject {}, configured.withChatTemplateKwargs(buildJsonObject {}).chatTemplateKwargs)
}

@Test
fun testExistingAdditionalPropertyIsRecognisedAndReplaced() {
val original = OpenAIChatParams(additionalProperties = mapOf("chat_template_kwargs" to kwargs))
assertEquals(kwargs, original.chatTemplateKwargs)
assertNull(OpenAIChatParams(additionalProperties = mapOf("chat_template_kwargs" to JsonPrimitive(true))).chatTemplateKwargs)
val replacement = buildJsonObject { put("enable_thinking", JsonPrimitive(true)) }
assertEquals(replacement, original.withChatTemplateKwargs(replacement).chatTemplateKwargs)
}

@Test
fun testRequestRoundTripPreservesArbitraryJson() =
runWithBothJsonConfigurations("chat template arguments") { json ->
val request = OpenAIChatCompletionRequest(
model = "provider/model",
messages = emptyList(),
additionalProperties = OpenAIChatParams().withChatTemplateKwargs(kwargs).additionalProperties,
)
val encoded = json.encodeToString(OpenAIChatCompletionRequestSerializer, request)
val body = json.parseToJsonElement(encoded).jsonObject
assertEquals(kwargs, body["chat_template_kwargs"])
assertFalse("chatTemplateKwargs" in body)
val decoded = json.decodeFromString(OpenAIChatCompletionRequestSerializer, encoded)
assertEquals(kwargs, decoded.additionalProperties?.get("chat_template_kwargs"))
}

@Test
fun testClientSendsArgumentsAndOmitsThemAfterRemoval() = runTest {
val bodies = mutableListOf<String>()
val engine = MockEngine { request ->
bodies.add((request.body as TextContent).text)
respond(
content = """{"id":"test","object":"chat.completion","created":0,"model":"gpt-4o","choices":[{"index":0,"message":{"role":"assistant","content":"Hello"},"finish_reason":"stop"}]}""",
headers = headersOf(HttpHeaders.ContentType, ContentType.Application.Json.toString()),
)
}
val client = OpenAILLMClient(
apiKey = "test-key",
httpClientFactory = KtorKoogHttpClient.Factory(HttpClient(engine)),
)
try {
val configured = OpenAIChatParams().withChatTemplateKwargs(kwargs)
for (params in listOf(configured, configured.withChatTemplateKwargs(null), OpenAIChatParams())) {
client.execute(Prompt.build("kwargs", params = params) { user("Hello") }, OpenAIModels.Chat.GPT4o)
}
assertEquals(kwargs, Json.parseToJsonElement(bodies[0]).jsonObject["chat_template_kwargs"])
for (body in bodies.drop(1)) {
assertFalse("chat_template_kwargs" in Json.parseToJsonElement(body).jsonObject)
}
} finally {
client.close()
}
}
}