88 lines
3.2 KiB
Kotlin
88 lines
3.2 KiB
Kotlin
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
|
||
|
||
/**
|
||
* 账号池:从 CredentialRepository 实时取号,负责轮询与会话粘性。
|
||
*
|
||
* - 多账号服务(Trae / WorkBuddy / Sub2API)按 ProviderKind 分组;
|
||
* - 会话粘性:同一 conversation 头持续命中同一账号,避免上下文错乱;
|
||
* - 区域匹配:优先选 credential.region 与路由 region 一致的账号,找不到时退回全部账号。
|
||
*/
|
||
@Singleton
|
||
class AccountPool @Inject constructor(
|
||
private val credentialRepository: CredentialRepository,
|
||
) {
|
||
|
||
data class PooledAccount(
|
||
val kind: ProviderKind,
|
||
val region: ProviderRegion? = null,
|
||
val accountId: String? = null,
|
||
val label: String? = null,
|
||
)
|
||
|
||
private val mutex = Mutex()
|
||
private val sessionSticky = mutableMapOf<String, String>()
|
||
private val cursor = mutableMapOf<ProviderKind, Int>()
|
||
|
||
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,
|
||
)
|
||
}
|
||
|
||
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) % ids.size
|
||
return ids[idx]
|
||
}
|
||
|
||
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
|
||
}
|
||
}
|