diff --git a/build.gradle.kts b/build.gradle.kts index d563aae..ce4fca1 100644 --- a/build.gradle.kts +++ b/build.gradle.kts @@ -12,10 +12,10 @@ repositories { } dependencies { - implementation("io.ktor:ktor-server-core:3.1.2") - implementation("io.ktor:ktor-server-cio:3.1.2") - implementation("io.ktor:ktor-server-content-negotiation:3.1.2") - implementation("io.ktor:ktor-serialization-kotlinx-json:3.1.2") + implementation("io.ktor:ktor-server-core:3.5.2") + implementation("io.ktor:ktor-server-cio:3.5.2") + implementation("io.ktor:ktor-server-content-negotiation:3.5.2") + implementation("io.ktor:ktor-serialization-kotlinx-json:3.5.2") implementation("org.jetbrains.kotlinx:kotlinx-serialization-json:1.8.1") implementation("ch.qos.logback:logback-classic:1.5.18") } diff --git a/src/main/kotlin/pw/binom/llmproxy/Main.kt b/src/main/kotlin/pw/binom/llmproxy/Main.kt index 11d4180..1d168dd 100644 --- a/src/main/kotlin/pw/binom/llmproxy/Main.kt +++ b/src/main/kotlin/pw/binom/llmproxy/Main.kt @@ -13,11 +13,19 @@ import io.ktor.server.response.respondOutputStream import io.ktor.server.routing.get import io.ktor.server.routing.post import io.ktor.server.routing.routing +import io.ktor.server.http.HttpRequestLifecycle import kotlinx.serialization.json.Json import kotlinx.serialization.json.JsonArray import kotlinx.serialization.json.JsonObject import kotlinx.serialization.json.JsonPrimitive import kotlinx.serialization.json.jsonArray +import kotlinx.coroutines.CancellationException +import kotlinx.coroutines.suspendCancellableCoroutine +import kotlinx.coroutines.currentCoroutineContext +import kotlinx.coroutines.ensureActive +import kotlin.coroutines.resume +import kotlin.coroutines.resumeWithException +import java.util.concurrent.CompletableFuture import kotlinx.serialization.json.jsonObject import kotlinx.serialization.json.jsonPrimitive import java.net.URI @@ -56,6 +64,9 @@ fun main() { .split(",").map { it.trim() }.filter { it.isNotEmpty() }.distinct() embeddedServer(CIO, port = port, host = "0.0.0.0") { + install(HttpRequestLifecycle) { + cancelCallOnClose = true + } proxyModule(upstream, apiKey, excluded, thinking) }.start(wait = true) } @@ -125,11 +136,45 @@ private suspend fun handleChat( } println("[llm-proxy] chat model=$model status=${resp.statusCode()} в ${System.currentTimeMillis() - start}ms stream=true") } else { - val resp = http.send(req, HttpResponse.BodyHandlers.ofByteArray()) + // Клиент ждёт non-stream ответ, но апстриму шлём stream=true: + // не-стрим генерация у llama.cpp НЕ отменяется обрывом соединения, + // а стрим — отменяется. Так отмена клиента реально рвёт генерацию. + val streamed = json.parseToJsonElement(patched).jsonObject.toMutableMap().apply { + this["stream"] = JsonPrimitive(true) + } + val req2 = HttpRequest.newBuilder() + .uri(URI.create(upstream.trimEnd('/') + "/chat/completions")) + .header("Authorization", "Bearer $apiKey") + .header("Content-Type", "application/json") + .POST(HttpRequest.BodyPublishers.ofString(JsonObject(streamed).toString())) + .build() + val future = http.sendAsync(req2, HttpResponse.BodyHandlers.ofInputStream()) + val resp = future.awaitOrCancel() + val body = resp.body() + // Читаем с проверкой отмены: при обрыве клиента ensureActive() бросит + // CancellationException, а finally закроет входной поток — это рвёт + // апстрим-соединение, и llama.cpp отменяет генерацию. + val full = try { + val sb = StringBuilder() + val reader = body.bufferedReader() + while (true) { + currentCoroutineContext().ensureActive() + val line = reader.readLine() ?: break + sb.append(line).append('\n') + } + sb.toString() + } finally { + body.close() + } val ct = resp.headers().firstValue("content-type").orElse("application/json") - call.respondBytes(resp.body(), ContentType.parse(ct), HttpStatusCode.fromValue(resp.statusCode())) + val out = if (ct.contains("text/event-stream")) rebuildFromChunks(full) else full + call.respondBytes(out.toByteArray(), ContentType.parse(ct), HttpStatusCode.fromValue(resp.statusCode())) println("[llm-proxy] chat model=$model status=${resp.statusCode()} в ${System.currentTimeMillis() - start}ms stream=false") } + } catch (e: CancellationException) { + // Клиент оборвал соединение: апстрим-запрос уже отменён через awaitOrCancel. + println("[llm-proxy] chat model=$model ОТМЕНЕНО клиентом в ${System.currentTimeMillis() - start}ms") + throw e } catch (e: Exception) { println("[llm-proxy] chat model=$model ОШИБКА: ${e.message} в ${System.currentTimeMillis() - start}ms") call.respondBytes( @@ -214,3 +259,65 @@ internal fun patchModelsCatalog(raw: String, thinking: List): String { root["data"] = JsonArray(data + copies) return JsonObject(root).toString() } + +/** + * Ожидание CompletableFuture с пробросом отмены корутины на апстрим-запрос: + * если клиент оборвал соединение (Ktor отменяет корутину), рвём и апстрим — + * upstream (llama.cpp/sglang) видит обрыв и отменяет генерацию (слот свободен). + */ +private suspend fun CompletableFuture.awaitOrCancel(): T = + suspendCancellableCoroutine { cont -> + this.whenComplete { res, err -> + if (err != null) cont.resumeWithException(err) else cont.resume(res) + } + cont.invokeOnCancellation { this.cancel(true) } + } + +/** + * Собрать полный chat.completion из SSE-чанков апстрима (для non-stream клиентов). + */ +private fun rebuildFromChunks(sse: String): String { + var content = StringBuilder() + var reasoning = StringBuilder() + var finish = "stop" + var id = "" + var model = "" + val created = System.currentTimeMillis() / 1000 + sse.lineSequence().forEach { line -> + if (!line.startsWith("data:")) return@forEach + val data = line.removePrefix("data:").trim() + if (data.isEmpty() || data == "[DONE]") return@forEach + try { + val obj = json.parseToJsonElement(data).jsonObject + if (id.isEmpty()) id = obj["id"]?.jsonPrimitive?.content ?: "" + if (model.isEmpty()) model = obj["model"]?.jsonPrimitive?.content ?: "" + val choice = obj["choices"]?.jsonArray?.firstOrNull()?.jsonObject + if (choice != null) { + choice["finish_reason"]?.jsonPrimitive?.content + ?.takeIf { it.isNotEmpty() && it != "null" }?.let { finish = it } + val delta = choice["delta"]?.jsonObject + delta?.get("content")?.jsonPrimitive?.content + ?.takeIf { it != "null" }?.let { content.append(it) } + delta?.get("reasoning_content")?.jsonPrimitive?.content + ?.takeIf { it != "null" }?.let { reasoning.append(it) } + } + } catch (_: Exception) {} + } + val msg = JsonObject(mutableMapOf( + "role" to JsonPrimitive("assistant"), + "content" to JsonPrimitive(content.toString()), + "reasoning" to JsonPrimitive(reasoning.toString()), + )) + val choice = JsonObject(mutableMapOf( + "index" to JsonPrimitive(0), + "message" to msg, + "finish_reason" to JsonPrimitive(finish), + )) + return JsonObject(mutableMapOf( + "id" to JsonPrimitive(id), + "object" to JsonPrimitive("chat.completion"), + "created" to JsonPrimitive(created), + "model" to JsonPrimitive(model), + "choices" to JsonArray(listOf(choice)), + )).toString() +}