259 lines
10 KiB
Kotlin
259 lines
10 KiB
Kotlin
package com.rainy.token.data.proxy
|
||
|
||
import com.rainy.token.data.repository.CredentialRepository
|
||
import com.rainy.token.domain.model.Credential
|
||
import com.rainy.token.domain.service.ServiceType
|
||
import java.util.UUID
|
||
import javax.inject.Inject
|
||
import javax.inject.Singleton
|
||
import kotlinx.coroutines.Dispatchers
|
||
import kotlinx.coroutines.withContext
|
||
import kotlinx.serialization.json.Json
|
||
import kotlinx.serialization.json.JsonArray
|
||
import kotlinx.serialization.json.JsonElement
|
||
import kotlinx.serialization.json.JsonNull
|
||
import kotlinx.serialization.json.JsonObject
|
||
import kotlinx.serialization.json.JsonPrimitive
|
||
import kotlinx.serialization.json.buildJsonArray
|
||
import kotlinx.serialization.json.buildJsonObject
|
||
import kotlinx.serialization.json.contentOrNull
|
||
import kotlinx.serialization.json.put
|
||
import okhttp3.Headers.Companion.toHeaders
|
||
import okhttp3.MediaType.Companion.toMediaType
|
||
import okhttp3.OkHttpClient
|
||
import okhttp3.Request
|
||
import okhttp3.RequestBody.Companion.toRequestBody
|
||
|
||
/**
|
||
* Trae CN 上游 Chat 代理。
|
||
*
|
||
* 上游为私有协议 POST {base}/api/agent/v3/llm_utils_chat。
|
||
* CN 模型网关固定走 trae-api-cn.mchost.guru(api.trae.cn 只是账号/OAuth 主机,返回 404)。
|
||
* 请求头/body 对齐 trae2api-cn 参考实现的 SOLO 协议。
|
||
*/
|
||
@Singleton
|
||
class TraeChatProxy @Inject constructor(
|
||
private val okHttpClient: OkHttpClient,
|
||
private val credentialRepository: CredentialRepository,
|
||
) {
|
||
|
||
private val json = Json { ignoreUnknownKeys = true }
|
||
|
||
suspend fun forwardChat(
|
||
requestBody: String,
|
||
accountId: String? = null,
|
||
forcedRegion: ProviderRegion? = null,
|
||
): ProxyUpstreamResponse? =
|
||
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://a0ai-api-sg.byteintlapi.com"
|
||
} else {
|
||
"https://trae-api-cn.mchost.guru"
|
||
}
|
||
val upstreamBody = buildUpstreamBody(requestBody, credential)
|
||
val requestId = UUID.randomUUID().toString()
|
||
|
||
val builder = Request.Builder()
|
||
.url(base + "/api/agent/v3/llm_utils_chat")
|
||
.headers(soloHeaders(jwt, credential, requestId).toHeaders())
|
||
.post(upstreamBody.toRequestBody("application/json".toMediaType()))
|
||
|
||
val response = try {
|
||
okHttpClient.newCall(builder.build()).execute()
|
||
} catch (e: java.io.IOException) {
|
||
throw e
|
||
}
|
||
val bytes = try { response.body?.bytes() ?: ByteArray(0) } catch (_: Throwable) { ByteArray(0) }
|
||
val contentType = response.header("Content-Type") ?: "application/json"
|
||
val status = response.code
|
||
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.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://a0ai-api-sg.byteintlapi.com"
|
||
} else {
|
||
"https://trae-api-cn.mchost.guru"
|
||
}
|
||
val upstreamBody = buildUpstreamBody(requestBody, credential)
|
||
val requestId = UUID.randomUUID().toString()
|
||
|
||
val builder = Request.Builder()
|
||
.url(base + "/api/agent/v3/llm_utils_chat")
|
||
.headers(soloHeaders(jwt, credential, requestId).toHeaders())
|
||
.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 soloHeaders(
|
||
jwt: String,
|
||
credential: Credential.TraeCredential,
|
||
requestId: String,
|
||
): Map<String, String> = linkedMapOf(
|
||
"Content-Type" to "application/json",
|
||
"Accept" to "text/event-stream",
|
||
"Connection" to "keep-alive",
|
||
"Authorization" to "Cloud-IDE-JWT $jwt",
|
||
"X-Cloudide-Token" to jwt,
|
||
"x-ide-token" to jwt,
|
||
"x-uid" to (credential.userId ?: ""),
|
||
"x-app-id" to "6eefa01c-1036-4c7e-9ca5-d891f63bfcd8",
|
||
"x-device-id" to (credential.deviceId ?: credential.checkinDeviceId ?: ""),
|
||
"x-machine-id" to (credential.deviceId ?: credential.checkinDeviceId ?: ""),
|
||
"x-request-id" to requestId,
|
||
"x-ide-version" to "0.1.52",
|
||
"x-ide-version-code" to "20260811",
|
||
"x-ide-version-type" to "stable",
|
||
"x-app-version" to "default",
|
||
"x-app-version-code" to "20260811",
|
||
"x-version-code" to "20260811",
|
||
"x-device-cpu" to "AMD",
|
||
"x-device-brand" to "83DG",
|
||
"x-device-type" to "windows",
|
||
"x-device-platform" to "windows",
|
||
"x-os-version" to "Windows 11 Pro",
|
||
"x-system-type" to "Windows",
|
||
"package-type" to "stable_cn",
|
||
"x-lscbd-aid" to "787976",
|
||
"x-lscbd-platform" to "windows",
|
||
"x-ss-dp" to "787976",
|
||
"x-plugin-channel" to "icube-ai",
|
||
"app-version" to "0.1.52",
|
||
"x-bridge-transport" to "aha",
|
||
"x-ahanet-timeout" to "86400",
|
||
"x-lgw-req-sdk-type" to "3",
|
||
"x-net-sdk-domain-dispatch" to "1",
|
||
"x-ttnet-bypass-decompression" to "1",
|
||
"x-ttnet-bypass-cookie" to "0",
|
||
"request-traffic-type" to "prod",
|
||
"User-Agent" to "Trae/0.1.52",
|
||
)
|
||
|
||
private fun convertNativeMessages(src: JsonElement?): List<JsonElement> {
|
||
val arr = src as? JsonArray ?: return emptyList()
|
||
return arr.mapNotNull { el ->
|
||
val m = el as? JsonObject ?: return@mapNotNull null
|
||
val rawRole = (m["role"] as? JsonPrimitive)?.contentOrNull?.lowercase() ?: "user"
|
||
val role = when (rawRole) {
|
||
"developer" -> "system"
|
||
"system", "user", "assistant", "tool", "function" -> rawRole
|
||
else -> "user"
|
||
}
|
||
buildJsonObject {
|
||
put("role", role)
|
||
val content = nativeContent(m["content"])
|
||
if (content != null) put("content", content)
|
||
m["name"]?.let { put("name", it) }
|
||
m["tool_call_id"]?.let { put("tool_call_id", it) }
|
||
m["tool_calls"]?.let { put("tool_calls", it) }
|
||
}
|
||
}
|
||
}
|
||
|
||
private fun nativeContent(content: JsonElement?): JsonElement? {
|
||
if (content == null || content is JsonNull) return null
|
||
if (content is JsonPrimitive) {
|
||
return buildJsonArray {
|
||
add(buildJsonObject {
|
||
put("type", "text")
|
||
put("text", content.content)
|
||
})
|
||
}
|
||
}
|
||
if (content is JsonArray) {
|
||
return buildJsonArray {
|
||
content.forEach { block ->
|
||
when (block) {
|
||
is JsonPrimitive -> add(buildJsonObject {
|
||
put("type", "text")
|
||
put("text", block.content)
|
||
})
|
||
is JsonObject -> {
|
||
val type = (block["type"] as? JsonPrimitive)?.contentOrNull?.lowercase()
|
||
if (type == "text" || type == "input_text") {
|
||
val text = (block["text"] as? JsonPrimitive)?.contentOrNull
|
||
?: (block["content"] as? JsonPrimitive)?.contentOrNull
|
||
?: ""
|
||
add(buildJsonObject {
|
||
put("type", "text")
|
||
put("text", text)
|
||
})
|
||
} else {
|
||
add(block)
|
||
}
|
||
}
|
||
else -> add(block)
|
||
}
|
||
}
|
||
}
|
||
}
|
||
return null
|
||
}
|
||
|
||
private fun buildUpstreamBody(raw: String, credential: Credential.TraeCredential): String {
|
||
val src = try {
|
||
json.parseToJsonElement(raw) as? JsonObject
|
||
} catch (_: Throwable) {
|
||
null
|
||
} ?: return raw
|
||
|
||
val model = (src["model"] as? JsonPrimitive)?.contentOrNull?.takeIf { it.isNotBlank() }
|
||
?: "glm-5.2"
|
||
val messages = convertNativeMessages(src["messages"])
|
||
val sessionId = UUID.randomUUID().toString().replace("-", "")
|
||
|
||
return buildJsonObject {
|
||
put("messages", JsonArray(messages))
|
||
put("config_name", model)
|
||
put("model", model)
|
||
put("function", "solo_work_lite")
|
||
put("stream", true)
|
||
put("request_id", sessionId)
|
||
put("session_id", sessionId)
|
||
src["tools"]?.let { put("tools", it) }
|
||
src["tool_choice"]?.let { put("tool_choice", it) }
|
||
(src["max_tokens"] as? JsonPrimitive)?.contentOrNull?.toIntOrNull()?.let {
|
||
put("max_tokens", it)
|
||
}
|
||
}.toString()
|
||
}
|
||
}
|