266 lines
11 KiB
Kotlin
266 lines
11 KiB
Kotlin
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<Boolean> = _isRunning.asStateFlow()
|
||
|
||
override fun start(config: ProxyServerConfig): Result<Unit> {
|
||
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<String, Any> =
|
||
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
|
||
}
|
||
}
|