package com.rainy.token.data.proxy import io.ktor.http.ContentType import io.ktor.http.HttpHeaders import io.ktor.http.HttpStatusCode import io.ktor.serialization.kotlinx.json.json import io.ktor.server.application.Application import io.ktor.server.application.ApplicationCall import io.ktor.server.application.install import io.ktor.server.cio.CIO import io.ktor.server.engine.EmbeddedServer import io.ktor.server.engine.embeddedServer import io.ktor.server.plugins.contentnegotiation.ContentNegotiation import io.ktor.server.request.receiveText import io.ktor.server.response.respond import io.ktor.server.response.respondBytes 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 java.io.IOException import javax.inject.Inject import javax.inject.Singleton import kotlinx.serialization.json.contentOrNull import kotlinx.serialization.json.put import kotlinx.serialization.json.buildJsonObject import kotlinx.coroutines.flow.MutableStateFlow import kotlinx.coroutines.flow.StateFlow import kotlinx.coroutines.flow.asStateFlow /** * 基于 Ktor CIO 的本地 HTTP 反代服务。 * * M1b 能力: * - GET /health * - GET /v1/models * - POST /v1/chat/completions * 按模型前缀路由:wbcn- → WorkBuddy 国内版;traecn- → Trae CN;其余 → Sub2API 透传 * * 安全: * - 只绑定 127.0.0.1 * - config.apiKey 非空时,所有 v1 业务路由要求 Bearer Key 一致,否则 401 */ @Singleton class KtorLocalProxyServer @Inject constructor( private val sub2ApiChatProxy: Sub2ApiChatProxy, private val workBuddyChatProxy: WorkBuddyChatProxy, private val traeChatProxy: TraeChatProxy, private val providerRouter: ProviderRouter, private val accountPool: AccountPool, ) : LocalProxyServer { private val lock = Any() @Volatile private var server: EmbeddedServer<*, *>? = null private val _isRunning = MutableStateFlow(false) override val isRunning: StateFlow = _isRunning.asStateFlow() override fun start(config: ProxyServerConfig): Result { synchronized(lock) { if (_isRunning.value) return Result.success(Unit) return try { val engine = embeddedServer(CIO, host = "127.0.0.1", port = config.port) { proxyModule(config.apiKey, sub2ApiChatProxy, workBuddyChatProxy, traeChatProxy, providerRouter, accountPool) } engine.start(wait = false) server = engine _isRunning.value = true Result.success(Unit) } catch (e: Throwable) { server = null _isRunning.value = false Result.failure(e) } } } override fun stop() { synchronized(lock) { runCatching { server?.stop(gracePeriodMillis = 500, timeoutMillis = 2000) } server = null _isRunning.value = false } } private fun Application.proxyModule( apiKey: String?, sub2Api: Sub2ApiChatProxy, workBuddy: WorkBuddyChatProxy, trae: TraeChatProxy, router: ProviderRouter, pool: AccountPool, ) { install(ContentNegotiation) { json() } routing { get("/health") { call.respond(mapOf("status" to "ok")) } get("/v1/models") { if (!authorized(call, apiKey)) { call.respond(HttpStatusCode.Unauthorized, errorBody("未授权")) return@get } val result = sub2Api.forwardModels() if (result == null) { call.respond(HttpStatusCode.BadRequest, errorBody("Sub2API 未配置或未登录,请在设置中填写 API Key")) } else { call.respondBytes(result.body, contentTypeOf(result.contentType), HttpStatusCode(result.status, "")) } } post("/v1/chat/completions") { if (!authorized(call, apiKey)) { call.respond(HttpStatusCode.Unauthorized, errorBody("未授权")) return@post } val rawBody = call.receiveText() if (rawBody.length > MAX_REQUEST_BYTES) { call.respond(HttpStatusCode(413, ""), errorBody("请求体过大")) return@post } val model = extractModel(rawBody) val route = router.route(model) val upstreamModel = stripModelPrefix(model) val body = if (upstreamModel != model) rewriteModelBody(rawBody, upstreamModel) else rawBody val conversationId = call.request.headers["X-Conversation-Id"] ?: extractUser(rawBody) val pooled = pool.next(route.kind, route.region, conversationId) val accountId = pooled?.accountId try { if (extractStream(body)) { val stream: ProxyUpstreamStream? = when (route.kind) { ProviderKind.WORKBUDDY_CN, ProviderKind.WORKBUDDY_INTL -> workBuddy.openStreamingChat(body, accountId, route.region) ProviderKind.TRAE_CN, ProviderKind.TRAE_INTL -> trae.openStreamingChat(body, accountId, route.region) else -> sub2Api.openStreamingChat(body, accountId) } if (stream == null) { call.respond( HttpStatusCode.BadRequest, errorBody("${route.kind.displayName} 未配置或未登录,请先在设置中配置") ) } else { call.respondOutputStream( contentType = contentTypeOf(stream.contentType), status = HttpStatusCode(stream.status, "") ) { try { val buffer = ByteArray(8192) val input = stream.input while (true) { val read = input.read(buffer) if (read < 0) break write(buffer, 0, read) flush() } } finally { stream.close() } } } } else { val result: ProxyUpstreamResponse? = when (route.kind) { ProviderKind.WORKBUDDY_CN, ProviderKind.WORKBUDDY_INTL -> workBuddy.forwardChat(body, accountId, route.region) ProviderKind.TRAE_CN, ProviderKind.TRAE_INTL -> trae.forwardChat(body, accountId, route.region) else -> sub2Api.forwardChat(body, accountId) } if (result == null) { call.respond( HttpStatusCode.BadRequest, errorBody("${route.kind.displayName} 未配置或未登录,请先在设置中配置") ) } else { call.respondBytes(result.body, contentTypeOf(result.contentType), HttpStatusCode(result.status, "")) } } } catch (e: IOException) { val detail = e.message ?: "未知错误" call.respond( HttpStatusCode.BadGateway, errorBody("上游网络错误:$detail") ) } catch (e: Exception) { if (e is kotlinx.coroutines.CancellationException) throw e val detail = e.message ?: "未知错误" call.respond( HttpStatusCode.InternalServerError, errorBody("网关内部错误:$detail") ) } } } } private fun authorized(call: ApplicationCall, apiKey: String?): Boolean { if (apiKey.isNullOrBlank()) return true val header = call.request.headers[HttpHeaders.Authorization] ?: return false val expect = "Bearer $apiKey" return header.trim() == expect } private fun errorBody(message: String): Map = mapOf("error" to mapOf("message" to message, "type" to "invalid_request_error")) private fun contentTypeOf(raw: String): ContentType = runCatching { ContentType.parse(raw) }.getOrDefault(ContentType.Application.Json) private fun extractStream(body: String): Boolean = runCatching { val el = kotlinx.serialization.json.Json.parseToJsonElement(body) (el as? kotlinx.serialization.json.JsonObject) ?.get("stream") ?.let { if (it is kotlinx.serialization.json.JsonPrimitive) it.contentOrNull?.toBooleanStrictOrNull() else null } }.getOrNull() ?: false private fun stripModelPrefix(model: String): String { val m = model.trim() val lower = m.lowercase() val prefix = listOf( "wbcn-", "workbuddy-cn", "codebuddy-", "wbintl-", "workbuddy-intl", "workbuddy-", "traeintl-", "trae-intl", "traecn-", "trae-cn", "trae-", "sub2api-", "openai-", ).firstOrNull { lower.startsWith(it) } return if (prefix != null) m.substring(prefix.length).ifBlank { m } else m } private fun rewriteModelBody(body: String, newModel: String): String = try { val obj = kotlinx.serialization.json.Json.parseToJsonElement(body) as? kotlinx.serialization.json.JsonObject ?: return body kotlinx.serialization.json.buildJsonObject { obj.forEach { (key, value) -> put(key, if (key == "model") kotlinx.serialization.json.JsonPrimitive(newModel) else value) } }.toString() } catch (_: Throwable) { body } private fun extractModel(body: String): String = runCatching { val el = kotlinx.serialization.json.Json.parseToJsonElement(body) (el as? kotlinx.serialization.json.JsonObject) ?.get("model") ?.let { if (it is kotlinx.serialization.json.JsonPrimitive) it.content else null } }.getOrNull() ?: "" private fun extractUser(body: String): String? = runCatching { val el = kotlinx.serialization.json.Json.parseToJsonElement(body) (el as? kotlinx.serialization.json.JsonObject) ?.get("user") ?.let { if (it is kotlinx.serialization.json.JsonPrimitive) it.content else null } }.getOrNull() companion object { private const val MAX_REQUEST_BYTES = 10 * 1024 * 1024 } }