feat(proxy): M2 账号池轮询 + 会话粘性 + 双区域路由

- ProviderRouter:前缀路由到 WorkBuddy CN/INTL、Trae CN/INTL、OpenAI 兼容
- AccountPool:从 CredentialRepository 实时取号,round-robin + 会话粘性(X-Conversation-Id / user)
- 账号按 credential.region 与路由区域匹配,找不到时回退全部账号
- WorkBuddy/Trae 代理支持 forcedRegion 覆盖
- Ktor 路由接入账号池
This commit is contained in:
Liuxinyu176 2026-10-08 23:34:35 +08:00
parent 5aa2d8f116
commit a1605e6993
6 changed files with 161 additions and 43 deletions

View File

@ -1,15 +1,24 @@
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 javax.inject.Inject
import javax.inject.Singleton
import kotlinx.coroutines.sync.Mutex
import kotlinx.coroutines.sync.withLock
/**
* 账号池:按 Provider 保存账号,负责轮询/冷却/会话粘性。
* 账号池:从 CredentialRepository 实时取号,负责轮询与会话粘性。
*
* 当前为骨架(round-robin)。后续把 workbuddy2api-panel 的
* 三因子加权、429 熔断冷却、会话粘性搬进来。
* - 多账号服务(Trae / WorkBuddy / Sub2API)按 ProviderKind 分组;
* - 会话粘性:同一 conversation 头持续命中同一账号,避免上下文错乱;
* - 区域匹配:优先选 credential.region 与路由 region 一致的账号,找不到时退回全部账号。
*/
class AccountPool {
@Singleton
class AccountPool @Inject constructor(
private val credentialRepository: CredentialRepository,
) {
data class PooledAccount(
val kind: ProviderKind,
@ -19,26 +28,60 @@ class AccountPool {
)
private val mutex = Mutex()
private val accounts = mutableMapOf<ProviderKind, MutableList<PooledAccount>>()
private val sessionSticky = mutableMapOf<String, String>()
private val cursor = mutableMapOf<ProviderKind, Int>()
fun register(account: PooledAccount) {
accounts.getOrPut(account.kind) { mutableListOf() }.add(account)
suspend fun next(
kind: ProviderKind,
region: ProviderRegion? = null,
sessionKey: String? = null,
): PooledAccount? = mutex.withLock {
val service = serviceFor(kind) ?: return null
val accounts = credentialRepository.accountsFor(service)
val allIds = accounts.map { it.id }
val matchedIds = if (region == null) {
allIds
} else {
allIds.filter { accountId -> regionMatches(service, accountId, region) }
.ifEmpty { allIds }
}
if (matchedIds.isEmpty()) return null
val accountId = if (sessionKey != null) {
sessionSticky[sessionKey]
?.takeIf { it in matchedIds }
?: pickRoundRobin(kind, matchedIds).also { sessionSticky[sessionKey] = it }
} else {
pickRoundRobin(kind, matchedIds)
}
PooledAccount(
kind = kind,
region = region,
accountId = accountId,
label = accounts.firstOrNull { it.id == accountId }?.label,
)
}
fun registerAll(list: List<PooledAccount>) {
list.forEach(::register)
private suspend fun regionMatches(service: ServiceType, accountId: String, region: ProviderRegion): Boolean {
val credential = credentialRepository.get(service, accountId) ?: return false
return when (credential) {
is Credential.TraeCredential -> credential.region.equals(region.name, ignoreCase = true)
is Credential.WorkBuddyCredential -> credential.region.equals(region.name, ignoreCase = true)
else -> true
}
}
suspend fun next(kind: ProviderKind, sessionKey: String? = null): PooledAccount? = mutex.withLock {
val list = accounts[kind]?.takeIf { it.isNotEmpty() } ?: return null
// TODO: 会话粘性:命中 sessionKey 对应账号时优先返回
private fun pickRoundRobin(kind: ProviderKind, ids: List<String>): String {
val idx = cursor[kind] ?: 0
cursor[kind] = (idx + 1) % list.size
list[idx]
cursor[kind] = (idx + 1) % ids.size
return ids[idx]
}
suspend fun snapshot(): Map<ProviderKind, List<PooledAccount>> = mutex.withLock {
accounts.mapValues { it.value.toList() }
private fun serviceFor(kind: ProviderKind): ServiceType? = when (kind) {
ProviderKind.WORKBUDDY_CN, ProviderKind.WORKBUDDY_INTL -> ServiceType.WORKBUDDY
ProviderKind.TRAE_CN, ProviderKind.TRAE_INTL -> ServiceType.TRAE
ProviderKind.OPENAI_COMPATIBLE -> ServiceType.SUB2API
else -> null
}
}

View File

@ -42,6 +42,8 @@ class KtorLocalProxyServer @Inject constructor(
private val sub2ApiChatProxy: Sub2ApiChatProxy,
private val workBuddyChatProxy: WorkBuddyChatProxy,
private val traeChatProxy: TraeChatProxy,
private val providerRouter: ProviderRouter,
private val accountPool: AccountPool,
) : LocalProxyServer {
private val lock = Any()
@ -57,7 +59,7 @@ class KtorLocalProxyServer @Inject constructor(
if (_isRunning.value) return Result.success(Unit)
return try {
val engine = embeddedServer(CIO, host = "127.0.0.1", port = config.port) {
proxyModule(config.apiKey, sub2ApiChatProxy, workBuddyChatProxy, traeChatProxy)
proxyModule(config.apiKey, sub2ApiChatProxy, workBuddyChatProxy, traeChatProxy, providerRouter, accountPool)
}
engine.start(wait = false)
server = engine
@ -84,6 +86,8 @@ class KtorLocalProxyServer @Inject constructor(
sub2Api: Sub2ApiChatProxy,
workBuddy: WorkBuddyChatProxy,
trae: TraeChatProxy,
router: ProviderRouter,
pool: AccountPool,
) {
install(ContentNegotiation) {
json()
@ -109,17 +113,24 @@ class KtorLocalProxyServer @Inject constructor(
return@post
}
val model = extractModel(body)
val provider = routeProvider(model)
val route = router.route(model)
val conversationId = call.request.headers["X-Conversation-Id"] ?: extractUser(body)
val pooled = pool.next(route.kind, route.region, conversationId)
val accountId = pooled?.accountId
try {
val result = when (provider) {
ChatRoute.WORKBUDDY_CN -> workBuddy.forwardChat(body)
ChatRoute.TRAE_CN -> trae.forwardChat(body)
ChatRoute.SUB2API -> sub2Api.forwardChat(body)
val result = 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)
}
if (result == null) {
call.respond(
HttpStatusCode.BadRequest,
errorBody("${provider.displayName} 未配置或未登录,请先在设置中配置")
errorBody("${route.kind.displayName} 未配置或未登录,请先在设置中配置")
)
} else {
call.respondBytes(result.body, contentTypeOf(result.contentType), HttpStatusCode(result.status, ""))
@ -162,20 +173,12 @@ class KtorLocalProxyServer @Inject constructor(
?.let { if (it is kotlinx.serialization.json.JsonPrimitive) it.content else null }
}.getOrNull() ?: ""
private fun routeProvider(model: String): ChatRoute {
val m = model.lowercase()
return when {
m.startsWith("wbcn-") -> ChatRoute.WORKBUDDY_CN
m.startsWith("traecn-") || m.startsWith("trae-") -> ChatRoute.TRAE_CN
else -> ChatRoute.SUB2API
}
}
private enum class ChatRoute(val displayName: String) {
WORKBUDDY_CN("WorkBuddy 国内版"),
TRAE_CN("Trae CN"),
SUB2API("Sub2API"),
}
private fun extractUser(body: String): String? = runCatching {
val el = kotlinx.serialization.json.Json.parseToJsonElement(body)
(el as? kotlinx.serialization.json.JsonObject)
?.get("user")
?.let { if (it is kotlinx.serialization.json.JsonPrimitive) it.content else null }
}.getOrNull()
companion object {
private const val MAX_REQUEST_BYTES = 10 * 1024 * 1024

View File

@ -0,0 +1,38 @@
package com.rainy.token.data.proxy
/**
* 模型名 → Provider + Region 路由。
*
* 约定前缀:
* - wbcn- / workbuddy-cn / codebuddy- → WorkBuddy 国内版
* - wbintl- / workbuddy-intl → WorkBuddy 国际版
* - traecn- / trae-cn / trae- → Trae CN
* - traeintl- / trae-intl → Trae INTL
* - 其余默认 OpenAI 兼容透传(Sub2API)
*/
class ProviderRouter {
data class Route(
val kind: ProviderKind,
val region: ProviderRegion?,
)
fun route(model: String): Route {
val m = model.trim().lowercase()
return when {
m.startsWith("wbcn-") || m.startsWith("workbuddy-cn") || m.startsWith("codebuddy-") ->
Route(ProviderKind.WORKBUDDY_CN, ProviderRegion.CN)
m.startsWith("wbintl-") || m.startsWith("workbuddy-intl") || m.startsWith("workbuddy-") ->
Route(ProviderKind.WORKBUDDY_INTL, ProviderRegion.INTL)
m.startsWith("traeintl-") || m.startsWith("trae-intl") ->
Route(ProviderKind.TRAE_INTL, ProviderRegion.INTL)
m.startsWith("traecn-") || m.startsWith("trae-cn") || m.startsWith("trae-") ->
Route(ProviderKind.TRAE_CN, ProviderRegion.CN)
else -> Route(ProviderKind.OPENAI_COMPATIBLE, null)
}
}
}

View File

@ -35,7 +35,11 @@ class TraeChatProxy @Inject constructor(
private val json = Json { ignoreUnknownKeys = true }
suspend fun forwardChat(requestBody: String, accountId: String? = null): ProxyUpstreamResponse? =
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
@ -43,7 +47,10 @@ class TraeChatProxy @Inject constructor(
val jwt = credential.jwt.trim().takeIf { it.isNotBlank() }
?: return@withContext null
val base = if (credential.region.equals("INTL", ignoreCase = true)) {
val effectiveRegion = forcedRegion ?: runCatching {
ProviderRegion.valueOf(credential.region.uppercase())
}.getOrNull()
val base = if (effectiveRegion == ProviderRegion.INTL) {
"https://grow-normal.trae.ai"
} else {
"https://api.trae.cn"

View File

@ -25,7 +25,11 @@ class WorkBuddyChatProxy @Inject constructor(
private val credentialRepository: CredentialRepository,
) {
suspend fun forwardChat(requestBody: String, accountId: String? = null): ProxyUpstreamResponse? =
suspend fun forwardChat(
requestBody: String,
accountId: String? = null,
forcedRegion: ProviderRegion? = null,
): ProxyUpstreamResponse? =
withContext(Dispatchers.IO) {
val credential = credentialRepository.get(ServiceType.WORKBUDDY, accountId)
?: return@withContext null
@ -33,7 +37,10 @@ class WorkBuddyChatProxy @Inject constructor(
val accessToken = credential.accessToken.trim().takeIf { it.isNotBlank() }
?: return@withContext null
val base = if (credential.region.equals("INTL", ignoreCase = true)) {
val effectiveRegion = forcedRegion ?: runCatching {
ProviderRegion.valueOf(credential.region.uppercase())
}.getOrNull()
val base = if (effectiveRegion == ProviderRegion.INTL) {
"https://www.workbuddy.ai"
} else {
"https://copilot.tencent.com"

View File

@ -20,6 +20,8 @@ import com.rainy.token.data.repository.OllamaRepository
import com.rainy.token.data.repository.Sub2ApiRepository
import com.rainy.token.data.repository.TraeRepository
import com.rainy.token.data.repository.UpdateRepository
import com.rainy.token.data.proxy.AccountPool
import com.rainy.token.data.proxy.ProviderRouter
import com.rainy.token.data.proxy.KtorLocalProxyServer
import com.rainy.token.data.proxy.LocalProxyServer
import com.rainy.token.data.proxy.Sub2ApiChatProxy
@ -251,13 +253,31 @@ object NetworkModule {
credentialRepository: CredentialRepository
): TraeChatProxy = TraeChatProxy(okHttpClient, credentialRepository)
@Provides
@Singleton
fun provideProviderRouter(): ProviderRouter = ProviderRouter()
@Provides
@Singleton
fun provideAccountPool(
credentialRepository: CredentialRepository
): AccountPool = AccountPool(credentialRepository)
@Provides
@Singleton
fun provideLocalProxyServer(
sub2ApiChatProxy: Sub2ApiChatProxy,
workBuddyChatProxy: WorkBuddyChatProxy,
traeChatProxy: TraeChatProxy
): LocalProxyServer = KtorLocalProxyServer(sub2ApiChatProxy, workBuddyChatProxy, traeChatProxy)
traeChatProxy: TraeChatProxy,
providerRouter: ProviderRouter,
accountPool: AccountPool
): LocalProxyServer = KtorLocalProxyServer(
sub2ApiChatProxy,
workBuddyChatProxy,
traeChatProxy,
providerRouter,
accountPool
)
/** 余额缓存 DataStore(计划 7.1) */
@Provides