profile: optimize conversation analysis concurrency

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