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 @@ -185,6 +185,9 @@ public class SpringAiLLMClient(
* [StreamFrame.ToolCallDelta] with a corresponding [StreamFrame.ToolCallComplete] and
* emits [StreamFrame.TextComplete] / [StreamFrame.ReasoningComplete] boundaries.
*
* The terminal [StreamFrame.End] carries the finish reason and the token usage collected over
* the whole stream, since a provider may report them in separate chunks.
*
* All blocking I/O runs on the configured [dispatcher] (default [Dispatchers.IO]).
*/
override fun executeStreaming(
Expand All @@ -201,10 +204,10 @@ public class SpringAiLLMClient(
} catch (e: Exception) {
throw LLMClientException(clientName, "ChatModel.stream() failed: ${e.message}", e)
}
var lastChatResponse: ChatResponse? = null
var finishReason: String? = null
var metaInfo: ResponseMetaInfo? = null
try {
flux.asFlow().collect { chatResponse ->
lastChatResponse = chatResponse
for ((generationIndex, generation) in chatResponse.results.withIndex()) {
val assistantMessage = generation.output
val text = assistantMessage.text
Expand All @@ -222,6 +225,23 @@ public class SpringAiLLMClient(
toolCallAssembler.accept(assistantMessage.toolCalls, generationIndex, this)
}
}
// Spring AI can report the finish reason and the token usage in different chunks:
// with stream-usage the last chunk carries the usage and no generations at all, so
// keep each as it arrives (see #2109). A chunk that has neither reports an empty
// finish reason and zero counts rather than nulls, so neither may overwrite a
// reported value.
chatResponse.results.firstOrNull()?.metadata?.finishReason
?.takeIf { it.isNotEmpty() }
?.let { finishReason = it }
val usage = chatResponse.metadata.usage
if (metaInfo == null || usage.totalTokens > 0) {
metaInfo = ResponseMetaInfo.create(
clock = clock,
totalTokensCount = usage.totalTokens,
inputTokensCount = usage.promptTokens,
outputTokensCount = usage.completionTokens
)
}
}
toolCallAssembler.flush(this)
} catch (e: CancellationException) {
Expand All @@ -232,18 +252,6 @@ public class SpringAiLLMClient(
throw LLMClientException(clientName, "ChatModel.stream() failed during collection: ${e.message}", e)
} finally {
// Always emit End frame so downstream consumers are not left hanging
val finishReason = lastChatResponse?.results?.firstOrNull()?.metadata?.finishReason
val usage = lastChatResponse?.metadata?.usage
val metaInfo = if (usage != null) {
ResponseMetaInfo.create(
clock = clock,
totalTokensCount = usage.totalTokens,
inputTokensCount = usage.promptTokens,
outputTokensCount = usage.completionTokens
)
} else {
null
}
emitEnd(finishReason = finishReason, metaInfo = metaInfo)
}
}.flowOn(dispatcher)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ import org.junit.jupiter.api.Assertions.assertTrue
import org.junit.jupiter.api.Test
import org.junit.jupiter.api.assertThrows
import org.springframework.ai.chat.messages.AssistantMessage
import org.springframework.ai.chat.metadata.ChatGenerationMetadata
import org.springframework.ai.chat.metadata.ChatResponseMetadata
import org.springframework.ai.chat.metadata.Usage
import org.springframework.ai.chat.model.ChatModel
Expand All @@ -45,6 +46,12 @@ class SpringAiLLMClientTest {

private fun requestMeta() = RequestMetaInfo.create(KoogClock.System)

private fun stubUsage(promptTokens: Int, completionTokens: Int) = object : Usage {
override fun getPromptTokens(): Int = promptTokens
override fun getCompletionTokens(): Int = completionTokens
override fun getNativeUsage(): Any = emptyMap<String, Any>()
}

// ---- llmProvider ----

@Test
Expand Down Expand Up @@ -153,12 +160,9 @@ class SpringAiLLMClientTest {

@Test
fun `execute maps usage metadata to response meta info`() = runBlocking {
val usage = object : Usage {
override fun getPromptTokens(): Int = 5
override fun getCompletionTokens(): Int = 15
override fun getNativeUsage(): Any = emptyMap<String, Any>()
}
val metadata = ChatResponseMetadata.builder().usage(usage).build()
val metadata = ChatResponseMetadata.builder()
.usage(stubUsage(promptTokens = 5, completionTokens = 15))
.build()
val client = SpringAiLLMClient.builder().chatModel(object : ChatModel {
override fun call(prompt: SpringPrompt) =
ChatResponse(listOf(Generation(AssistantMessage("Done"))), metadata)
Expand Down Expand Up @@ -340,6 +344,86 @@ class SpringAiLLMClientTest {
assertEquals("Hi", textFrames[0].text)
}

@Test
fun testExecuteStreamingKeepsFinishReasonWhenLastChunkCarriesOnlyUsage() = runBlocking {
val finishMetadata = ChatGenerationMetadata.builder().finishReason("STOP").build()
val usageMetadata = ChatResponseMetadata.builder()
.usage(stubUsage(promptTokens = 5, completionTokens = 10))
.build()
val client = SpringAiLLMClient.builder().chatModel(object : ChatModel {
override fun call(prompt: SpringPrompt) = throw UnsupportedOperationException()
override fun stream(prompt: SpringPrompt) = Flux.just(
ChatResponse(listOf(Generation(AssistantMessage("Hello"), finishMetadata))),
// With stream-usage enabled the terminal chunk reports the usage and no generations
ChatResponse(emptyList(), usageMetadata)
)
}).build()
val prompt = createPrompt(Message.User("Hi", requestMeta()))
val frames = client.executeStreaming(prompt, testModel, emptyList()).toList()

val end = frames.filterIsInstance<StreamFrame.End>().single()
assertEquals("STOP", end.finishReason)
assertEquals(15, end.metaInfo.totalTokensCount)
}

@Test
fun testExecuteStreamingKeepsTokenUsageWhenALaterChunkReportsNone() = runBlocking {
val usageMetadata = ChatResponseMetadata.builder()
.usage(stubUsage(promptTokens = 5, completionTokens = 10))
.build()
val finishMetadata = ChatGenerationMetadata.builder().finishReason("STOP").build()
val client = SpringAiLLMClient.builder().chatModel(object : ChatModel {
override fun call(prompt: SpringPrompt) = throw UnsupportedOperationException()
override fun stream(prompt: SpringPrompt) = Flux.just(
ChatResponse(listOf(Generation(AssistantMessage("Hello"))), usageMetadata),
ChatResponse(listOf(Generation(AssistantMessage(" world"), finishMetadata)))
)
}).build()
val prompt = createPrompt(Message.User("Hi", requestMeta()))
val frames = client.executeStreaming(prompt, testModel, emptyList()).toList()

val end = frames.filterIsInstance<StreamFrame.End>().single()
assertEquals("STOP", end.finishReason)
assertEquals(5, end.metaInfo.inputTokensCount)
assertEquals(10, end.metaInfo.outputTokensCount)
assertEquals(15, end.metaInfo.totalTokensCount)
}

@Test
fun testExecuteStreamingReportsZeroTokenCountsWhenNoChunkReportsUsage() = runBlocking {
val client = SpringAiLLMClient.builder().chatModel(object : ChatModel {
override fun call(prompt: SpringPrompt) = throw UnsupportedOperationException()
override fun stream(prompt: SpringPrompt) = Flux.just(
ChatResponse(listOf(Generation(AssistantMessage("Hello"))))
)
}).build()
val prompt = createPrompt(Message.User("Hi", requestMeta()))
val frames = client.executeStreaming(prompt, testModel, emptyList()).toList()

// A response without usage reports zero counts, not unknown ones, as it did before
val end = frames.filterIsInstance<StreamFrame.End>().single()
assertEquals(0, end.metaInfo.totalTokensCount)
}

@Test
fun testExecuteStreamingIgnoresAnEmptyFinishReasonFromALaterChunk() = runBlocking {
val finishMetadata = ChatGenerationMetadata.builder().finishReason("STOP").build()
// Mistral AI and DeepSeek pad every streamed chunk with an empty finish reason
val paddedMetadata = ChatGenerationMetadata.builder().finishReason("").build()
val client = SpringAiLLMClient.builder().chatModel(object : ChatModel {
override fun call(prompt: SpringPrompt) = throw UnsupportedOperationException()
override fun stream(prompt: SpringPrompt) = Flux.just(
ChatResponse(listOf(Generation(AssistantMessage("Hello"), finishMetadata))),
ChatResponse(listOf(Generation(AssistantMessage(" world"), paddedMetadata)))
)
}).build()
val prompt = createPrompt(Message.User("Hi", requestMeta()))
val frames = client.executeStreaming(prompt, testModel, emptyList()).toList()

val end = frames.filterIsInstance<StreamFrame.End>().single()
assertEquals("STOP", end.finishReason)
}

// ---- moderate ----

@Test
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -184,6 +184,9 @@ public class SpringAiLLMClient(
* [StreamFrame.ToolCallDelta] with a corresponding [StreamFrame.ToolCallComplete] and
* emits [StreamFrame.TextComplete] / [StreamFrame.ReasoningComplete] boundaries.
*
* The terminal [StreamFrame.End] carries the finish reason and the token usage collected over
* the whole stream, since a provider may report them in separate chunks.
*
* All blocking I/O runs on the configured [dispatcher] (default [Dispatchers.IO]).
*/
override fun executeStreaming(
Expand All @@ -200,10 +203,10 @@ public class SpringAiLLMClient(
} catch (e: Exception) {
throw LLMClientException(clientName, "ChatModel.stream() failed: ${e.message}", e)
}
var lastChatResponse: ChatResponse? = null
var finishReason: String? = null
var metaInfo: ResponseMetaInfo? = null
try {
flux.asFlow().collect { chatResponse ->
lastChatResponse = chatResponse
for ((generationIndex, generation) in chatResponse.results.withIndex()) {
val assistantMessage = generation.output
val text = assistantMessage.text
Expand All @@ -221,6 +224,23 @@ public class SpringAiLLMClient(
toolCallAssembler.accept(assistantMessage.toolCalls, generationIndex, this)
}
}
// Spring AI can report the finish reason and the token usage in different chunks:
// with stream-usage the last chunk carries the usage and no generations at all, so
// keep each as it arrives (see #2109). A chunk that has neither reports an empty
// finish reason and zero counts rather than nulls, so neither may overwrite a
// reported value.
chatResponse.results.firstOrNull()?.metadata?.finishReason
?.takeIf { it.isNotEmpty() }
?.let { finishReason = it }
val usage = chatResponse.metadata.usage
if (metaInfo == null || usage.totalTokens > 0) {
metaInfo = ResponseMetaInfo.create(
clock = clock,
totalTokensCount = usage.totalTokens,
inputTokensCount = usage.promptTokens,
outputTokensCount = usage.completionTokens
)
}
}
toolCallAssembler.flush(this)
} catch (e: CancellationException) {
Expand All @@ -231,18 +251,6 @@ public class SpringAiLLMClient(
throw LLMClientException(clientName, "ChatModel.stream() failed during collection: ${e.message}", e)
} finally {
// Always emit End frame so downstream consumers are not left hanging
val finishReason = lastChatResponse?.results?.firstOrNull()?.metadata?.finishReason
val usage = lastChatResponse?.metadata?.usage
val metaInfo = if (usage != null) {
ResponseMetaInfo.create(
clock = clock,
totalTokensCount = usage.totalTokens,
inputTokensCount = usage.promptTokens,
outputTokensCount = usage.completionTokens
)
} else {
null
}
emitEnd(finishReason = finishReason, metaInfo = metaInfo)
}
}.flowOn(dispatcher)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@ import org.junit.jupiter.api.Assertions.assertTrue
import org.junit.jupiter.api.Test
import org.junit.jupiter.api.assertThrows
import org.springframework.ai.chat.messages.AssistantMessage
import org.springframework.ai.chat.metadata.ChatGenerationMetadata
import org.springframework.ai.chat.metadata.ChatResponseMetadata
import org.springframework.ai.chat.metadata.Usage
import org.springframework.ai.chat.model.ChatModel
Expand All @@ -44,6 +45,12 @@ class SpringAiLLMClientTest {

private fun requestMeta() = RequestMetaInfo.create(KoogClock.System)

private fun stubUsage(promptTokens: Int, completionTokens: Int) = object : Usage {
override fun getPromptTokens(): Int = promptTokens
override fun getCompletionTokens(): Int = completionTokens
override fun getNativeUsage(): Any = emptyMap<String, Any>()
}

// ---- llmProvider ----

@Test
Expand Down Expand Up @@ -152,12 +159,9 @@ class SpringAiLLMClientTest {

@Test
fun `execute maps usage metadata to response meta info`() = runBlocking {
val usage = object : Usage {
override fun getPromptTokens(): Int = 5
override fun getCompletionTokens(): Int = 15
override fun getNativeUsage(): Any = emptyMap<String, Any>()
}
val metadata = ChatResponseMetadata.builder().usage(usage).build()
val metadata = ChatResponseMetadata.builder()
.usage(stubUsage(promptTokens = 5, completionTokens = 15))
.build()
val client = SpringAiLLMClient.builder().chatModel(object : ChatModel {
override fun call(prompt: SpringPrompt) =
ChatResponse(listOf(Generation(AssistantMessage("Done"))), metadata)
Expand Down Expand Up @@ -336,6 +340,86 @@ class SpringAiLLMClientTest {
assertEquals("Hi", textFrames[0].text)
}

@Test
fun testExecuteStreamingKeepsFinishReasonWhenLastChunkCarriesOnlyUsage() = runBlocking {
val finishMetadata = ChatGenerationMetadata.builder().finishReason("STOP").build()
val usageMetadata = ChatResponseMetadata.builder()
.usage(stubUsage(promptTokens = 5, completionTokens = 10))
.build()
val client = SpringAiLLMClient.builder().chatModel(object : ChatModel {
override fun call(prompt: SpringPrompt) = throw UnsupportedOperationException()
override fun stream(prompt: SpringPrompt) = Flux.just(
ChatResponse(listOf(Generation(AssistantMessage("Hello"), finishMetadata))),
// With stream-usage enabled the terminal chunk reports the usage and no generations
ChatResponse(emptyList(), usageMetadata)
)
}).build()
val prompt = createPrompt(Message.User("Hi", requestMeta()))
val frames = client.executeStreaming(prompt, testModel, emptyList()).toList()

val end = frames.filterIsInstance<StreamFrame.End>().single()
assertEquals("STOP", end.finishReason)
assertEquals(15, end.metaInfo.totalTokensCount)
}

@Test
fun testExecuteStreamingKeepsTokenUsageWhenALaterChunkReportsNone() = runBlocking {
val usageMetadata = ChatResponseMetadata.builder()
.usage(stubUsage(promptTokens = 5, completionTokens = 10))
.build()
val finishMetadata = ChatGenerationMetadata.builder().finishReason("STOP").build()
val client = SpringAiLLMClient.builder().chatModel(object : ChatModel {
override fun call(prompt: SpringPrompt) = throw UnsupportedOperationException()
override fun stream(prompt: SpringPrompt) = Flux.just(
ChatResponse(listOf(Generation(AssistantMessage("Hello"))), usageMetadata),
ChatResponse(listOf(Generation(AssistantMessage(" world"), finishMetadata)))
)
}).build()
val prompt = createPrompt(Message.User("Hi", requestMeta()))
val frames = client.executeStreaming(prompt, testModel, emptyList()).toList()

val end = frames.filterIsInstance<StreamFrame.End>().single()
assertEquals("STOP", end.finishReason)
assertEquals(5, end.metaInfo.inputTokensCount)
assertEquals(10, end.metaInfo.outputTokensCount)
assertEquals(15, end.metaInfo.totalTokensCount)
}

@Test
fun testExecuteStreamingReportsZeroTokenCountsWhenNoChunkReportsUsage() = runBlocking {
val client = SpringAiLLMClient.builder().chatModel(object : ChatModel {
override fun call(prompt: SpringPrompt) = throw UnsupportedOperationException()
override fun stream(prompt: SpringPrompt) = Flux.just(
ChatResponse(listOf(Generation(AssistantMessage("Hello"))))
)
}).build()
val prompt = createPrompt(Message.User("Hi", requestMeta()))
val frames = client.executeStreaming(prompt, testModel, emptyList()).toList()

// A response without usage reports zero counts, not unknown ones, as it did before
val end = frames.filterIsInstance<StreamFrame.End>().single()
assertEquals(0, end.metaInfo.totalTokensCount)
}

@Test
fun testExecuteStreamingIgnoresAnEmptyFinishReasonFromALaterChunk() = runBlocking {
val finishMetadata = ChatGenerationMetadata.builder().finishReason("STOP").build()
// Mistral AI and DeepSeek pad every streamed chunk with an empty finish reason
val paddedMetadata = ChatGenerationMetadata.builder().finishReason("").build()
val client = SpringAiLLMClient.builder().chatModel(object : ChatModel {
override fun call(prompt: SpringPrompt) = throw UnsupportedOperationException()
override fun stream(prompt: SpringPrompt) = Flux.just(
ChatResponse(listOf(Generation(AssistantMessage("Hello"), finishMetadata))),
ChatResponse(listOf(Generation(AssistantMessage(" world"), paddedMetadata)))
)
}).build()
val prompt = createPrompt(Message.User("Hi", requestMeta()))
val frames = client.executeStreaming(prompt, testModel, emptyList()).toList()

val end = frames.filterIsInstance<StreamFrame.End>().single()
assertEquals("STOP", end.finishReason)
}

// ---- moderate ----

@Test
Expand Down