Repository navigation
Expand file tree
/
Copy pathChatCompletionAccumulator.kt
More file actions
425 lines (373 loc) · 19.7 KB
/
Copy pathChatCompletionAccumulator.kt
File metadata and controls
425 lines (373 loc) · 19.7 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
package com.openai.helpers
import com.openai.core.JsonNull
import com.openai.core.JsonValue
import com.openai.errors.OpenAIInvalidDataException
import com.openai.models.chat.completions.*
import java.util.Optional
import kotlin.jvm.optionals.getOrNull
/**
* An accumulator that constructs a [ChatCompletion] from a sequence of streamed chunks. Pass all
* chunks to [accumulate] and then call [chatCompletion] to get the final accumulated chat
* completion. The final [ChatCompletion] will be similar to what would have been received had the
* non-streaming API been used.
*
* A [ChatCompletionAccumulator] may only be used to accumulate _one_ chat completion. To accumulate
* another chat completion, create another instance of [ChatCompletionAccumulator].
*/
class ChatCompletionAccumulator private constructor() {
/**
* The final [ChatCompletion] accumulated from the chunks. This is created when the final chunk
* carrying the `finishReason` is accumulated. It may be re-created if an optional `usage` chunk
* is accumulated after the presumed final chunk.
*/
private var chatCompletion: ChatCompletion? = null
/**
* The builder for the [chatCompletion]. This is updated until the final chunk is accumulated.
* It may be used to build the [chatCompletion] again if a final usage chunk is accumulated.
*/
private var chatCompletionBuilder: ChatCompletion.Builder? = null
/**
* The builders for the [ChatCompletion.Choice] elements that are updated for the choices in
* each chunk. When the last chunk is accumulated, the choices are finally built and added to
* the [chatCompletionBuilder]. The keys are the `index` values from each chunk delta.
*/
private val choiceBuilders = mutableMapOf<Long, ChatCompletion.Choice.Builder>()
/**
* The builders for the [ChatCompletionMessage] elements that are added to their respective
* choices when all chunks have been accumulated. The keys are the `index` values from each
* chunk delta.
*/
private val messageBuilders = mutableMapOf<Long, ChatCompletionMessage.Builder>()
/**
* The accumulated content for each message. The keys correspond to the indexes in
* [messageBuilders].
*/
private val messageContents = mutableMapOf<Long, StringBuilder>()
/**
* The accumulated refusal for each message. The keys correspond to the indexes in
* [messageBuilders].
*/
private val messageRefusals = mutableMapOf<Long, StringBuilder>()
private val messageAudio = mutableMapOf<Long, AudioState>()
private class AudioState {
var id: String? = null
var expiresAt: Long? = null
var data: StringBuilder? = null
var transcript: StringBuilder? = null
val additionalProperties = mutableMapOf<String, JsonValue>()
fun isComplete() = id != null && expiresAt != null && data != null && transcript != null
fun build(): ChatCompletionAudio {
if (!isComplete()) throw OpenAIInvalidDataException("Incomplete streamed audio.")
return ChatCompletionAudio.builder()
.id(id!!)
.data(data.toString())
.transcript(transcript.toString())
.expiresAt(expiresAt!!)
.putAllAdditionalProperties(additionalProperties)
.build()
}
}
/**
* The builders for the [ChatCompletion.Choice.Logprobs] of each choice. These are only
* accumulated if the option to report log probabilities was enabled in the request. When the
* last chunk is accumulated, these are added to their respective [chatCompletionBuilder]. The
* keys are the `index` values from each chunk delta.
*/
private val logprobsBuilders = mutableMapOf<Long, ChatCompletion.Choice.Logprobs.Builder>()
/**
* The accumulated tool call builders for each message. The "outer" keys correspond to the
* indexes in [messageBuilders] (the choice index). The "inner" keys correspond to the position
* of each tool call in the message's list of tool calls (the tool call index).
*/
private val toolCallBuilders =
mutableMapOf<Long, MutableMap<Long, ChatCompletionMessageFunctionToolCall.Builder>>()
/**
* The accumulated tool call function builders for the tool call builders of each message. The
* entries correspond to those in [toolCallBuilders].
*/
private val toolCallFunctionBuilders =
mutableMapOf<
Long,
MutableMap<Long, ChatCompletionMessageFunctionToolCall.Function.Builder>,
>()
/**
* The accumulated tool call function arguments that will be set on the function builders when
* completed. The entries correspond to those in [toolCallFunctionBuilders].
*/
private val toolCallFunctionArgs = mutableMapOf<Long, MutableMap<Long, StringBuilder>>()
/**
* The finished status of each of the `n` completions. When a chunk with a `finishReason` is
* encountered, its index is recorded against a `true` value. When a `true` has been recorded
* for all indexes, the chat completion is finished and no further chunks (except perhaps a
* usage chunk) are expected to be accumulated. The keys are the `index` values from each chunk
* delta.
*/
private val isFinished = mutableMapOf<Long, Boolean>()
companion object {
@JvmStatic fun create() = ChatCompletionAccumulator()
@JvmSynthetic
internal fun convertFunctionCall(
chunkFunctionCall: ChatCompletionChunk.Choice.Delta.FunctionCall
) =
ChatCompletionMessage.FunctionCall.builder()
.name(chunkFunctionCall._name())
.arguments(chunkFunctionCall._arguments())
.additionalProperties(chunkFunctionCall._additionalProperties())
.build()
}
/**
* Gets the final accumulated chat completion. Until the last chunk has been accumulated, a
* [ChatCompletion] will not be available. Wait until all chunks have been handled by
* [accumulate] before calling this method. See that method for more details on how the last
* chunk is detected.
*
* @throws IllegalStateException If called before the last chunk has been accumulated.
*/
fun chatCompletion(): ChatCompletion =
checkNotNull(chatCompletion) { "Final chat completion chunk(s) not yet received." }
/**
* Gets the final accumulated chat completion with support for structured outputs. Until the
* last chunk has been accumulated, a [StructuredChatCompletion] will not be available. Wait
* until all chunks have been handled by [accumulate] before calling this method. See that
* method for more details on how the last chunk is detected. See the SDK documentation on
* _Structured Outputs_ for more details and example code.
*
* @param responseType The Java class from which the JSON schema in the request was derived. The
* output JSON conforming to that schema can be converted automatically back to an instance of
* that Java class by the [StructuredChatCompletion].
* @throws IllegalStateException If called before the last chunk has been accumulated.
* @throws OpenAIInvalidDataException If the JSON data cannot be parsed to an instance of the
* [responseType] class.
*/
fun <T : Any> chatCompletion(responseType: Class<T>) =
StructuredChatCompletion(responseType, chatCompletion())
/**
* Accumulates a streamed chunk and uses it to construct a [ChatCompletion]. When all chunks
* have been accumulated, the chat completion can be retrieved by calling [chatCompletion].
*
* The last chunk provides a `finishReason`, or an expiry-only audio update completing the
* accumulated audio. A finish reason received before the complete audio is retained. At that
* point, the [ChatCompletion] is created and available. However, if the request was configured
* to include the usage details, one additional chunk may be accumulated which provides those
* usage details and the [ChatCompletion] will be recreated to reflect them. After that, no more
* chunks of any kind may be accumulated.
*
* @return The given [chunk] for convenience, such as when chaining method calls.
* @throws IllegalStateException If [accumulate] is called again after the last chunk has been
* accumulated. A [ChatCompletionAccumulator] can only be used to accumulate a single
* [ChatCompletion].
* @throws OpenAIInvalidDataException If the [chunk] contains invalid or unexpected data.
*/
fun accumulate(chunk: ChatCompletionChunk): ChatCompletionChunk {
val localChatCompletion = chatCompletion
check(
localChatCompletion == null ||
(chunk.usage().getOrNull() != null &&
localChatCompletion.usage().getOrNull() == null)
) {
"Already accumulated the final chunk(s)."
}
val chatCompletionBuilder = ensureChatCompletionBuilder()
chunk.usage().getOrNull()?.let {
chatCompletionBuilder.usage(it)
chatCompletion = chatCompletionBuilder.build()
// This final `usage` chunk will have no `choices`, so stop here.
return@accumulate chunk
}
// Each chunk in the stream will repeat these values. They could be copied just once, but it
// is simpler copy them each time they are encountered.
chatCompletionBuilder
// These are the mandatory fields that must be set to something.
.id(chunk.id())
.created(chunk.created())
.model(chunk.model())
// The other fields are optional and `object` has an appropriate default value.
.systemFingerprint(chunk._systemFingerprint())
.apply {
chunk.serviceTier().ifPresent {
serviceTier(ChatCompletion.ServiceTier.of(it.asString()))
}
}
.putAllAdditionalProperties(chunk._additionalProperties())
// Each chunk contains a `choices` list (the size depends on `n` in the request). Each
// chunked `choice` has an `index`, that identifies the corresponding `Choice` within the
// accumulated list of choices; and a `Delta`, that contains additional content to be
// accumulated in that `Choice`, mostly within its `message`.
chunk.choices().forEach { choice ->
val index = choice.index()
val choiceBuilder =
choiceBuilders.getOrPut(index) { ChatCompletion.Choice.builder().index(index) }
accumulateLogprobs(index, choice.logprobs())
accumulateMessage(index, choice.delta())
choiceBuilder.additionalProperties(choice._additionalProperties())
if (choice.finishReason().isPresent) {
choiceBuilder.finishReason(
ChatCompletion.Choice.FinishReason.of(choice.finishReason().get().asString())
)
// A non-null `finishReason` signals that this is the final `Chunk` for its index.
// For each of the `n` messages completed, there will be one chunk carrying only a
// `finishReason` (its `Choice` will have no `content`). These `finalReason` chunks
// will all arrive after the non-empty chunks have been received for all `n`
// choices. Therefore, the value of `n` can be inferred from the number of distinct
// indexes observed thus far when building choices. When all `n` `finishReason`
// chunks have been received, presumptively build the `ChatCompletion`. (One further
// usage chunk _might_ be notified.)
isFinished[index] = true
} else if (isTerminalAudioUpdate(choice) && messageAudio[index]?.isComplete() == true) {
if (isFinished[index] != true) {
choiceBuilder.finishReason(ChatCompletion.Choice.FinishReason.STOP)
}
isFinished[index] = true
} else if (isFinished[index] != true) {
choiceBuilder.finishReason(JsonNull.of())
}
}
if (
chunk.choices().any { it.finishReason().isPresent || isTerminalAudioUpdate(it) } &&
choiceBuilders.keys.all { isFinished[it] == true } &&
messageAudio.values.all { it.isComplete() }
) {
chatCompletion = chatCompletionBuilder.choices(buildChoices()).build()
// Release mutable storage only after the whole chunk and final build succeed.
toolCallFunctionArgs.clear()
messageContents.clear()
messageRefusals.clear()
messageAudio.clear()
}
return chunk
}
private fun ensureChatCompletionBuilder(): ChatCompletion.Builder =
chatCompletionBuilder ?: ChatCompletion.builder().also { chatCompletionBuilder = it }
@JvmSynthetic
internal fun accumulateLogprobs(
index: Long,
logprobs: Optional<ChatCompletionChunk.Choice.Logprobs>,
) {
val chunkLogprobs = logprobs.getOrNull() ?: return
val logprobsBuilder =
logprobsBuilders.getOrPut(index) {
ChatCompletion.Choice.Logprobs.builder()
// These must be set explicitly to avoid an error when `build()` is called in
// the event that no `content` or `refusal` is added later.
.content(mutableListOf())
.refusal(mutableListOf())
}
chunkLogprobs.content().ifPresent { chunkContentList ->
chunkContentList.forEach(logprobsBuilder::addContent)
}
chunkLogprobs.refusal().ifPresent { chunkRefusalsList ->
chunkRefusalsList.forEach(logprobsBuilder::addRefusal)
}
logprobsBuilder.putAllAdditionalProperties(chunkLogprobs._additionalProperties())
}
@JvmSynthetic
internal fun accumulateMessage(index: Long, delta: ChatCompletionChunk.Choice.Delta) {
val messageBuilder = messageBuilders.getOrPut(index) { ChatCompletionMessage.builder() }
delta.audio().ifPresent { audio ->
val state = messageAudio.getOrPut(index) { AudioState() }
audio.id().ifPresent { state.id = it }
audio.expiresAt().ifPresent { state.expiresAt = it }
audio.data().ifPresent {
val buffer = state.data ?: StringBuilder().also { state.data = it }
buffer.append(it)
}
audio.transcript().ifPresent {
val buffer = state.transcript ?: StringBuilder().also { state.transcript = it }
buffer.append(it)
}
state.additionalProperties.putAll(audio._additionalProperties())
}
delta.content().ifPresent { messageContents.getOrPut(index) { StringBuilder() }.append(it) }
delta.refusal().ifPresent { messageRefusals.getOrPut(index) { StringBuilder() }.append(it) }
// The `role` defaults to "assistant", so if no other `role` is set on the delta, there is
// no need to set it explicitly to anything else.
delta.role().ifPresent { messageBuilder.role(JsonValue.from(it.asString())) }
delta.functionCall().ifPresent { messageBuilder.functionCall(convertFunctionCall(it)) }
delta.toolCalls().ifPresent {
it.map { deltaToolCall ->
// The first chunk delta will carry the tool call ID and the function name. Later
// deltas will carry only fragments of the function arguments, but the tool call
// index will identify the function to which those argument fragments belong.
val messageToolCallBuilders = toolCallBuilders.getOrPut(index) { mutableMapOf() }
messageToolCallBuilders.getOrPut(deltaToolCall.index()) {
ChatCompletionMessageFunctionToolCall.builder()
.id(deltaToolCall._id())
.additionalProperties(deltaToolCall._additionalProperties())
// Must wait until the `function` is accumulated and built before adding it to
// the tool call later when `buildChoices` is called.
}
val messageToolCallFunctionBuilders =
toolCallFunctionBuilders.getOrPut(index) { mutableMapOf() }
messageToolCallFunctionBuilders.getOrPut(deltaToolCall.index()) {
ChatCompletionMessageFunctionToolCall.Function.builder()
.name(ensureFunction(deltaToolCall.function())._name())
.additionalProperties(deltaToolCall._additionalProperties())
}
val messageToolCallFunctionArgs =
toolCallFunctionArgs.getOrPut(index) { mutableMapOf() }
val fragment =
ensureFunction(deltaToolCall.function()).arguments().getOrNull() ?: ""
messageToolCallFunctionArgs
.getOrPut(deltaToolCall.index()) { StringBuilder() }
.append(fragment)
}
}
messageBuilder.putAllAdditionalProperties(delta._additionalProperties())
}
private fun ensureFunction(
function: Optional<ChatCompletionChunk.Choice.Delta.ToolCall.Function>
): ChatCompletionChunk.Choice.Delta.ToolCall.Function =
function.orElseThrow { OpenAIInvalidDataException("Tool call chunk missing function.") }
private fun isTerminalAudioUpdate(choice: ChatCompletionChunk.Choice): Boolean =
choice.delta().audio().getOrNull()?.let {
it.expiresAt().isPresent &&
!it.id().isPresent &&
!it.data().isPresent &&
!it.transcript().isPresent
} ?: false
private fun buildChoices() =
choiceBuilders.entries
.sortedBy { it.key }
.map {
it.value
.message(buildMessage(it.key))
.logprobs(logprobsBuilders[it.key]?.build())
.build()
}
private fun buildMessage(index: Long) =
messageBuilders
.getOrElse(index) {
throw OpenAIInvalidDataException("Missing message for index $index.")
}
.content(messageContents[index]?.toString())
.refusal(messageRefusals[index]?.toString())
.toolCalls(buildToolCalls(index))
.apply { messageAudio[index]?.let { audio(it.build()) } }
.build()
private fun buildToolCalls(index: Long): List<ChatCompletionMessageToolCall> =
// It is OK for a message not to have any tool calls; most will not and an empty list will
// be returned. An entry (if it exists) will be a collection of tool call builders and each
// has a function that needs to be set.
toolCallBuilders[index]
?.entries
?.sortedBy { it.key }
?.map { (toolCallIndex, toolCallBuilder) ->
ChatCompletionMessageToolCall.ofFunction(
toolCallBuilder.function(buildFunction(index, toolCallIndex)).build()
)
} ?: listOf()
private fun buildFunction(index: Long, toolCallIndex: Long) =
// Every tool call is expected to have a function with arguments.
toolCallFunctionBuilders[index]
?.get(toolCallIndex)
?.arguments(
toolCallFunctionArgs[index]?.get(toolCallIndex)?.toString()
?: throw OpenAIInvalidDataException(
"Missing function arguments for index $index.$toolCallIndex."
)
)
?.build()
?: throw OpenAIInvalidDataException(
"Missing function builder for index $index.$toolCallIndex."
)
}