From 94f303ec72bca966d85a3df75315ff6c449dc1c9 Mon Sep 17 00:00:00 2001 From: jie65535 Date: Mon, 3 Aug 2026 16:39:19 +0800 Subject: [PATCH] profile: optimize conversation analysis concurrency --- README.md | 9 +- src/main/kotlin/command/PluginCommands.kt | 11 +- src/main/kotlin/config/PluginConfig.kt | 3 + src/main/kotlin/llm/LargeLanguageModels.kt | 1 + src/main/kotlin/llm/ModelService.kt | 41 +++++- .../kotlin/profile/ProfileHistoryReader.kt | 30 +++- .../profile/UserProfileAnalysisService.kt | 136 +++++++++++++----- src/main/kotlin/profile/UserProfileStore.kt | 35 ++++- src/test/kotlin/llm/ModelServiceTest.kt | 58 ++++++++ .../profile/ProfileHistoryReaderTest.kt | 3 +- .../profile/UserProfileAnalysisServiceTest.kt | 67 +++++++-- .../kotlin/profile/UserProfileStoreTest.kt | 2 + 12 files changed, 328 insertions(+), 68 deletions(-) diff --git a/README.md b/README.md index e94d1ee..6b4c2f1 100644 --- a/README.md +++ b/README.md @@ -105,6 +105,8 @@ profileModelToken: '' profileModel: '' profileModelTemperature: null profileModelExtraBody: '' +# 画像模型同时执行的请求上限,取值1~512;超出后在应用层排队,不计入首块超时 +profileMaxConcurrentRequests: 128 # 留空使用插件自己的聊天库;本地实验可填写外部SQLite历史库的绝对路径 profileHistoryDatabasePath: '' # 异步刷新好友、群、群成员联系人快照,写入 chat-history.sqlite @@ -232,8 +234,11 @@ searchHistoryMaxRecords: 5000 模型只使用单次请求内有效的临时编号,不接触画像条目的内部 UUID;不合规建议会被跳过,不阻断其他有效更新。 临时用户别名不会写入最终画像。`profileCompact` 可独立清理重复或低价值条目,且不推进历史水位线。 旧历史可通过 `profileAnalyze` 或 `profileAnalyzeGroup` 手动分批推进。执行不带参数的 `profileAnalyzeGroup` 时, -插件会从画像历史库枚举所有含有效群消息的群,并为每个群同时启动一个分析任务,不设置额外的应用层并发上限; -显式传入群号时仍只分析指定群。全量模式仅发送启动和最终汇总,逐群进度与结果写入日志,避免回执刷屏。 +插件会从画像历史库枚举所有含有效群消息的群,并在启动任务前排除游标已覆盖最新历史快照的群,只并发推进仍有 +历史待处理的群。画像模型通过应用层信号量按 `profileMaxConcurrentRequests` 限制同时执行的请求数,默认 128; +等待信号量的时间不计入首块响应超时,OkHttp 使用相同上限兜底。群批次只在读取画像快照和提交结果时短暂 +持有用户锁,模型请求在锁外执行;提交前若发现画像已变化,会加载最新版并重新分析,避免共享群友把网络请求 +串行化。显式传入群号时仍只分析指定群。全量模式仅发送启动和最终汇总,逐群进度与结果写入日志,避免回执刷屏。 下一次正常群聊会自动携带触发者和最近发言者的认识。现有好感度、Bot 代号、标签和主观印象会与证据驱动的 长期画像按同一个人合并渲染,并明确给出长期画像条目数;私聊也会携带对方的可靠画像摘要和条目数。 diff --git a/src/main/kotlin/command/PluginCommands.kt b/src/main/kotlin/command/PluginCommands.kt index 67b3904..2c46f86 100644 --- a/src/main/kotlin/command/PluginCommands.kt +++ b/src/main/kotlin/command/PluginCommands.kt @@ -20,6 +20,7 @@ import top.jie65535.mirai.data.PluginData import top.jie65535.mirai.data.SkillStore import top.jie65535.mirai.data.TokenUsageStore import top.jie65535.mirai.llm.LargeLanguageModels +import top.jie65535.mirai.llm.normalizeMaxConcurrentRequests import top.jie65535.mirai.profile.GroupProfileAnalysisReport import top.jie65535.mirai.profile.ProfileAnalysisReport import top.jie65535.mirai.profile.ProfileAutoMaintenance @@ -90,18 +91,22 @@ object PluginCommands : CompositeCommand( require(batches > 0) { "batches 必须是正数" } val analyzeAllGroups = groupIds.isBlank() val parsedGroupIds = if (analyzeAllGroups) { - UserProfileAnalysisService.listHistoryGroupIds() + UserProfileAnalysisService.listPendingHistoryGroupIds() } else { parseProfileGroupIds(groupIds) } if (parsedGroupIds.isEmpty()) { - sendMessage("聊天记录中没有可分析的群消息。") + sendMessage("没有待推进的群画像:历史库中无有效群消息,或所有群均已追平当前快照。") return } val runToken = UserProfileAnalysisService.newRunToken() sendMessage( "已启动 ${parsedGroupIds.size} 个群的画像分析,每群最多推进 $batches 个批次。" + - if (analyzeAllGroups) " 全量模式不设并发上限。" else "" + if (analyzeAllGroups) { + " 画像请求并发上限 ${normalizeMaxConcurrentRequests(PluginConfig.profileMaxConcurrentRequests)}。" + } else { + "" + } ) val completedGroups = AtomicInteger() val successfulGroups = AtomicInteger() diff --git a/src/main/kotlin/config/PluginConfig.kt b/src/main/kotlin/config/PluginConfig.kt index fdc6105..0ce8e4a 100644 --- a/src/main/kotlin/config/PluginConfig.kt +++ b/src/main/kotlin/config/PluginConfig.kt @@ -75,6 +75,9 @@ object PluginConfig : AutoSavePluginConfig("Config") { @ValueDescription("画像分析模型额外请求体JSON。留空时继承聊天模型额外请求体") val profileModelExtraBody: String by value("") + @ValueDescription("画像模型同时执行的请求上限,取值1~512,默认128;应用层排队且不计入首块超时,仅影响画像模型") + val profileMaxConcurrentRequests: Int by value(128) + @ValueDescription("画像分析使用的聊天记录SQLite路径。留空时使用插件自己的chat-history.sqlite;本地实验可填写历史库绝对路径") val profileHistoryDatabasePath: String by value("") diff --git a/src/main/kotlin/llm/LargeLanguageModels.kt b/src/main/kotlin/llm/LargeLanguageModels.kt index abe78ac..1a668aa 100644 --- a/src/main/kotlin/llm/LargeLanguageModels.kt +++ b/src/main/kotlin/llm/LargeLanguageModels.kt @@ -171,6 +171,7 @@ object LargeLanguageModels { timeout = maxOf(timeout, profileFirstChunk), firstChunkTimeout = profileFirstChunk, extraBody = parseExtraBody(extraBody), + maxConcurrentRequests = PluginConfig.profileMaxConcurrentRequests, ), model = model, temperature = PluginConfig.profileModelTemperature ?: PluginConfig.chatTemperature, diff --git a/src/main/kotlin/llm/ModelService.kt b/src/main/kotlin/llm/ModelService.kt index 611c908..704b616 100644 --- a/src/main/kotlin/llm/ModelService.kt +++ b/src/main/kotlin/llm/ModelService.kt @@ -12,10 +12,14 @@ import io.ktor.http.* import io.ktor.utils.io.* import kotlinx.coroutines.currentCoroutineContext import kotlinx.coroutines.flow.Flow +import kotlinx.coroutines.flow.collect import kotlinx.coroutines.flow.flow import kotlinx.coroutines.isActive +import kotlinx.coroutines.sync.Semaphore +import kotlinx.coroutines.sync.withPermit import kotlinx.coroutines.withTimeout import kotlinx.serialization.json.* +import okhttp3.Dispatcher as OkHttpDispatcher import kotlin.time.Duration class ModelService( @@ -23,10 +27,21 @@ class ModelService( val token: String, val timeout: Duration, val firstChunkTimeout: Duration, - val extraBody: JsonObject? = null + val extraBody: JsonObject? = null, + maxConcurrentRequests: Int? = null, ) { + private val maxConcurrentRequests = maxConcurrentRequests?.let(::normalizeMaxConcurrentRequests) + private val requestSemaphore = this.maxConcurrentRequests?.let(::Semaphore) + val httpClient: HttpClient by lazy { HttpClient(OkHttp) { + this@ModelService.maxConcurrentRequests?.let { concurrencyLimit -> + engine { + config { + dispatcher(createRequestDispatcher(concurrencyLimit)) + } + } + } install(HttpTimeout) { // 流式响应的「首 token」与「token 间隔」超时统一由应用层 withTimeout 管控(见 chatCompletions)。 // 这里特意不设 requestTimeoutMillis:否则正常但耗时较长的流式输出会被 Ktor 在中途整体掐断。 @@ -79,7 +94,7 @@ class ModelService( } val body = JsonObject(requestJson).toString() - return flow { + val responseFlow: Flow = flow { // 关键:服务器繁忙时会拖住「响应头」,使 httpClient.post() 自身阻塞在等待响应的阶段, // 因此必须把 post() 连同首个 data 块的读取一起包进 withTimeout。 // 否则首 token 超时永远不会触发(post() 还没返回,根本进不到读取循环), @@ -137,5 +152,27 @@ class ModelService( channel?.cancel() } } + return responseFlow.withConcurrencyLimit(requestSemaphore) + } +} + +internal const val MAX_MODEL_CONCURRENT_REQUESTS = 512 + +internal fun normalizeMaxConcurrentRequests(value: Int): Int = + value.coerceIn(1, MAX_MODEL_CONCURRENT_REQUESTS) + +internal fun createRequestDispatcher(maxConcurrentRequests: Int): OkHttpDispatcher = + OkHttpDispatcher().apply { + val concurrencyLimit = normalizeMaxConcurrentRequests(maxConcurrentRequests) + maxRequests = concurrencyLimit + maxRequestsPerHost = concurrencyLimit + } + +internal fun Flow.withConcurrencyLimit(semaphore: Semaphore?): Flow { + semaphore ?: return this + return flow { + semaphore.withPermit { + this@withConcurrencyLimit.collect { value -> emit(value) } + } } } diff --git a/src/main/kotlin/profile/ProfileHistoryReader.kt b/src/main/kotlin/profile/ProfileHistoryReader.kt index cc2b06b..be03a8d 100644 --- a/src/main/kotlin/profile/ProfileHistoryReader.kt +++ b/src/main/kotlin/profile/ProfileHistoryReader.kt @@ -13,7 +13,12 @@ import java.sql.ResultSet class ProfileHistoryReader(private val databaseFile: File) { data class TimeBounds(val startTime: Int, val endTime: Int) - data class GroupTimeBounds(val botId: Long, val startTime: Int, val endTime: Int) + data class GroupTimeBounds( + val botId: Long, + val groupId: Long, + val startTime: Int, + val endTime: Int, + ) private data class Episode( val index: Int, @@ -63,6 +68,7 @@ class ProfileHistoryReader(private val databaseFile: File) { if (!results.next()) return@use null GroupTimeBounds( botId = results.getLong("bot_id"), + groupId = groupId, startTime = results.getInt("min_time"), endTime = results.getInt("max_time").safeNextSecond(), ) @@ -70,20 +76,32 @@ class ProfileHistoryReader(private val databaseFile: File) { } } - fun listGroupIds(): List = openReadConnection().use { connection -> + fun listGroupTimeBounds(): List = openReadConnection().use { connection -> connection.prepareStatement( """ - SELECT target_id, MAX(time) AS last_message_time + SELECT bot_id, target_id, MIN(time) AS min_time, MAX(time) AS max_time FROM message_record WHERE kind = ? AND recalled = 0 AND target_id > 0 - GROUP BY target_id - ORDER BY last_message_time DESC, target_id ASC + GROUP BY bot_id, target_id + ORDER BY max_time DESC, target_id ASC, bot_id ASC """.trimIndent() ).use { statement -> statement.setInt(1, MessageSourceKind.GROUP.ordinal) statement.executeQuery().use { results -> + val seenGroupIds = hashSetOf() buildList { - while (results.next()) add(results.getLong("target_id")) + while (results.next()) { + val groupId = results.getLong("target_id") + if (!seenGroupIds.add(groupId)) continue + add( + GroupTimeBounds( + botId = results.getLong("bot_id"), + groupId = groupId, + startTime = results.getInt("min_time"), + endTime = results.getInt("max_time").safeNextSecond(), + ) + ) + } } } } diff --git a/src/main/kotlin/profile/UserProfileAnalysisService.kt b/src/main/kotlin/profile/UserProfileAnalysisService.kt index bc8bd75..40b8d0d 100644 --- a/src/main/kotlin/profile/UserProfileAnalysisService.kt +++ b/src/main/kotlin/profile/UserProfileAnalysisService.kt @@ -14,6 +14,8 @@ import java.security.MessageDigest import java.util.concurrent.ConcurrentHashMap object UserProfileAnalysisService { + private const val MAX_CONVERSATION_CONFLICT_RETRIES = 3 + private val runningUsers = ConcurrentHashMap.newKeySet() private val runningGroups = ConcurrentHashMap.newKeySet() private val runningCompactions = ConcurrentHashMap.newKeySet() @@ -32,11 +34,19 @@ object UserProfileAnalysisService { return report } - suspend fun listHistoryGroupIds(): List { + suspend fun listPendingHistoryGroupIds(): List { check(PluginConfig.profileEnabled) { "历史用户画像分析未启用" } check(UserProfileStore.isAvailable) { "用户画像数据库不可用" } return withContext(Dispatchers.IO) { - ProfileHistoryReader(resolveHistoryFile()).listGroupIds() + val historyBounds = ProfileHistoryReader(resolveHistoryFile()).listGroupTimeBounds() + val cursors = UserProfileStore.loadGroupCursors() + .associateBy { cursor -> cursor.botId to cursor.groupId } + historyBounds.asSequence() + .filter { bounds -> + isGroupAnalysisPending(bounds, cursors[bounds.botId to bounds.groupId]) + } + .map(ProfileHistoryReader.GroupTimeBounds::groupId) + .toList() } } @@ -274,10 +284,6 @@ object UserProfileAnalysisService { val reader = withContext(Dispatchers.IO) { ProfileHistoryReader(resolveHistoryFile()) } val bounds = withContext(Dispatchers.IO) { reader.findGroupTimeBounds(groupId) } ?: return emptyGroupReport(groupId) - val endpoint = checkNotNull(LargeLanguageModels.profile) { - "画像分析模型未配置,请设置 profileModelApi/profileModelToken,或配置可继承的聊天模型接入点" - } - val model: ConversationProfileModel = ProfileModelClient(endpoint) var cursor = withContext(Dispatchers.IO) { UserProfileStore.loadGroupCursor(bounds.botId, groupId) } ?: GroupProfileCursor( @@ -297,6 +303,12 @@ object UserProfileAnalysisService { var skippedOperations = 0 var totalUsage = ProfileTokenUsage() var caughtUp = cursor.cursorTime >= cursor.snapshotEndTime + val model: ConversationProfileModel by lazy { + val endpoint = checkNotNull(LargeLanguageModels.profile) { + "画像分析模型未配置,请设置 profileModelApi/profileModelToken,或配置可继承的聊天模型接入点" + } + ProfileModelClient(endpoint) + } while (processedBatches < maxBatches && !caughtUp && runGate.canContinue(runToken)) { val batch = withContext(Dispatchers.IO) { @@ -421,20 +433,16 @@ object UserProfileAnalysisService { .filterValues { it >= minAuthoredTextChars.coerceAtLeast(1) } .keys if (eligibleUserIds.isEmpty()) return null - return userLocks.withUserLocks(eligibleUserIds) locked@{ - if (withContext(Dispatchers.IO) { UserProfileStore.isConversationProcessed(batch.inputHash) }) { - return@locked null - } - val profiles = withContext(Dispatchers.IO) { - eligibleUserIds.associateWith { userId -> - UserProfileStore.load(userId) ?: UserProfileSnapshot( - userId = userId, - cursorTime = 0, - snapshotEndTime = 0, - ) - } + var profiles = userLocks.withUserLocks(eligibleUserIds) { + withContext(Dispatchers.IO) { + if (UserProfileStore.isConversationProcessed(batch.inputHash)) null + else loadConversationProfiles(eligibleUserIds) } + } ?: return null + var conflictRetries = 0 + var totalUsage = ProfileTokenUsage() + while (true) { val (result, reductions) = analyzeConversationWithRetry( model = model, profiles = profiles, @@ -444,25 +452,72 @@ object UserProfileAnalysisService { summaryMaxLength = summaryMaxLength, onRetryFailure = onRetryFailure, ) - withContext(Dispatchers.IO) { - UserProfileStore.commitConversation( - reductions = reductions.map { reduction -> reduction to batch.forUser(reduction.profile.userId) }, - usage = result.usage, - ) + totalUsage += result.usage + val commitOutcome = userLocks.withUserLocks(eligibleUserIds) { + withContext(Dispatchers.IO) { + if (UserProfileStore.isConversationProcessed(batch.inputHash)) { + ConversationCommitOutcome.AlreadyProcessed + } else { + val latestProfiles = loadConversationProfiles(eligibleUserIds) + if (hasProfileVersionConflict(profiles, latestProfiles)) { + ConversationCommitOutcome.Conflict(latestProfiles) + } else { + UserProfileStore.commitConversation( + reductions = reductions.map { reduction -> + reduction to batch.forUser(reduction.profile.userId) + }, + usage = totalUsage, + ) + ConversationCommitOutcome.Committed + } + } + } } - onCommittedOperations( - "source=CONVERSATION bot=${batch.botId} group=${batch.groupId} " + - "batch=[${batch.startTime},${batch.endTime})", - reductions, - ) - ConversationProfileAnalysisReport( - analyzedUsers = eligibleUserIds.size, - processedMessages = batch.messages.size, - appliedOperations = reductions.sumOf { it.operations.size }, - skippedOperations = reductions.sumOf { it.skippedOperations.size }, - usage = result.usage, + + when (commitOutcome) { + ConversationCommitOutcome.AlreadyProcessed -> return null + ConversationCommitOutcome.Committed -> { + onCommittedOperations( + "source=CONVERSATION bot=${batch.botId} group=${batch.groupId} " + + "batch=[${batch.startTime},${batch.endTime})", + reductions, + ) + return ConversationProfileAnalysisReport( + analyzedUsers = eligibleUserIds.size, + processedMessages = batch.messages.size, + appliedOperations = reductions.sumOf { it.operations.size }, + skippedOperations = reductions.sumOf { it.skippedOperations.size }, + usage = totalUsage, + ) + } + + is ConversationCommitOutcome.Conflict -> { + conflictRetries++ + if (conflictRetries > MAX_CONVERSATION_CONFLICT_RETRIES) { + throw IllegalStateException( + "群 ${batch.groupId} 会话画像提交连续冲突 $conflictRetries 次,未提交结果" + ) + } + profiles = commitOutcome.latestProfiles + } + } + } + } + + private fun loadConversationProfiles(userIds: Set): Map = + userIds.associateWith { userId -> + UserProfileStore.load(userId) ?: UserProfileSnapshot( + userId = userId, + cursorTime = 0, + snapshotEndTime = 0, ) } + + private fun hasProfileVersionConflict( + expectedProfiles: Map, + latestProfiles: Map, + ): Boolean = expectedProfiles.any { (userId, expected) -> + latestProfiles[userId]?.version != expected.version } private suspend fun analyzeConversationWithRetry( @@ -698,4 +753,17 @@ object UserProfileAnalysisService { if (skipped.size > 8) ";其余 ${skipped.size - 8} 项已省略" else "" ) } + + private sealed class ConversationCommitOutcome { + object Committed : ConversationCommitOutcome() + object AlreadyProcessed : ConversationCommitOutcome() + data class Conflict( + val latestProfiles: Map, + ) : ConversationCommitOutcome() + } } + +internal fun isGroupAnalysisPending( + bounds: ProfileHistoryReader.GroupTimeBounds, + cursor: GroupProfileCursor?, +): Boolean = cursor == null || cursor.cursorTime < maxOf(cursor.snapshotEndTime, bounds.endTime) diff --git a/src/main/kotlin/profile/UserProfileStore.kt b/src/main/kotlin/profile/UserProfileStore.kt index 63de85b..a152584 100644 --- a/src/main/kotlin/profile/UserProfileStore.kt +++ b/src/main/kotlin/profile/UserProfileStore.kt @@ -192,13 +192,26 @@ object UserProfileStore { statement.setLong(2, groupId) statement.executeQuery().use { results -> if (!results.next()) return@use null - GroupProfileCursor( - botId = results.getLong("bot_id"), - groupId = results.getLong("group_id"), - cursorTime = results.getInt("cursor_time"), - snapshotEndTime = results.getInt("snapshot_end_time"), - updatedAt = results.getLong("updated_at"), - ) + results.toGroupProfileCursor() + } + } + } + } + + fun loadGroupCursors(): List { + check(initialized) { "用户画像数据库尚未初始化" } + return openReadConnection().use { connection -> + connection.prepareStatement( + """ + SELECT bot_id, group_id, cursor_time, snapshot_end_time, updated_at + FROM profile_group_cursor + ORDER BY bot_id, group_id + """.trimIndent() + ).use { statement -> + statement.executeQuery().use { results -> + buildList { + while (results.next()) add(results.toGroupProfileCursor()) + } } } } @@ -602,6 +615,14 @@ object UserProfileStore { lastConfirmedAt = getInt("last_confirmed_at"), ) + private fun ResultSet.toGroupProfileCursor() = GroupProfileCursor( + botId = getLong("bot_id"), + groupId = getLong("group_id"), + cursorTime = getInt("cursor_time"), + snapshotEndTime = getInt("snapshot_end_time"), + updatedAt = getLong("updated_at"), + ) + private fun ProfileCategory.toStorageValue(): String = name.lowercase() private fun ProfileConfidence.toStorageValue(): String = name.lowercase() private fun Int.safeNextSecond(): Int = if (this == Int.MAX_VALUE) this else this + 1 diff --git a/src/test/kotlin/llm/ModelServiceTest.kt b/src/test/kotlin/llm/ModelServiceTest.kt index 71fd27b..906082b 100644 --- a/src/test/kotlin/llm/ModelServiceTest.kt +++ b/src/test/kotlin/llm/ModelServiceTest.kt @@ -1,5 +1,15 @@ package top.jie65535.mirai.llm +import kotlinx.coroutines.CompletableDeferred +import kotlinx.coroutines.async +import kotlinx.coroutines.awaitAll +import kotlinx.coroutines.coroutineScope +import kotlinx.coroutines.flow.collect +import kotlinx.coroutines.flow.flow +import kotlinx.coroutines.runBlocking +import kotlinx.coroutines.sync.Semaphore +import kotlinx.coroutines.withTimeout +import java.util.concurrent.atomic.AtomicInteger import kotlin.test.Test import kotlin.test.assertEquals import kotlin.test.assertNull @@ -35,4 +45,52 @@ class ModelServiceTest { fun ignoresResponsesWithoutCacheDetails() { assertNull(service.extractCacheUsage("""{"usage":{"prompt_tokens":100}}""")) } + + @Test + fun configuresGlobalAndPerHostConcurrencyTogether() { + val dispatcher = createRequestDispatcher(128) + + assertEquals(128, dispatcher.maxRequests) + assertEquals(128, dispatcher.maxRequestsPerHost) + } + + @Test + fun clampsConcurrencyToSafetyRange() { + assertEquals(1, normalizeMaxConcurrentRequests(0)) + assertEquals(MAX_MODEL_CONCURRENT_REQUESTS, normalizeMaxConcurrentRequests(Int.MAX_VALUE)) + } + + @Test + fun queuesFlowsBeforeStartingWork() = runBlocking { + val active = AtomicInteger() + val maximumActive = AtomicInteger() + val started = AtomicInteger() + val firstWaveStarted = CompletableDeferred() + val release = CompletableDeferred() + val source = flow { + val currentActive = active.incrementAndGet() + maximumActive.accumulateAndGet(currentActive, ::maxOf) + if (started.incrementAndGet() == 2) firstWaveStarted.complete(Unit) + try { + release.await() + emit(Unit) + } finally { + active.decrementAndGet() + } + } + val semaphore = Semaphore(2) + + coroutineScope { + val collectors = (1..6).map { + async { source.withConcurrencyLimit(semaphore).collect() } + } + withTimeout(1_000) { firstWaveStarted.await() } + assertEquals(2, active.get()) + release.complete(Unit) + collectors.awaitAll() + } + + assertEquals(2, maximumActive.get()) + assertEquals(6, started.get()) + } } diff --git a/src/test/kotlin/profile/ProfileHistoryReaderTest.kt b/src/test/kotlin/profile/ProfileHistoryReaderTest.kt index da164fc..80aacf2 100644 --- a/src/test/kotlin/profile/ProfileHistoryReaderTest.kt +++ b/src/test/kotlin/profile/ProfileHistoryReaderTest.kt @@ -68,7 +68,7 @@ class ProfileHistoryReaderTest { } val reader = ProfileHistoryReader(database.toFile()) - assertEquals(listOf(10L, 20L), reader.listGroupIds()) + assertEquals(listOf(10L, 20L), reader.listGroupTimeBounds().map { it.groupId }) val bounds = assertNotNull(reader.findUserTimeBounds(TARGET)) assertEquals(100, bounds.startTime) assertEquals(201, bounds.endTime) @@ -125,6 +125,7 @@ class ProfileHistoryReaderTest { val groupBounds = assertNotNull(reader.findGroupTimeBounds(10)) assertEquals(1, groupBounds.botId) + assertEquals(10L, groupBounds.groupId) assertEquals(100, groupBounds.startTime) assertEquals(201, groupBounds.endTime) val firstGroupBatch = assertNotNull( diff --git a/src/test/kotlin/profile/UserProfileAnalysisServiceTest.kt b/src/test/kotlin/profile/UserProfileAnalysisServiceTest.kt index d2a2ef4..a4262d9 100644 --- a/src/test/kotlin/profile/UserProfileAnalysisServiceTest.kt +++ b/src/test/kotlin/profile/UserProfileAnalysisServiceTest.kt @@ -5,7 +5,6 @@ import kotlinx.coroutines.async import kotlinx.coroutines.coroutineScope import kotlinx.coroutines.runBlocking import kotlinx.coroutines.withTimeout -import kotlinx.coroutines.withTimeoutOrNull import net.mamoe.mirai.message.data.MessageSourceKind import top.jie65535.mirai.data.ChatMessageRecord import java.io.IOException @@ -20,6 +19,36 @@ import kotlin.test.assertNull import kotlin.test.assertTrue class UserProfileAnalysisServiceTest { + @Test + fun onlySchedulesGroupsWhoseCursorDoesNotCoverLatestHistory() { + val bounds = ProfileHistoryReader.GroupTimeBounds( + botId = BOT, + groupId = GROUP, + startTime = 100, + endTime = 500, + ) + + assertTrue(isGroupAnalysisPending(bounds, null)) + assertTrue( + isGroupAnalysisPending( + bounds, + GroupProfileCursor(BOT, GROUP, cursorTime = 300, snapshotEndTime = 500), + ) + ) + assertTrue( + isGroupAnalysisPending( + bounds, + GroupProfileCursor(BOT, GROUP, cursorTime = 400, snapshotEndTime = 400), + ) + ) + assertFalse( + isGroupAnalysisPending( + bounds, + GroupProfileCursor(BOT, GROUP, cursorTime = 500, snapshotEndTime = 500), + ) + ) + } + @Test fun analyzesAllEligibleUsersWithOneModelCall() = withProfileStore { val model = FakeConversationProfileModel { successfulResult() } @@ -120,11 +149,13 @@ class UserProfileAnalysisServiceTest { } @Test - fun sharedUserBatchWaitsAndLoadsLatestCommittedProfile() = withProfileStore { + fun sharedUserRequestsRunConcurrentlyAndConflictReloadsLatestProfile() = withProfileStore { val firstEntered = CompletableDeferred() val releaseFirst = CompletableDeferred() - val secondStarted = CompletableDeferred() - val secondEntered = CompletableDeferred() + val secondFirstEntered = CompletableDeferred() + val releaseSecondFirst = CompletableDeferred() + val secondRetried = CompletableDeferred() + val secondCalls = AtomicInteger() val firstModel = InspectingConversationProfileModel { profiles -> assertEquals(0, profiles.getValue(USER_A).version) firstEntered.complete(Unit) @@ -132,26 +163,36 @@ class UserProfileAnalysisServiceTest { result(responseFor("U1", "第一群归纳的信息", evidenceRef = 1)) } val secondModel = InspectingConversationProfileModel { profiles -> - assertEquals(1, profiles.getValue(USER_A).version) - secondEntered.complete(Unit) + when (secondCalls.incrementAndGet()) { + 1 -> { + assertEquals(0, profiles.getValue(USER_A).version) + secondFirstEntered.complete(Unit) + releaseSecondFirst.await() + } + + 2 -> { + assertEquals(1, profiles.getValue(USER_A).version) + secondRetried.complete(Unit) + } + + else -> error("共享用户冲突只应触发一次重算") + } result(responseFor("U1", "第二群归纳的信息", evidenceRef = 1)) } coroutineScope { val first = async { analyze(singleUserBatch(USER_A, 10, "shared-a"), firstModel) } firstEntered.await() - val second = async { - secondStarted.complete(Unit) - analyze(singleUserBatch(USER_A, 20, "shared-b"), secondModel) - } - secondStarted.await() - assertNull(withTimeoutOrNull(100) { secondEntered.await() }) + val second = async { analyze(singleUserBatch(USER_A, 20, "shared-b"), secondModel) } + withTimeout(1_000) { secondFirstEntered.await() } releaseFirst.complete(Unit) assertNotNull(first.await()) - withTimeout(1_000) { secondEntered.await() } + releaseSecondFirst.complete(Unit) + withTimeout(1_000) { secondRetried.await() } assertNotNull(second.await()) } + assertEquals(2, secondCalls.get()) val profile = assertNotNull(UserProfileStore.load(USER_A)) assertEquals(2, profile.version) assertEquals(2, profile.items.size) diff --git a/src/test/kotlin/profile/UserProfileStoreTest.kt b/src/test/kotlin/profile/UserProfileStoreTest.kt index db35845..402a230 100644 --- a/src/test/kotlin/profile/UserProfileStoreTest.kt +++ b/src/test/kotlin/profile/UserProfileStoreTest.kt @@ -175,10 +175,12 @@ class UserProfileStoreTest { UserProfileStore.saveGroupCursor(cursor) assertEquals(cursor, UserProfileStore.loadGroupCursor(1, 300)) + assertEquals(listOf(cursor), UserProfileStore.loadGroupCursors()) val advanced = cursor.copy(cursorTime = 400, updatedAt = 456) UserProfileStore.saveGroupCursor(advanced) assertEquals(advanced, UserProfileStore.loadGroupCursor(1, 300)) + assertEquals(listOf(advanced), UserProfileStore.loadGroupCursors()) } finally { UserProfileStore.close() directory.toFile().deleteRecursively()