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 = 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 { 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() } }