diff --git a/src/main/kotlin/JChatGPT.kt b/src/main/kotlin/JChatGPT.kt index 5c68c75..89c3e67 100644 --- a/src/main/kotlin/JChatGPT.kt +++ b/src/main/kotlin/JChatGPT.kt @@ -28,6 +28,8 @@ import net.mamoe.mirai.message.data.content import net.mamoe.mirai.utils.info import top.jie65535.mirai.command.PluginCommands import top.jie65535.mirai.config.PluginConfig +import top.jie65535.mirai.config.ModelConfig +import top.jie65535.mirai.config.ModelConfigMigration import top.jie65535.mirai.conversation.ConversationContext import top.jie65535.mirai.conversation.ConversationEngine import top.jie65535.mirai.data.ChatHistoryStore @@ -62,8 +64,19 @@ object JChatGPT : KotlinPlugin( override fun onEnable() { PermissionService.INSTANCE.register(chatPermission, "JChatGPT Chat Permission") PluginConfig.reload() + ModelConfig.reload() + ModelConfigMigration.migrateLoadedConfig().takeIf { it.changed }?.let { migration -> + logger.info( + "已将旧模型配置自动迁移到 Models.yml:新增 ${migration.addedProviders} 个提供商、" + + "${migration.addedModels} 个模型别名" + ) + } PluginData.reload() - TokenUsageStore.init(dataFolder) + runCatching { + TokenUsageStore.init(dataFolder) { message, cause -> + if (cause == null) logger.warning(message) else logger.warning(message, cause) + } + }.onFailure { logger.error("初始化 SQLite 模型用量记录失败,用量统计将暂时禁用", it) } SkillStore.init(dataFolder) includeHistory = try { @@ -139,6 +152,7 @@ object JChatGPT : KotlinPlugin( ContactSnapshotRefresher.clear() UserProfileStore.close() ContactSnapshotStore.close() + TokenUsageStore.close() ChatHistoryStore.close() } diff --git a/src/main/kotlin/command/PluginCommands.kt b/src/main/kotlin/command/PluginCommands.kt index 7196388..fd70514 100644 --- a/src/main/kotlin/command/PluginCommands.kt +++ b/src/main/kotlin/command/PluginCommands.kt @@ -16,6 +16,8 @@ import net.mamoe.mirai.contact.User import top.jie65535.mirai.JChatGPT import top.jie65535.mirai.JChatGPT.reload import top.jie65535.mirai.config.PluginConfig +import top.jie65535.mirai.config.ModelConfig +import top.jie65535.mirai.config.ModelConfigMigration import top.jie65535.mirai.data.PluginData import top.jie65535.mirai.data.SkillStore import top.jie65535.mirai.data.TokenUsageStore @@ -45,6 +47,8 @@ object PluginCommands : CompositeCommand( @SubCommand suspend fun CommandSender.reload() { PluginConfig.reload() + ModelConfig.reload() + ModelConfigMigration.migrateLoadedConfig() PluginData.reload() LargeLanguageModels.reload() ProfileDailyMaintenance.reload() @@ -371,89 +375,93 @@ object PluginCommands : CompositeCommand( suspend fun CommandSender.tokens(days: Int = 7) { validateDays(days) - if (TokenUsageStore.all.isEmpty()) { - sendMessage("暂无 Token 使用记录") + if (!TokenUsageStore.isAvailable) { + sendMessage("Token SQLite 尚未初始化") return } val cutoff = calculateCutoffDate(days) - val today = LocalDate.now().toString() - - val windowed = TokenUsageStore.all.filter { it.date >= cutoff } - if (windowed.isEmpty()) { - sendMessage("最近 $days 天无 Token 使用记录") + val summary = runCatching { TokenUsageStore.summary(cutoff, rankingLimit = 100) } + .getOrElse { + sendMessage("读取 Token 使用记录失败:${it.message ?: it::class.simpleName}") + return + } + if (!TokenUsageStore.hasAny(cutoff)) { + sendMessage("暂无 Token 使用记录") return } - // 窗口汇总 - var prompt = 0L; var completion = 0L; var total = 0L; var cached = 0L - var calls = 0; var todayTotal = 0L - val users = HashSet() - for (r in windowed) { - prompt += r.promptTokens - completion += r.completionTokens - total += r.totalTokens - cached += r.cachedTokens - calls += r.callCount - users.add(r.userId) - if (r.date == today) todayTotal += r.totalTokens - } - val hitRate = if (prompt > 0) cached * 100.0 / prompt else 0.0 - - // 每日趋势 - val daily = windowed.groupBy { it.date } - .mapValues { (_, rs) -> rs.sumOf { it.totalTokens } } - .toSortedMap() - - // Top 用户 - val topUsers = windowed.groupBy { it.userId } - .map { (_, rs) -> - val name = rs.maxByOrNull { it.date }!!.userNickname - name to rs.sumOf { it.totalTokens } - } - .sortedByDescending { it.second } - .take(TOP_LIMIT) - - // Top 群组:只显示群名,绝不暴露群号(避免被误判宣群) - val topGroups = windowed.filter { it.groupId != null } - .groupBy { it.groupId!! } - .map { (gid, rs) -> - val name = rs.firstNotNullOfOrNull { r -> r.groupName?.takeIf { it.isNotBlank() } } - ?: resolveGroupName(gid) - name to rs.sumOf { it.totalTokens } - } - .sortedByDescending { it.second } - .take(TOP_LIMIT) + val hitRate = if (summary.promptTokens > 0) { + summary.cachedTokens * 100.0 / summary.promptTokens + } else 0.0 val response = buildString { appendLine("📊 Token 简报 · 最近 $days 天") appendLine() - appendLine("输入 ${formatCompact(prompt)}(缓存命中 ${"%.1f".format(hitRate)}%,省 ${formatCompact(cached)})") - appendLine("输出 ${formatCompact(completion)}") - appendLine("总计 ${formatCompact(total)} | 调用 ${formatNumber(calls)} 次 | 活跃 ${users.size} 人") - appendLine("今日 ${formatCompact(todayTotal)}") + appendLine("输入 ${formatCompact(summary.promptTokens)}(缓存命中 ${"%.1f".format(hitRate)}%,省 ${formatCompact(summary.cachedTokens)})") + appendLine("输出 ${formatCompact(summary.completionTokens)}") + appendLine("总计 ${formatCompact(summary.totalTokens)} | 调用 ${formatNumber(summary.callCount)} 次 | 活跃 ${summary.activeUsers} 人") + if (summary.allCallCount != summary.callCount) { + appendLine("全部模型调用 ${formatNumber(summary.allCallCount)} 次") + } + appendLine("今日 ${formatCompact(summary.todayTotal)}") - if (daily.size > 1) { + if (summary.daily.size > 1) { appendLine() appendLine("📈 每日趋势") - daily.forEach { (date, t) -> - appendLine(" ${date.substring(5)} ${formatCompact(t)}") + summary.daily.forEach { daily -> + appendLine(" ${daily.date.substring(5)} ${formatCompact(daily.totalTokens)}") } } - if (topUsers.isNotEmpty()) { + if (summary.topUsers.isNotEmpty()) { appendLine() appendLine("👤 Top 用户") - topUsers.forEachIndexed { i, (name, t) -> - appendLine(" ${i + 1}. $name ${formatCompact(t)}") + summary.topUsers.forEachIndexed { i, ranking -> + appendLine(" ${i + 1}. ${ranking.name.ifBlank { ranking.id.toString() }} ${formatCompact(ranking.totalTokens)}") } } - if (topGroups.isNotEmpty()) { + if (summary.topGroups.isNotEmpty()) { appendLine() appendLine("👥 Top 群组") - topGroups.forEachIndexed { i, (name, t) -> - appendLine(" ${i + 1}. $name ${formatCompact(t)}") + summary.topGroups.forEachIndexed { i, ranking -> + val name = ranking.name.ifBlank { resolveGroupName(ranking.id) } + appendLine(" ${i + 1}. $name ${formatCompact(ranking.totalTokens)}") + } + } + + if (summary.models.size > 1) { + appendLine() + appendLine("🤖 模型") + summary.models.take(TOP_LIMIT).forEach { model -> + appendLine(" ${model.provider}/${model.model} ${formatCompact(model.totalTokens)}") + } + } + + val tokenUsageByKind = summary.breakdown.asSequence() + .filter { it.unit == "tokens" } + .groupBy { it.usageKind } + .mapValues { (_, usage) -> usage.sumOf { it.totalUnits } } + .entries + .sortedByDescending { it.value } + if (tokenUsageByKind.size > 1 || tokenUsageByKind.firstOrNull()?.key != "chat") { + appendLine() + appendLine("Token 用途") + tokenUsageByKind.take(TOP_LIMIT).forEach { (kind, total) -> + appendLine(" $kind ${formatCompact(total)}") + } + } + + val otherUsage = summary.breakdown.filter { it.unit != "tokens" } + if (otherUsage.isNotEmpty()) { + appendLine() + appendLine("其他模型用量") + otherUsage.take(TOP_LIMIT).forEach { usage -> + appendLine( + " ${usage.provider}/${usage.model} ${usage.usageKind} " + + "${formatCompact(usage.totalUnits)} ${usage.unit}" + ) } } } diff --git a/src/main/kotlin/config/ModelConfig.kt b/src/main/kotlin/config/ModelConfig.kt new file mode 100644 index 0000000..2b11f7f --- /dev/null +++ b/src/main/kotlin/config/ModelConfig.kt @@ -0,0 +1,59 @@ +package top.jie65535.mirai.config + +import kotlinx.serialization.Serializable +import net.mamoe.mirai.console.data.AutoSavePluginConfig +import net.mamoe.mirai.console.data.ValueDescription +import net.mamoe.mirai.console.data.value + +/** A credential/API definition shared by one or more model aliases. */ +@Serializable +data class ModelProviderDefinition( + val name: String = "", + val type: String = "openai", + val api: String = "", + val token: String = "", +) + +/** A model alias used by role bindings in [PluginConfig]. */ +@Serializable +data class ModelDefinition( + val name: String = "", + val provider: String = "", + val model: String = "", + val extraBody: String = "", +) + +/** + * Shared model catalog. Credentials live here so roles can reference one alias + * without duplicating API URLs and keys in the main plugin configuration. + */ +object ModelConfig : AutoSavePluginConfig("Models") { + @ValueDescription( + """模型提供商与凭据;type 可用 openai 或 dashscope,token 对应 API Token/Key。 +新安装可参考: +- name: deepseek + type: openai + api: 'https://api.deepseek.com/v1/' + token: 'sk-xxxx' +- name: dashscope-native + type: dashscope + api: '' + token: 'sk-xxxx'""" + ) + var providers: List by value() + + @ValueDescription( + """模型别名;每项通过 provider 绑定一个提供商,并填写实际模型名。 +配置后还需在 Config.yml 将 chatModelAlias 等用途字段设为对应 name。 +示例: +- name: chat-main + provider: deepseek + model: deepseek-chat + extraBody: '' +- name: image-main + provider: dashscope-native + model: qwen-image-2.0 + extraBody: ''""" + ) + var models: List by value() +} diff --git a/src/main/kotlin/config/ModelConfigMigration.kt b/src/main/kotlin/config/ModelConfigMigration.kt new file mode 100644 index 0000000..4bc59ee --- /dev/null +++ b/src/main/kotlin/config/ModelConfigMigration.kt @@ -0,0 +1,254 @@ +package top.jie65535.mirai.config + +import java.net.URI + +internal data class ModelRoleBindings( + val chat: String = "", + val chatFallbacks: List = emptyList(), + val profile: String = "", + val reasoning: String = "", + val visual: String = "", + val webSummary: String = "", + val image: String = "", + val tts: String = "", +) + +internal data class LegacyOpenAiModel( + val api: String, + val token: String, + val model: String, + val extraBody: String = "", +) + +internal data class LegacyModelSettings( + val chat: LegacyOpenAiModel, + val chatFallbacks: List, + val profile: LegacyOpenAiModel, + val reasoning: LegacyOpenAiModel, + val visual: LegacyOpenAiModel, + val webSummary: LegacyOpenAiModel, + val dashScopeToken: String, + val imageModel: String, + val ttsModel: String, +) + +internal data class ModelConfigMigrationResult( + val providers: List, + val models: List, + val bindings: ModelRoleBindings, + val addedProviders: Int, + val addedModels: Int, + val bindingsChanged: Boolean, +) { + val changed: Boolean + get() = addedProviders > 0 || addedModels > 0 || bindingsChanged +} + +internal object ModelConfigMigration { + fun migrateLoadedConfig(): ModelConfigMigrationResult { + val currentBindings = ModelRoleBindings( + chat = PluginConfig.chatModelAlias, + chatFallbacks = PluginConfig.chatFallbackModelAliases, + profile = PluginConfig.profileModelAlias, + reasoning = PluginConfig.reasoningModelAlias, + visual = PluginConfig.visualModelAlias, + webSummary = PluginConfig.webSummaryModelAlias, + image = PluginConfig.imageModelAlias, + tts = PluginConfig.ttsModelAlias, + ) + val chat = LegacyOpenAiModel( + api = PluginConfig.openAiApi, + token = PluginConfig.openAiToken, + model = PluginConfig.chatModel, + extraBody = PluginConfig.chatModelExtraBody, + ) + val legacy = LegacyModelSettings( + chat = chat, + chatFallbacks = PluginConfig.chatFallbacks.map { fallback -> + LegacyOpenAiModel( + api = fallback.api.ifBlank { chat.api }, + token = fallback.token.ifBlank { chat.token }, + model = fallback.model.ifBlank { chat.model }, + extraBody = fallback.extraBody.ifBlank { chat.extraBody }, + ) + }, + profile = LegacyOpenAiModel( + api = PluginConfig.profileModelApi.ifBlank { chat.api }, + token = PluginConfig.profileModelToken.ifBlank { chat.token }, + model = PluginConfig.profileModel.ifBlank { chat.model }, + extraBody = PluginConfig.profileModelExtraBody.ifBlank { chat.extraBody }, + ), + reasoning = LegacyOpenAiModel( + api = PluginConfig.reasoningModelApi, + token = PluginConfig.reasoningModelToken, + model = PluginConfig.reasoningModel, + extraBody = PluginConfig.reasoningModelExtraBody, + ), + visual = LegacyOpenAiModel( + api = PluginConfig.visualModelApi, + token = PluginConfig.visualModelToken, + model = PluginConfig.visualModel, + extraBody = PluginConfig.visualModelExtraBody, + ), + webSummary = LegacyOpenAiModel( + api = PluginConfig.webSummaryModelApi, + token = PluginConfig.webSummaryModelToken, + model = PluginConfig.webSummaryModel, + extraBody = PluginConfig.webSummaryModelExtraBody, + ), + dashScopeToken = PluginConfig.dashScopeApiKey, + imageModel = PluginConfig.imageModel, + ttsModel = PluginConfig.ttsModel, + ) + val result = migrate(ModelConfig.providers, ModelConfig.models, currentBindings, legacy) + if (result.providers != ModelConfig.providers) ModelConfig.providers = result.providers + if (result.models != ModelConfig.models) ModelConfig.models = result.models + if (result.bindings.chat != PluginConfig.chatModelAlias) PluginConfig.chatModelAlias = result.bindings.chat + if (result.bindings.chatFallbacks != PluginConfig.chatFallbackModelAliases) { + PluginConfig.chatFallbackModelAliases = result.bindings.chatFallbacks + } + if (result.bindings.profile != PluginConfig.profileModelAlias) PluginConfig.profileModelAlias = result.bindings.profile + if (result.bindings.reasoning != PluginConfig.reasoningModelAlias) { + PluginConfig.reasoningModelAlias = result.bindings.reasoning + } + if (result.bindings.visual != PluginConfig.visualModelAlias) PluginConfig.visualModelAlias = result.bindings.visual + if (result.bindings.webSummary != PluginConfig.webSummaryModelAlias) { + PluginConfig.webSummaryModelAlias = result.bindings.webSummary + } + if (result.bindings.image != PluginConfig.imageModelAlias) PluginConfig.imageModelAlias = result.bindings.image + if (result.bindings.tts != PluginConfig.ttsModelAlias) PluginConfig.ttsModelAlias = result.bindings.tts + return result + } + + fun migrate( + existingProviders: List, + existingModels: List, + bindings: ModelRoleBindings, + legacy: LegacyModelSettings, + ): ModelConfigMigrationResult { + val providers = existingProviders.toMutableList() + val models = existingModels.toMutableList() + val initialProviderCount = providers.size + val initialModelCount = models.size + + fun bindOpenAi(current: String, preferredAlias: String, legacyModel: LegacyOpenAiModel): String = + current.ifBlank { + addModel(providers, models, preferredAlias, "openai", legacyModel) + } + + fun bindDashScope(current: String, preferredAlias: String, model: String): String = + current.ifBlank { + addModel( + providers = providers, + models = models, + preferredAlias = preferredAlias, + providerType = "dashscope", + legacyModel = LegacyOpenAiModel("", legacy.dashScopeToken, model), + ) + } + + val chat = bindOpenAi(bindings.chat, "chat-main", legacy.chat) + val chatFallbacks = if (bindings.chatFallbacks.isNotEmpty()) { + bindings.chatFallbacks + } else { + legacy.chatFallbacks.mapIndexedNotNull { index, fallback -> + bindOpenAi("", "chat-fallback-${index + 1}", fallback).takeIf(String::isNotBlank) + } + } + val migratedBindings = ModelRoleBindings( + chat = chat, + chatFallbacks = chatFallbacks, + profile = bindOpenAi(bindings.profile, "profile-main", legacy.profile), + reasoning = bindOpenAi(bindings.reasoning, "reasoning-main", legacy.reasoning), + visual = bindOpenAi(bindings.visual, "visual-main", legacy.visual), + webSummary = bindOpenAi(bindings.webSummary, "web-summary-main", legacy.webSummary), + image = bindDashScope(bindings.image, "image-main", legacy.imageModel), + tts = bindDashScope(bindings.tts, "tts-main", legacy.ttsModel), + ) + return ModelConfigMigrationResult( + providers = providers, + models = models, + bindings = migratedBindings, + addedProviders = providers.size - initialProviderCount, + addedModels = models.size - initialModelCount, + bindingsChanged = migratedBindings != bindings, + ) + } + + private fun addModel( + providers: MutableList, + models: MutableList, + preferredAlias: String, + providerType: String, + legacyModel: LegacyOpenAiModel, + ): String { + val token = legacyModel.token.trim() + val modelName = legacyModel.model.trim() + val api = legacyModel.api.trim() + val extraBody = legacyModel.extraBody.trim() + if (token.isEmpty() || modelName.isEmpty() || providerType == "openai" && api.isEmpty()) return "" + + val providerName = findOrAddProvider(providers, providerType, api, token) + val reusable = models.firstOrNull { candidate -> + candidate.name.isNotBlank() && + models.count { it.name.trim() == candidate.name.trim() } == 1 && + candidate.provider.trim() == providerName && + candidate.model.trim() == modelName && + candidate.extraBody.trim() == extraBody + } + if (reusable != null) return reusable.name.trim() + + val alias = uniqueName(preferredAlias, models.mapTo(HashSet()) { it.name.trim() }) + models += ModelDefinition( + name = alias, + provider = providerName, + model = modelName, + extraBody = extraBody, + ) + return alias + } + + private fun findOrAddProvider( + providers: MutableList, + type: String, + api: String, + token: String, + ): String { + val reusable = providers.firstOrNull { candidate -> + candidate.name.isNotBlank() && + providers.count { it.name.trim() == candidate.name.trim() } == 1 && + normalizedType(candidate.type) == type && + normalizedApi(candidate.api) == normalizedApi(api) && + candidate.token.trim() == token + } + if (reusable != null) return reusable.name.trim() + + val name = uniqueName(providerBaseName(type, api), providers.mapTo(HashSet()) { it.name.trim() }) + providers += ModelProviderDefinition(name = name, type = type, api = api, token = token) + return name + } + + private fun normalizedType(type: String): String = when (type.trim().lowercase()) { + "openai-compatible", "openai_compatible" -> "openai" + else -> type.trim().lowercase() + } + + private fun normalizedApi(api: String): String = api.trim().trimEnd('/') + + private fun providerBaseName(type: String, api: String): String { + if (type == "dashscope") return "dashscope-native" + val host = runCatching { URI.create(api).host?.lowercase() }.getOrNull().orEmpty() + if (host.contains("deepseek")) return "deepseek" + if (host.contains("dashscope")) return "dashscope-openai" + if (host.contains("openai")) return "openai" + val segment = host.split('.').firstOrNull { it !in setOf("", "api", "www", "v1") }.orEmpty() + return segment.replace(Regex("[^a-z0-9]+"), "-").trim('-').ifBlank { "openai" } + } + + private fun uniqueName(preferred: String, occupied: Set): String { + if (preferred !in occupied) return preferred + var suffix = 2 + while ("$preferred-$suffix" in occupied) suffix++ + return "$preferred-$suffix" + } +} diff --git a/src/main/kotlin/config/PluginConfig.kt b/src/main/kotlin/config/PluginConfig.kt index a65ba07..c2f3c11 100644 --- a/src/main/kotlin/config/PluginConfig.kt +++ b/src/main/kotlin/config/PluginConfig.kt @@ -33,6 +33,12 @@ object PluginConfig : AutoSavePluginConfig("Config") { @ValueDescription("Chat模型温度,默认为null") var chatTemperature: Double? by value(null) + @ValueDescription("主聊天模型别名;填写后优先从 Models.yml 解析,留空时兼容旧的 openAiApi/openAiToken/chatModel") + var chatModelAlias: String by value("") + + @ValueDescription("聊天备用模型别名列表;主模型失败时按顺序切换") + var chatFallbackModelAliases: List by value() + @ValueDescription("推理模型API") var reasoningModelApi: String by value("https://dashscope.aliyuncs.com/compatible-mode/v1/") @@ -42,6 +48,9 @@ object PluginConfig : AutoSavePluginConfig("Config") { @ValueDescription("推理模型") var reasoningModel: String by value("qwq-plus") + @ValueDescription("推理模型别名;留空时兼容旧的推理模型配置") + var reasoningModelAlias: String by value("") + @ValueDescription("视觉模型API") var visualModelApi: String by value("https://dashscope.aliyuncs.com/compatible-mode/v1/") @@ -51,6 +60,9 @@ object PluginConfig : AutoSavePluginConfig("Config") { @ValueDescription("视觉模型") var visualModel: String by value("qwen-vl-plus") + @ValueDescription("视觉模型别名;留空时兼容旧的视觉模型配置") + var visualModelAlias: String by value("") + @ValueDescription("聊天模型额外请求体JSON,会合并到请求体中。例如DeepSeek关闭思维: {\"thinking\": {\"type\": \"disabled\"}}") val chatModelExtraBody: String by value("") @@ -69,6 +81,9 @@ object PluginConfig : AutoSavePluginConfig("Config") { @ValueDescription("画像分析模型。留空时继承聊天模型") val profileModel: String by value("") + @ValueDescription("画像分析模型别名;留空时继承主聊天模型别名或兼容旧配置") + var profileModelAlias: String by value("") + @ValueDescription("画像分析模型额外请求体JSON。留空时继承聊天模型额外请求体") val profileModelExtraBody: String by value("") @@ -174,12 +189,18 @@ object PluginConfig : AutoSavePluginConfig("Config") { @ValueDescription("百炼平台图像模型,支持文生图与图像编辑。可选:qwen-image-2.0 / qwen-image-2.0-pro / qwen-image-edit-max / qwen-image-edit-plus 等") val imageModel: String by value("qwen-image-2.0") + @ValueDescription("图像模型别名;留空时兼容旧的 dashScopeApiKey/imageModel") + var imageModelAlias: String by value("") + @ValueDescription("是否在生成的图片右下角添加 Qwen-Image 水印") val imageWatermark: Boolean by value(false) @ValueDescription("百炼平台TTS模型。qwen3-tts-instruct-flash 支持 instructions 指令控制;纯发音可用 qwen3-tts-flash 或 qwen-tts") val ttsModel: String by value("qwen3-tts-instruct-flash") + @ValueDescription("TTS 模型别名;留空时兼容旧的 dashScopeApiKey/ttsModel") + var ttsModelAlias: String by value("") + @ValueDescription("Jina API Key") val jinaApiKey by value("") @@ -195,6 +216,9 @@ object PluginConfig : AutoSavePluginConfig("Config") { @ValueDescription("网页摘要模型名称") val webSummaryModel: String by value("") + @ValueDescription("网页摘要模型别名;留空时兼容旧的网页摘要模型配置") + var webSummaryModelAlias: String by value("") + @ValueDescription("网页摘要模型额外请求体JSON,会合并到请求体中") val webSummaryModelExtraBody: String by value("") diff --git a/src/main/kotlin/conversation/ConversationEngine.kt b/src/main/kotlin/conversation/ConversationEngine.kt index f055ffa..e20369e 100644 --- a/src/main/kotlin/conversation/ConversationEngine.kt +++ b/src/main/kotlin/conversation/ConversationEngine.kt @@ -6,6 +6,7 @@ import com.aallam.openai.api.chat.ChatMessage import com.aallam.openai.api.chat.ChatRole import com.aallam.openai.api.chat.ToolCall import com.aallam.openai.api.chat.ToolChoice +import com.aallam.openai.api.chat.StreamOptions import com.aallam.openai.api.core.Usage import com.aallam.openai.api.model.ModelId import kotlinx.coroutines.CancellationException @@ -21,7 +22,7 @@ import net.mamoe.mirai.event.events.MessageEvent import net.mamoe.mirai.message.data.source import top.jie65535.mirai.JChatGPT import top.jie65535.mirai.config.PluginConfig -import top.jie65535.mirai.data.TokenUsageStore +import top.jie65535.mirai.data.ModelUsageRecorder import top.jie65535.mirai.llm.LargeLanguageModels import top.jie65535.mirai.llm.ModelService import top.jie65535.mirai.profile.ProfileAutoMaintenance @@ -50,6 +51,7 @@ import top.jie65535.mirai.tools.VisitWeb import top.jie65535.mirai.tools.VisualAgent import top.jie65535.mirai.tools.WeatherService import top.jie65535.mirai.tools.WebSearch +import top.jie65535.mirai.tools.QueryTokenUsageAgent import top.jie65535.mirai.util.RetryBackoff import java.time.OffsetDateTime import java.time.format.DateTimeFormatter @@ -84,6 +86,7 @@ internal object ConversationEngine { AdjustUserFavorabilityAgent(), RequestOwner(), GroupManageAgent(), + QueryTokenUsageAgent(), ) fun clear() { @@ -202,7 +205,8 @@ internal object ConversationEngine { var lastTokenUsage: Usage? = null responseFlow.collect { chunk -> - val delta = chunk.choices[0].delta ?: return@collect + chunk.usage?.let { lastTokenUsage = it } + val delta = chunk.choices.firstOrNull()?.delta ?: return@collect delta.reasoningContent?.let { content -> if (reasoningContent == null) reasoningContent = StringBuilder(content) else reasoningContent.append(content) @@ -238,7 +242,6 @@ internal object ConversationEngine { responseToolCalls[index] = current.copy(function = updated) } } - chunk.usage?.let { lastTokenUsage = it } } streamingOk = true @@ -252,7 +255,7 @@ internal object ConversationEngine { toolCalls = responseToolCalls.ifEmpty { null }, reasoningContent = if (responseToolCalls.isNotEmpty()) reasoningContent?.toString() else null, ) - recordUsage(roundEvent, lastTokenUsage, lastCacheUsage) + recordUsage(roundEvent, endpoint, lastTokenUsage, lastCacheUsage) completedRounds++ if (responseToolCalls.size > toolCallTasks.size) { @@ -455,6 +458,7 @@ internal object ConversationEngine { messages = history, tools = availableTools, toolChoice = ToolChoice.Required, + streamOptions = StreamOptions(includeUsage = true), ) JChatGPT.logger.info("API Requesting... Model=${endpoint.model} [${endpoint.label}]") return endpoint.service.chatCompletions(request, onCacheUsage) @@ -534,21 +538,19 @@ internal object ConversationEngine { private fun recordUsage( event: MessageEvent, + endpoint: LargeLanguageModels.ChatEndpoint, usage: Usage?, cacheUsage: ModelService.CacheUsage?, ) { - usage ?: return - val group = (event as? GroupMessageEvent)?.group - TokenUsageStore.record( - timestamp = OffsetDateTime.now().toEpochSecond(), - userId = event.sender.id, - userNickname = event.senderName, - groupId = group?.id, - groupName = group?.name, - promptTokens = usage.promptTokens ?: 0, - completionTokens = usage.completionTokens ?: 0, - totalTokens = usage.totalTokens ?: 0, - cachedTokens = cacheUsage?.hitTokens ?: 0, + ModelUsageRecorder.recordTokens( + event = event, + endpointLabel = endpoint.label, + modelAlias = endpoint.alias, + provider = endpoint.provider, + model = endpoint.model, + usageKind = "chat", + usage = usage, + cacheUsage = cacheUsage, ) } diff --git a/src/main/kotlin/data/ModelUsageRecorder.kt b/src/main/kotlin/data/ModelUsageRecorder.kt new file mode 100644 index 0000000..3729cf5 --- /dev/null +++ b/src/main/kotlin/data/ModelUsageRecorder.kt @@ -0,0 +1,154 @@ +package top.jie65535.mirai.data + +import com.aallam.openai.api.core.Usage +import net.mamoe.mirai.event.events.GroupMessageEvent +import net.mamoe.mirai.event.events.MessageEvent +import top.jie65535.mirai.llm.ModelService +import java.time.OffsetDateTime + +data class ModelUsageAttribution( + val botId: Long = 0, + val userId: Long = 0, + val userNickname: String = "", + val groupId: Long? = null, + val groupName: String? = null, +) { + companion object { + fun from(event: MessageEvent): ModelUsageAttribution { + val group = (event as? GroupMessageEvent)?.group + return ModelUsageAttribution( + botId = event.bot.id, + userId = event.sender.id, + userNickname = event.senderName, + groupId = group?.id, + groupName = group?.name, + ) + } + } +} + +object ModelUsageRecorder { + fun recordTokens( + event: MessageEvent, + endpointLabel: String, + modelAlias: String, + provider: String, + model: String, + usageKind: String, + usage: Usage?, + cacheUsage: ModelService.CacheUsage? = null, + ) { + usage ?: return + val promptTokens = usage.promptTokens ?: 0 + val completionTokens = usage.completionTokens ?: 0 + recordTokenValues( + attribution = ModelUsageAttribution.from(event), + endpointLabel = endpointLabel, + modelAlias = modelAlias, + provider = provider, + model = model, + usageKind = usageKind, + promptTokens = promptTokens.toLong(), + completionTokens = completionTokens.toLong(), + totalTokens = (usage.totalTokens ?: (promptTokens + completionTokens)).toLong(), + cachedTokens = (cacheUsage?.hitTokens ?: 0).toLong(), + ) + } + + fun recordTokenValues( + attribution: ModelUsageAttribution, + endpointLabel: String, + modelAlias: String, + provider: String, + model: String, + usageKind: String, + promptTokens: Long, + completionTokens: Long, + totalTokens: Long = promptTokens + completionTokens, + cachedTokens: Long = 0, + ) { + record( + attribution = attribution, + endpointLabel = endpointLabel, + modelAlias = modelAlias, + provider = provider, + model = model, + usageKind = usageKind, + unit = "tokens", + inputUnits = promptTokens, + outputUnits = completionTokens, + totalUnits = totalTokens, + promptTokens = promptTokens, + completionTokens = completionTokens, + totalTokens = totalTokens, + cachedTokens = cachedTokens, + ) + } + + fun recordUnits( + event: MessageEvent, + endpointLabel: String, + modelAlias: String, + provider: String, + model: String, + usageKind: String, + unit: String, + inputUnits: Long = 0, + outputUnits: Long = 0, + totalUnits: Long = inputUnits + outputUnits, + ) { + record( + attribution = ModelUsageAttribution.from(event), + endpointLabel = endpointLabel, + modelAlias = modelAlias, + provider = provider, + model = model, + usageKind = usageKind, + unit = unit, + inputUnits = inputUnits, + outputUnits = outputUnits, + totalUnits = totalUnits, + ) + } + + private fun record( + attribution: ModelUsageAttribution, + endpointLabel: String, + modelAlias: String, + provider: String, + model: String, + usageKind: String, + unit: String, + inputUnits: Long, + outputUnits: Long, + totalUnits: Long, + promptTokens: Long = 0, + completionTokens: Long = 0, + totalTokens: Long = 0, + cachedTokens: Long = 0, + ) { + TokenUsageStore.recordUsage( + ModelUsageEvent( + timestamp = OffsetDateTime.now().toEpochSecond(), + botId = attribution.botId, + userId = attribution.userId, + userNickname = attribution.userNickname, + groupId = attribution.groupId, + groupName = attribution.groupName, + endpointLabel = endpointLabel, + modelAlias = modelAlias, + provider = provider, + model = model, + usageKind = usageKind, + unit = unit, + inputUnits = inputUnits.coerceAtLeast(0), + outputUnits = outputUnits.coerceAtLeast(0), + totalUnits = totalUnits.coerceAtLeast(0), + promptTokens = promptTokens.coerceAtLeast(0), + completionTokens = completionTokens.coerceAtLeast(0), + totalTokens = totalTokens.coerceAtLeast(0), + cachedTokens = cachedTokens.coerceAtLeast(0), + ) + ) + } +} diff --git a/src/main/kotlin/data/PluginData.kt b/src/main/kotlin/data/PluginData.kt index 0a561c9..b2288d0 100644 --- a/src/main/kotlin/data/PluginData.kt +++ b/src/main/kotlin/data/PluginData.kt @@ -40,7 +40,7 @@ data class FavorabilityInfo( } /** - * Token使用日聚合记录。按 (date, userId, groupId) 维度合并。由 [TokenUsageStore] 持久化到独立 JSON 文件。 + * 旧版 Token 使用日聚合记录。仅用于将 token_usage.json 迁移到 [TokenUsageStore] 的 SQLite 明细表。 * @param date 本地时区下的日期,格式 yyyy-MM-dd * @param userId QQ * @param userNickname 最近一次记录到的昵称 diff --git a/src/main/kotlin/data/TokenUsageModels.kt b/src/main/kotlin/data/TokenUsageModels.kt new file mode 100644 index 0000000..41932ec --- /dev/null +++ b/src/main/kotlin/data/TokenUsageModels.kt @@ -0,0 +1,114 @@ +package top.jie65535.mirai.data + +data class ModelUsageEvent( + val timestamp: Long, + val botId: Long = 0, + val userId: Long = 0, + val userNickname: String = "", + val groupId: Long? = null, + val groupName: String? = null, + val endpointLabel: String? = null, + val modelAlias: String? = null, + val provider: String? = null, + val model: String? = null, + val usageKind: String = "chat", + val unit: String = "tokens", + val inputUnits: Long = 0, + val outputUnits: Long = 0, + val totalUnits: Long = inputUnits + outputUnits, + val promptTokens: Long = 0, + val completionTokens: Long = 0, + val totalTokens: Long = 0, + val cachedTokens: Long = 0, +) + +data class TokenUsageRecord( + val id: Long, + val timestamp: Long, + val date: String, + val botId: Long?, + val userId: Long, + val userNickname: String, + val groupId: Long?, + val groupName: String?, + val endpointLabel: String?, + val provider: String?, + val model: String?, + val modelAlias: String? = null, + val usageKind: String = "chat", + val unit: String = "tokens", + val inputUnits: Long = 0, + val outputUnits: Long = 0, + val totalUnits: Long = 0, + val promptTokens: Long, + val completionTokens: Long, + val totalTokens: Long, + val cachedTokens: Long, + val callCount: Int, + val detailed: Boolean, +) + +data class TokenUsageRanking( + val id: Long, + val name: String, + val totalTokens: Long, +) + +data class TokenUsageModelTotal( + val provider: String, + val model: String, + val totalTokens: Long, + val callCount: Int, +) + +data class TokenUsageBreakdown( + val provider: String, + val model: String, + val usageKind: String, + val unit: String, + val inputUnits: Long, + val outputUnits: Long, + val totalUnits: Long, + val callCount: Int, +) + +data class ModelUsageUserTotal( + val userId: Long, + val name: String, + val usageKind: String, + val unit: String, + val totalUnits: Long, + val callCount: Int, +) + +data class ModelUsageDailyTotal( + val date: String, + val usageKind: String, + val unit: String, + val totalUnits: Long, + val callCount: Int, +) + +data class TokenUsageDailyTotal( + val date: String, + val totalTokens: Long, +) + +data class TokenUsageSummary( + val promptTokens: Long, + val completionTokens: Long, + val totalTokens: Long, + val cachedTokens: Long, + val callCount: Int, + val activeUsers: Int, + val todayTotal: Long, + val daily: List, + val topUsers: List, + val topGroups: List, + val models: List, + val allCallCount: Int = callCount, + val allActiveUsers: Int = activeUsers, + val breakdown: List = emptyList(), + val userUsage: List = emptyList(), + val usageDaily: List = emptyList(), +) diff --git a/src/main/kotlin/data/TokenUsageStore.kt b/src/main/kotlin/data/TokenUsageStore.kt index bfd482c..5d7a47a 100644 --- a/src/main/kotlin/data/TokenUsageStore.kt +++ b/src/main/kotlin/data/TokenUsageStore.kt @@ -2,119 +2,1009 @@ package top.jie65535.mirai.data import kotlinx.serialization.builtins.ListSerializer import kotlinx.serialization.json.Json +import org.sqlite.SQLiteConfig import java.io.File +import java.sql.Connection +import java.sql.DriverManager +import java.sql.ResultSet +import java.security.MessageDigest import java.time.Instant import java.time.LocalDate import java.time.ZoneId import java.time.format.DateTimeFormatter /** - * Token使用日聚合存储。独立于 mamoe 的 plugin data 系统,直接管 JSON 文件, - * 避免 yamlkt 在大数据量下编/解码不互通的 bug。 + * SQLite model-usage ledger. New calls are stored individually; old daily JSON + * aggregates are imported as rows marked [TokenUsageRecord.detailed] = false. */ object TokenUsageStore { + private const val SCHEMA_VERSION = 2 + private const val BUSY_TIMEOUT_MS = 30_000 + private const val DATABASE_NAME = "chat-history.sqlite" + private const val LEGACY_FILE_NAME = "token_usage.json" + private const val LEGACY_FINGERPRINT_KEY = "legacy_json_sha256" + private const val TOP_LIMIT = 5 + private const val MAX_RANKING_LIMIT = 100 + + private val lifecycleLock = Any() + private val writeLock = Any() + private val dateFmt = DateTimeFormatter.ISO_LOCAL_DATE + private val zone = ZoneId.systemDefault() private val json = Json { - prettyPrint = true ignoreUnknownKeys = true encodeDefaults = true } - private val dateFmt = DateTimeFormatter.ISO_LOCAL_DATE - private val listSerializer = ListSerializer(TokenUsageDailyRecord.serializer()) + private val legacySerializer = ListSerializer(TokenUsageDailyRecord.serializer()) - private lateinit var file: File - private val records = mutableListOf() + @Volatile + private var initialized = false + private lateinit var databaseFile: File + private var writeConnection: Connection? = null + private var warningLogger: ((String, Throwable?) -> Unit)? = null - /** - * 在 onEnable 中调用一次,传入插件数据目录。 - */ - fun init(dataFolder: File) { - file = File(dataFolder, "token_usage.json") - records.clear() - if (file.exists() && file.length() > 0) { + val isAvailable: Boolean + get() = initialized + + fun init( + dataFolder: File, + onWarning: (String, Throwable?) -> Unit = { _, _ -> }, + ) { + synchronized(lifecycleLock) { + if (initialized) return + Class.forName("org.sqlite.JDBC") + dataFolder.mkdirs() + databaseFile = dataFolder.resolve(DATABASE_NAME) + warningLogger = onWarning + + val connection = openConnection(databaseFile) try { - records.addAll(json.decodeFromString(listSerializer, file.readText())) - } catch (_: Exception) { - // 加载失败不阻塞插件启动,备份原文件后从空开始 - val backup = File(file.parentFile, "token_usage.json.broken-${System.currentTimeMillis()}") - file.copyTo(backup, overwrite = true) + configureWriteConnection(connection) + createSchema(connection) + writeConnection = connection + initialized = true + importLegacyJson(dataFolder.resolve(LEGACY_FILE_NAME)) + } catch (cause: Throwable) { + connection.close() + warningLogger = null + throw cause } } } - val all: List get() = records + fun close() { + synchronized(lifecycleLock) { + if (!initialized) return + synchronized(writeLock) { + writeConnection?.let { connection -> + runCatching { + connection.createStatement().use { it.execute("PRAGMA wal_checkpoint(TRUNCATE)") } + }.onFailure { warn("Token SQLite WAL checkpoint 失败", it) } + connection.close() + } + writeConnection = null + initialized = false + warningLogger = null + } + } + } - /** - * 将一次调用的 token 用量累加到当日聚合行;若不存在则创建。写盘失败不抛。 - */ - @Synchronized + /** Records one successful chat-model response. Storage failures are non-fatal. */ fun record( timestamp: Long, + botId: Long, userId: Long, userNickname: String, groupId: Long?, groupName: String?, + endpointLabel: String?, + apiBaseUrl: String?, + model: String?, promptTokens: Int, completionTokens: Int, totalTokens: Int, - cachedTokens: Int + cachedTokens: Int, ) { - val date = LocalDate.ofInstant(Instant.ofEpochSecond(timestamp), ZoneId.systemDefault()) - .format(dateFmt) - val nickname = sanitizeNickname(userNickname) - val groupNameClean = groupName?.let { sanitizeNickname(it) } - val idx = records.indexOfFirst { - it.date == date && it.userId == userId && it.groupId == groupId - } - if (idx >= 0) { - val r = records[idx] - records[idx] = r.copy( - userNickname = nickname.ifEmpty { r.userNickname }, - groupName = groupNameClean?.ifEmpty { null } ?: r.groupName, - promptTokens = r.promptTokens + promptTokens, - completionTokens = r.completionTokens + completionTokens, - totalTokens = r.totalTokens + totalTokens, - cachedTokens = r.cachedTokens + cachedTokens, - callCount = r.callCount + 1 + recordUsage( + ModelUsageEvent( + timestamp = timestamp, + botId = botId, + userId = userId, + userNickname = userNickname, + groupId = groupId, + groupName = groupName, + endpointLabel = endpointLabel, + modelAlias = endpointLabel, + provider = providerFor(apiBaseUrl), + model = model, + usageKind = "chat", + inputUnits = promptTokens.toLong().coerceAtLeast(0), + outputUnits = completionTokens.toLong().coerceAtLeast(0), + totalUnits = totalTokens.toLong().coerceAtLeast(0), + promptTokens = promptTokens.toLong().coerceAtLeast(0), + completionTokens = completionTokens.toLong().coerceAtLeast(0), + totalTokens = totalTokens.toLong().coerceAtLeast(0), + cachedTokens = cachedTokens.toLong().coerceAtLeast(0), ) - } else { - records.add( - TokenUsageDailyRecord( - date = date, - userId = userId, - userNickname = nickname, - groupId = groupId, - groupName = groupNameClean?.ifEmpty { null }, - promptTokens = promptTokens.toLong(), - completionTokens = completionTokens.toLong(), - totalTokens = totalTokens.toLong(), - cachedTokens = cachedTokens.toLong(), - callCount = 1 + ) + } + + fun recordUsage(event: ModelUsageEvent) { + if (!initialized) return + runCatching { + withWriteConnection { connection -> + connection.prepareStatement( + """ + INSERT INTO token_usage_record( + occurred_at, usage_date, bot_id, user_id, user_nickname, + group_id, group_name, endpoint_label, model_alias, provider, model, + usage_kind, unit, input_units, output_units, total_units, + prompt_tokens, completion_tokens, total_tokens, cached_tokens, + call_count, detailed + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 1, 1) + """.trimIndent() + ).use { statement -> + statement.setLong(1, event.timestamp) + statement.setString(2, dateFor(event.timestamp)) + statement.setLong(3, event.botId) + statement.setLong(4, event.userId) + statement.setString(5, sanitizeNickname(event.userNickname)) + if (event.groupId == null) statement.setNull(6, java.sql.Types.INTEGER) else statement.setLong(6, event.groupId) + if (event.groupName == null) statement.setNull(7, java.sql.Types.VARCHAR) + else statement.setString(7, sanitizeNickname(event.groupName)) + statement.setString(8, event.endpointLabel?.takeIf(String::isNotBlank)) + statement.setString(9, event.modelAlias?.takeIf(String::isNotBlank)) + statement.setString(10, event.provider?.takeIf(String::isNotBlank)) + statement.setString(11, event.model?.takeIf(String::isNotBlank)) + statement.setString(12, event.usageKind.ifBlank { "other" }) + statement.setString(13, event.unit.ifBlank { "units" }) + statement.setLong(14, event.inputUnits.coerceAtLeast(0)) + statement.setLong(15, event.outputUnits.coerceAtLeast(0)) + statement.setLong(16, event.totalUnits.coerceAtLeast(0)) + statement.setLong(17, event.promptTokens.coerceAtLeast(0)) + statement.setLong(18, event.completionTokens.coerceAtLeast(0)) + statement.setLong(19, event.totalTokens.coerceAtLeast(0)) + statement.setLong(20, event.cachedTokens.coerceAtLeast(0)) + statement.executeUpdate() + } + } + }.onFailure { warn("保存模型用量记录失败", it) } + } + + fun summary( + startDate: String, + endDate: String = LocalDate.now(zone).format(dateFmt), + botId: Long? = null, + userId: Long? = null, + groupId: Long? = null, + privateOnly: Boolean = false, + model: String? = null, + usageKind: String? = null, + rankingLimit: Int = TOP_LIMIT, + ): TokenUsageSummary { + check(initialized) { "Token SQLite 尚未初始化" } + require(groupId == null || !privateOnly) { "groupId and privateOnly cannot be used together" } + require(rankingLimit > 0) { "rankingLimit must be positive" } + val normalizedRankingLimit = rankingLimit.coerceAtMost(MAX_RANKING_LIMIT) + val filter = Filter( + startDate, + endDate, + botId, + userId, + groupId, + privateOnly, + model?.trim()?.takeIf(String::isNotBlank), + usageKind?.trim()?.takeIf(String::isNotBlank), + ) + return openReadConnection(databaseFile).use { connection -> + connection.autoCommit = false + try { + val aggregate = queryAggregate(connection, filter, tokenOnly = true) + val allAggregate = queryAggregate(connection, filter, tokenOnly = false) + val daily = queryDaily(connection, filter) + val topUsers = queryTopUsers(connection, filter, normalizedRankingLimit) + val topGroups = queryTopGroups(connection, filter, normalizedRankingLimit) + val models = queryModels(connection, filter, normalizedRankingLimit) + val breakdown = queryBreakdown(connection, filter, normalizedRankingLimit) + val userUsage = queryUserUsage(connection, filter, normalizedRankingLimit) + val usageDaily = queryUsageDaily(connection, filter) + TokenUsageSummary( + promptTokens = aggregate.promptTokens, + completionTokens = aggregate.completionTokens, + totalTokens = aggregate.totalTokens, + cachedTokens = aggregate.cachedTokens, + callCount = aggregate.callCount, + activeUsers = aggregate.activeUsers, + todayTotal = queryTodayTotal(connection, filter), + daily = daily, + topUsers = topUsers, + topGroups = topGroups, + models = models, + allCallCount = allAggregate.callCount, + allActiveUsers = allAggregate.activeUsers, + breakdown = breakdown, + userUsage = userUsage, + usageDaily = usageDaily, ) + } finally { + connection.rollback() + } + } + } + + fun recent( + limit: Int = 20, + startDate: String? = null, + endDate: String? = null, + botId: Long? = null, + userId: Long? = null, + groupId: Long? = null, + privateOnly: Boolean = false, + model: String? = null, + usageKind: String? = null, + ): List { + check(initialized) { "Token SQLite 尚未初始化" } + require(limit > 0) { "limit must be positive" } + require(groupId == null || !privateOnly) { "groupId and privateOnly cannot be used together" } + val filter = Filter( + startDate, + endDate, + botId, + userId, + groupId, + privateOnly, + model?.trim()?.takeIf(String::isNotBlank), + usageKind?.trim()?.takeIf(String::isNotBlank), + ) + return openReadConnection(databaseFile).use { connection -> + val (where, args) = whereClause(filter) + val sql = "SELECT * FROM token_usage_record WHERE $where ORDER BY occurred_at DESC, id DESC LIMIT ?" + connection.prepareStatement(sql).use { statement -> + bind(statement, args) + statement.setInt(args.size + 1, limit.coerceAtMost(100)) + statement.executeQuery().use { results -> + buildList { while (results.next()) add(results.toTokenUsageRecord()) } + } + } + } + } + + /** Best-effort access for the command during a partial startup. */ + fun hasAny(startDate: String, endDate: String = LocalDate.now(zone).format(dateFmt)): Boolean { + if (!initialized) return false + return openReadConnection(databaseFile).use { connection -> + val (where, args) = whereClause(Filter(startDate, endDate, null, null, null, false, null, null)) + connection.prepareStatement("SELECT 1 FROM token_usage_record WHERE $where LIMIT 1").use { statement -> + bind(statement, args) + statement.executeQuery().use(ResultSet::next) + } + } + } + + private data class Filter( + val startDate: String?, + val endDate: String?, + val botId: Long?, + val userId: Long?, + val groupId: Long?, + val privateOnly: Boolean, + val model: String?, + val usageKind: String?, + ) + + private data class Aggregate( + val promptTokens: Long, + val completionTokens: Long, + val totalTokens: Long, + val cachedTokens: Long, + val callCount: Int, + val activeUsers: Int, + ) + + private data class UserUsageAggregate( + val userId: Long, + val usageKind: String, + val unit: String, + val totalUnits: Long, + val callCount: Int, + ) + + private data class LegacyImportState( + val occurredAt: Long, + val date: String, + val userId: Long, + val userNickname: String, + val groupId: Long?, + val groupName: String?, + val inputUnits: Long, + val outputUnits: Long, + val totalUnits: Long, + val promptTokens: Long, + val completionTokens: Long, + val totalTokens: Long, + val cachedTokens: Long, + val callCount: Int, + ) + + private fun queryAggregate(connection: Connection, filter: Filter, tokenOnly: Boolean): Aggregate { + val (where, args) = whereClause(filter, tokenOnly = tokenOnly) + val sql = """ + SELECT COALESCE(SUM(prompt_tokens), 0), COALESCE(SUM(completion_tokens), 0), + COALESCE(SUM(total_tokens), 0), COALESCE(SUM(cached_tokens), 0), + COALESCE(SUM(call_count), 0), + COUNT(DISTINCT CASE WHEN user_id > 0 THEN user_id END) + FROM token_usage_record WHERE $where + """.trimIndent() + return connection.prepareStatement(sql).use { statement -> + bind(statement, args) + statement.executeQuery().use { results -> + check(results.next()) + Aggregate( + promptTokens = results.getLong(1), + completionTokens = results.getLong(2), + totalTokens = results.getLong(3), + cachedTokens = results.getLong(4), + callCount = results.getLong(5).toInt(), + activeUsers = results.getInt(6), + ) + } + } + } + + private fun queryTodayTotal(connection: Connection, filter: Filter): Long { + val today = LocalDate.now(zone).format(dateFmt) + if (filter.startDate != null && today < filter.startDate || filter.endDate != null && today > filter.endDate) return 0 + val todayFilter = filter.copy(startDate = today, endDate = today) + val (where, args) = whereClause(todayFilter, tokenOnly = true) + return connection.prepareStatement("SELECT COALESCE(SUM(total_tokens), 0) FROM token_usage_record WHERE $where") + .use { statement -> + bind(statement, args) + statement.executeQuery().use { results -> check(results.next()); results.getLong(1) } + } + } + + private fun queryDaily(connection: Connection, filter: Filter): List { + val (where, args) = whereClause(filter, tokenOnly = true) + return connection.prepareStatement( + "SELECT usage_date, SUM(total_tokens) FROM token_usage_record WHERE $where GROUP BY usage_date ORDER BY usage_date" + ).use { statement -> + bind(statement, args) + statement.executeQuery().use { results -> + buildList { + while (results.next()) add(TokenUsageDailyTotal(results.getString(1), results.getLong(2))) + } + } + } + } + + private fun queryTopUsers( + connection: Connection, + filter: Filter, + limit: Int, + ): List { + val (where, args) = whereClause(filter, tokenOnly = true) + val totals = connection.prepareStatement( + "SELECT user_id, SUM(total_tokens) AS total FROM token_usage_record " + + "WHERE $where AND user_id > 0 GROUP BY user_id ORDER BY total DESC LIMIT ?" + ).use { statement -> + bind(statement, args) + statement.setInt(args.size + 1, limit) + statement.executeQuery().use { results -> + buildList { while (results.next()) add(results.getLong("user_id") to results.getLong("total")) } + } + } + val (nameWhere, nameArgs) = whereClause(filter.copy(model = null, usageKind = null)) + return totals.map { (id, total) -> + val name = connection.prepareStatement( + "SELECT user_nickname FROM token_usage_record WHERE $nameWhere AND user_id = ? " + + "AND user_nickname <> '' ORDER BY occurred_at DESC, id DESC LIMIT 1" + ).use { statement -> + bind(statement, nameArgs) + statement.setLong(nameArgs.size + 1, id) + statement.executeQuery().use { results -> if (results.next()) results.getString(1).orEmpty() else "" } + } + TokenUsageRanking(id, name, total) + } + } + + private fun queryTopGroups( + connection: Connection, + filter: Filter, + limit: Int, + ): List { + val (where, args) = whereClause(filter, requireGroup = true, tokenOnly = true) + val totals = connection.prepareStatement( + "SELECT group_id, SUM(total_tokens) AS total FROM token_usage_record WHERE $where GROUP BY group_id ORDER BY total DESC LIMIT ?" + ).use { statement -> + bind(statement, args) + statement.setInt(args.size + 1, limit) + statement.executeQuery().use { results -> + buildList { while (results.next()) add(results.getLong("group_id") to results.getLong("total")) } + } + } + val (nameWhere, nameArgs) = whereClause( + filter.copy(model = null, usageKind = null), + requireGroup = true, + ) + return totals.map { (id, total) -> + val name = connection.prepareStatement( + "SELECT group_name FROM token_usage_record WHERE $nameWhere AND group_id = ? " + + "AND group_name IS NOT NULL AND group_name <> '' ORDER BY occurred_at DESC, id DESC LIMIT 1" + ).use { statement -> + bind(statement, nameArgs) + statement.setLong(nameArgs.size + 1, id) + statement.executeQuery().use { results -> if (results.next()) results.getString(1).orEmpty() else "" } + } + TokenUsageRanking(id, name, total) + } + } + + private fun queryModels( + connection: Connection, + filter: Filter, + limit: Int, + ): List { + val (where, args) = whereClause(filter, tokenOnly = true) + return connection.prepareStatement( + """ + SELECT COALESCE(provider, ''), COALESCE(model, ''), SUM(total_units), SUM(call_count) + FROM token_usage_record WHERE $where + GROUP BY provider, model ORDER BY SUM(total_tokens) DESC LIMIT ? + """.trimIndent() + ).use { statement -> + bind(statement, args) + statement.setInt(args.size + 1, limit) + statement.executeQuery().use { results -> + buildList { + while (results.next()) { + add( + TokenUsageModelTotal( + provider = results.getString(1).ifBlank { "unknown" }, + model = results.getString(2).ifBlank { "unknown" }, + totalTokens = results.getLong(3), + callCount = results.getLong(4).toInt(), + ) + ) + } + } + } + } + } + + private fun queryBreakdown( + connection: Connection, + filter: Filter, + limit: Int, + ): List { + val (where, args) = whereClause(filter) + return connection.prepareStatement( + """ + SELECT COALESCE(provider, ''), COALESCE(model, ''), + COALESCE(usage_kind, 'other'), COALESCE(unit, 'units'), + SUM(input_units), SUM(output_units), SUM(total_units), SUM(call_count) + FROM token_usage_record WHERE $where + GROUP BY provider, model, usage_kind, unit + ORDER BY SUM(total_units) DESC LIMIT ? + """.trimIndent() + ).use { statement -> + bind(statement, args) + statement.setInt(args.size + 1, limit) + statement.executeQuery().use { results -> + buildList { + while (results.next()) { + add( + TokenUsageBreakdown( + provider = results.getString(1).ifBlank { "unknown" }, + model = results.getString(2).ifBlank { "unknown" }, + usageKind = results.getString(3).ifBlank { "other" }, + unit = results.getString(4).ifBlank { "units" }, + inputUnits = results.getLong(5), + outputUnits = results.getLong(6), + totalUnits = results.getLong(7), + callCount = results.getLong(8).toInt(), + ) + ) + } + } + } + } + } + + private fun queryUserUsage( + connection: Connection, + filter: Filter, + limit: Int, + ): List { + val (where, args) = whereClause(filter) + val totals = connection.prepareStatement( + """ + SELECT user_id, COALESCE(usage_kind, 'other'), COALESCE(unit, 'units'), + SUM(total_units), SUM(call_count) + FROM token_usage_record + WHERE $where AND user_id > 0 + GROUP BY user_id, usage_kind, unit + ORDER BY usage_kind, unit, SUM(total_units) DESC + LIMIT ? + """.trimIndent() + ).use { statement -> + bind(statement, args) + statement.setInt(args.size + 1, limit) + statement.executeQuery().use { results -> + buildList { + while (results.next()) { + add( + UserUsageAggregate( + userId = results.getLong(1), + usageKind = results.getString(2), + unit = results.getString(3), + totalUnits = results.getLong(4), + callCount = results.getLong(5).toInt(), + ) + ) + } + } + } + } + val (nameWhere, nameArgs) = whereClause(filter.copy(model = null, usageKind = null)) + return totals.map { row -> + val name = connection.prepareStatement( + "SELECT user_nickname FROM token_usage_record WHERE $nameWhere AND user_id = ? AND user_nickname <> '' " + + "ORDER BY occurred_at DESC, id DESC LIMIT 1" + ).use { statement -> + bind(statement, nameArgs) + statement.setLong(nameArgs.size + 1, row.userId) + statement.executeQuery().use { results -> if (results.next()) results.getString(1).orEmpty() else "" } + } + ModelUsageUserTotal( + userId = row.userId, + name = name, + usageKind = row.usageKind, + unit = row.unit, + totalUnits = row.totalUnits, + callCount = row.callCount, ) + } + } + + private fun queryUsageDaily( + connection: Connection, + filter: Filter, + ): List { + val (where, args) = whereClause(filter) + return connection.prepareStatement( + """ + SELECT usage_date, COALESCE(usage_kind, 'other'), COALESCE(unit, 'units'), + SUM(total_units), SUM(call_count) + FROM token_usage_record + WHERE $where + GROUP BY usage_date, usage_kind, unit + ORDER BY usage_date, usage_kind, unit + """.trimIndent() + ).use { statement -> + bind(statement, args) + statement.executeQuery().use { results -> + buildList { + while (results.next()) { + add( + ModelUsageDailyTotal( + date = results.getString(1), + usageKind = results.getString(2), + unit = results.getString(3), + totalUnits = results.getLong(4), + callCount = results.getLong(5).toInt(), + ) + ) + } + } + } + } + } + + private fun whereClause( + filter: Filter, + requireGroup: Boolean = false, + tokenOnly: Boolean = false, + ): Pair> { + val conditions = mutableListOf() + val args = mutableListOf() + if (tokenOnly) conditions += "unit = 'tokens'" + filter.startDate?.let { conditions += "usage_date >= ?"; args += it } + filter.endDate?.let { conditions += "usage_date <= ?"; args += it } + filter.botId?.let { conditions += "bot_id = ?"; args += it } + filter.userId?.let { conditions += "user_id = ?"; args += it } + if (requireGroup) conditions += "group_id IS NOT NULL" + if (filter.privateOnly) conditions += "group_id IS NULL" + filter.groupId?.let { conditions += "group_id = ?"; args += it } + filter.model?.let { conditions += "model = ?"; args += it } + filter.usageKind?.let { conditions += "usage_kind = ?"; args += it } + return (conditions.takeIf { it.isNotEmpty() }?.joinToString(" AND ") ?: "1 = 1") to args + } + + private fun bind(statement: java.sql.PreparedStatement, args: List) { + args.forEachIndexed { index, value -> + when (value) { + is Long -> statement.setLong(index + 1, value) + is Int -> statement.setInt(index + 1, value) + else -> statement.setString(index + 1, value.toString()) + } + } + } + + private fun importLegacyJson(file: File) { + if (!file.isFile || file.length() == 0L) return + val fingerprint = runCatching { sha256(file) }.getOrElse { cause -> + warn("计算旧 Token 使用文件指纹失败,已跳过迁移", cause) + return + } + val fingerprintMatches = runCatching { + withWriteConnection { connection -> readMeta(connection, LEGACY_FINGERPRINT_KEY) == fingerprint } + }.getOrElse { cause -> + warn("读取旧 Token 迁移状态失败,将继续校验现有数据", cause) + false + } + if (fingerprintMatches) return + + val records = try { + json.decodeFromString(legacySerializer, file.readText()) + } catch (cause: Throwable) { + val backup = File(file.parentFile, "${file.name}.broken-${System.currentTimeMillis()}") + runCatching { file.copyTo(backup, overwrite = true) } + warn("读取旧 Token 使用文件失败,已保留原文件并跳过迁移", cause) + return + } + if (records.isEmpty()) return + runCatching { + var imported = false + withWriteConnection { connection -> + val oldAutoCommit = connection.autoCommit + connection.autoCommit = false + try { + if (!legacyRowsMatch(connection, records)) { + connection.prepareStatement( + """ + INSERT INTO token_usage_record( + occurred_at, usage_date, user_id, user_nickname, group_id, group_name, + usage_kind, unit, input_units, output_units, total_units, + prompt_tokens, completion_tokens, total_tokens, cached_tokens, + call_count, detailed, legacy_key + ) VALUES (?, ?, ?, ?, ?, ?, 'chat', 'tokens', ?, ?, ?, ?, ?, ?, ?, ?, 0, ?) + ON CONFLICT(legacy_key) DO UPDATE SET + occurred_at = excluded.occurred_at, + usage_date = excluded.usage_date, + user_nickname = excluded.user_nickname, + group_name = excluded.group_name, + usage_kind = excluded.usage_kind, + unit = excluded.unit, + input_units = excluded.input_units, + output_units = excluded.output_units, + total_units = excluded.total_units, + prompt_tokens = excluded.prompt_tokens, + completion_tokens = excluded.completion_tokens, + total_tokens = excluded.total_tokens, + cached_tokens = excluded.cached_tokens, + call_count = excluded.call_count + """.trimIndent() + ).use { statement -> + records.forEach { record -> + val state = record.toLegacyImportState() + statement.setLong(1, state.occurredAt) + statement.setString(2, state.date) + statement.setLong(3, state.userId) + statement.setString(4, state.userNickname) + if (state.groupId == null) statement.setNull(5, java.sql.Types.INTEGER) + else statement.setLong(5, state.groupId) + if (state.groupName == null) statement.setNull(6, java.sql.Types.VARCHAR) + else statement.setString(6, state.groupName) + statement.setLong(7, state.inputUnits) + statement.setLong(8, state.outputUnits) + statement.setLong(9, state.totalUnits) + statement.setLong(10, state.promptTokens) + statement.setLong(11, state.completionTokens) + statement.setLong(12, state.totalTokens) + statement.setLong(13, state.cachedTokens) + statement.setInt(14, state.callCount) + statement.setString(15, legacyKey(record)) + statement.addBatch() + } + statement.executeBatch() + } + imported = true + } + writeMeta(connection, LEGACY_FINGERPRINT_KEY, fingerprint) + connection.commit() + } catch (cause: Throwable) { + connection.rollback() + throw cause + } finally { + connection.autoCommit = oldAutoCommit + } + } + if (imported) warn("已将 ${records.size} 条旧 Token 聚合记录迁移到 SQLite", null) + }.onFailure { warn("迁移旧 Token 使用记录失败,保留原 JSON 以便下次重试", it) } + } + + private fun legacyRowsMatch(connection: Connection, records: List): Boolean { + val expected = records.associate { legacyKey(it) to it.toLegacyImportState() } + val actual = connection.prepareStatement( + """ + SELECT legacy_key, occurred_at, usage_date, user_id, user_nickname, group_id, group_name, + input_units, output_units, total_units, prompt_tokens, completion_tokens, + total_tokens, cached_tokens, call_count + FROM token_usage_record WHERE legacy_key IS NOT NULL + """.trimIndent() + ).use { statement -> + statement.executeQuery().use { results -> + buildMap { + while (results.next()) { + put( + results.getString("legacy_key"), + LegacyImportState( + occurredAt = results.getLong("occurred_at"), + date = results.getString("usage_date"), + userId = results.getLong("user_id"), + userNickname = results.getString("user_nickname").orEmpty(), + groupId = results.getLong("group_id").takeUnless { results.wasNull() }, + groupName = results.getString("group_name"), + inputUnits = results.getLong("input_units"), + outputUnits = results.getLong("output_units"), + totalUnits = results.getLong("total_units"), + promptTokens = results.getLong("prompt_tokens"), + completionTokens = results.getLong("completion_tokens"), + totalTokens = results.getLong("total_tokens"), + cachedTokens = results.getLong("cached_tokens"), + callCount = results.getInt("call_count"), + ) + ) + } + } + } + } + return actual == expected + } + + private fun TokenUsageDailyRecord.toLegacyImportState(): LegacyImportState = LegacyImportState( + occurredAt = LocalDate.parse(date, dateFmt).atStartOfDay(zone).toEpochSecond(), + date = date, + userId = userId, + userNickname = sanitizeNickname(userNickname), + groupId = groupId, + groupName = groupName?.let(::sanitizeNickname), + inputUnits = promptTokens.coerceAtLeast(0), + outputUnits = completionTokens.coerceAtLeast(0), + totalUnits = totalTokens.coerceAtLeast(0), + promptTokens = promptTokens.coerceAtLeast(0), + completionTokens = completionTokens.coerceAtLeast(0), + totalTokens = totalTokens.coerceAtLeast(0), + cachedTokens = cachedTokens.coerceAtLeast(0), + callCount = callCount.coerceAtLeast(0), + ) + + private fun readMeta(connection: Connection, key: String): String? = + connection.prepareStatement("SELECT value FROM token_usage_meta WHERE key = ?").use { statement -> + statement.setString(1, key) + statement.executeQuery().use { results -> if (results.next()) results.getString(1) else null } + } + + private fun writeMeta(connection: Connection, key: String, value: String) { + connection.prepareStatement( + "INSERT INTO token_usage_meta(key, value) VALUES (?, ?) " + + "ON CONFLICT(key) DO UPDATE SET value = excluded.value" + ).use { statement -> + statement.setString(1, key) + statement.setString(2, value) + statement.executeUpdate() + } + } + + private fun sha256(file: File): String { + val digest = MessageDigest.getInstance("SHA-256") + file.inputStream().buffered().use { input -> + val buffer = ByteArray(DEFAULT_BUFFER_SIZE) + while (true) { + val count = input.read(buffer) + if (count < 0) break + digest.update(buffer, 0, count) + } + } + val hex = "0123456789abcdef" + return buildString(64) { + digest.digest().forEach { byte -> + val value = byte.toInt() and 0xff + append(hex[value ushr 4]) + append(hex[value and 0x0f]) + } + } + } + + private fun legacyKey(record: TokenUsageDailyRecord): String = + listOf(record.date, record.userId, record.groupId ?: "private").joinToString("|") + + private fun dateFor(timestamp: Long): String = + LocalDate.ofInstant(Instant.ofEpochSecond(timestamp), zone).format(dateFmt) + + private fun providerFor(apiBaseUrl: String?): String? { + val host = runCatching { apiBaseUrl?.trim()?.let(java.net.URI::create)?.host?.lowercase() }.getOrNull() ?: return null + return when { + host == "api.deepseek.com" || host.endsWith(".deepseek.com") -> "deepseek" + else -> host } - save() } - /** 把控制字符压成空格,避免昵称里的换行/零宽字符把 JSON/展示弄乱。 */ private fun sanitizeNickname(s: String): String { if (s.isEmpty()) return s - val cleaned = buildString(s.length) { + return buildString(s.length) { for (c in s) { if (c == ' ' || (!c.isISOControl() && c.category != CharCategory.FORMAT)) append(c) else append(' ') } - } - return cleaned.trim().replace(Regex(" {2,}"), " ") + }.trim().replace(Regex(" {2,}"), " ") } - private fun save() { - try { - val tmp = File(file.parentFile, "${file.name}.tmp") - tmp.writeText(json.encodeToString(listSerializer, records)) - tmp.copyTo(file, overwrite = true) - tmp.delete() - } catch (_: Exception) { - // 写盘失败由日志/上层关心,这里不抛断对话流程 + private fun withWriteConnection(block: (Connection) -> T): T { + synchronized(writeLock) { + check(initialized) { "Token SQLite 尚未初始化" } + val connection = writeConnection?.takeUnless(Connection::isClosed) + ?: openConnection(databaseFile).also { + configureWriteConnection(it) + writeConnection = it + } + return block(connection) } } + + private fun createSchema(connection: Connection) { + val oldAutoCommit = connection.autoCommit + connection.autoCommit = false + try { + connection.createStatement().use { statement -> + statement.executeUpdate( + """ + CREATE TABLE IF NOT EXISTS token_usage_record( + id INTEGER PRIMARY KEY AUTOINCREMENT, + occurred_at INTEGER NOT NULL, + usage_date TEXT NOT NULL, + bot_id INTEGER, + user_id INTEGER NOT NULL, + user_nickname TEXT NOT NULL DEFAULT '', + group_id INTEGER, + group_name TEXT, + endpoint_label TEXT, + model_alias TEXT, + provider TEXT, + model TEXT, + usage_kind TEXT NOT NULL DEFAULT 'chat', + unit TEXT NOT NULL DEFAULT 'tokens', + input_units INTEGER NOT NULL DEFAULT 0, + output_units INTEGER NOT NULL DEFAULT 0, + total_units INTEGER NOT NULL DEFAULT 0, + prompt_tokens INTEGER NOT NULL DEFAULT 0, + completion_tokens INTEGER NOT NULL DEFAULT 0, + total_tokens INTEGER NOT NULL DEFAULT 0, + cached_tokens INTEGER NOT NULL DEFAULT 0, + call_count INTEGER NOT NULL DEFAULT 1, + detailed INTEGER NOT NULL DEFAULT 1, + legacy_key TEXT UNIQUE + ) + """.trimIndent() + ) + ensureColumn(connection, "token_usage_record", "model_alias", "TEXT") + ensureColumn(connection, "token_usage_record", "usage_kind", "TEXT NOT NULL DEFAULT 'chat'") + ensureColumn(connection, "token_usage_record", "unit", "TEXT NOT NULL DEFAULT 'tokens'") + ensureColumn(connection, "token_usage_record", "input_units", "INTEGER NOT NULL DEFAULT 0") + ensureColumn(connection, "token_usage_record", "output_units", "INTEGER NOT NULL DEFAULT 0") + ensureColumn(connection, "token_usage_record", "total_units", "INTEGER NOT NULL DEFAULT 0") + statement.executeUpdate( + """ + UPDATE token_usage_record + SET input_units = prompt_tokens, + output_units = completion_tokens, + total_units = total_tokens + WHERE unit = 'tokens' + AND input_units = 0 AND output_units = 0 AND total_units = 0 + """.trimIndent() + ) + statement.executeUpdate( + "UPDATE token_usage_record SET usage_kind = 'chat' WHERE usage_kind = 'tokens'" + ) + statement.executeUpdate("CREATE INDEX IF NOT EXISTS idx_token_usage_date ON token_usage_record(usage_date)") + statement.executeUpdate("CREATE INDEX IF NOT EXISTS idx_token_usage_user_date ON token_usage_record(user_id, usage_date)") + statement.executeUpdate("CREATE INDEX IF NOT EXISTS idx_token_usage_group_date ON token_usage_record(group_id, usage_date)") + statement.executeUpdate("CREATE INDEX IF NOT EXISTS idx_token_usage_model_date ON token_usage_record(model, usage_date)") + statement.executeUpdate( + "CREATE INDEX IF NOT EXISTS idx_token_usage_kind_date ON token_usage_record(usage_kind, unit, usage_date)" + ) + statement.executeUpdate( + """ + CREATE TABLE IF NOT EXISTS token_usage_meta( + key TEXT PRIMARY KEY, + value TEXT NOT NULL + ) + """.trimIndent() + ) + } + connection.prepareStatement( + "INSERT INTO token_usage_meta(key, value) VALUES ('schema_version', ?) " + + "ON CONFLICT(key) DO UPDATE SET value = excluded.value" + ).use { statement -> + statement.setString(1, SCHEMA_VERSION.toString()) + statement.executeUpdate() + } + connection.commit() + } catch (cause: Throwable) { + connection.rollback() + throw cause + } finally { + connection.autoCommit = oldAutoCommit + } + } + + private fun ensureColumn( + connection: Connection, + table: String, + column: String, + definition: String, + ) { + val exists = connection.createStatement().use { statement -> + statement.executeQuery("PRAGMA table_info($table)").use { results -> + generateSequence { if (results.next()) results.getString("name") else null } + .any { it == column } + } + } + if (!exists) { + connection.createStatement().use { statement -> + statement.executeUpdate("ALTER TABLE $table ADD COLUMN $column $definition") + } + } + } + + private fun openConnection(databaseFile: File): Connection = + DriverManager.getConnection("jdbc:sqlite:${databaseFile.absolutePath}") + + private fun openReadConnection(databaseFile: File): Connection { + val config = SQLiteConfig().apply { + setReadOnly(true) + setBusyTimeout(BUSY_TIMEOUT_MS) + } + return DriverManager.getConnection("jdbc:sqlite:${databaseFile.absolutePath}", config.toProperties()).also { + it.createStatement().use { statement -> statement.execute("PRAGMA query_only=ON") } + } + } + + private fun configureWriteConnection(connection: Connection) { + connection.createStatement().use { statement -> + statement.execute("PRAGMA journal_mode=WAL") + statement.execute("PRAGMA synchronous=NORMAL") + statement.execute("PRAGMA busy_timeout=$BUSY_TIMEOUT_MS") + statement.execute("PRAGMA wal_autocheckpoint=1000") + } + } + + private fun ResultSet.toTokenUsageRecord(): TokenUsageRecord = TokenUsageRecord( + id = getLong("id"), + timestamp = getLong("occurred_at"), + date = getString("usage_date"), + botId = getLong("bot_id").takeUnless { wasNull() }, + userId = getLong("user_id"), + userNickname = getString("user_nickname").orEmpty(), + groupId = getLong("group_id").takeUnless { wasNull() }, + groupName = getString("group_name"), + endpointLabel = getString("endpoint_label"), + modelAlias = getString("model_alias"), + provider = getString("provider"), + model = getString("model"), + usageKind = getString("usage_kind") ?: "chat", + unit = getString("unit") ?: "tokens", + inputUnits = getLong("input_units"), + outputUnits = getLong("output_units"), + totalUnits = getLong("total_units"), + promptTokens = getLong("prompt_tokens"), + completionTokens = getLong("completion_tokens"), + totalTokens = getLong("total_tokens"), + cachedTokens = getLong("cached_tokens"), + callCount = getInt("call_count"), + detailed = getInt("detailed") != 0, + ) + + private fun warn(message: String, cause: Throwable?) { + warningLogger?.invoke(message, cause) + } } diff --git a/src/main/kotlin/llm/LargeLanguageModels.kt b/src/main/kotlin/llm/LargeLanguageModels.kt index b629f41..3d2c97c 100644 --- a/src/main/kotlin/llm/LargeLanguageModels.kt +++ b/src/main/kotlin/llm/LargeLanguageModels.kt @@ -4,6 +4,7 @@ import kotlinx.serialization.json.Json import kotlinx.serialization.json.JsonObject import kotlinx.serialization.json.jsonObject import top.jie65535.mirai.JChatGPT +import top.jie65535.mirai.config.ModelConfig import top.jie65535.mirai.config.PluginConfig import kotlin.time.Duration.Companion.milliseconds @@ -25,16 +26,29 @@ object LargeLanguageModels { val temperature: Double?, /** 唯一标识,用于健康状态跟踪与日志 */ val label: String, + val alias: String = "", + val provider: String = "", ) data class ProfileEndpoint( val service: ModelService, val model: String, + val alias: String = "", + val provider: String = "", ) data class WebSummaryEndpoint( val service: ModelService, val model: String, + val alias: String = "", + val provider: String = "", + ) + + data class AuxiliaryEndpoint( + val service: ModelService, + val model: String, + val alias: String = "", + val provider: String = "", ) /** @@ -52,12 +66,12 @@ object LargeLanguageModels { /** * 推理模型 */ - var reasoning: ModelService? = null + var reasoning: AuxiliaryEndpoint? = null /** * 视觉模型 */ - var visual: ModelService? = null + var visual: AuxiliaryEndpoint? = null /** 历史用户画像分析模型。 */ var profile: ProfileEndpoint? = null @@ -116,28 +130,68 @@ object LargeLanguageModels { } fun reload() { + (ModelCatalog.validationIssues() + ModelCatalog.roleValidationIssues()) + .distinct() + .forEach(JChatGPT.logger::warning) val timeout = PluginConfig.timeout.milliseconds val firstChunkTimeout = PluginConfig.firstChunkTimeout.milliseconds // 初始化聊天接入点(主 + 备用),并重置健康状态 cooldownUntil.clear() val endpoints = mutableListOf() - if (PluginConfig.openAiApi.isNotBlank() && PluginConfig.openAiToken.isNotBlank()) { - endpoints.add( - ChatEndpoint( - service = ModelService( - baseUrl = PluginConfig.openAiApi, - token = PluginConfig.openAiToken, - timeout = timeout, - firstChunkTimeout = firstChunkTimeout, - extraBody = parseExtraBody(PluginConfig.chatModelExtraBody) - ), - model = PluginConfig.chatModel, + val primaryAlias = PluginConfig.chatModelAlias.trim() + if (primaryAlias.isNotEmpty()) { + resolveOpenAi(primaryAlias)?.let { definition -> + endpoints += ChatEndpoint( + service = modelService(definition, timeout, firstChunkTimeout), + model = definition.model, temperature = PluginConfig.chatTemperature, - label = "primary", + label = "primary:$primaryAlias", + alias = primaryAlias, + provider = definition.provider, ) + } + } + var legacyPrimaryUsed = false + if (endpoints.isEmpty() && PluginConfig.openAiApi.isNotBlank() && PluginConfig.openAiToken.isNotBlank()) { + endpoints += ChatEndpoint( + service = ModelService( + baseUrl = ModelCatalog.normalizeOpenAiApi(PluginConfig.openAiApi), + token = PluginConfig.openAiToken, + timeout = timeout, + firstChunkTimeout = firstChunkTimeout, + extraBody = parseExtraBody(PluginConfig.chatModelExtraBody) + ), + model = PluginConfig.chatModel, + temperature = PluginConfig.chatTemperature, + label = "primary", + alias = "legacy-primary", + provider = providerName(PluginConfig.openAiApi), ) + legacyPrimaryUsed = true + } + PluginConfig.chatFallbackModelAliases.map(String::trim) + .filter(String::isNotEmpty) + .forEach { alias -> + resolveOpenAi(alias)?.let { definition -> + val label = if (endpoints.isEmpty()) { + "primary:$alias" + } else { + "fallback${endpoints.size - 1}:$alias" + } + endpoints += ChatEndpoint( + service = modelService(definition, timeout, firstChunkTimeout), + model = definition.model, + temperature = PluginConfig.chatTemperature, + label = label, + alias = alias, + provider = definition.provider, + ) + } + } + + if (legacyPrimaryUsed) { // 备用接入点:留空字段继承主接入点配置 PluginConfig.chatFallbacks.forEachIndexed { i, fb -> val api = fb.api.ifBlank { PluginConfig.openAiApi } @@ -145,19 +199,19 @@ object LargeLanguageModels { val model = fb.model.ifBlank { PluginConfig.chatModel } val extraBody = fb.extraBody.ifBlank { PluginConfig.chatModelExtraBody } if (api.isNotBlank() && token.isNotBlank()) { - endpoints.add( - ChatEndpoint( - service = ModelService( - baseUrl = api, - token = token, - timeout = timeout, - firstChunkTimeout = firstChunkTimeout, - extraBody = parseExtraBody(extraBody) - ), - model = model, - temperature = PluginConfig.chatTemperature, - label = "fallback$i:$model", - ) + endpoints += ChatEndpoint( + service = ModelService( + baseUrl = ModelCatalog.normalizeOpenAiApi(api), + token = token, + timeout = timeout, + firstChunkTimeout = firstChunkTimeout, + extraBody = parseExtraBody(extraBody) + ), + model = model, + temperature = PluginConfig.chatTemperature, + label = "fallback${endpoints.size - 1}:legacy-$i:$model", + alias = "legacy-fallback$i", + provider = providerName(api), ) } } @@ -166,69 +220,144 @@ object LargeLanguageModels { profile = null if (PluginConfig.profileEnabled) { - val api = PluginConfig.profileModelApi.ifBlank { PluginConfig.openAiApi } - val token = PluginConfig.profileModelToken.ifBlank { PluginConfig.openAiToken } - val model = PluginConfig.profileModel.ifBlank { PluginConfig.chatModel } - val extraBody = PluginConfig.profileModelExtraBody.ifBlank { PluginConfig.chatModelExtraBody } - if (api.isNotBlank() && token.isNotBlank() && model.isNotBlank()) { + val profileAlias = PluginConfig.profileModelAlias.ifBlank { PluginConfig.chatModelAlias } + val definition = profileAlias.trim().takeIf(String::isNotEmpty)?.let(::resolveOpenAi) + if (definition != null) { val profileFirstChunk = PluginConfig.profileFirstChunkTimeout.milliseconds profile = ProfileEndpoint( - service = ModelService( - baseUrl = api, - token = token, + service = modelService( + definition, timeout = maxOf(timeout, profileFirstChunk), firstChunkTimeout = profileFirstChunk, - extraBody = parseExtraBody(extraBody), maxConcurrentRequests = PluginConfig.profileMaxConcurrentRequests, - ), - model = model, - ) + ), + model = definition.model, + alias = definition.alias, + provider = definition.provider, + ) + } else { + val api = PluginConfig.profileModelApi.ifBlank { PluginConfig.openAiApi } + val token = PluginConfig.profileModelToken.ifBlank { PluginConfig.openAiToken } + val model = PluginConfig.profileModel.ifBlank { PluginConfig.chatModel } + val extraBody = PluginConfig.profileModelExtraBody.ifBlank { PluginConfig.chatModelExtraBody } + if (api.isNotBlank() && token.isNotBlank() && model.isNotBlank()) { + val profileFirstChunk = PluginConfig.profileFirstChunkTimeout.milliseconds + profile = ProfileEndpoint( + service = ModelService( + baseUrl = ModelCatalog.normalizeOpenAiApi(api), + token = token, + timeout = maxOf(timeout, profileFirstChunk), + firstChunkTimeout = profileFirstChunk, + extraBody = parseExtraBody(extraBody), + maxConcurrentRequests = PluginConfig.profileMaxConcurrentRequests, + ), + model = model, + alias = "legacy-profile", + provider = providerName(api), + ) + } } } webSummary = null - if (PluginConfig.webSummaryModelApi.isNotBlank() && + val webSummaryDefinition = PluginConfig.webSummaryModelAlias.trim().takeIf(String::isNotEmpty) + ?.let(::resolveOpenAi) + if (webSummaryDefinition != null) { + val webSummaryFirstChunk = PluginConfig.webSummaryFirstChunkTimeout.milliseconds + webSummary = WebSummaryEndpoint( + service = modelService( + webSummaryDefinition, + timeout = maxOf(timeout, webSummaryFirstChunk), + firstChunkTimeout = webSummaryFirstChunk, + ), + model = webSummaryDefinition.model, + alias = webSummaryDefinition.alias, + provider = webSummaryDefinition.provider, + ) + } else if (PluginConfig.webSummaryModelApi.isNotBlank() && PluginConfig.webSummaryModelToken.isNotBlank() && PluginConfig.webSummaryModel.isNotBlank() ) { val webSummaryFirstChunk = PluginConfig.webSummaryFirstChunkTimeout.milliseconds webSummary = WebSummaryEndpoint( service = ModelService( - baseUrl = PluginConfig.webSummaryModelApi, + baseUrl = ModelCatalog.normalizeOpenAiApi(PluginConfig.webSummaryModelApi), token = PluginConfig.webSummaryModelToken, timeout = maxOf(timeout, webSummaryFirstChunk), firstChunkTimeout = webSummaryFirstChunk, extraBody = parseExtraBody(PluginConfig.webSummaryModelExtraBody), ), model = PluginConfig.webSummaryModel, + alias = "legacy-web-summary", + provider = providerName(PluginConfig.webSummaryModelApi), ) } // 初始化推理模型 - if (PluginConfig.reasoningModelApi.isNotBlank() && PluginConfig.reasoningModelToken.isNotBlank()) { + reasoning = null + val reasoningDefinition = PluginConfig.reasoningModelAlias.trim().takeIf(String::isNotEmpty) + ?.let(::resolveOpenAi) + if (reasoningDefinition != null) { + val reasoningFirstChunk = PluginConfig.reasoningFirstChunkTimeout.milliseconds + reasoning = AuxiliaryEndpoint( + service = modelService( + reasoningDefinition, + timeout = maxOf(timeout, reasoningFirstChunk), + firstChunkTimeout = reasoningFirstChunk, + ), + model = reasoningDefinition.model, + alias = reasoningDefinition.alias, + provider = reasoningDefinition.provider, + ) + } else if (PluginConfig.reasoningModelApi.isNotBlank() && PluginConfig.reasoningModelToken.isNotBlank()) { // 推理模型出首块前常有思考预热,比对话慢,使用单独放宽的首块超时; // socket 超时(两次读间隔,等首块时也归它管)不能小于首块预算,否则首块超时形同虚设 val reasoningFirstChunk = PluginConfig.reasoningFirstChunkTimeout.milliseconds - reasoning = ModelService( - baseUrl = PluginConfig.reasoningModelApi, - token = PluginConfig.reasoningModelToken, - timeout = maxOf(timeout, reasoningFirstChunk), - firstChunkTimeout = reasoningFirstChunk, - extraBody = parseExtraBody(PluginConfig.reasoningModelExtraBody) + reasoning = AuxiliaryEndpoint( + service = ModelService( + baseUrl = ModelCatalog.normalizeOpenAiApi(PluginConfig.reasoningModelApi), + token = PluginConfig.reasoningModelToken, + timeout = maxOf(timeout, reasoningFirstChunk), + firstChunkTimeout = reasoningFirstChunk, + extraBody = parseExtraBody(PluginConfig.reasoningModelExtraBody) + ), + model = PluginConfig.reasoningModel, + alias = "legacy-reasoning", + provider = providerName(PluginConfig.reasoningModelApi), ) } // 初始化视觉模型 - if (PluginConfig.visualModelApi.isNotBlank() && PluginConfig.visualModelToken.isNotBlank()) { + visual = null + val visualDefinition = PluginConfig.visualModelAlias.trim().takeIf(String::isNotEmpty) + ?.let(::resolveOpenAi) + if (visualDefinition != null) { + val visualFirstChunk = PluginConfig.visualFirstChunkTimeout.milliseconds + visual = AuxiliaryEndpoint( + service = modelService( + visualDefinition, + timeout = maxOf(timeout, visualFirstChunk), + firstChunkTimeout = visualFirstChunk, + ), + model = visualDefinition.model, + alias = visualDefinition.alias, + provider = visualDefinition.provider, + ) + } else if (PluginConfig.visualModelApi.isNotBlank() && PluginConfig.visualModelToken.isNotBlank()) { // 视觉模型需服务端先下载图片再出首块,比对话天然慢,使用单独放宽的首块超时; // socket 超时(两次读间隔,等首块时也归它管)不能小于首块预算,否则首块超时形同虚设 val visualFirstChunk = PluginConfig.visualFirstChunkTimeout.milliseconds - visual = ModelService( - baseUrl = PluginConfig.visualModelApi, - token = PluginConfig.visualModelToken, - timeout = maxOf(timeout, visualFirstChunk), - firstChunkTimeout = visualFirstChunk, - extraBody = parseExtraBody(PluginConfig.visualModelExtraBody) + visual = AuxiliaryEndpoint( + service = ModelService( + baseUrl = ModelCatalog.normalizeOpenAiApi(PluginConfig.visualModelApi), + token = PluginConfig.visualModelToken, + timeout = maxOf(timeout, visualFirstChunk), + firstChunkTimeout = visualFirstChunk, + extraBody = parseExtraBody(PluginConfig.visualModelExtraBody) + ), + model = PluginConfig.visualModel, + alias = "legacy-visual", + provider = providerName(PluginConfig.visualModelApi), ) } @@ -249,4 +378,38 @@ object LargeLanguageModels { } } } + + private fun resolveOpenAi(alias: String): ResolvedModelDefinition? { + val definition = ModelCatalog.resolve(alias) + if (definition == null) { + JChatGPT.logger.warning("Models.yml 中不存在模型别名:$alias") + return null + } + if (definition.providerType !in setOf("openai", "openai-compatible", "openai_compatible")) { + JChatGPT.logger.warning("模型别名 $alias 的提供商类型 ${definition.providerType} 不能用于 OpenAI 兼容客户端") + return null + } + if (definition.api.isBlank() || definition.token.isBlank() || definition.model.isBlank()) { + JChatGPT.logger.warning("模型别名 $alias 的 provider/api/token/model 配置不完整") + return null + } + return definition + } + + private fun modelService( + definition: ResolvedModelDefinition, + timeout: kotlin.time.Duration, + firstChunkTimeout: kotlin.time.Duration, + maxConcurrentRequests: Int? = null, + ): ModelService = ModelService( + baseUrl = ModelCatalog.normalizeOpenAiApi(definition.api), + token = definition.token, + timeout = timeout, + firstChunkTimeout = firstChunkTimeout, + extraBody = parseExtraBody(definition.extraBody), + maxConcurrentRequests = maxConcurrentRequests, + ) + + private fun providerName(api: String): String = + runCatching { java.net.URI.create(api.trim()).host.orEmpty() }.getOrDefault("") } diff --git a/src/main/kotlin/llm/ModelCatalog.kt b/src/main/kotlin/llm/ModelCatalog.kt new file mode 100644 index 0000000..e32e0bd --- /dev/null +++ b/src/main/kotlin/llm/ModelCatalog.kt @@ -0,0 +1,187 @@ +package top.jie65535.mirai.llm + +import top.jie65535.mirai.config.ModelConfig +import top.jie65535.mirai.config.ModelDefinition +import top.jie65535.mirai.config.ModelProviderDefinition +import top.jie65535.mirai.config.PluginConfig + +data class ResolvedModelDefinition( + val alias: String, + val provider: String, + val providerType: String, + val api: String, + val token: String, + val model: String, + val extraBody: String, +) + +/** Resolves shared model aliases and keeps legacy Config.yml fallback logic in one place. */ +object ModelCatalog { + private const val DEFAULT_DASHSCOPE_IMAGE_API = + "https://dashscope.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation" + + private val openAiProviderTypes = setOf( + "openai", + "openai-compatible", + "openai_compatible", + ) + private val supportedProviderTypes = openAiProviderTypes + setOf( + "dashscope", + ) + + fun resolve(alias: String): ResolvedModelDefinition? = + resolve(alias, ModelConfig.providers, ModelConfig.models) + + internal fun resolve( + alias: String, + providers: List, + models: List, + ): ResolvedModelDefinition? { + val normalizedAlias = alias.trim() + if (normalizedAlias.isEmpty()) return null + val model = models.filter { it.name.trim() == normalizedAlias }.singleOrNull() ?: return null + val provider = providers.filter { it.name.trim() == model.provider.trim() }.singleOrNull() + ?: return null + return model.resolve(provider, normalizedAlias) + } + + fun validationIssues(): List = validationIssues(ModelConfig.providers, ModelConfig.models) + + fun roleValidationIssues(): List { + val bindings = buildList { + add(Triple("主聊天", PluginConfig.chatModelAlias, openAiProviderTypes)) + PluginConfig.chatFallbackModelAliases.forEachIndexed { index, alias -> + add(Triple("聊天备用 ${index + 1}", alias, openAiProviderTypes)) + } + add( + Triple( + "画像", + PluginConfig.profileModelAlias.ifBlank { PluginConfig.chatModelAlias }, + openAiProviderTypes, + ) + ) + add(Triple("推理", PluginConfig.reasoningModelAlias, openAiProviderTypes)) + add(Triple("视觉", PluginConfig.visualModelAlias, openAiProviderTypes)) + add(Triple("网页摘要", PluginConfig.webSummaryModelAlias, openAiProviderTypes)) + add(Triple("图像", PluginConfig.imageModelAlias, setOf("dashscope"))) + add(Triple("TTS", PluginConfig.ttsModelAlias, setOf("dashscope"))) + } + return bindings.mapNotNull { (role, alias, allowedTypes) -> + bindingValidationIssue( + role = role, + alias = alias, + allowedTypes = allowedTypes, + providers = ModelConfig.providers, + models = ModelConfig.models, + ) + }.distinct() + } + + internal fun bindingValidationIssue( + role: String, + alias: String, + allowedTypes: Set, + providers: List, + models: List, + ): String? { + val normalizedAlias = alias.trim() + if (normalizedAlias.isEmpty()) return null + val definition = resolve(normalizedAlias, providers, models) + ?: return "Models.yml 的 $role 角色引用了无效模型别名:$normalizedAlias" + return if (definition.providerType !in allowedTypes) { + "Models.yml 的 $role 角色不能使用 provider type ${definition.providerType}:$normalizedAlias" + } else { + null + } + } + + internal fun validationIssues( + providers: List, + models: List, + ): List = buildList { + val providerNames = providers.map { it.name.trim() } + providerNames.filter(String::isEmpty).forEach { add("Models.yml 存在空 provider 名称") } + providerNames.groupingBy(String::toString).eachCount() + .filterValues { it > 1 } + .keys + .forEach { add("Models.yml provider 名称重复:$it") } + + val modelNames = models.map { it.name.trim() } + modelNames.filter(String::isEmpty).forEach { add("Models.yml 存在空模型别名") } + modelNames.groupingBy(String::toString).eachCount() + .filterValues { it > 1 } + .keys + .forEach { add("Models.yml 模型别名重复:$it") } + + providers.forEach { provider -> + val name = provider.name.trim().ifBlank { "" } + val type = provider.type.trim().lowercase().ifBlank { "openai" } + if (type !in supportedProviderTypes) add("Models.yml provider $name 的 type 不受支持:$type") + if (type != "dashscope" && provider.api.isBlank()) add("Models.yml provider $name 未配置 api") + if (provider.token.isBlank()) add("Models.yml provider $name 未配置 token") + } + models.forEach { model -> + val alias = model.name.trim().ifBlank { "" } + val providerName = model.provider.trim() + if (providerName.isEmpty() || providerNames.count { it == providerName } != 1) { + add("Models.yml 模型 $alias 引用的 provider 无效:${providerName.ifBlank { "" }}") + } + if (model.model.isBlank()) add("Models.yml 模型 $alias 未配置实际模型名") + } + }.distinct() + + fun resolveImage(): ResolvedModelDefinition? = + resolveDashScope(PluginConfig.imageModelAlias) + ?: legacy( + alias = "legacy-image", + provider = "dashscope", + providerType = "dashscope", + api = DEFAULT_DASHSCOPE_IMAGE_API, + token = PluginConfig.dashScopeApiKey, + model = PluginConfig.imageModel, + ) + + fun resolveTts(): ResolvedModelDefinition? = + resolveDashScope(PluginConfig.ttsModelAlias) + ?: legacy( + alias = "legacy-tts", + provider = "dashscope", + providerType = "dashscope", + api = DEFAULT_DASHSCOPE_IMAGE_API, + token = PluginConfig.dashScopeApiKey, + model = PluginConfig.ttsModel, + ) + + fun normalizeOpenAiApi(api: String): String = api.trim().trimEnd('/') + "/" + + private fun resolveDashScope(alias: String): ResolvedModelDefinition? = + resolve(alias) + ?.takeIf { it.providerType == "dashscope" && it.token.isNotBlank() && it.model.isNotBlank() } + ?.let { it.copy(api = it.api.ifBlank { DEFAULT_DASHSCOPE_IMAGE_API }) } + + private fun ModelDefinition.resolve( + provider: ModelProviderDefinition, + alias: String, + ): ResolvedModelDefinition = ResolvedModelDefinition( + alias = alias, + provider = provider.name.trim(), + providerType = provider.type.trim().lowercase().ifBlank { "openai" }, + api = provider.api.trim(), + token = provider.token.trim(), + model = model.trim(), + extraBody = extraBody, + ) + + private fun legacy( + alias: String, + provider: String, + providerType: String, + api: String, + token: String, + model: String, + extraBody: String = "", + ): ResolvedModelDefinition? { + if (api.isBlank() || token.isBlank() || model.isBlank()) return null + return ResolvedModelDefinition(alias, provider, providerType, api, token, model.trim(), extraBody) + } +} diff --git a/src/main/kotlin/profile/ProfileModelClient.kt b/src/main/kotlin/profile/ProfileModelClient.kt index 77be439..5cd5788 100644 --- a/src/main/kotlin/profile/ProfileModelClient.kt +++ b/src/main/kotlin/profile/ProfileModelClient.kt @@ -8,6 +8,9 @@ import com.aallam.openai.api.core.Usage import com.aallam.openai.api.model.ModelId import kotlinx.serialization.SerializationException import kotlinx.serialization.json.Json +import net.mamoe.mirai.message.data.MessageSourceKind +import top.jie65535.mirai.data.ModelUsageAttribution +import top.jie65535.mirai.data.ModelUsageRecorder import top.jie65535.mirai.llm.LargeLanguageModels import top.jie65535.mirai.llm.ModelService @@ -86,6 +89,8 @@ class ProfileModelClient( ), ) ) + recordUsage(batch.usageAttribution(), completion) + require(completion.content.isNotBlank()) { "模型流式响应没有文本内容" } val raw = completion.content.replace(THINK_REGEX, "").trim() val response = parseResponse(raw) return ProfileModelResult( @@ -125,6 +130,15 @@ class ProfileModelClient( ), ) ) + recordUsage( + ModelUsageAttribution( + botId = batch.botId, + userId = 0, + groupId = batch.groupId, + ), + completion, + ) + require(completion.content.isNotBlank()) { "模型流式响应没有文本内容" } val raw = completion.content.replace(THINK_REGEX, "").trim() return ConversationProfileModelResult( response = parseObject(raw), @@ -148,6 +162,13 @@ class ProfileModelClient( ), ) ) + recordUsage( + ModelUsageAttribution( + userId = profile.userId, + ), + completion, + ) + require(completion.content.isNotBlank()) { "模型流式响应没有文本内容" } val raw = completion.content.replace(THINK_REGEX, "").trim() return ProfileCompactionModelResult( response = json.decodeFromString(extractObject(raw)), @@ -164,10 +185,10 @@ class ProfileModelClient( chunk.choices.firstOrNull()?.delta?.content?.let(content::append) chunk.usage?.let { lastUsage = it } } - require(content.isNotBlank()) { "模型流式响应没有文本内容" } return CompletedProfileResponse( content = content.toString(), usage = lastUsage.toProfileUsage(cacheUsage), + usageAvailable = lastUsage != null, ) } @@ -177,6 +198,31 @@ class ProfileModelClient( cachedTokens = cacheUsage?.hitTokens ?: 0, ) + private fun recordUsage(attribution: ModelUsageAttribution, completion: CompletedProfileResponse) { + if (!completion.usageAvailable) return + val usage = completion.usage + ModelUsageRecorder.recordTokenValues( + attribution = attribution, + endpointLabel = "profile", + modelAlias = endpoint.alias, + provider = endpoint.provider, + model = endpoint.model, + usageKind = "profile", + promptTokens = usage.promptTokens.toLong(), + completionTokens = usage.completionTokens.toLong(), + cachedTokens = usage.cachedTokens.toLong(), + ) + } + + private fun ProfileHistoryBatch.usageAttribution(): ModelUsageAttribution { + val record = messages.firstOrNull()?.record + return ModelUsageAttribution( + botId = record?.botId ?: 0, + userId = userId, + groupId = record?.targetId?.takeIf { record.kind == MessageSourceKind.GROUP }, + ) + } + private fun parseResponse(raw: String): ProfileModelResponse { return parseObject(raw) } @@ -206,5 +252,6 @@ class ProfileModelClient( private data class CompletedProfileResponse( val content: String, val usage: ProfileTokenUsage, + val usageAvailable: Boolean, ) } diff --git a/src/main/kotlin/tools/ImageAgent.kt b/src/main/kotlin/tools/ImageAgent.kt index 34402e6..c55b608 100644 --- a/src/main/kotlin/tools/ImageAgent.kt +++ b/src/main/kotlin/tools/ImageAgent.kt @@ -17,12 +17,15 @@ import kotlinx.serialization.json.int import kotlinx.serialization.json.jsonArray import kotlinx.serialization.json.jsonObject import kotlinx.serialization.json.jsonPrimitive +import kotlinx.serialization.json.longOrNull import kotlinx.serialization.json.put import kotlinx.serialization.json.putJsonArray import kotlinx.serialization.json.putJsonObject import net.mamoe.mirai.event.events.MessageEvent import top.jie65535.mirai.JChatGPT import top.jie65535.mirai.config.PluginConfig +import top.jie65535.mirai.data.ModelUsageRecorder +import top.jie65535.mirai.llm.ModelCatalog class ImageAgent : BaseAgent( tool = Tool.function( @@ -59,13 +62,15 @@ class ImageAgent : BaseAgent( } override val isEnabled: Boolean - get() = PluginConfig.dashScopeApiKey.isNotEmpty() + get() = ModelCatalog.resolveImage() != null override val loadingMessage: String get() = "作图中..." override suspend fun execute(args: JsonObject?, event: MessageEvent): String { requireNotNull(args) + val modelDefinition = ModelCatalog.resolveImage() + ?: return "未配置图像模型,无法生成图片。" val prompt = args.getValue("prompt").jsonPrimitive.content val imageIndices = args["image_indices"]?.jsonArray ?.map { it.jsonPrimitive.int } @@ -76,11 +81,11 @@ class ImageAgent : BaseAgent( ?: throw IllegalArgumentException("图片编号[$imageIndex]不存在或已失效") } - val response = httpClient.post(API_URL) { + val response = httpClient.post(modelDefinition.api.ifBlank { API_URL }) { contentType(ContentType("application", "json")) - header("Authorization", "Bearer " + PluginConfig.dashScopeApiKey) + header("Authorization", "Bearer " + modelDefinition.token) setBody(buildJsonObject { - put("model", PluginConfig.imageModel) + put("model", modelDefinition.model) putJsonObject("input") { putJsonArray("messages") { addJsonObject { @@ -115,6 +120,21 @@ class ImageAgent : BaseAgent( .getValue("message").jsonObject .getValue("content").jsonArray[0].jsonObject .getValue("image").jsonPrimitive.content + val outputImages = (responseObject["usage"] as? JsonObject) + ?.get("image_count")?.jsonPrimitive?.longOrNull + ?.coerceAtLeast(1) + ?: 1L + ModelUsageRecorder.recordUnits( + event = event, + endpointLabel = "image", + modelAlias = modelDefinition.alias, + provider = modelDefinition.provider, + model = modelDefinition.model, + usageKind = "image", + unit = "images", + inputUnits = imageUrls.size.toLong(), + outputUnits = outputImages, + ) "图片已生成,发送时请务必包含完整的url和查询参数,因为下载地址存在鉴权:![图片]($url)" } catch (e: Throwable) { JChatGPT.logger.error("图像生成结果解析异常", e) diff --git a/src/main/kotlin/tools/QueryTokenUsageAgent.kt b/src/main/kotlin/tools/QueryTokenUsageAgent.kt new file mode 100644 index 0000000..164ba28 --- /dev/null +++ b/src/main/kotlin/tools/QueryTokenUsageAgent.kt @@ -0,0 +1,453 @@ +package top.jie65535.mirai.tools + +import com.aallam.openai.api.chat.Tool +import com.aallam.openai.api.core.Parameters +import io.ktor.client.plugins.timeout +import io.ktor.client.request.get +import io.ktor.client.request.header +import io.ktor.client.statement.bodyAsText +import io.ktor.http.HttpHeaders +import io.ktor.http.isSuccess +import kotlinx.serialization.json.Json +import kotlinx.serialization.json.JsonObject +import kotlinx.serialization.json.add +import kotlinx.serialization.json.addJsonObject +import kotlinx.serialization.json.buildJsonObject +import kotlinx.serialization.json.contentOrNull +import kotlinx.serialization.json.intOrNull +import kotlinx.serialization.json.jsonArray +import kotlinx.serialization.json.jsonObject +import kotlinx.serialization.json.jsonPrimitive +import kotlinx.serialization.json.longOrNull +import kotlinx.serialization.json.put +import kotlinx.serialization.json.putJsonArray +import kotlinx.serialization.json.putJsonObject +import net.mamoe.mirai.event.events.GroupMessageEvent +import net.mamoe.mirai.event.events.MessageEvent +import top.jie65535.mirai.config.ModelConfig +import top.jie65535.mirai.config.ModelDefinition +import top.jie65535.mirai.config.ModelProviderDefinition +import top.jie65535.mirai.config.PluginConfig +import top.jie65535.mirai.data.TokenUsageRecord +import top.jie65535.mirai.data.TokenUsageStore +import top.jie65535.mirai.data.TokenUsageSummary +import top.jie65535.mirai.llm.ModelCatalog +import java.net.URI +import java.time.LocalDate +import java.time.ZoneId + +class QueryTokenUsageAgent : BaseAgent( + tool = Tool.function( + name = "queryTokenUsage", + description = "查询当前会话的模型用量或提供商余额。", + parameters = Parameters.buildJsonObject { + put("type", "object") + putJsonObject("properties") { + putJsonObject("operation") { + put("type", "string") + putJsonArray("enum") { + add("summary") + add("details") + add("balance") + add("overview") + } + put("description", "summary用量,details明细,balance余额,overview为用量加余额") + } + putJsonObject("days") { + put("type", "integer") + put("description", "统计最近多少天,包含今天,默认7,范围1到3650") + } + putJsonObject("limit") { + put("type", "integer") + put("description", "排名或最近明细条数,默认20,最多100") + } + putJsonObject("userId") { + put("type", "integer") + put("description", "群聊中可选,仅统计当前群内指定用户QQ号;私聊中忽略此参数并固定为当前私聊对象") + } + putJsonObject("usageType") { + put("type", "string") + putJsonArray("enum") { + add("chat") + add("profile") + add("reasoning") + add("visual") + add("web_summary") + add("image") + add("tts") + } + put("description", "可选,按模型用途筛选") + } + } + } + ) +) { + companion object { + private const val MAX_DAYS = 3650 + private const val MAX_DETAILS = 100 + private const val DEEPSEEK_BALANCE_URL = "https://api.deepseek.com/user/balance" + private val OPERATIONS = setOf("summary", "details", "balance", "overview") + private val USAGE_TYPES = setOf("chat", "profile", "reasoning", "visual", "web_summary", "image", "tts") + private val json = Json { ignoreUnknownKeys = true; explicitNulls = false } + + internal data class QueryScope( + val botId: Long, + val userId: Long?, + val groupId: Long?, + val privateOnly: Boolean, + ) + + internal data class BalanceAccount( + val name: String, + val api: String, + val token: String, + ) + + internal fun queryScope( + botId: Long, + senderId: Long, + currentGroupId: Long?, + requestedUserId: Long?, + ): QueryScope = if (currentGroupId == null) { + QueryScope(botId, senderId, null, privateOnly = true) + } else { + QueryScope(botId, requestedUserId, currentGroupId, privateOnly = false) + } + + internal fun isDeepSeekApi(api: String): Boolean { + val host = runCatching { URI.create(api.trim()).host?.lowercase() }.getOrNull() ?: return false + return host == "api.deepseek.com" || host.endsWith(".deepseek.com") + } + + internal fun parseDeepSeekBalance(body: String): JsonObject { + val root = json.parseToJsonElement(body).jsonObject + return buildJsonObject { + put("available", root["is_available"]?.jsonPrimitive?.contentOrNull?.toBooleanStrictOrNull() ?: false) + putJsonArray("balances") { + root["balance_infos"]?.jsonArray?.forEach { element -> + val info = element.jsonObject + addJsonObject { + put("currency", info["currency"]?.jsonPrimitive?.contentOrNull.orEmpty()) + put("total", info["total_balance"]?.jsonPrimitive?.contentOrNull.orEmpty()) + put("granted", info["granted_balance"]?.jsonPrimitive?.contentOrNull.orEmpty()) + put("toppedUp", info["topped_up_balance"]?.jsonPrimitive?.contentOrNull.orEmpty()) + } + } + } + } + } + + internal fun collectDeepSeekBalanceAccounts( + providers: List, + models: List, + legacyAccounts: List = emptyList(), + ): List { + val referencedProviders = models.mapTo(HashSet()) { it.provider.trim() } + val accounts = providers.asSequence() + .filter { it.name.trim() in referencedProviders } + .filter { it.token.isNotBlank() && isDeepSeekApi(it.api) } + .map { BalanceAccount(it.name.trim().ifBlank { "deepseek" }, it.api.trim(), it.token.trim()) } + .toMutableList() + accounts += legacyAccounts.filter { it.token.isNotBlank() && isDeepSeekApi(it.api) } + .map { it.copy(api = it.api.trim(), token = it.token.trim()) } + return accounts.distinctBy { it.token } + } + } + + override val isEnabled: Boolean + get() = TokenUsageStore.isAvailable || + configuredDeepSeekAccounts().isNotEmpty() + + override val loadingMessage: String + get() = "查询模型用量中..." + + override suspend fun execute(args: JsonObject?, event: MessageEvent): String { + val operation = args?.get("operation")?.jsonPrimitive?.contentOrNull?.lowercase() ?: "summary" + require(operation in OPERATIONS) { "不支持的模型用量查询操作:$operation" } + val days = (args?.get("days")?.jsonPrimitive?.intOrNull ?: 7).coerceIn(1, MAX_DAYS) + val limit = (args?.get("limit")?.jsonPrimitive?.intOrNull ?: 20).coerceIn(1, MAX_DETAILS) + val requestedUserId = args?.get("userId")?.jsonPrimitive?.longOrNull?.takeIf { it > 0 } + val usageType = args?.get("usageType")?.jsonPrimitive?.contentOrNull?.lowercase() + require(usageType == null || usageType in USAGE_TYPES) { "不支持的模型用途:$usageType" } + val scope = queryScope( + botId = event.bot.id, + senderId = event.sender.id, + currentGroupId = (event as? GroupMessageEvent)?.group?.id, + requestedUserId = requestedUserId, + ) + + val result = buildJsonObject { + put("operation", operation) + if (operation == "summary" || operation == "details" || operation == "overview") { + if (!TokenUsageStore.isAvailable) { + put("usageError", "模型用量 SQLite 尚未初始化") + } else { + val startDate = LocalDate.now(ZoneId.systemDefault()).minusDays((days - 1).toLong()).toString() + val summary = TokenUsageStore.summary( + startDate = startDate, + botId = scope.botId, + userId = scope.userId, + groupId = scope.groupId, + privateOnly = scope.privateOnly, + usageKind = usageType, + rankingLimit = limit, + ) + putJsonObject("usage") { writeSummary(summary, days) } + if (operation == "details") { + putJsonArray("details") { + TokenUsageStore.recent( + limit = limit, + startDate = startDate, + botId = scope.botId, + userId = scope.userId, + groupId = scope.groupId, + privateOnly = scope.privateOnly, + usageKind = usageType, + ).forEach { record -> addRecord(record) } + } + } + } + } + if (operation == "balance" || operation == "overview") { + put("balance", queryBalances()) + } + } + return result.toString() + } + + private fun kotlinx.serialization.json.JsonObjectBuilder.writeSummary(summary: TokenUsageSummary, days: Int) { + put("days", days) + put("promptTokens", summary.promptTokens) + put("completionTokens", summary.completionTokens) + put("totalTokens", summary.totalTokens) + put("cachedTokens", summary.cachedTokens) + put("tokenCallCount", summary.callCount) + put("allCallCount", summary.allCallCount) + put("activeUsers", summary.activeUsers) + put("allActiveUsers", summary.allActiveUsers) + put("todayTotalTokens", summary.todayTotal) + put("cacheHitRatePercent", if (summary.promptTokens > 0) summary.cachedTokens * 100.0 / summary.promptTokens else 0.0) + putJsonArray("daily") { + summary.daily.forEach { daily -> + addJsonObject { + put("date", daily.date) + put("totalTokens", daily.totalTokens) + } + } + } + putJsonArray("usageDaily") { + summary.usageDaily.forEach { daily -> + addJsonObject { + put("date", daily.date) + put("usageKind", daily.usageKind) + put("unit", daily.unit) + put("totalUnits", daily.totalUnits) + put("callCount", daily.callCount) + } + } + } + putJsonArray("models") { + summary.models.forEach { model -> + addJsonObject { + put("provider", model.provider) + put("model", model.model) + put("totalTokens", model.totalTokens) + put("callCount", model.callCount) + } + } + } + putJsonArray("breakdown") { + summary.breakdown.forEach { item -> + addJsonObject { + put("provider", item.provider) + put("model", item.model) + put("usageKind", item.usageKind) + put("unit", item.unit) + put("inputUnits", item.inputUnits) + put("outputUnits", item.outputUnits) + put("totalUnits", item.totalUnits) + put("callCount", item.callCount) + } + } + } + putJsonArray("userUsage") { + summary.userUsage.forEach { item -> + addJsonObject { + put("userId", item.userId) + put("name", item.name) + put("usageKind", item.usageKind) + put("unit", item.unit) + put("totalUnits", item.totalUnits) + put("callCount", item.callCount) + } + } + } + putJsonArray("topUsers") { + summary.topUsers.forEach { ranking -> + addJsonObject { + put("userId", ranking.id) + put("name", ranking.name) + put("totalTokens", ranking.totalTokens) + } + } + } + putJsonArray("topGroups") { + summary.topGroups.forEach { ranking -> + addJsonObject { + put("groupId", ranking.id) + put("name", ranking.name) + put("totalTokens", ranking.totalTokens) + } + } + } + } + + private fun kotlinx.serialization.json.JsonArrayBuilder.addRecord(record: TokenUsageRecord) { + addJsonObject { + put("date", record.date) + put("timestamp", record.timestamp) + put("userId", record.userId) + put("userName", record.userNickname) + record.groupId?.let { put("groupId", it) } + record.groupName?.let { put("groupName", it) } + record.provider?.let { put("provider", it) } + record.model?.let { put("model", it) } + put("usageKind", record.usageKind) + put("unit", record.unit) + put("inputUnits", record.inputUnits) + put("outputUnits", record.outputUnits) + put("totalUnits", record.totalUnits) + put("promptTokens", record.promptTokens) + put("completionTokens", record.completionTokens) + put("totalTokens", record.totalTokens) + put("cachedTokens", record.cachedTokens) + put("callCount", record.callCount) + put("detailed", record.detailed) + } + } + + private fun configuredDeepSeekAccounts(): List { + val openAiTypes = setOf("openai", "openai-compatible", "openai_compatible") + fun resolveOpenAiAlias(alias: String) = ModelCatalog.resolve(alias)?.takeIf { + it.providerType in openAiTypes && it.api.isNotBlank() && it.token.isNotBlank() && it.model.isNotBlank() + } + val boundAliases = buildSet { + add(PluginConfig.chatModelAlias) + addAll(PluginConfig.chatFallbackModelAliases) + add(PluginConfig.profileModelAlias.ifBlank { PluginConfig.chatModelAlias }) + add(PluginConfig.reasoningModelAlias) + add(PluginConfig.visualModelAlias) + add(PluginConfig.webSummaryModelAlias) + }.mapTo(HashSet(), String::trim).filterTo(HashSet(), String::isNotEmpty) + val boundModels = boundAliases.mapNotNull { alias -> + resolveOpenAiAlias(alias)?.let { definition -> + ModelDefinition( + name = definition.alias, + provider = definition.provider, + model = definition.model, + ) + } + } + val legacyAccounts = buildList { + val primaryResolved = resolveOpenAiAlias(PluginConfig.chatModelAlias) != null + if (!primaryResolved) { + add(BalanceAccount("legacy-chat", PluginConfig.openAiApi, PluginConfig.openAiToken)) + PluginConfig.chatFallbacks.forEachIndexed { index, fallback -> + add( + BalanceAccount( + name = "legacy-chat-fallback-${index + 1}", + api = fallback.api.ifBlank { PluginConfig.openAiApi }, + token = fallback.token.ifBlank { PluginConfig.openAiToken }, + ) + ) + } + } + val profileAlias = PluginConfig.profileModelAlias.ifBlank { PluginConfig.chatModelAlias } + if (PluginConfig.profileEnabled && resolveOpenAiAlias(profileAlias) == null) { + add( + BalanceAccount( + "legacy-profile", + PluginConfig.profileModelApi.ifBlank { PluginConfig.openAiApi }, + PluginConfig.profileModelToken.ifBlank { PluginConfig.openAiToken }, + ) + ) + } + if (resolveOpenAiAlias(PluginConfig.reasoningModelAlias) == null) { + add( + BalanceAccount( + "legacy-reasoning", + PluginConfig.reasoningModelApi, + PluginConfig.reasoningModelToken, + ) + ) + } + if (resolveOpenAiAlias(PluginConfig.visualModelAlias) == null) { + add(BalanceAccount("legacy-visual", PluginConfig.visualModelApi, PluginConfig.visualModelToken)) + } + if (resolveOpenAiAlias(PluginConfig.webSummaryModelAlias) == null) { + add( + BalanceAccount( + "legacy-web-summary", + PluginConfig.webSummaryModelApi, + PluginConfig.webSummaryModelToken, + ) + ) + } + } + return collectDeepSeekBalanceAccounts( + providers = ModelConfig.providers, + models = boundModels, + legacyAccounts = legacyAccounts, + ) + } + + private suspend fun queryBalances(): JsonObject { + val accounts = configuredDeepSeekAccounts() + if (accounts.isEmpty()) { + return buildJsonObject { + put("supported", false) + put("message", "未配置可查询余额的官方 DeepSeek 账号") + } + } + val results = accounts.map { account -> queryDeepSeekBalance(account) } + return buildJsonObject { + put("supported", true) + putJsonArray("accounts") { + results.forEach(::add) + } + } + } + + private suspend fun queryDeepSeekBalance(account: BalanceAccount): JsonObject = try { + val response = httpClient.get(DEEPSEEK_BALANCE_URL) { + header(HttpHeaders.Authorization, "Bearer ${account.token}") + timeout { + requestTimeoutMillis = 20_000 + connectTimeoutMillis = 10_000 + socketTimeoutMillis = 20_000 + } + } + val body = response.bodyAsText() + if (!response.status.isSuccess()) { + buildJsonObject { + put("provider", account.name) + put("service", "deepseek") + put("error", "HTTP ${response.status.value}") + put("message", body.take(300)) + } + } else { + buildJsonObject { + put("provider", account.name) + put("service", "deepseek") + put("data", parseDeepSeekBalance(body)) + } + } + } catch (cause: Throwable) { + buildJsonObject { + put("provider", account.name) + put("service", "deepseek") + put("error", cause.message ?: cause::class.simpleName.orEmpty()) + } + } + +} diff --git a/src/main/kotlin/tools/ReasoningAgent.kt b/src/main/kotlin/tools/ReasoningAgent.kt index c3a8c55..0281e57 100644 --- a/src/main/kotlin/tools/ReasoningAgent.kt +++ b/src/main/kotlin/tools/ReasoningAgent.kt @@ -2,12 +2,16 @@ package top.jie65535.mirai.tools import com.aallam.openai.api.chat.ChatCompletionRequest import com.aallam.openai.api.chat.ChatMessage +import com.aallam.openai.api.chat.StreamOptions import com.aallam.openai.api.chat.Tool import com.aallam.openai.api.core.Parameters +import com.aallam.openai.api.core.Usage import com.aallam.openai.api.model.ModelId import kotlinx.serialization.json.* -import top.jie65535.mirai.config.PluginConfig +import net.mamoe.mirai.event.events.MessageEvent +import top.jie65535.mirai.data.ModelUsageRecorder import top.jie65535.mirai.llm.LargeLanguageModels +import top.jie65535.mirai.llm.ModelService class ReasoningAgent : BaseAgent( tool = Tool.function( @@ -33,17 +37,23 @@ class ReasoningAgent : BaseAgent( override val isEnabled: Boolean get() = LargeLanguageModels.reasoning != null - override suspend fun execute(args: JsonObject?): String { + override suspend fun execute(args: JsonObject?, event: MessageEvent): String { requireNotNull(args) - val llm = LargeLanguageModels.reasoning ?: return "未配置llm,无法进行推理。" + val endpoint = LargeLanguageModels.reasoning ?: return "未配置llm,无法进行推理。" val prompt = args.getValue("prompt").jsonPrimitive.content val answerContent = StringBuilder() val reasoningContent = StringBuilder() - llm.chatCompletions(ChatCompletionRequest( - model = ModelId(PluginConfig.reasoningModel), - messages = listOf(ChatMessage.User(prompt)) - )).collect { + var lastUsage: Usage? = null + var cacheUsage: ModelService.CacheUsage? = null + endpoint.service.chatCompletions( + ChatCompletionRequest( + model = ModelId(endpoint.model), + messages = listOf(ChatMessage.User(prompt)), + streamOptions = StreamOptions(includeUsage = true), + ) + ) { cacheUsage = it }.collect { + it.usage?.let { usage -> lastUsage = usage } if (it.choices.isNotEmpty()) { val delta = it.choices[0].delta ?: return@collect if (!delta.reasoningContent.isNullOrEmpty()) { @@ -57,10 +67,21 @@ class ReasoningAgent : BaseAgent( val result = answerContent.toString() val reasoning = reasoningContent.toString() - return when { + ModelUsageRecorder.recordTokens( + event = event, + endpointLabel = "reasoning", + modelAlias = endpoint.alias, + provider = endpoint.provider, + model = endpoint.model, + usageKind = "reasoning", + usage = lastUsage, + cacheUsage = cacheUsage, + ) + val output = when { result.isNotEmpty() -> result reasoning.isNotEmpty() -> reasoning else -> "推理出错,结果为空" } + return output } } diff --git a/src/main/kotlin/tools/SendVoiceMessage.kt b/src/main/kotlin/tools/SendVoiceMessage.kt index 0deffa9..28a0313 100644 --- a/src/main/kotlin/tools/SendVoiceMessage.kt +++ b/src/main/kotlin/tools/SendVoiceMessage.kt @@ -11,6 +11,8 @@ import net.mamoe.mirai.event.events.MessageEvent import net.mamoe.mirai.utils.ExternalResource.Companion.toExternalResource import top.jie65535.mirai.JChatGPT import top.jie65535.mirai.config.PluginConfig +import top.jie65535.mirai.data.ModelUsageRecorder +import top.jie65535.mirai.llm.ModelCatalog import java.io.File import java.util.concurrent.TimeUnit import kotlin.time.measureTime @@ -48,22 +50,24 @@ class SendVoiceMessage : BaseAgent( get() = "录音中..." override val isEnabled: Boolean - get() = PluginConfig.dashScopeApiKey.isNotEmpty() + get() = ModelCatalog.resolveTts() != null override suspend fun execute(args: JsonObject?, event: MessageEvent): String { requireNotNull(args) if (event.subject !is AudioSupported) return "当前聊天环境不支持发送语音!" + val modelDefinition = ModelCatalog.resolveTts() + ?: return "未配置 TTS 模型,无法生成语音。" val content = args.getValue("content").jsonPrimitive.content val instructions = args["instructions"]?.jsonPrimitive?.content?.takeIf { it.isNotBlank() } // https://help.aliyun.com/zh/model-studio/qwen-tts - val response = httpClient.post(API_URL) { + val response = httpClient.post(modelDefinition.api.ifBlank { API_URL }) { contentType(ContentType("application", "json")) - header("Authorization", "Bearer " + PluginConfig.dashScopeApiKey) + header("Authorization", "Bearer " + modelDefinition.token) setBody(buildJsonObject { - put("model", PluginConfig.ttsModel) + put("model", modelDefinition.model) putJsonObject("input") { put("text", content) put("voice", "Chelsie") // Chelsie(女) Cherry(女) Ethan(男) Serena(女) @@ -82,6 +86,24 @@ class SendVoiceMessage : BaseAgent( .getValue("output").jsonObject .getValue("audio").jsonObject .getValue("url").jsonPrimitive.content + val inputCharacters = (responseObject["usage"] as? JsonObject) + ?.let { usage -> + usage["input_characters"]?.jsonPrimitive?.longOrNull + ?: usage["characters"]?.jsonPrimitive?.longOrNull + } + ?.coerceAtLeast(0) + ?: content.codePointCount(0, content.length).toLong() + ModelUsageRecorder.recordUnits( + event = event, + endpointLabel = "tts", + modelAlias = modelDefinition.alias, + provider = modelDefinition.provider, + model = modelDefinition.model, + usageKind = "tts", + unit = "characters", + inputUnits = inputCharacters, + totalUnits = inputCharacters, + ) val voiceFolder = JChatGPT.resolveDataFile("voice") voiceFolder.mkdir() diff --git a/src/main/kotlin/tools/VisitWeb.kt b/src/main/kotlin/tools/VisitWeb.kt index c75883e..282c8f3 100644 --- a/src/main/kotlin/tools/VisitWeb.kt +++ b/src/main/kotlin/tools/VisitWeb.kt @@ -2,8 +2,10 @@ package top.jie65535.mirai.tools import com.aallam.openai.api.chat.ChatCompletionRequest import com.aallam.openai.api.chat.ChatMessage +import com.aallam.openai.api.chat.StreamOptions import com.aallam.openai.api.chat.Tool import com.aallam.openai.api.core.Parameters +import com.aallam.openai.api.core.Usage import com.aallam.openai.api.model.ModelId import io.ktor.client.request.* import io.ktor.client.statement.* @@ -16,9 +18,12 @@ import kotlinx.coroutines.coroutineScope import kotlinx.coroutines.flow.collect import kotlinx.coroutines.withContext import kotlinx.serialization.json.* +import net.mamoe.mirai.event.events.MessageEvent import top.jie65535.mirai.JChatGPT import top.jie65535.mirai.config.PluginConfig +import top.jie65535.mirai.data.ModelUsageRecorder import top.jie65535.mirai.llm.LargeLanguageModels +import top.jie65535.mirai.llm.ModelService import java.net.InetAddress import java.net.URI import java.net.UnknownHostException @@ -142,7 +147,7 @@ class VisitWeb : BaseAgent( override val loadingMessage: String get() = "上网中..." - override suspend fun execute(args: JsonObject?): String { + override suspend fun execute(args: JsonObject?, event: MessageEvent): String { requireNotNull(args) val urlJson = args.getValue("url") val instruction = args["instruction"] @@ -162,12 +167,17 @@ class VisitWeb : BaseAgent( return coroutineScope { urls.map { url -> - async(Dispatchers.IO) { jinaReadPage(url, instruction, outputLimit) } + async(Dispatchers.IO) { jinaReadPage(url, instruction, outputLimit, event) } }.awaitAll().joinToString("\n\n---\n\n") } } - private suspend fun jinaReadPage(url: String, instruction: String, outputLimit: Int): String { + private suspend fun jinaReadPage( + url: String, + instruction: String, + outputLimit: Int, + event: MessageEvent, + ): String { return try { val request = withContext(Dispatchers.IO) { createReaderRequest( @@ -182,7 +192,7 @@ class VisitWeb : BaseAgent( } val body = response.bodyAsText() if (response.status.isSuccess()) { - summarizeOrExcerpt(url, body, instruction, outputLimit) + summarizeOrExcerpt(url, body, instruction, outputLimit, event) } else { "Error fetching \"$url\": HTTP ${response.status.value} ${body.take(500)}" } @@ -198,6 +208,7 @@ class VisitWeb : BaseAgent( body: String, instruction: String, outputLimit: Int, + event: MessageEvent, ): String { val endpoint = LargeLanguageModels.webSummary if (endpoint == null) { @@ -207,6 +218,8 @@ class VisitWeb : BaseAgent( val input = prepareWebContent(body, PluginConfig.webSummaryMaxInputChars) val prompt = buildSummaryUserPrompt(url, instruction, input) val rawSummary = StringBuilder() + var lastUsage: Usage? = null + var cacheUsage: ModelService.CacheUsage? = null return try { endpoint.service.chatCompletions( ChatCompletionRequest( @@ -215,14 +228,26 @@ class VisitWeb : BaseAgent( ChatMessage.System(WEB_SUMMARY_SYSTEM_PROMPT), ChatMessage.User(prompt), ), + streamOptions = StreamOptions(includeUsage = true), ) - ).collect { chunk -> + ) { cacheUsage = it }.collect { chunk -> + chunk.usage?.let { usage -> lastUsage = usage } chunk.choices.firstOrNull()?.delta?.content?.let { content -> if (rawSummary.length < MAX_RAW_SUMMARY_CHARS) { rawSummary.append(content.take(MAX_RAW_SUMMARY_CHARS - rawSummary.length)) } } } + ModelUsageRecorder.recordTokens( + event = event, + endpointLabel = "web-summary", + modelAlias = endpoint.alias, + provider = endpoint.provider, + model = endpoint.model, + usageKind = "web_summary", + usage = lastUsage, + cacheUsage = cacheUsage, + ) val summary = limitSummaryOutput(cleanSummary(rawSummary.toString()), outputLimit) if (summary.isBlank()) { logSummaryFallback(url, body.length, "模型返回为空") diff --git a/src/main/kotlin/tools/VisualAgent.kt b/src/main/kotlin/tools/VisualAgent.kt index 77f45f2..5961499 100644 --- a/src/main/kotlin/tools/VisualAgent.kt +++ b/src/main/kotlin/tools/VisualAgent.kt @@ -4,9 +4,11 @@ import com.aallam.openai.api.chat.ChatCompletionRequest import com.aallam.openai.api.chat.ChatMessage import com.aallam.openai.api.chat.ContentPart import com.aallam.openai.api.chat.ImagePart +import com.aallam.openai.api.chat.StreamOptions import com.aallam.openai.api.chat.TextPart import com.aallam.openai.api.chat.Tool import com.aallam.openai.api.core.Parameters +import com.aallam.openai.api.core.Usage import com.aallam.openai.api.model.ModelId import io.ktor.client.plugins.ClientRequestException import kotlinx.serialization.json.JsonObject @@ -24,7 +26,9 @@ import kotlinx.coroutines.sync.withPermit import net.mamoe.mirai.event.events.MessageEvent import top.jie65535.mirai.JChatGPT import top.jie65535.mirai.config.PluginConfig +import top.jie65535.mirai.data.ModelUsageRecorder import top.jie65535.mirai.llm.LargeLanguageModels +import top.jie65535.mirai.llm.ModelService import top.jie65535.mirai.util.RetryBackoff import java.net.URI @@ -68,7 +72,7 @@ class VisualAgent : BaseAgent( override suspend fun execute(args: JsonObject?, event: MessageEvent): String { requireNotNull(args) - val llm = LargeLanguageModels.visual ?: return "未配置llm,无法进行识别。" + val endpoint = LargeLanguageModels.visual ?: return "未配置llm,无法进行识别。" val imageIndices = args["image_indices"]?.jsonArray ?.map { it.jsonPrimitive.int } ?.ifEmpty { null } @@ -124,16 +128,20 @@ class VisualAgent : BaseAgent( repeat(maxAttempts) { attempt -> try { val answerContent = StringBuilder() - llm.chatCompletions( + var lastUsage: Usage? = null + var cacheUsage: ModelService.CacheUsage? = null + endpoint.service.chatCompletions( ChatCompletionRequest( - model = ModelId(PluginConfig.visualModel), + model = ModelId(endpoint.model), messages = listOf( ChatMessage.User( content = messageContent ) - ) + ), + streamOptions = StreamOptions(includeUsage = true), ) - ).collect { + ) { cacheUsage = it }.collect { + it.usage?.let { usage -> lastUsage = usage } if (it.choices.isNotEmpty()) { val delta = it.choices[0].delta ?: return@collect if (!delta.content.isNullOrEmpty()) { @@ -142,6 +150,16 @@ class VisualAgent : BaseAgent( } } + ModelUsageRecorder.recordTokens( + event = event, + endpointLabel = "visual", + modelAlias = endpoint.alias, + provider = endpoint.provider, + model = endpoint.model, + usageKind = "visual", + usage = lastUsage, + cacheUsage = cacheUsage, + ) if (answerContent.isNotEmpty()) { return@withPermit answerContent.toString() } diff --git a/src/test/kotlin/config/ModelConfigMigrationTest.kt b/src/test/kotlin/config/ModelConfigMigrationTest.kt new file mode 100644 index 0000000..1b58d41 --- /dev/null +++ b/src/test/kotlin/config/ModelConfigMigrationTest.kt @@ -0,0 +1,120 @@ +package top.jie65535.mirai.config + +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertTrue + +class ModelConfigMigrationTest { + @Test + fun migratesLegacySettingsAndDeduplicatesProvidersAndModels() { + val legacy = LegacyModelSettings( + chat = openAi("https://api.deepseek.com/v1/", "token-a", "deepseek-chat"), + chatFallbacks = listOf( + openAi("https://api.deepseek.com/v1", "token-b", "deepseek-chat"), + ), + profile = openAi("https://api.deepseek.com/v1/", "token-a", "deepseek-chat"), + reasoning = openAi("https://api.deepseek.com/v1/", "token-a", "deepseek-reasoner"), + visual = openAi("https://dashscope.aliyuncs.com/compatible-mode/v1/", "token-c", "qwen-vl-plus"), + webSummary = openAi("", "", ""), + dashScopeToken = "dashscope-token", + imageModel = "qwen-image-2.0", + ttsModel = "qwen3-tts-flash", + ) + + val result = ModelConfigMigration.migrate(emptyList(), emptyList(), ModelRoleBindings(), legacy) + + assertEquals(4, result.providers.size) + assertEquals(6, result.models.size) + assertEquals("chat-main", result.bindings.chat) + assertEquals("chat-main", result.bindings.profile) + assertEquals(listOf("chat-fallback-1"), result.bindings.chatFallbacks) + assertEquals("reasoning-main", result.bindings.reasoning) + assertEquals("visual-main", result.bindings.visual) + assertEquals("", result.bindings.webSummary) + assertEquals("image-main", result.bindings.image) + assertEquals("tts-main", result.bindings.tts) + assertEquals(1, result.providers.count { it.type == "dashscope" }) + } + + @Test + fun preservesExistingEntriesAndBindingsAndIsIdempotent() { + val existingProvider = ModelProviderDefinition("custom", "openai", "https://api.deepseek.com/v1", "token-a") + val existingModel = ModelDefinition("my-chat", "custom", "deepseek-chat") + val bindings = ModelRoleBindings(chat = "my-chat", visual = "manual-visual") + val legacy = LegacyModelSettings( + chat = openAi("https://api.deepseek.com/v1/", "token-a", "deepseek-chat"), + chatFallbacks = emptyList(), + profile = openAi("https://api.deepseek.com/v1/", "token-a", "deepseek-chat"), + reasoning = openAi("", "", ""), + visual = openAi("https://example.com/v1", "other", "vision"), + webSummary = openAi("", "", ""), + dashScopeToken = "", + imageModel = "qwen-image-2.0", + ttsModel = "qwen3-tts-flash", + ) + + val first = ModelConfigMigration.migrate(listOf(existingProvider), listOf(existingModel), bindings, legacy) + assertEquals(listOf(existingProvider), first.providers) + assertEquals(listOf(existingModel), first.models) + assertEquals("my-chat", first.bindings.chat) + assertEquals("my-chat", first.bindings.profile) + assertEquals("manual-visual", first.bindings.visual) + + val second = ModelConfigMigration.migrate(first.providers, first.models, first.bindings, legacy) + assertEquals(first.providers, second.providers) + assertEquals(first.models, second.models) + assertEquals(first.bindings, second.bindings) + assertEquals(0, second.addedProviders) + assertEquals(0, second.addedModels) + assertTrue(!second.changed) + } + + @Test + fun leavesFreshInstallEmptyWhenNoCredentialsExist() { + val empty = openAi("", "", "") + val legacy = LegacyModelSettings( + chat = empty, + chatFallbacks = emptyList(), + profile = empty, + reasoning = empty, + visual = empty, + webSummary = empty, + dashScopeToken = "", + imageModel = "qwen-image-2.0", + ttsModel = "qwen3-tts-flash", + ) + + val result = ModelConfigMigration.migrate(emptyList(), emptyList(), ModelRoleBindings(), legacy) + + assertTrue(result.providers.isEmpty()) + assertTrue(result.models.isEmpty()) + assertEquals(ModelRoleBindings(), result.bindings) + assertTrue(!result.changed) + } + + @Test + fun allocatesNewNamesWithoutOverwritingConflicts() { + val providers = listOf(ModelProviderDefinition("deepseek", "openai", "https://other.example/v1", "other")) + val models = listOf(ModelDefinition("chat-main", "deepseek", "other-model")) + val empty = openAi("", "", "") + val legacy = LegacyModelSettings( + chat = openAi("https://api.deepseek.com/v1", "token-a", "deepseek-chat"), + chatFallbacks = emptyList(), + profile = empty, + reasoning = empty, + visual = empty, + webSummary = empty, + dashScopeToken = "", + imageModel = "", + ttsModel = "", + ) + + val result = ModelConfigMigration.migrate(providers, models, ModelRoleBindings(), legacy) + + assertEquals("deepseek-2", result.providers.last().name) + assertEquals("chat-main-2", result.models.last().name) + assertEquals("chat-main-2", result.bindings.chat) + } + + private fun openAi(api: String, token: String, model: String) = LegacyOpenAiModel(api, token, model) +} diff --git a/src/test/kotlin/data/TokenUsageStoreTest.kt b/src/test/kotlin/data/TokenUsageStoreTest.kt new file mode 100644 index 0000000..e50611e --- /dev/null +++ b/src/test/kotlin/data/TokenUsageStoreTest.kt @@ -0,0 +1,403 @@ +package top.jie65535.mirai.data + +import java.nio.file.Files +import java.sql.DriverManager +import java.time.Instant +import java.time.LocalDate +import kotlin.io.path.absolutePathString +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertFalse +import kotlin.test.assertTrue + +class TokenUsageStoreTest { + @Test + fun recordsDetailedUsageAndScopesQueriesToCurrentConversation() { + val directory = Files.createTempDirectory("jchatgpt-token-usage-test-") + try { + TokenUsageStore.init(directory.toFile()) + val now = Instant.now().epochSecond + record(now, botId = 1, userId = 10, groupId = 100, total = 15, model = "deepseek-chat") + record(now + 1, botId = 1, userId = 11, groupId = 100, total = 30, model = "deepseek-chat") + record(now + 2, botId = 1, userId = 10, groupId = 200, total = 40, model = "model-b") + record(now + 3, botId = 2, userId = 12, groupId = 100, total = 50, model = "model-c") + record(now + 4, botId = 1, userId = 10, groupId = null, total = 60, model = "deepseek-chat") + + val startDate = LocalDate.now().minusDays(1).toString() + val global = TokenUsageStore.summary(startDate, rankingLimit = 20) + assertEquals(195, global.totalTokens) + assertEquals(5, global.callCount) + assertEquals(3, global.activeUsers) + + val currentGroup = TokenUsageStore.summary( + startDate = startDate, + botId = 1, + groupId = 100, + rankingLimit = 20, + ) + assertEquals(45, currentGroup.totalTokens) + assertEquals(2, currentGroup.callCount) + assertEquals(45, currentGroup.todayTotal) + assertEquals(45, currentGroup.usageDaily.sumOf { it.totalUnits }) + assertEquals(listOf(11L, 10L), currentGroup.topUsers.map(TokenUsageRanking::id)) + + val oneGroupUser = TokenUsageStore.summary( + startDate = startDate, + botId = 1, + userId = 11, + groupId = 100, + ) + assertEquals(30, oneGroupUser.totalTokens) + assertEquals(1, oneGroupUser.callCount) + + val currentPrivateChat = TokenUsageStore.summary( + startDate = startDate, + botId = 1, + userId = 10, + privateOnly = true, + ) + assertEquals(60, currentPrivateChat.totalTokens) + assertEquals(1, currentPrivateChat.callCount) + assertTrue(currentPrivateChat.topGroups.isEmpty()) + + val details = TokenUsageStore.recent( + startDate = startDate, + botId = 1, + groupId = 100, + ) + assertEquals(listOf(11L, 10L), details.map(TokenUsageRecord::userId)) + assertTrue(details.all(TokenUsageRecord::detailed)) + assertEquals(setOf("deepseek"), details.mapNotNull(TokenUsageRecord::provider).toSet()) + } finally { + TokenUsageStore.close() + directory.toFile().deleteRecursively() + } + } + + @Test + fun importsLegacyJsonIdempotentlyAndKeepsSourceFile() { + val directory = Files.createTempDirectory("jchatgpt-token-legacy-test-") + val legacy = directory.resolve("token_usage.json") + legacy.toFile().writeText( + """ + [ + { + "date": "2026-08-01", + "userId": 123, + "userNickname": "legacy-user", + "groupId": 456, + "groupName": "legacy-group", + "promptTokens": 100, + "completionTokens": 20, + "totalTokens": 120, + "cachedTokens": 40, + "callCount": 3 + } + ] + """.trimIndent() + ) + val warnings = mutableListOf() + val database = directory.resolve("chat-history.sqlite") + try { + TokenUsageStore.init(directory.toFile()) { message, _ -> warnings += message } + assertLegacyImport() + TokenUsageStore.close() + + DriverManager.getConnection("jdbc:sqlite:${database.absolutePathString()}").use { connection -> + connection.createStatement().use { statement -> + statement.executeUpdate("DELETE FROM token_usage_meta WHERE key = 'legacy_json_sha256'") + } + } + + TokenUsageStore.init(directory.toFile()) { message, _ -> warnings += message } + assertLegacyImport() + assertTrue(Files.isRegularFile(legacy)) + assertEquals(1, warnings.count { it.startsWith("已将 1 条旧 Token 聚合记录迁移") }) + + DriverManager.getConnection("jdbc:sqlite:${database.absolutePathString()}").use { connection -> + connection.createStatement().use { statement -> + statement.executeQuery("PRAGMA journal_mode").use { results -> + assertTrue(results.next()) + assertEquals("wal", results.getString(1).lowercase()) + } + statement.executeQuery("SELECT value FROM token_usage_meta WHERE key = 'schema_version'").use { results -> + assertTrue(results.next()) + assertEquals("2", results.getString(1)) + } + statement.executeQuery( + "SELECT value FROM token_usage_meta WHERE key = 'legacy_json_sha256'" + ).use { results -> + assertTrue(results.next()) + assertEquals(64, results.getString(1).length) + } + statement.executeQuery("PRAGMA table_info(token_usage_record)").use { results -> + val columns = buildSet { + while (results.next()) add(results.getString("name")) + } + assertEquals( + setOf( + "id", + "occurred_at", + "usage_date", + "bot_id", + "user_id", + "user_nickname", + "group_id", + "group_name", + "endpoint_label", + "model_alias", + "provider", + "model", + "usage_kind", + "unit", + "input_units", + "output_units", + "total_units", + "prompt_tokens", + "completion_tokens", + "total_tokens", + "cached_tokens", + "call_count", + "detailed", + "legacy_key", + ), + columns, + ) + } + statement.executeQuery("SELECT COUNT(*) FROM token_usage_record").use { results -> + assertTrue(results.next()) + assertEquals(1, results.getInt(1)) + } + statement.executeQuery("PRAGMA index_list(token_usage_record)").use { results -> + val indexes = buildSet { + while (results.next()) add(results.getString("name")) + } + assertTrue("idx_token_usage_kind_date" in indexes) + } + } + } + } finally { + TokenUsageStore.close() + directory.toFile().deleteRecursively() + } + } + + @Test + fun keepsNonTokenUnitsOutOfTokenTotalsAndIncludesThemInBreakdown() { + val directory = Files.createTempDirectory("jchatgpt-model-usage-test-") + try { + TokenUsageStore.init(directory.toFile()) + val now = Instant.now().epochSecond + record(now, botId = 1, userId = 10, groupId = 100, total = 100, model = "chat-model") + TokenUsageStore.recordUsage( + ModelUsageEvent( + timestamp = now + 1, + botId = 1, + userId = 10, + userNickname = "user-10", + groupId = 100, + groupName = "group-100", + endpointLabel = "image", + modelAlias = "image-main", + provider = "dashscope", + model = "qwen-image", + usageKind = "image", + unit = "images", + outputUnits = 1, + totalUnits = 1, + promptTokens = 0, + completionTokens = 0, + totalTokens = 0, + ) + ) + TokenUsageStore.recordUsage( + ModelUsageEvent( + timestamp = now + 2, + botId = 1, + userId = 10, + userNickname = "user-10", + groupId = 100, + groupName = "group-100", + endpointLabel = "tts", + modelAlias = "tts-main", + provider = "dashscope", + model = "qwen-tts", + usageKind = "tts", + unit = "characters", + inputUnits = 50, + totalUnits = 50, + promptTokens = 0, + completionTokens = 0, + totalTokens = 0, + ) + ) + + val summary = TokenUsageStore.summary( + startDate = LocalDate.now().minusDays(1).toString(), + botId = 1, + groupId = 100, + rankingLimit = 20, + ) + assertEquals(100, summary.totalTokens) + assertEquals(1, summary.callCount) + assertEquals(3, summary.allCallCount) + assertEquals(setOf("tokens", "images", "characters"), summary.breakdown.map { it.unit }.toSet()) + assertEquals(1, summary.breakdown.single { it.unit == "images" }.outputUnits) + assertEquals(50, summary.breakdown.single { it.unit == "characters" }.inputUnits) + assertEquals(1, summary.userUsage.single { it.usageKind == "image" }.totalUnits) + assertEquals(50, summary.userUsage.single { it.usageKind == "tts" }.totalUnits) + + val imageOnly = TokenUsageStore.summary( + startDate = LocalDate.now().minusDays(1).toString(), + botId = 1, + groupId = 100, + usageKind = "image", + ) + assertEquals(0, imageOnly.totalTokens) + assertEquals(0, imageOnly.callCount) + assertEquals(1, imageOnly.allCallCount) + assertEquals(1, imageOnly.allActiveUsers) + assertEquals("images", imageOnly.breakdown.single().unit) + assertEquals(10, imageOnly.userUsage.single().userId) + assertEquals(1, imageOnly.usageDaily.single().totalUnits) + + val details = TokenUsageStore.recent(botId = 1, groupId = 100) + assertEquals(listOf("tts", "image", "chat"), details.map { it.usageKind }) + assertEquals("tts-main", details.first().modelAlias) + } finally { + TokenUsageStore.close() + directory.toFile().deleteRecursively() + } + } + + @Test + fun migratesVersionOneTableAndBackfillsGenericTokenUnits() { + val directory = Files.createTempDirectory("jchatgpt-token-v1-migration-test-") + val database = directory.resolve("chat-history.sqlite") + try { + DriverManager.getConnection("jdbc:sqlite:${database.absolutePathString()}").use { connection -> + connection.createStatement().use { statement -> + statement.executeUpdate( + """ + CREATE TABLE token_usage_record( + id INTEGER PRIMARY KEY AUTOINCREMENT, + occurred_at INTEGER NOT NULL, + usage_date TEXT NOT NULL, + bot_id INTEGER, + user_id INTEGER NOT NULL, + user_nickname TEXT NOT NULL DEFAULT '', + group_id INTEGER, + group_name TEXT, + endpoint_label TEXT, + provider TEXT, + model TEXT, + prompt_tokens INTEGER NOT NULL DEFAULT 0, + completion_tokens INTEGER NOT NULL DEFAULT 0, + total_tokens INTEGER NOT NULL DEFAULT 0, + cached_tokens INTEGER NOT NULL DEFAULT 0, + call_count INTEGER NOT NULL DEFAULT 1, + detailed INTEGER NOT NULL DEFAULT 1, + legacy_key TEXT UNIQUE + ) + """.trimIndent() + ) + statement.executeUpdate( + "CREATE TABLE token_usage_meta(key TEXT PRIMARY KEY, value TEXT NOT NULL)" + ) + statement.executeUpdate( + "INSERT INTO token_usage_meta(key, value) VALUES ('schema_version', '1')" + ) + statement.executeUpdate( + """ + INSERT INTO token_usage_record( + occurred_at, usage_date, bot_id, user_id, user_nickname, + group_id, group_name, endpoint_label, provider, model, + prompt_tokens, completion_tokens, total_tokens, cached_tokens + ) VALUES (1, '2026-08-01', 1, 10, 'user', 100, 'group', + 'primary', 'deepseek', 'deepseek-chat', 80, 20, 100, 40) + """.trimIndent() + ) + } + } + + TokenUsageStore.init(directory.toFile()) + val record = TokenUsageStore.recent( + startDate = "2026-08-01", + endDate = "2026-08-01", + ).single() + assertEquals("chat", record.usageKind) + assertEquals("tokens", record.unit) + assertEquals(80, record.inputUnits) + assertEquals(20, record.outputUnits) + assertEquals(100, record.totalUnits) + assertEquals(100, TokenUsageStore.summary("2026-08-01", "2026-08-01").totalTokens) + + TokenUsageStore.close() + DriverManager.getConnection("jdbc:sqlite:${database.absolutePathString()}").use { connection -> + connection.createStatement().use { statement -> + statement.executeQuery("SELECT value FROM token_usage_meta WHERE key = 'schema_version'").use { results -> + assertTrue(results.next()) + assertEquals("2", results.getString(1)) + } + } + } + } finally { + TokenUsageStore.close() + directory.toFile().deleteRecursively() + } + } + + private fun assertLegacyImport() { + val summary = TokenUsageStore.summary("2026-08-01", "2026-08-01") + assertEquals(100, summary.promptTokens) + assertEquals(20, summary.completionTokens) + assertEquals(120, summary.totalTokens) + assertEquals(40, summary.cachedTokens) + assertEquals(3, summary.callCount) + assertEquals( + 0, + TokenUsageStore.summary( + startDate = "2026-08-01", + endDate = "2026-08-01", + botId = 1, + groupId = 456, + ).totalTokens, + ) + + val record = TokenUsageStore.recent( + limit = 10, + startDate = "2026-08-01", + endDate = "2026-08-01", + ).single() + assertFalse(record.detailed) + assertEquals(3, record.callCount) + assertEquals(100, record.inputUnits) + assertEquals(20, record.outputUnits) + assertEquals(120, record.totalUnits) + } + + private fun record( + timestamp: Long, + botId: Long, + userId: Long, + groupId: Long?, + total: Int, + model: String, + ) { + TokenUsageStore.record( + timestamp = timestamp, + botId = botId, + userId = userId, + userNickname = "user-$userId", + groupId = groupId, + groupName = groupId?.let { "group-$it" }, + endpointLabel = "primary", + apiBaseUrl = "https://api.deepseek.com/v1/", + model = model, + promptTokens = total - 5, + completionTokens = 5, + totalTokens = total, + cachedTokens = 2, + ) + } +} diff --git a/src/test/kotlin/llm/ModelCatalogTest.kt b/src/test/kotlin/llm/ModelCatalogTest.kt new file mode 100644 index 0000000..1346544 --- /dev/null +++ b/src/test/kotlin/llm/ModelCatalogTest.kt @@ -0,0 +1,107 @@ +package top.jie65535.mirai.llm + +import top.jie65535.mirai.config.ModelDefinition +import top.jie65535.mirai.config.ModelProviderDefinition +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertNull +import kotlin.test.assertTrue + +class ModelCatalogTest { + @Test + fun resolvesModelAliasThroughProvider() { + val resolved = ModelCatalog.resolve( + alias = "chat-main", + providers = listOf( + ModelProviderDefinition( + name = "deepseek", + type = "openai", + api = "https://api.deepseek.com/v1/", + token = "secret", + ) + ), + models = listOf( + ModelDefinition( + name = "chat-main", + provider = "deepseek", + model = "deepseek-chat", + extraBody = "{\"thinking\":false}", + ) + ), + ) + + requireNotNull(resolved) + assertEquals("chat-main", resolved.alias) + assertEquals("deepseek", resolved.provider) + assertEquals("openai", resolved.providerType) + assertEquals("deepseek-chat", resolved.model) + assertEquals("secret", resolved.token) + } + + @Test + fun rejectsAmbiguousAliasesAndProviders() { + val provider = ModelProviderDefinition("deepseek", "openai", "https://api.deepseek.com/v1/", "secret") + val model = ModelDefinition("chat-main", "deepseek", "deepseek-chat") + + assertNull(ModelCatalog.resolve("chat-main", listOf(provider), listOf(model, model))) + assertNull(ModelCatalog.resolve("chat-main", listOf(provider, provider), listOf(model))) + assertNull( + ModelCatalog.resolve( + "chat-main", + listOf(provider), + listOf(model.copy(provider = "missing")), + ) + ) + } + + @Test + fun reportsInvalidCatalogEntries() { + val issues = ModelCatalog.validationIssues( + providers = listOf( + ModelProviderDefinition("duplicate", "openai", "https://example.com/v1/", "token"), + ModelProviderDefinition("duplicate", "unknown", "", ""), + ), + models = listOf( + ModelDefinition("same", "duplicate", "model-a"), + ModelDefinition("same", "missing", ""), + ), + ) + + assertTrue(issues.any { it.contains("provider 名称重复") }) + assertTrue(issues.any { it.contains("模型别名重复") }) + assertTrue(issues.any { it.contains("type 不受支持") }) + assertTrue(issues.any { it.contains("未配置 token") }) + assertTrue(issues.any { it.contains("引用的 provider 无效") }) + assertTrue(issues.any { it.contains("未配置实际模型名") }) + } + + @Test + fun reportsRoleProtocolMismatch() { + val providers = listOf( + ModelProviderDefinition( + name = "dashscope-native", + type = "dashscope", + token = "secret", + ) + ) + val models = listOf(ModelDefinition("image-main", "dashscope-native", "qwen-image")) + + val issue = ModelCatalog.bindingValidationIssue( + role = "主聊天", + alias = "image-main", + allowedTypes = setOf("openai"), + providers = providers, + models = models, + ) + assertTrue(issue?.contains("不能使用 provider type dashscope") == true) + assertNull( + ModelCatalog.bindingValidationIssue( + role = "图像", + alias = "image-main", + allowedTypes = setOf("dashscope"), + providers = providers, + models = models, + ) + ) + } +} diff --git a/src/test/kotlin/tools/QueryTokenUsageAgentTest.kt b/src/test/kotlin/tools/QueryTokenUsageAgentTest.kt new file mode 100644 index 0000000..c9c4ee4 --- /dev/null +++ b/src/test/kotlin/tools/QueryTokenUsageAgentTest.kt @@ -0,0 +1,101 @@ +package top.jie65535.mirai.tools + +import kotlinx.serialization.json.boolean +import kotlinx.serialization.json.jsonArray +import kotlinx.serialization.json.jsonObject +import kotlinx.serialization.json.jsonPrimitive +import top.jie65535.mirai.config.ModelDefinition +import top.jie65535.mirai.config.ModelProviderDefinition +import kotlin.test.Test +import kotlin.test.assertEquals +import kotlin.test.assertFalse +import kotlin.test.assertTrue + +class QueryTokenUsageAgentTest { + @Test + fun scopesUsageToCurrentGroupOrPrivateSender() { + val group = QueryTokenUsageAgent.queryScope( + botId = 1, + senderId = 10, + currentGroupId = 100, + requestedUserId = 20, + ) + assertEquals(1, group.botId) + assertEquals(100, group.groupId) + assertEquals(20, group.userId) + assertFalse(group.privateOnly) + + val private = QueryTokenUsageAgent.queryScope( + botId = 1, + senderId = 10, + currentGroupId = null, + requestedUserId = 999, + ) + assertEquals(1, private.botId) + assertEquals(null, private.groupId) + assertEquals(10, private.userId) + assertTrue(private.privateOnly) + } + + @Test + fun recognizesOnlyOfficialDeepSeekHosts() { + assertTrue(QueryTokenUsageAgent.isDeepSeekApi("https://api.deepseek.com/v1/")) + assertTrue(QueryTokenUsageAgent.isDeepSeekApi("https://cn.api.deepseek.com/v1")) + assertFalse(QueryTokenUsageAgent.isDeepSeekApi("https://deepseek.example.com/v1")) + assertFalse(QueryTokenUsageAgent.isDeepSeekApi("not a url")) + } + + @Test + fun parsesDeepSeekBalanceResponse() { + val result = QueryTokenUsageAgent.parseDeepSeekBalance( + """ + { + "is_available": true, + "balance_infos": [ + { + "currency": "CNY", + "total_balance": "12.34", + "granted_balance": "2.00", + "topped_up_balance": "10.34" + } + ] + } + """.trimIndent() + ) + + assertTrue(result.getValue("available").jsonPrimitive.boolean) + val balance = result.getValue("balances").jsonArray.single().jsonObject + assertEquals("CNY", balance.getValue("currency").jsonPrimitive.content) + assertEquals("12.34", balance.getValue("total").jsonPrimitive.content) + assertEquals("2.00", balance.getValue("granted").jsonPrimitive.content) + assertEquals("10.34", balance.getValue("toppedUp").jsonPrimitive.content) + } + + @Test + fun collectsReferencedDeepSeekAccountsAndDeduplicatesCredentials() { + val accounts = QueryTokenUsageAgent.collectDeepSeekBalanceAccounts( + providers = listOf( + ModelProviderDefinition("deepseek-a", "openai", "https://api.deepseek.com/v1/", "same-token"), + ModelProviderDefinition("deepseek-b", "openai", "https://cn.api.deepseek.com/v1/", "same-token"), + ModelProviderDefinition("unused", "openai", "https://api.deepseek.com/v1/", "unused-token"), + ModelProviderDefinition("dashscope", "dashscope", "https://dashscope.aliyuncs.com/api/v1/", "dash-token"), + ), + models = listOf( + ModelDefinition("chat", "deepseek-a", "deepseek-chat"), + ModelDefinition("reasoning", "deepseek-b", "deepseek-reasoner"), + ModelDefinition("image", "dashscope", "qwen-image"), + ), + legacyAccounts = listOf( + QueryTokenUsageAgent.Companion.BalanceAccount( + "legacy-chat", + "https://api.deepseek.com/v1/", + "same-token", + ) + ), + ) + + assertEquals(1, accounts.size) + assertEquals("deepseek-a", accounts.single().name) + assertEquals("same-token", accounts.single().token) + } +}