diff --git a/app/src/main/java/com/rainy/token/data/proxy/KtorLocalProxyServer.kt b/app/src/main/java/com/rainy/token/data/proxy/KtorLocalProxyServer.kt index 778fa75..aacdc13 100644 --- a/app/src/main/java/com/rainy/token/data/proxy/KtorLocalProxyServer.kt +++ b/app/src/main/java/com/rainy/token/data/proxy/KtorLocalProxyServer.kt @@ -20,6 +20,9 @@ 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 @@ -97,7 +100,10 @@ class KtorLocalProxyServer @Inject constructor( call.respond(mapOf("status" to "ok")) } get("/v1/models") { - if (!authorized(call, apiKey)) return@get + 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")) @@ -106,19 +112,24 @@ class KtorLocalProxyServer @Inject constructor( } } post("/v1/chat/completions") { - if (!authorized(call, apiKey)) return@post - val body = call.receiveText() - if (body.length > MAX_REQUEST_BYTES) { + 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(body) + val model = extractModel(rawBody) val route = router.route(model) - val conversationId = call.request.headers["X-Conversation-Id"] ?: extractUser(body) + 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 { - val result = when (route.kind) { + val result: ProxyUpstreamResponse? = when (route.kind) { ProviderKind.WORKBUDDY_CN, ProviderKind.WORKBUDDY_INTL -> workBuddy.forwardChat(body, accountId, route.region) @@ -166,6 +177,37 @@ class KtorLocalProxyServer @Inject constructor( 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) diff --git a/app/src/main/java/com/rainy/token/data/proxy/ProxyUpstreamStream.kt b/app/src/main/java/com/rainy/token/data/proxy/ProxyUpstreamStream.kt new file mode 100644 index 0000000..2718439 --- /dev/null +++ b/app/src/main/java/com/rainy/token/data/proxy/ProxyUpstreamStream.kt @@ -0,0 +1,15 @@ +package com.rainy.token.data.proxy + +import java.io.InputStream + +/** + * 上游流式响应句柄:由本地网关在收到客户端 stream=true 请求时创建, + * 把上游 SSE/原始字节流实时转发给本地客户端。 + * 使用方必须在 finally 中调用 [close]。 + */ +class ProxyUpstreamStream( + val status: Int, + val contentType: String, + val input: InputStream, + val close: () -> Unit, +) diff --git a/app/src/main/java/com/rainy/token/data/proxy/Sub2ApiChatProxy.kt b/app/src/main/java/com/rainy/token/data/proxy/Sub2ApiChatProxy.kt index 2048d13..34f4480 100644 --- a/app/src/main/java/com/rainy/token/data/proxy/Sub2ApiChatProxy.kt +++ b/app/src/main/java/com/rainy/token/data/proxy/Sub2ApiChatProxy.kt @@ -34,6 +34,36 @@ class Sub2ApiChatProxy @Inject constructor( return forward(base = base, path = "/v1/chat/completions", requestBody = requestBody, accountId = accountId) } + /** 转发流式 POST /v1/chat/completions(上游 SSE 原样转发)。 */ + suspend fun openStreamingChat( + requestBody: String, + accountId: String? = null, + ): ProxyUpstreamStream? = withContext(Dispatchers.IO) { + val credential = credentialRepository.get(ServiceType.SUB2API, accountId) + ?: return@withContext null + if (credential !is Credential.Sub2ApiCredential) return@withContext null + val auth = resolveAuth(credential) ?: return@withContext null + val base = normalizeBase(credential.baseUrl) ?: return@withContext null + + val builder = Request.Builder() + .url(base + "/v1/chat/completions") + .addHeader("Authorization", auth) + .addHeader("Content-Type", "application/json") + .post(requestBody.toRequestBody("application/json".toMediaType())) + + val response = okHttpClient.newCall(builder.build()).execute() + val input = response.body?.byteStream() ?: run { + response.close() + return@withContext null + } + ProxyUpstreamStream( + status = response.code, + contentType = response.header("Content-Type") ?: "text/event-stream", + input = input, + close = { response.close() }, + ) + } + /** 转发 GET /v1/models。 */ suspend fun forwardModels(accountId: String? = null): ProxyUpstreamResponse? { val base = resolveBase(accountId) ?: return null @@ -68,12 +98,8 @@ class Sub2ApiChatProxy @Inject constructor( .post(body.toRequestBody("application/json".toMediaType())) } - val response = try { - okHttpClient.newCall(builder.build()).execute() - } catch (_: Throwable) { - return@withContext null - } - val bytes = try { response.body?.bytes() ?: ByteArray(0) } catch (_: Throwable) { ByteArray(0) } + val response = okHttpClient.newCall(builder.build()).execute() + val bytes = response.body?.bytes() ?: ByteArray(0) val contentType = response.header("Content-Type") ?: "application/json" val status = response.code response.close() diff --git a/app/src/main/java/com/rainy/token/data/proxy/TraeChatProxy.kt b/app/src/main/java/com/rainy/token/data/proxy/TraeChatProxy.kt index 3dd8a19..8f3f8af 100644 --- a/app/src/main/java/com/rainy/token/data/proxy/TraeChatProxy.kt +++ b/app/src/main/java/com/rainy/token/data/proxy/TraeChatProxy.kt @@ -63,7 +63,7 @@ class TraeChatProxy @Inject constructor( .url(base + "/api/agent/v3/llm_utils_chat") .addHeader("Content-Type", "application/json") .addHeader("Accept", "text/event-stream") - .addHeader("Authorization", "Cloud-IDE-JWT \$jwt") + .addHeader("Authorization", "Cloud-IDE-JWT $jwt") .addHeader("X-Cloudide-Token", jwt) .addHeader("x-uid", credential.userId ?: "") .addHeader("x-device-id", credential.deviceId ?: credential.checkinDeviceId ?: "") @@ -90,6 +90,65 @@ class TraeChatProxy @Inject constructor( ProxyUpstreamResponse(status, contentType, bytes) } + suspend fun openStreamingChat( + requestBody: String, + accountId: String? = null, + forcedRegion: ProviderRegion? = null, + ): ProxyUpstreamStream? = withContext(Dispatchers.IO) { + val credential = credentialRepository.get(ServiceType.TRAE, accountId) + ?: return@withContext null + if (credential !is Credential.TraeCredential) return@withContext null + val jwt = credential.jwt.trim().takeIf { it.isNotBlank() } + ?: return@withContext null + + val effectiveRegion = forcedRegion ?: runCatching { + ProviderRegion.valueOf(credential.region.uppercase()) + }.getOrNull() + val base = if (effectiveRegion == ProviderRegion.INTL) { + "https://grow-normal.trae.ai" + } else { + "https://api.trae.cn" + } + val upstreamBody = buildUpstreamBody(requestBody, credential) + val requestId = UUID.randomUUID().toString() + val sessionId = UUID.randomUUID().toString().replace("-", "") + + val builder = Request.Builder() + .url(base + "/api/agent/v3/llm_utils_chat") + .addHeader("Content-Type", "application/json") + .addHeader("Accept", "text/event-stream") + .addHeader("Authorization", "Cloud-IDE-JWT $jwt") + .addHeader("X-Cloudide-Token", jwt) + .addHeader("x-uid", credential.userId ?: "") + .addHeader("x-device-id", credential.deviceId ?: credential.checkinDeviceId ?: "") + .addHeader("x-machine-id", credential.deviceId ?: credential.checkinDeviceId ?: "") + .addHeader("x-request-id", requestId) + .addHeader("x-app-id", "trae") + .addHeader("x-ide-version", "1.0.0") + .addHeader("x-ide-version-code", "1000000") + .addHeader("x-ide-version-type", "stable") + .addHeader("x-os-version", "Android") + .addHeader("x-system-type", "Android") + .addHeader("User-Agent", "RainyToken/1.0") + .post(upstreamBody.toRequestBody("application/json".toMediaType())) + + val response = try { + okHttpClient.newCall(builder.build()).execute() + } catch (e: java.io.IOException) { + throw e + } + val input = response.body?.byteStream() ?: run { + response.close() + return@withContext null + } + ProxyUpstreamStream( + status = response.code, + contentType = response.header("Content-Type") ?: "text/event-stream", + input = input, + close = { response.close() }, + ) + } + private fun buildUpstreamBody(raw: String, credential: Credential.TraeCredential): String { val src = try { json.parseToJsonElement(raw) as? JsonObject diff --git a/app/src/main/java/com/rainy/token/data/proxy/WorkBuddyChatProxy.kt b/app/src/main/java/com/rainy/token/data/proxy/WorkBuddyChatProxy.kt index 031b363..c22a7cc 100644 --- a/app/src/main/java/com/rainy/token/data/proxy/WorkBuddyChatProxy.kt +++ b/app/src/main/java/com/rainy/token/data/proxy/WorkBuddyChatProxy.kt @@ -50,7 +50,7 @@ class WorkBuddyChatProxy @Inject constructor( .url(base + "/v2/chat/completions") .addHeader("Content-Type", "application/json") .addHeader("Accept", "text/event-stream") - .addHeader("Authorization", "Bearer \$accessToken") + .addHeader("Authorization", "Bearer $accessToken") .addHeader("User-Agent", "RainyToken/1.0") .addHeader("X-IDE-Type", "VSCode") .addHeader("X-IDE-Name", "CodeBuddy") @@ -72,4 +72,56 @@ class WorkBuddyChatProxy @Inject constructor( response.close() ProxyUpstreamResponse(status, contentType, bytes) } + suspend fun openStreamingChat( + requestBody: String, + accountId: String? = null, + forcedRegion: ProviderRegion? = null, + ): ProxyUpstreamStream? = withContext(Dispatchers.IO) { + val credential = credentialRepository.get(ServiceType.WORKBUDDY, accountId) + ?: return@withContext null + if (credential !is Credential.WorkBuddyCredential) return@withContext null + val accessToken = credential.accessToken.trim().takeIf { it.isNotBlank() } + ?: return@withContext null + + val effectiveRegion = forcedRegion ?: runCatching { + ProviderRegion.valueOf(credential.region.uppercase()) + }.getOrNull() + val base = if (effectiveRegion == ProviderRegion.INTL) { + "https://www.workbuddy.ai" + } else { + "https://copilot.tencent.com" + } + val requestId = UUID.randomUUID().toString() + val builder = Request.Builder() + .url(base + "/v2/chat/completions") + .addHeader("Content-Type", "application/json") + .addHeader("Accept", "text/event-stream") + .addHeader("Authorization", "Bearer $accessToken") + .addHeader("User-Agent", "RainyToken/1.0") + .addHeader("X-IDE-Type", "VSCode") + .addHeader("X-IDE-Name", "CodeBuddy") + .addHeader("X-IDE-Version", "3.0.0") + .addHeader("X-Product", "CodeBuddy") + .addHeader("X-Agent-Intent", "craft") + .addHeader("X-Request-ID", requestId) + .addHeader("X-Conv-Request-ID", requestId) + .post(requestBody.toRequestBody("application/json".toMediaType())) + + val response = try { + okHttpClient.newCall(builder.build()).execute() + } catch (e: java.io.IOException) { + throw e + } + val input = response.body?.byteStream() ?: run { + response.close() + return@withContext null + } + ProxyUpstreamStream( + status = response.code, + contentType = response.header("Content-Type") ?: "text/event-stream", + input = input, + close = { response.close() }, + ) + } + }