diff --git a/app/src/main/java/com/rainy/token/data/proxy/AccountPool.kt b/app/src/main/java/com/rainy/token/data/proxy/AccountPool.kt index ae71135..4de1876 100644 --- a/app/src/main/java/com/rainy/token/data/proxy/AccountPool.kt +++ b/app/src/main/java/com/rainy/token/data/proxy/AccountPool.kt @@ -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>() + private val sessionSticky = mutableMapOf() private val cursor = mutableMapOf() - 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) { - 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 { 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> = 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 } } diff --git a/app/src/main/java/com/rainy/token/data/proxy/KtorLocalProxyServer.kt b/app/src/main/java/com/rainy/token/data/proxy/KtorLocalProxyServer.kt index 4e685ec..778fa75 100644 --- a/app/src/main/java/com/rainy/token/data/proxy/KtorLocalProxyServer.kt +++ b/app/src/main/java/com/rainy/token/data/proxy/KtorLocalProxyServer.kt @@ -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 diff --git a/app/src/main/java/com/rainy/token/data/proxy/ProviderRouter.kt b/app/src/main/java/com/rainy/token/data/proxy/ProviderRouter.kt new file mode 100644 index 0000000..126de77 --- /dev/null +++ b/app/src/main/java/com/rainy/token/data/proxy/ProviderRouter.kt @@ -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) + } + } +} diff --git a/app/src/main/java/com/rainy/token/data/proxy/TraeChatProxy.kt b/app/src/main/java/com/rainy/token/data/proxy/TraeChatProxy.kt index be4f8b8..3dd8a19 100644 --- a/app/src/main/java/com/rainy/token/data/proxy/TraeChatProxy.kt +++ b/app/src/main/java/com/rainy/token/data/proxy/TraeChatProxy.kt @@ -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" diff --git a/app/src/main/java/com/rainy/token/data/proxy/WorkBuddyChatProxy.kt b/app/src/main/java/com/rainy/token/data/proxy/WorkBuddyChatProxy.kt index 7ad1071..031b363 100644 --- a/app/src/main/java/com/rainy/token/data/proxy/WorkBuddyChatProxy.kt +++ b/app/src/main/java/com/rainy/token/data/proxy/WorkBuddyChatProxy.kt @@ -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" diff --git a/app/src/main/java/com/rainy/token/di/NetworkModule.kt b/app/src/main/java/com/rainy/token/di/NetworkModule.kt index 755416f..3c3124a 100644 --- a/app/src/main/java/com/rainy/token/di/NetworkModule.kt +++ b/app/src/main/java/com/rainy/token/di/NetworkModule.kt @@ -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