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 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.Mutex
import kotlinx.coroutines.sync.withLock import kotlinx.coroutines.sync.withLock
/** /**
* 账号池:按 Provider 保存账号,负责轮询/冷却/会话粘性。 * 账号池:从 CredentialRepository 实时取号,负责轮询与会话粘性。
* *
* 当前为骨架(round-robin)。后续把 workbuddy2api-panel 的 * - 多账号服务(Trae / WorkBuddy / Sub2API)按 ProviderKind 分组;
* 三因子加权、429 熔断冷却、会话粘性搬进来。 * - 会话粘性:同一 conversation 头持续命中同一账号,避免上下文错乱;
* - 区域匹配:优先选 credential.region 与路由 region 一致的账号,找不到时退回全部账号。
*/ */
class AccountPool { @Singleton
class AccountPool @Inject constructor(
private val credentialRepository: CredentialRepository,
) {
data class PooledAccount( data class PooledAccount(
val kind: ProviderKind, val kind: ProviderKind,
@ -19,26 +28,60 @@ class AccountPool {
) )
private val mutex = Mutex() private val mutex = Mutex()
private val accounts = mutableMapOf<ProviderKind, MutableList<PooledAccount>>() private val sessionSticky = mutableMapOf<String, String>()
private val cursor = mutableMapOf<ProviderKind, Int>() private val cursor = mutableMapOf<ProviderKind, Int>()
fun register(account: PooledAccount) { suspend fun next(
accounts.getOrPut(account.kind) { mutableListOf() }.add(account) 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>) { private suspend fun regionMatches(service: ServiceType, accountId: String, region: ProviderRegion): Boolean {
list.forEach(::register) 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 { private fun pickRoundRobin(kind: ProviderKind, ids: List<String>): String {
val list = accounts[kind]?.takeIf { it.isNotEmpty() } ?: return null
// TODO: 会话粘性:命中 sessionKey 对应账号时优先返回
val idx = cursor[kind] ?: 0 val idx = cursor[kind] ?: 0
cursor[kind] = (idx + 1) % list.size cursor[kind] = (idx + 1) % ids.size
list[idx] return ids[idx]
} }
suspend fun snapshot(): Map<ProviderKind, List<PooledAccount>> = mutex.withLock { private fun serviceFor(kind: ProviderKind): ServiceType? = when (kind) {
accounts.mapValues { it.value.toList() } 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 sub2ApiChatProxy: Sub2ApiChatProxy,
private val workBuddyChatProxy: WorkBuddyChatProxy, private val workBuddyChatProxy: WorkBuddyChatProxy,
private val traeChatProxy: TraeChatProxy, private val traeChatProxy: TraeChatProxy,
private val providerRouter: ProviderRouter,
private val accountPool: AccountPool,
) : LocalProxyServer { ) : LocalProxyServer {
private val lock = Any() private val lock = Any()
@ -57,7 +59,7 @@ class KtorLocalProxyServer @Inject constructor(
if (_isRunning.value) return Result.success(Unit) if (_isRunning.value) return Result.success(Unit)
return try { return try {
val engine = embeddedServer(CIO, host = "127.0.0.1", port = config.port) { 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) engine.start(wait = false)
server = engine server = engine
@ -84,6 +86,8 @@ class KtorLocalProxyServer @Inject constructor(
sub2Api: Sub2ApiChatProxy, sub2Api: Sub2ApiChatProxy,
workBuddy: WorkBuddyChatProxy, workBuddy: WorkBuddyChatProxy,
trae: TraeChatProxy, trae: TraeChatProxy,
router: ProviderRouter,
pool: AccountPool,
) { ) {
install(ContentNegotiation) { install(ContentNegotiation) {
json() json()
@ -109,17 +113,24 @@ class KtorLocalProxyServer @Inject constructor(
return@post return@post
} }
val model = extractModel(body) 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 { try {
val result = when (provider) { val result = when (route.kind) {
ChatRoute.WORKBUDDY_CN -> workBuddy.forwardChat(body) ProviderKind.WORKBUDDY_CN, ProviderKind.WORKBUDDY_INTL ->
ChatRoute.TRAE_CN -> trae.forwardChat(body) workBuddy.forwardChat(body, accountId, route.region)
ChatRoute.SUB2API -> sub2Api.forwardChat(body)
ProviderKind.TRAE_CN, ProviderKind.TRAE_INTL ->
trae.forwardChat(body, accountId, route.region)
else -> sub2Api.forwardChat(body, accountId)
} }
if (result == null) { if (result == null) {
call.respond( call.respond(
HttpStatusCode.BadRequest, HttpStatusCode.BadRequest,
errorBody("${provider.displayName} 未配置或未登录,请先在设置中配置") errorBody("${route.kind.displayName} 未配置或未登录,请先在设置中配置")
) )
} else { } else {
call.respondBytes(result.body, contentTypeOf(result.contentType), HttpStatusCode(result.status, "")) 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 } ?.let { if (it is kotlinx.serialization.json.JsonPrimitive) it.content else null }
}.getOrNull() ?: "" }.getOrNull() ?: ""
private fun routeProvider(model: String): ChatRoute { private fun extractUser(body: String): String? = runCatching {
val m = model.lowercase() val el = kotlinx.serialization.json.Json.parseToJsonElement(body)
return when { (el as? kotlinx.serialization.json.JsonObject)
m.startsWith("wbcn-") -> ChatRoute.WORKBUDDY_CN ?.get("user")
m.startsWith("traecn-") || m.startsWith("trae-") -> ChatRoute.TRAE_CN ?.let { if (it is kotlinx.serialization.json.JsonPrimitive) it.content else null }
else -> ChatRoute.SUB2API }.getOrNull()
}
}
private enum class ChatRoute(val displayName: String) {
WORKBUDDY_CN("WorkBuddy 国内版"),
TRAE_CN("Trae CN"),
SUB2API("Sub2API"),
}
companion object { companion object {
private const val MAX_REQUEST_BYTES = 10 * 1024 * 1024 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 } 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) { withContext(Dispatchers.IO) {
val credential = credentialRepository.get(ServiceType.TRAE, accountId) val credential = credentialRepository.get(ServiceType.TRAE, accountId)
?: return@withContext null ?: return@withContext null
@ -43,7 +47,10 @@ class TraeChatProxy @Inject constructor(
val jwt = credential.jwt.trim().takeIf { it.isNotBlank() } val jwt = credential.jwt.trim().takeIf { it.isNotBlank() }
?: return@withContext null ?: 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" "https://grow-normal.trae.ai"
} else { } else {
"https://api.trae.cn" "https://api.trae.cn"

View File

@ -25,7 +25,11 @@ class WorkBuddyChatProxy @Inject constructor(
private val credentialRepository: CredentialRepository, 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) { withContext(Dispatchers.IO) {
val credential = credentialRepository.get(ServiceType.WORKBUDDY, accountId) val credential = credentialRepository.get(ServiceType.WORKBUDDY, accountId)
?: return@withContext null ?: return@withContext null
@ -33,7 +37,10 @@ class WorkBuddyChatProxy @Inject constructor(
val accessToken = credential.accessToken.trim().takeIf { it.isNotBlank() } val accessToken = credential.accessToken.trim().takeIf { it.isNotBlank() }
?: return@withContext null ?: 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" "https://www.workbuddy.ai"
} else { } else {
"https://copilot.tencent.com" "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.Sub2ApiRepository
import com.rainy.token.data.repository.TraeRepository import com.rainy.token.data.repository.TraeRepository
import com.rainy.token.data.repository.UpdateRepository 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.KtorLocalProxyServer
import com.rainy.token.data.proxy.LocalProxyServer import com.rainy.token.data.proxy.LocalProxyServer
import com.rainy.token.data.proxy.Sub2ApiChatProxy import com.rainy.token.data.proxy.Sub2ApiChatProxy
@ -251,13 +253,31 @@ object NetworkModule {
credentialRepository: CredentialRepository credentialRepository: CredentialRepository
): TraeChatProxy = TraeChatProxy(okHttpClient, credentialRepository) ): TraeChatProxy = TraeChatProxy(okHttpClient, credentialRepository)
@Provides
@Singleton
fun provideProviderRouter(): ProviderRouter = ProviderRouter()
@Provides
@Singleton
fun provideAccountPool(
credentialRepository: CredentialRepository
): AccountPool = AccountPool(credentialRepository)
@Provides @Provides
@Singleton @Singleton
fun provideLocalProxyServer( fun provideLocalProxyServer(
sub2ApiChatProxy: Sub2ApiChatProxy, sub2ApiChatProxy: Sub2ApiChatProxy,
workBuddyChatProxy: WorkBuddyChatProxy, workBuddyChatProxy: WorkBuddyChatProxy,
traeChatProxy: TraeChatProxy traeChatProxy: TraeChatProxy,
): LocalProxyServer = KtorLocalProxyServer(sub2ApiChatProxy, workBuddyChatProxy, traeChatProxy) providerRouter: ProviderRouter,
accountPool: AccountPool
): LocalProxyServer = KtorLocalProxyServer(
sub2ApiChatProxy,
workBuddyChatProxy,
traeChatProxy,
providerRouter,
accountPool
)
/** 余额缓存 DataStore(计划 7.1) */ /** 余额缓存 DataStore(计划 7.1) */
@Provides @Provides