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:
parent
5aa2d8f116
commit
a1605e6993
@ -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)
|
||||
}
|
||||
|
||||
fun registerAll(list: List<PooledAccount>) {
|
||||
list.forEach(::register)
|
||||
PooledAccount(
|
||||
kind = kind,
|
||||
region = region,
|
||||
accountId = accountId,
|
||||
label = accounts.firstOrNull { it.id == accountId }?.label,
|
||||
)
|
||||
}
|
||||
|
||||
suspend fun next(kind: ProviderKind, sessionKey: String? = null): PooledAccount? = mutex.withLock {
|
||||
val list = accounts[kind]?.takeIf { it.isNotEmpty() } ?: return null
|
||||
// TODO: 会话粘性:命中 sessionKey 对应账号时优先返回
|
||||
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
|
||||
}
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
}
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@ -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"
|
||||
|
||||
@ -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"
|
||||
|
||||
@ -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
|
||||
|
||||
Loading…
Reference in New Issue
Block a user