diff --git a/app/build.gradle.kts b/app/build.gradle.kts index 5ea6e18..acdca91 100644 --- a/app/build.gradle.kts +++ b/app/build.gradle.kts @@ -16,8 +16,8 @@ android { applicationId = "com.rainy.token" minSdk = 31 targetSdk = 35 - versionCode = 26 - versionName = "1.7.6" + versionCode = 27 + versionName = "1.7.7" testInstrumentationRunner = "androidx.test.runner.AndroidJUnitRunner" vectorDrawables { 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 48be89a..0f26125 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 @@ -122,28 +122,14 @@ class KtorLocalProxyServer @Inject constructor( 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 (extractStream(rawBody)) { + val stream = router.openStreamingChat(rawBody, conversationId) if (stream == null) { call.respond( HttpStatusCode.BadRequest, - errorBody("${route.kind.displayName} 未配置或未登录,请先在设置中配置") + errorBody("所有上游均不可用或未配置,请检查设置") ) } else { call.respondOutputStream( @@ -165,19 +151,11 @@ class KtorLocalProxyServer @Inject constructor( } } } 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) - } + val result = router.forwardChat(rawBody, conversationId) if (result == null) { call.respond( HttpStatusCode.BadRequest, - errorBody("${route.kind.displayName} 未配置或未登录,请先在设置中配置") + errorBody("所有上游均不可用或未配置,请检查设置") ) } else { call.respondBytes(result.body, contentTypeOf(result.contentType), HttpStatusCode(result.status, "")) diff --git a/app/src/main/java/com/rainy/token/data/proxy/ProviderRouter.kt b/app/src/main/java/com/rainy/token/data/proxy/ProviderRouter.kt index 126de77..bc99dd0 100644 --- a/app/src/main/java/com/rainy/token/data/proxy/ProviderRouter.kt +++ b/app/src/main/java/com/rainy/token/data/proxy/ProviderRouter.kt @@ -1,38 +1,293 @@ package com.rainy.token.data.proxy +import java.io.IOException +import java.util.concurrent.ConcurrentHashMap +import javax.inject.Inject +import javax.inject.Singleton +import kotlinx.serialization.json.Json +import kotlinx.serialization.json.JsonObject +import kotlinx.serialization.json.JsonPrimitive + /** - * 模型名 → Provider + Region 路由。 + * 聚合路由:一个本地 API Key 通吃所有供应商。 * - * 约定前缀: - * - wbcn- / workbuddy-cn / codebuddy- → WorkBuddy 国内版 - * - wbintl- / workbuddy-intl → WorkBuddy 国际版 - * - traecn- / trae-cn / trae- → Trae CN - * - traeintl- / trae-intl → Trae INTL - * - 其余默认 OpenAI 兼容透传(Sub2API) + * 规则: + * 1. 模型名带显式前缀(traecn- / wbintl- 等)→ 强制指定供应商; + * 2. 否则按【模型名】找所有支持的供应商,按历史速度排序逐个尝试; + * 3. 上游不可用/HTTP>=400/抛错 → 自动切换到下一个供应商; + * 4. Sub2API 作为通用兜底(配置了账号时)。 */ -class ProviderRouter { +@Singleton +class ProviderRouter @Inject constructor( + private val traeChatProxy: TraeChatProxy, + private val workBuddyChatProxy: WorkBuddyChatProxy, + private val sub2ApiChatProxy: Sub2ApiChatProxy, + private val traeModelProvider: TraeModelProvider, + private val workBuddyModelProvider: WorkBuddyModelProvider, + private val accountPool: AccountPool, +) { data class Route( val kind: ProviderKind, val region: ProviderRegion?, ) - fun route(model: String): Route { + /** 网关内供应商目标。 */ + enum class ProviderTarget( + val kind: ProviderKind, + val region: ProviderRegion?, + val displayName: String, + ) { + TRAE_CN(ProviderKind.TRAE_CN, ProviderRegion.CN, "Trae CN"), + TRAE_INTL(ProviderKind.TRAE_INTL, ProviderRegion.INTL, "Trae INTL"), + WORKBUDDY_CN(ProviderKind.WORKBUDDY_CN, ProviderRegion.CN, "WorkBuddy CN"), + WORKBUDDY_INTL(ProviderKind.WORKBUDDY_INTL, ProviderRegion.INTL, "WorkBuddy INTL"), + SUB2API(ProviderKind.OPENAI_COMPATIBLE, null, "Sub2API"), + } + + private data class Candidate( + val target: ProviderTarget, + val accountId: String?, + ) + + // ---- 速度 / 健康度统计(内存态) ---- + + private val avgLatency = ConcurrentHashMap() + private val cooldownUntil = ConcurrentHashMap() + private val failCount = ConcurrentHashMap() + + private fun recordSuccess(target: ProviderTarget, startedMs: Long, status: Int) { + val ms = (System.currentTimeMillis() - startedMs).coerceAtLeast(1L) + val old = avgLatency[target] + avgLatency[target] = if (old == null) ms else (old * 3 + ms) / 4 + cooldownUntil.remove(target) + failCount.remove(target) + } + + private fun recordFail(target: ProviderTarget, startedMs: Long) { + val ms = (System.currentTimeMillis() - startedMs).coerceAtLeast(1L) + val old = avgLatency[target] + avgLatency[target] = if (old == null) ms + 5000L else old + 5000L + val fails = (failCount[target] ?: 0) + 1 + failCount[target] = fails + cooldownUntil[target] = System.currentTimeMillis() + 30_000L * fails + } + + private fun cooling(target: ProviderTarget): Boolean { + val until = cooldownUntil[target] ?: return false + if (System.currentTimeMillis() >= until) { + cooldownUntil.remove(target) + return false + } + return true + } + + private fun orderBySpeed(input: List): List = input.sortedWith( + compareBy( + { cooling(it.target) }, + { avgLatency[it.target] ?: Long.MAX_VALUE }, + { if (it.target == ProviderTarget.SUB2API) 1 else 0 }, + ) + ) + + // ---- 流式转发(带失败切换) ---- + + suspend fun openStreamingChat( + requestBody: String, + conversationId: String? = null, + ): ProxyUpstreamStream? { + val model = extractModel(requestBody) + val forced = explicitTarget(model) + val upstreamBody = if (forced != null) stripModelInBody(requestBody, model) else requestBody + val candidates = candidatesFor(model, conversationId, forced) + for (candidate in candidates) { + val started = System.currentTimeMillis() + val stream = try { + when (candidate.target.kind) { + ProviderKind.WORKBUDDY_CN, ProviderKind.WORKBUDDY_INTL -> + workBuddyChatProxy.openStreamingChat( + upstreamBody, candidate.accountId, candidate.target.region, + ) + + ProviderKind.TRAE_CN, ProviderKind.TRAE_INTL -> + traeChatProxy.openStreamingChat( + upstreamBody, candidate.accountId, candidate.target.region, + ) + + else -> + sub2ApiChatProxy.openStreamingChat(upstreamBody, candidate.accountId) + } + } catch (e: java.io.IOException) { + recordFail(candidate.target, started) + null + } catch (e: Exception) { + if (e is kotlinx.coroutines.CancellationException) throw e + recordFail(candidate.target, started) + null + } + + if (stream == null) { + recordFail(candidate.target, started) + continue + } + if (stream.status >= 400) { + stream.close() + recordFail(candidate.target, started) + continue + } + recordSuccess(candidate.target, started, stream.status) + return stream + } + return null + } + + // ---- 非流式转发(带失败切换) ---- + + suspend fun forwardChat( + requestBody: String, + conversationId: String? = null, + ): ProxyUpstreamResponse? { + val model = extractModel(requestBody) + val forced = explicitTarget(model) + val upstreamBody = if (forced != null) stripModelInBody(requestBody, model) else requestBody + val candidates = candidatesFor(model, conversationId, forced) + for (candidate in candidates) { + val started = System.currentTimeMillis() + val result = try { + when (candidate.target.kind) { + ProviderKind.WORKBUDDY_CN, ProviderKind.WORKBUDDY_INTL -> + workBuddyChatProxy.forwardChat( + upstreamBody, candidate.accountId, candidate.target.region, + ) + + ProviderKind.TRAE_CN, ProviderKind.TRAE_INTL -> + traeChatProxy.forwardChat( + upstreamBody, candidate.accountId, candidate.target.region, + ) + + else -> + sub2ApiChatProxy.forwardChat(upstreamBody, candidate.accountId) + } + } catch (e: java.io.IOException) { + recordFail(candidate.target, started) + null + } catch (e: Exception) { + if (e is kotlinx.coroutines.CancellationException) throw e + recordFail(candidate.target, started) + null + } + + if (result == null) { + recordFail(candidate.target, started) + continue + } + if (result.status >= 400) { + recordFail(candidate.target, started) + continue + } + recordSuccess(candidate.target, started, result.status) + return result + } + return null + } + + // ---- 候选构建 ---- + + private suspend fun candidatesFor( + model: String, + conversationId: String?, + forced: ProviderTarget?, + ): List { + val targets = if (forced != null) { + listOf(forced) + } else { + capableTargets(model) + } + val built = targets.mapNotNull { target -> + val pooled = accountPool.next(target.kind, target.region, conversationId) + pooled?.accountId?.let { Candidate(target, it) } + } + return orderBySpeed(built) + } + + private fun capableTargets(model: String): List { + val m = model.trim().lowercase() + if (m.isBlank() || m == "auto") { + return listOf( + ProviderTarget.TRAE_CN, + ProviderTarget.WORKBUDDY_CN, + ProviderTarget.TRAE_INTL, + ProviderTarget.WORKBUDDY_INTL, + ProviderTarget.SUB2API, + ) + } + val out = mutableListOf() + if (traeModelProvider.supports(model, ProviderRegion.CN)) out += ProviderTarget.TRAE_CN + if (traeModelProvider.supports(model, ProviderRegion.INTL)) out += ProviderTarget.TRAE_INTL + if (workBuddyModelProvider.supports(model, ProviderRegion.CN)) out += ProviderTarget.WORKBUDDY_CN + if (workBuddyModelProvider.supports(model, ProviderRegion.INTL)) out += ProviderTarget.WORKBUDDY_INTL + // 通用兜底:Sub2API 是 OpenAI 兼容实例,什么模型都可能支持 + out += ProviderTarget.SUB2API + return out.distinct().ifEmpty { listOf(ProviderTarget.TRAE_CN) } + } + + private fun explicitTarget(model: String): ProviderTarget? { val m = model.trim().lowercase() return when { m.startsWith("wbcn-") || m.startsWith("workbuddy-cn") || m.startsWith("codebuddy-") -> - Route(ProviderKind.WORKBUDDY_CN, ProviderRegion.CN) + ProviderTarget.WORKBUDDY_CN m.startsWith("wbintl-") || m.startsWith("workbuddy-intl") || m.startsWith("workbuddy-") -> - Route(ProviderKind.WORKBUDDY_INTL, ProviderRegion.INTL) + ProviderTarget.WORKBUDDY_INTL m.startsWith("traeintl-") || m.startsWith("trae-intl") -> - Route(ProviderKind.TRAE_INTL, ProviderRegion.INTL) + ProviderTarget.TRAE_INTL m.startsWith("traecn-") || m.startsWith("trae-cn") || m.startsWith("trae-") -> - Route(ProviderKind.TRAE_CN, ProviderRegion.CN) + ProviderTarget.TRAE_CN - else -> Route(ProviderKind.OPENAI_COMPATIBLE, null) + m.startsWith("sub2api-") || m.startsWith("openai-") -> + ProviderTarget.SUB2API + + else -> null } } + + /** 兼容旧调用:只返回模型名前缀对应的路由。 */ + fun route(model: String): Route { + val target = explicitTarget(model) + return if (target == null) { + Route(ProviderKind.OPENAI_COMPATIBLE, null) + } else { + Route(target.kind, target.region) + } + } + + private fun stripModelInBody(body: String, originalModel: String): String { + val m = originalModel.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 body + val newModel = m.substring(prefix.length).ifBlank { m } + if (newModel == m) return body + return runCatching { + val obj = Json.parseToJsonElement(body) as? JsonObject ?: return@runCatching body + JsonObject( + obj.entries.associate { (key, value) -> + key to (if (key == "model") JsonPrimitive(newModel) else value) + } + ).toString() + }.getOrDefault(body) + } + + private fun extractModel(body: String): String = runCatching { + val el = Json.parseToJsonElement(body) + (el as? JsonObject) + ?.get("model") + ?.let { if (it is JsonPrimitive) it.content else null } + }.getOrNull() ?: "" } 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 34f4480..9087cec 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 @@ -117,6 +117,12 @@ class Sub2ApiChatProxy @Inject constructor( while (s.endsWith("/")) s = s.dropLast(1) return s.takeIf { it.isNotBlank() } } + + /** 是否已配置 Sub2API 凭据(用于聚合网关兜底)。 */ + suspend fun hasCredential(): Boolean { + val credential = credentialRepository.get(ServiceType.SUB2API, null) + return credential != null + } } /** 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 9d27c84..7a84c27 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 @@ -266,6 +266,9 @@ class TraeChatProxy @Inject constructor( } companion object { + /** 供路由判断模型是否可能由 Trae 消化。 */ + internal fun resolveAlias(modelName: String): String? = + MODEL_ALIASES[modelName.trim().lowercase()] /** OpenAI/Claude 常用名 -> Trae CN 内部模型名(参考 trae2api-cn)。 */ private val MODEL_ALIASES = mapOf( "auto" to "glm-5.2", diff --git a/app/src/main/java/com/rainy/token/data/proxy/TraeModelProvider.kt b/app/src/main/java/com/rainy/token/data/proxy/TraeModelProvider.kt index 7f6fec9..b54b258 100644 --- a/app/src/main/java/com/rainy/token/data/proxy/TraeModelProvider.kt +++ b/app/src/main/java/com/rainy/token/data/proxy/TraeModelProvider.kt @@ -57,6 +57,14 @@ class TraeModelProvider @Inject constructor( prefs.edit().putString(key, id).apply() } + /** 该区域是否可能支持此模型(内置列表/在线列表/别名)。auto 视为支持。 */ + fun supports(modelId: String, region: ProviderRegion): Boolean { + val want = modelId.trim().lowercase() + if (want.isBlank() || want == "auto") return true + if (models.value.any { it.id.equals(modelId, ignoreCase = true) }) return true + return TraeChatProxy.resolveAlias(modelId) != null + } + /** 拉取账号在线模型列表;失败时保留内置列表并返回 false。 */ suspend fun refreshFor(forcedRegion: ProviderRegion?): Boolean = withContext(Dispatchers.IO) { val credential = credentialRepository.get(ServiceType.TRAE, null) diff --git a/app/src/main/java/com/rainy/token/di/NetworkModule.kt b/app/src/main/java/com/rainy/token/di/NetworkModule.kt index 7c8bb96..4f9219d 100644 --- a/app/src/main/java/com/rainy/token/di/NetworkModule.kt +++ b/app/src/main/java/com/rainy/token/di/NetworkModule.kt @@ -28,6 +28,7 @@ import com.rainy.token.data.proxy.Sub2ApiChatProxy import com.rainy.token.data.proxy.TraeChatProxy import com.rainy.token.data.proxy.TraeModelProvider import com.rainy.token.data.proxy.WorkBuddyChatProxy +import com.rainy.token.data.proxy.WorkBuddyModelProvider import com.rainy.token.data.repository.WorkBuddyRepository import dagger.Module import dagger.Provides @@ -257,7 +258,21 @@ object NetworkModule { @Provides @Singleton - fun provideProviderRouter(): ProviderRouter = ProviderRouter() + fun provideProviderRouter( + sub2ApiChatProxy: Sub2ApiChatProxy, + workBuddyChatProxy: WorkBuddyChatProxy, + traeChatProxy: TraeChatProxy, + traeModelProvider: TraeModelProvider, + workBuddyModelProvider: WorkBuddyModelProvider, + accountPool: AccountPool, + ): ProviderRouter = ProviderRouter( + traeChatProxy, + workBuddyChatProxy, + sub2ApiChatProxy, + traeModelProvider, + workBuddyModelProvider, + accountPool, + ) @Provides @Singleton