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
|
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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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 }
|
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"
|
||||||
|
|||||||
@ -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"
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user