- ProviderRouter 重构为聚合路由:请求只写 model,网关自动匹配供应商 - 候选按历史速度排序,上游不可用/HTTP>=400/抛错自动切换下一个 - Sub2API 作为通用兜底;显式前缀 traecn-/wbintl- 仍可强制指定 - 显式前缀转发前自动剥离,避免上游收到 traecn-glm-5.2 这类模型名 - 内存速度/冷却统计:失败供应商冷却 30s 起步、指数退避 - 版本升至 1.7.7
This commit is contained in:
parent
73ff6b2ee7
commit
597632dad3
@ -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 {
|
||||
|
||||
@ -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, ""))
|
||||
|
||||
@ -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() ?: ""
|
||||
}
|
||||
|
||||
@ -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
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
|
||||
@ -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",
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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
|
||||
|
||||
Loading…
Reference in New Issue
Block a user