feat(proxy): 聚合网关——按模型自动选供应商并失败切换
Some checks are pending
Release / build (push) Waiting to run

- ProviderRouter 重构为聚合路由:请求只写 model,网关自动匹配供应商
- 候选按历史速度排序,上游不可用/HTTP>=400/抛错自动切换下一个
- Sub2API 作为通用兜底;显式前缀 traecn-/wbintl- 仍可强制指定
- 显式前缀转发前自动剥离,避免上游收到 traecn-glm-5.2 这类模型名
- 内存速度/冷却统计:失败供应商冷却 30s 起步、指数退避
- 版本升至 1.7.7
This commit is contained in:
Liuxinyu176 2026-10-09 11:49:38 +08:00
parent 73ff6b2ee7
commit 597632dad3
7 changed files with 309 additions and 44 deletions

View File

@ -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 {

View File

@ -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, ""))

View File

@ -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<ProviderTarget, Long>()
private val cooldownUntil = ConcurrentHashMap<ProviderTarget, Long>()
private val failCount = ConcurrentHashMap<ProviderTarget, Int>()
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<Candidate>): List<Candidate> = input.sortedWith(
compareBy<Candidate>(
{ 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<Candidate> {
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<ProviderTarget> {
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<ProviderTarget>()
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() ?: ""
}

View File

@ -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
}
}
/**

View File

@ -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",

View File

@ -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)

View File

@ -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