diff --git a/infrastructure/inference/llama_cpp/src/main/kotlin/com/correx/infrastructure/inference/llama/cpp/LlamaCppInferenceProvider.kt b/infrastructure/inference/llama_cpp/src/main/kotlin/com/correx/infrastructure/inference/llama/cpp/LlamaCppInferenceProvider.kt index d4a498d3..8a29a8e2 100644 --- a/infrastructure/inference/llama_cpp/src/main/kotlin/com/correx/infrastructure/inference/llama/cpp/LlamaCppInferenceProvider.kt +++ b/infrastructure/inference/llama_cpp/src/main/kotlin/com/correx/infrastructure/inference/llama/cpp/LlamaCppInferenceProvider.kt @@ -13,6 +13,7 @@ import com.correx.core.inference.TokenUsage import com.correx.core.inference.Tokenizer import com.correx.core.inference.ToolCallFunction import com.correx.core.inference.ToolCallRequest +import com.correx.core.inference.ToolDefinition import com.correx.infrastructure.inference.commons.ModelDescriptor import io.ktor.client.HttpClient import io.ktor.client.call.body @@ -148,6 +149,9 @@ class LlamaCppInferenceProvider( val tools = request.tools.takeIf { it.isNotEmpty() } + // #291: clamp the stage completion cap to the window's real runway (see clampMaxTokens). + val effectiveMaxTokens = clampMaxTokens(request.generationConfig.maxTokens, messages, tools) + // llama.cpp rejects requests that carry BOTH a custom grammar and tools // ("Cannot use custom grammar constraints with tools"). When a stage has tools, // we drop the grammar and rely on post-hoc schema validation + retry to keep the @@ -167,7 +171,7 @@ class LlamaCppInferenceProvider( messages = messages, temperature = request.generationConfig.temperature, topP = request.generationConfig.topP, - maxTokens = request.generationConfig.maxTokens, + maxTokens = effectiveMaxTokens, stopSequences = request.generationConfig.stopSequences, seed = request.generationConfig.seed, topK = request.generationConfig.topK, @@ -222,6 +226,47 @@ class LlamaCppInferenceProvider( ) } + // Per-message chat-template framing (role tags, delimiters) the client can't see but the server + // adds; a small fixed estimate is enough headroom for a safety clamp. ponytail: constant, not a + // per-family table — bump if a template proves heavier. + private val templateTokensPerMessage = 8 + private val safetyReserveTokens = 512 + private val minCompletionTokens = 256 + + /** + * #291: returns min(requestedCap, contextSize − promptTokens − toolSchema − templateOverhead − + * reserve), floored so a nearly-full window still asks for *something* rather than a negative or + * absurd allowance. Prompt/tool tokens are counted with the model's own tokenizer (one extra + * /tokenize round-trip); a char/4 estimate is the fallback if that call fails. Live-path only — + * deterministic replay never calls [infer], so nothing here needs to be an event. + */ + private suspend fun clampMaxTokens( + requestedCap: Int, + messages: List, + tools: List?, + ): Int { + val countable = buildString { + messages.forEach { m -> + append(m.role).append('\n') + m.content?.let { append(it).append('\n') } + m.reasoningContent?.let { append(it).append('\n') } + m.toolCalls.forEach { append(it.function.name).append(it.function.arguments).append('\n') } + } + tools?.let { append(json.encodeToString(it)) } + } + val promptTokens = runCatching { tokenizer.countTokens(countable) } + .getOrElse { countable.length / 4 } + val overhead = messages.size * templateTokensPerMessage + val runway = descriptor.contextSize - promptTokens - overhead - safetyReserveTokens + if (runway < minCompletionTokens) { + log.warn( + "context runway low: ctx={} prompt~{} overhead={} reserve={} runway={} (cap was {})", + descriptor.contextSize, promptTokens, overhead, safetyReserveTokens, runway, requestedCap, + ) + } + return requestedCap.coerceAtMost(runway).coerceAtLeast(0) + } + override suspend fun healthCheck(): ProviderHealth = try { val response = httpClient.get("$baseUrl/health") if (response.status.value in 200..299) { diff --git a/infrastructure/inference/llama_cpp/src/test/kotlin/com/correx/infrastructure/inference/llama/cpp/LlamaCppMaxTokensClampTest.kt b/infrastructure/inference/llama_cpp/src/test/kotlin/com/correx/infrastructure/inference/llama/cpp/LlamaCppMaxTokensClampTest.kt new file mode 100644 index 00000000..bec880dc --- /dev/null +++ b/infrastructure/inference/llama_cpp/src/test/kotlin/com/correx/infrastructure/inference/llama/cpp/LlamaCppMaxTokensClampTest.kt @@ -0,0 +1,111 @@ +package com.correx.infrastructure.inference.llama.cpp + +import com.correx.core.context.model.ContextEntry +import com.correx.core.context.model.ContextLayer +import com.correx.core.context.model.ContextPack +import com.correx.core.context.model.EntryRole +import com.correx.core.events.types.ContextEntryId +import com.correx.core.events.types.ContextPackId +import com.correx.core.events.types.InferenceRequestId +import com.correx.core.events.types.SessionId +import com.correx.core.events.types.StageId +import com.correx.core.inference.GenerationConfig +import com.correx.core.inference.InferenceRequest +import com.correx.infrastructure.inference.commons.ModelDescriptor +import com.correx.infrastructure.inference.commons.ResidencyMode +import io.ktor.client.HttpClient +import io.ktor.client.engine.mock.MockEngine +import io.ktor.client.engine.mock.respond +import io.ktor.client.plugins.contentnegotiation.ContentNegotiation +import io.ktor.serialization.kotlinx.json.json +import io.ktor.http.ContentType +import io.ktor.http.HttpHeaders +import io.ktor.http.content.TextContent +import io.ktor.http.headersOf +import kotlinx.coroutines.runBlocking +import org.junit.jupiter.api.Test +import kotlin.test.assertEquals + +// #291: the provider must clamp the requested completion cap against the model context window so a +// nearly-full prompt never asks for an impossible allowance (the traced Gemma4 truncation). +class LlamaCppMaxTokensClampTest { + + private val promptTokenCount = 1000 + private val reserve = 512 + private val overheadPerMessage = 8 + + private fun request(cap: Int) = InferenceRequest( + requestId = InferenceRequestId("req-1"), + sessionId = SessionId("s-1"), + stageId = StageId("stage-1"), + contextPack = ContextPack( + id = ContextPackId("pack-1"), + sessionId = SessionId("s-1"), + stageId = StageId("stage-1"), + layers = mapOf( + ContextLayer.L1 to listOf( + ContextEntry( + id = ContextEntryId("e-1"), + layer = ContextLayer.L1, + content = "do the thing", + sourceType = "stagePrompt", + sourceId = "e-1", + tokenEstimate = 4, + role = EntryRole.USER, + ), + ), + ), + budgetUsed = 4, + budgetLimit = 24_576, + ), + generationConfig = GenerationConfig(temperature = 1.0, topP = 1.0, maxTokens = cap), + ) + + private fun provider(contextSize: Int, capture: (Int) -> Unit): LlamaCppInferenceProvider { + val engine = MockEngine { req -> + when { + req.url.encodedPath.endsWith("/tokenize") -> respond( + content = """{"tokens":[${(1..promptTokenCount).joinToString(",")}]}""", + headers = headersOf(HttpHeaders.ContentType, ContentType.Application.Json.toString()), + ) + else -> { + val body = (req.body as TextContent).text + val maxTokens = Regex("\"max_tokens\":(\\d+)").find(body)!!.groupValues[1].toInt() + capture(maxTokens) + val usage = + """"usage":{"prompt_tokens":$promptTokenCount,"completion_tokens":1,"total_tokens":1}""" + respond( + content = """{"id":"x","choices":[{"message":{"role":"assistant",""" + + """"content":"ok"},"finish_reason":"stop"}],$usage}""", + headers = headersOf(HttpHeaders.ContentType, ContentType.Application.Json.toString()), + ) + } + } + } + val descriptor = ModelDescriptor( + modelId = "test", + modelPath = "/dev/null", + residencyMode = ResidencyMode.EPHEMERAL, + contextSize = contextSize, + ) + val client = HttpClient(engine) { + install(ContentNegotiation) { json(kotlinx.serialization.json.Json { ignoreUnknownKeys = true }) } + } + return LlamaCppInferenceProvider(descriptor, "http://localhost:10000", client) + } + + @Test + fun `clamps completion cap to remaining context runway`(): Unit = runBlocking { + var sent = -1 + // 1 message -> overhead 8. runway = 2000 - 1000 - 8 - 512 = 480, below the 24_576 cap. + provider(contextSize = 2000) { sent = it }.infer(request(cap = 24_576)) + assertEquals(2000 - promptTokenCount - overheadPerMessage - reserve, sent) + } + + @Test + fun `keeps the stage cap when the window has ample runway`(): Unit = runBlocking { + var sent = -1 + provider(contextSize = 100_000) { sent = it }.infer(request(cap = 24_576)) + assertEquals(24_576, sent) + } +}