diff --git a/README.md b/README.md index aa9d3200..9fac1310 100644 --- a/README.md +++ b/README.md @@ -84,6 +84,7 @@ Seira正在活跃开发中,在使用的过程中可能会有一些Bug,也会 | `/sa` | `/sa ` | 获取指定成绩分析图 | | `/ma` | `/ma [id/locId/rsN/bpN] [n/#n]` | 获取指定或最近目标成绩的Miss分析;省略目标时用`#n`指定Miss | | `/u` | `/u [uid/username/@user]` | 获取指定用户信息 | +| `/@` | `/@[someone]` | 获取自己或指定用户的文字版信息 | | `/r` | `/r [id/locId/rsN/bpN] [[mm:ss]-[mm:ss]]` | 生成并发送指定或最近目标的回放视频。省略范围时自动识别高光,使用`-`渲染整个回放 | | `/rg` | `/rg ` | 猜 Rank 游戏及个人战绩查询 | | `/rsc` | `/rsc [id/locId/rsN/bpN] [+,...]` | 生成并发送指定或最近目标的成绩同屏回放视频;追加用户和范围顺序不限 | @@ -94,10 +95,15 @@ Seira正在活跃开发中,在使用的过程中可能会有一些Bug,也会 | `/sms` | `/sms ` | 搜索谱面集 | | `/lb` | `/lb [id] [,...]` | 列出指定谱面排行或表现分排行 | | `/daily` | `/daily` | 每日挑战信息 | +| `/roll` | `/roll [dice]` | 掷骰子 | | `/luck` | `/luck` | 今日人品 | +| `/rbp` | `/rbp [uid/username/@user]` | 获取随机BP | | `/mp` | `/mp` | 多人房间列表 | | `/watch` | `/watch add/del/list [目标]` | 添加/删除/列出监视任务 | | `/mpwatch` | `/mpwatch start/stop/status [目标]` | 按群成员添加、停止或查看多人房间监视;`stop all` 可停止本群全部监视 | +| `/romai` | `/romai [目标]` | 开始监视自己或目标正在进行的RomAI比赛 | +| `/mc` | `/mc <服务器地址>` | 获取指定MC服务器状态 | +| `/ai` | `/ai [on/off/reset/reset all]` | 开启/关闭/重置本群的AI对话 | | `/wx` | `/wx start <谱面ID列表>` / `/wx stop` | 监视指定玩家在指定谱面取得的成绩,重启后自动恢复 | | `/dcs` | `/dcs start .` / `/dcs stop` | 开启或解除当前 QQ 群与 Discord 频道的双向消息同步 | | `/stat` | `/stat` | 服务状态和统计信息文本 | @@ -108,17 +114,21 @@ Seira正在活跃开发中,在使用的过程中可能会有一些Bug,也会 `/r`和`/rsc`(回放渲染)会先返回“生成请求正在等待中,队列位置:N”,随后返回请求状态,最后在渲染完成后再发送回放视频。 -上传 osu! 服务器上不存在的 `.osr` 回放后,机器人会返回形如 `loc123456789` 的本地成绩ID;该ID可用于 `/s`、`/sa`、`/ma`、`/r`、`/rsc` 等成绩目标指令。 +上传 osu! 服务器上不存在的 `.osr` 回放后,机器人会返回形如 `loc123456789` 的本地成绩ID;该ID可用于 `/s`、`/sa`、`/ma`、`/r`、 +`/rsc` 等成绩目标指令。 ### 快捷查询 对于一些需要指定谱面ID或成绩ID的指令(如 `/m`、`/s`、`/ms` 等),支持快捷查询写法,格式为 `rs5`、`bp3`、`rp2`。 -快捷查询也可以直接作为指令使用,并在后面指定玩家,例如 `/bp5 @用户`;紧凑写法同样支持列表范围,例如 `/bp21-30`、`/rp6-10 @用户`。 +快捷查询也可以直接作为指令使用,并在后面指定玩家,例如 `/bp5 @用户`;紧凑写法同样支持列表范围,例如 `/bp21-30`、 +`/rp6-10 @用户`。 -这些指令会共享最近一次显式指定的查询目标。已有最近目标时,`/r`、`/rsc`、`/ma` 可以省略目标,例如 `/r 01:00-01:30`、`/rsc +12345,67890 -`、`/ma #2`。 +这些指令会共享最近一次显式指定的查询目标。已有最近目标时,`/r`、`/rsc`、`/ma` 可以省略目标,例如 `/r 01:00-01:30`、 +`/rsc +12345,67890 -`、`/ma #2`。 -也可以在前面写上玩家ID、用户名或@用户,例如 `123456 rs5`、`peppy bp3`、`@ABC bp3`,表示查询指定玩家的最近成绩第 5 条或最好成绩第 3 条。 +也可以在前面写上玩家ID、用户名或@用户,例如 `123456 rs5`、`peppy bp3`、`@ABC bp3`,表示查询指定玩家的最近成绩第 5 条或最好成绩第 +3 条。 - `rs5`:使用你已绑定的玩家ID,查询“最近成绩第 5 条” - `rp1`:使用你已绑定的玩家ID,查询“最近通过成绩第 1 条” @@ -129,8 +139,6 @@ Seira正在活跃开发中,在使用的过程中可能会有一些Bug,也会 省略玩家目标时,使用快捷查询前需要先执行 `/bind`,否则会提示无法使用快捷查询。 -> 由于QQ业务调整,暂时无法使用`@用户`查询绑定信息,请改用直接输入uid的方式。 - ### 回放上传 你可以在私聊中直接发送你的回放文件,Seira会将其转存。 @@ -139,11 +147,11 @@ Seira正在活跃开发中,在使用的过程中可能会有一些Bug,也会 ## 调试命令 > [!WARNING] -> +> > **警告:危险区域!** -> +> > 以下命令仅用于调试,使用不当可能会造成数据丢失、账号封禁等后果,且仅在调试模式启用且发送者为bpt管理员的情况下可用。 -> +> > 这些指令可能会随时进行添加、修改或删除,且不保证向后兼容。 | 命令 | 结果 | diff --git a/pom.xml b/pom.xml index c003a53f..d3b44572 100644 --- a/pom.xml +++ b/pom.xml @@ -6,7 +6,7 @@ com.github.BotSeira SeiraCore - 1.12.0 + 1.13.0 25 diff --git a/src/main/java/xyz/zcraft/seira/SeiraApplication.java b/src/main/java/xyz/zcraft/seira/SeiraApplication.java index dbf49881..b0d1be9b 100644 --- a/src/main/java/xyz/zcraft/seira/SeiraApplication.java +++ b/src/main/java/xyz/zcraft/seira/SeiraApplication.java @@ -11,10 +11,7 @@ import xyz.zcraft.seira.console.JLineConsole; import xyz.zcraft.seira.console.UserDataConsoleAccess; import xyz.zcraft.seira.db.SqliteDatabase; -import xyz.zcraft.seira.services.BindingService; -import xyz.zcraft.seira.services.BotStat; -import xyz.zcraft.seira.services.DailyLuck; -import xyz.zcraft.seira.services.NoticeStore; +import xyz.zcraft.seira.services.*; import xyz.zcraft.seira.util.AdminRegistry; import xyz.zcraft.seira.util.ApplicationExecutors; @@ -47,6 +44,7 @@ public SeiraApplication(AppConfig config) { DailyLuck.initialize(config.qq().appId()); BotStat.initialize(); NoticeStore.initialize(); + AiPermission.initialize(); ApplicationExecutors createdExecutors = new ApplicationExecutors(); BindingService createdBindingService = new BindingService( @@ -97,6 +95,7 @@ public void close() { bot.close(); DailyLuck.saveToFile(); BotStat.shutdown(); + AiPermission.saveToFile(); LOG.info("Shutdown complete"); } } diff --git a/src/main/java/xyz/zcraft/seira/ai/AiChatHandler.java b/src/main/java/xyz/zcraft/seira/ai/AiChatHandler.java new file mode 100644 index 00000000..0097e4b8 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/ai/AiChatHandler.java @@ -0,0 +1,279 @@ +package xyz.zcraft.seira.ai; + +import com.google.gson.Gson; +import com.google.gson.JsonObject; +import xyz.zcraft.seira.ai.data.AgentFile; +import xyz.zcraft.seira.ai.provider.ChatProvider; +import xyz.zcraft.seira.bot.data.Attachment; +import xyz.zcraft.seira.bot.data.GroupBotState; +import xyz.zcraft.seira.bot.data.MsgElem; +import xyz.zcraft.seira.bot.data.PendingMessage; +import xyz.zcraft.seira.command.Context; +import xyz.zcraft.seira.command.parse.Resolver; +import xyz.zcraft.seira.db.UserDataStore; +import xyz.zcraft.seira.services.AiPermission; + +import java.util.*; +import java.util.concurrent.ConcurrentHashMap; +import java.util.function.Function; +import java.util.function.Predicate; + +import static xyz.zcraft.seira.command.reply.ReplyFactory.ExternalUrls.CHANNEL; +import static xyz.zcraft.seira.command.reply.ReplyFactory.ExternalUrls.PERMISSION; +import static xyz.zcraft.seira.command.reply.ReplyFactory.at; +import static xyz.zcraft.seira.command.reply.ReplyFactory.cmd; + +public class AiChatHandler { + private static final Gson GSON = new Gson(); + private final Resolver resolver; + private final Predicate adminAuthorizer; + private final ChatProvider chatProvider; + private final Function botStateGetter; + private final Map> groupMentionedHistory = new ConcurrentHashMap<>(); + + public AiChatHandler( + Resolver resolver, ChatProvider chatProvider, Predicate isAdmin, + Function botStateGetter + ) { + this.resolver = resolver; + this.adminAuthorizer = isAdmin; + this.chatProvider = chatProvider; + this.botStateGetter = botStateGetter; + } + + public void handleAi(Context ctx) { + if (!ctx.inGroup()) { + ctx.sendReply(at(ctx) + "AI对话仅在群组中可用喵。"); + return; + } + + if (ctx.argumentCount() == 0) { + final boolean b = AiPermission.permits(ctx.groupId()); + final boolean c = AiPermission.isActivated(ctx.groupId()); + ctx.sendReply(at(ctx) + "目前AI对话在本群状态\n" + + "> 已授权:" + (b ? "√" : "×") + "\n" + + "> 已启用:" + (c ? "√" : "×") + ); + return; + } else if (ctx.argumentCount() == 1 + && List.of("on", "off", "grant", "revoke", "reset", "stop").contains(ctx.argument(0).toLowerCase(Locale.ROOT))) { + if ("on".equalsIgnoreCase(ctx.argument(0))) { + final GroupBotState apply = botStateGetter.apply(ctx.groupId()); + if (apply.allowProactiveMsg() && apply.receiveMsgSetting() == GroupBotState.ReceiveMsgSetting.ALL) { + if (AiPermission.permits(ctx.groupId())) { + AiPermission.activate(ctx.groupId()); + ctx.sendReply(at(ctx) + "已启用本群AI对话喵。"); + } else { + if (adminAuthorizer.test(ctx.senderUserId())) { + ctx.sendReply(at(ctx) + "已授权并启用本群AI对话喵。"); + } else { + ctx.sendReply(at(ctx) + "本群无此功能权限,请 [联系 Bot 管理员](%s) 喵。".formatted(CHANNEL)); + } + } + } else { + ctx.sendReply(PendingMessage.ofMarkdownRaw( + at(ctx) + "由于本群未配置权限或配置不完整,暂无法启用本群AI对话喵。" + + "权限配置见[这里](" + PERMISSION + ")~") + ); + } + return; + } else if ("off".equalsIgnoreCase(ctx.argument(0))) { + if (AiPermission.isActivated(ctx.groupId())) { + AiPermission.deactivate(ctx.groupId()); + ctx.sendReply(at(ctx) + "已禁用本群AI对话喵。"); + } else { + ctx.sendReply(at(ctx) + "本群AI对话还未启用喵。"); + } + return; + } else if ("grant".equalsIgnoreCase(ctx.argument(0))) { + if (adminAuthorizer.test(ctx.senderUserId())) { + AiPermission.grant(ctx.groupId()); + ctx.sendReply(at(ctx) + "已授予本群AI对话权限喵。"); + } else { + ctx.sendReply(at(ctx) + "你无权使用该命令喵。"); + } + return; + } else if ("revoke".equalsIgnoreCase(ctx.argument(0))) { + if (adminAuthorizer.test(ctx.senderUserId())) { + AiPermission.revoke(ctx.groupId()); + if (AiPermission.isActivated(ctx.groupId())) { + ctx.sendReply(at(ctx) + "已停用并取消本群AI对话权限喵。"); + } else { + ctx.sendReply(at(ctx) + "已撤销本群AI对话权限喵。"); + } + } else { + ctx.sendReply(at(ctx) + "你无权使用该命令喵。"); + } + return; + } else if ("reset".equalsIgnoreCase(ctx.argument(0))) { + chatProvider.clearState(ctx.groupId(), ctx.senderUserId()); + ctx.sendReply(at(ctx) + "已重置你在本群的AI对话状态喵。"); + return; + } else if ("stop".equalsIgnoreCase(ctx.argument(0))) { + final ChatProvider.StopStatus stopStatus = chatProvider.requireStop(ctx.groupId(), ctx.senderUserId()); + ctx.sendReply(at(ctx) + switch (stopStatus) { + case SUCCESS -> "已停止你在本群的AI对话喵。"; + case FAILED -> "停止AI对话失败了喵。"; + case NO_CONVERSATION -> "目前没有运行中的对话喵。"; + case NOT_SUPPORTED -> "当前不支持停止AI对话。"; + }); + + return; + } + } else if (ctx.argumentCount() == 2 && ctx.argument(0).equalsIgnoreCase("reset")) { + if ("group".equalsIgnoreCase(ctx.argument(1))) { + if (!adminAuthorizer.test(ctx.senderUserId())) { + ctx.sendReply(at(ctx) + "你无权使用该命令喵。"); + return; + } + final int i = chatProvider.clearStateOfGroup(ctx.groupId()); + ctx.sendReply(at(ctx) + "已重置本群" + i + "个用户的AI对话状态喵。"); + return; + } else if ("all".equalsIgnoreCase(ctx.argument(1))) { + final int i = chatProvider.clearStateOfUser(ctx.senderUserId()); + ctx.sendReply(at(ctx) + "已重置你在" + i + "个群中的AI对话状态喵。"); + return; + } + } else if (ctx.argumentCount() == 2 && ctx.argument(0).equalsIgnoreCase("parallel")) { + if (!adminAuthorizer.test(ctx.senderUserId())) { + ctx.sendReply(at(ctx) + "你无权使用该命令喵。"); + return; + } + + final Integer n = resolver.parsePositiveInt(ctx.argument(1)); + if (n != null) { + AiPermission.setParallel(ctx.groupId(), n); + ctx.sendReply(at(ctx) + "已设置本群的最大并行AI对话数量为" + n + "个喵。"); + return; + } + } + + ctx.sendReply(at(ctx) + "用法:/ai [on|off|reset|stop]"); + } + + public void handleChat(Context ctx, String message, List elems) { + if (!AiPermission.isActivated(ctx.groupId())) { + return; + } + + final String at = at(ctx); + if (chatProvider.isRunning(ctx.groupId(), ctx.senderUserId())) { + ctx.sendReply( + at + "你已有一轮对话正在进行中了喵,请稍作等待或" + cmd("/ai stop", "取消对话") + "~" + ); + return; + } + + if (chatProvider.runningCount(ctx.groupId()) >= AiPermission.getParallel(ctx.groupId())) { + ctx.sendReply( + at + "当前群聊中的同时运行对话数量已达到上限,无法开始新的对话喵,请稍作等待。" + ); + return; + } + + final List attachments = elems.stream() + .filter(e -> e.attachments() != null) + .flatMap(e -> e.attachments().stream()) + .map(e -> new AgentFile(e.filename(), null, e.size(), e.url())) + .filter(AgentFile::isValid) + .toList(); + + final var refContent = elems.stream() + .filter(e -> e.content() != null && !e.content().isBlank()) + .findFirst() + .map(MsgElem::content) + .orElse(null); + + chatProvider.input( + ctx.groupId(), + ctx.senderUserId(), + message, + input -> generateVar(ctx, input), + new StreamHandler() { + @Override + public void onText(String message) { + if (!message.trim().startsWith(at.trim())) { + message = at + message; + } + ctx.send(true, PendingMessage.ofMarkdownRaw(message), true); + } + + @Override + public void onComplete(String fullText) { + // Do nothing + } + + @Override + public void onError(String errorCode, String errorMsg) { + ctx.send( + true, + PendingMessage.ofMarkdownRaw( + at + "回复生成失败了喵。\n" + + "> " + errorCode + ": " + errorMsg + "\n" + + "> 若重复出现错误,请尝试" + cmd("/ai reset", "重置会话") + ), + true + ); + } + }, + attachments, + refContent + ); + } + + private String generateVar(Context ctx, String input) { + JsonObject qqContext = new JsonObject(); + qqContext.addProperty("in_group", ctx.inGroup()); + qqContext.addProperty("sender_open_id", ctx.senderUserId()); + qqContext.addProperty("group_id", ctx.groupId()); + + Map> bindings = new HashMap<>(); + + final Deque mentionedHistory = groupMentionedHistory.computeIfAbsent(ctx.groupId(), _ -> new ArrayDeque<>(100)); + + final Set ids = resolver.extractAllMentionedIds(input); + + for (String id : ids) { + if (!mentionedHistory.contains(id)) { + mentionedHistory.push(id); + } + } + + while (mentionedHistory.size() > 64) { + mentionedHistory.removeFirst(); + } + + ids.add(ctx.senderUserId()); + ids.addAll(mentionedHistory); + + for (String openId : ids) { + final Long uid = resolver.resolveBoundUid(openId); + if (uid != null) { + final String username = UserDataStore.findUsername(uid).orElse(""); + bindings.put(openId, Map.of("uid", uid.toString(), "username", username)); + } + } + + qqContext.add("bindings", GSON.toJsonTree(bindings)); + + return qqContext.toString(); + } + + public void recordHistory(String groupId, String userId, String rawContent) { + if (rawContent == null || rawContent.isBlank()) { + return; + } + + recordHistory(groupId, userId, rawContent.trim(), List.of()); + } + + public void recordHistory(String groupId, String userId, String rawContent, List attachments) { + StringBuilder sb = new StringBuilder(rawContent.trim()); + if (attachments != null && !attachments.isEmpty()) { + for (Attachment attachment : attachments) { + sb.append("\n").append("![%s](%s)".formatted(attachment.filename(), attachment.url())); + } + } + chatProvider.recordHistory(groupId, userId, sb.toString()); + } +} diff --git a/src/main/java/xyz/zcraft/seira/ai/StreamHandler.java b/src/main/java/xyz/zcraft/seira/ai/StreamHandler.java new file mode 100644 index 00000000..5179d64e --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/ai/StreamHandler.java @@ -0,0 +1,9 @@ +package xyz.zcraft.seira.ai; + +public interface StreamHandler { + void onText(String message); + + void onComplete(String fullText); + + void onError(String errorCode, String errorMsg); +} diff --git a/src/main/java/xyz/zcraft/seira/ai/data/AgentFile.java b/src/main/java/xyz/zcraft/seira/ai/data/AgentFile.java new file mode 100644 index 00000000..8615e2d2 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/ai/data/AgentFile.java @@ -0,0 +1,14 @@ +package xyz.zcraft.seira.ai.data; + +import com.google.gson.annotations.SerializedName; + +public record AgentFile( + @SerializedName("Name") String name, + @SerializedName("Path") String path, + @SerializedName("Size") Long size, + @SerializedName("Url") String url +) { + public boolean isValid() { + return name != null && url != null; + } +} diff --git a/src/main/java/xyz/zcraft/seira/ai/data/AppConversationBrief.java b/src/main/java/xyz/zcraft/seira/ai/data/AppConversationBrief.java new file mode 100644 index 00000000..902f4c33 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/ai/data/AppConversationBrief.java @@ -0,0 +1,24 @@ +package xyz.zcraft.seira.ai.data; + +import com.google.gson.annotations.SerializedName; + + +/* +{ + "AppConversationID" : "danob6v71098vo7jfsq0", + "ConversationName" : "新的会话", + "CreateTime" : "2026-09-20 15:04:59", + "CreateTimestamp" : 1789887899, + "LastChatTime" : "", + "LastChatTimestamp" : 0, + "EmptyConversation" : false, + "IsPinned" : false, + "ConversationID" : "01M2YT3SNDMQ34FZF2WJD8FTKF" + } + */ +public record AppConversationBrief( + @SerializedName("AppConversationID") String appConversationID, + @SerializedName("ConversationID") String conversationID, + @SerializedName("ConversationName") String conversationName +) { +} diff --git a/src/main/java/xyz/zcraft/seira/ai/data/ChatQueryResponse.java b/src/main/java/xyz/zcraft/seira/ai/data/ChatQueryResponse.java new file mode 100644 index 00000000..3e85e447 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/ai/data/ChatQueryResponse.java @@ -0,0 +1,42 @@ +package xyz.zcraft.seira.ai.data; + + +import com.google.gson.annotations.SerializedName; + +import java.util.List; + +/* +{ + "total_tokens": 993, + "event": "message", + "task_id": "01K0P6NXPHD62CXY80J3HPJCZC", + "id": "01K0P6NXPHD62CXY80J3HPJCZC", + "conversation_id": "01K0BXVWX8CV3HX81V2SPCRP8K", + "answer": "有点遗憾知识库没找到相关内容。不过我很想知道你说“还可以”是在评价什么呀,是一部电影、一顿美食,还是其他方面呢?能多给我些提示,这样我们就能更畅快地交流啦。 ", + "created_at": 1753091868, + "latency": 4.05, + "input_tokens": 937, + "output_tokens": 56, + "start_time_first_resp": 1753091868324, + "latency_first_resp": 4050, + "think_messages": ["推理过程消息内容"], + "tool_messages": ["工具消息输出内容"] +} + */ +public record ChatQueryResponse ( + @SerializedName("total_tokens") Long totalTokens, + @SerializedName("event") String event, + @SerializedName("task_id") String taskId, + @SerializedName("id") String id, + @SerializedName("conversation_id") String conversationId, + @SerializedName("answer") String answer, + @SerializedName("created_at") Long createdAt, + @SerializedName("latency") Double latency, + @SerializedName("input_tokens") Long inputTokens, + @SerializedName("output_tokens") Long outputTokens, + @SerializedName("start_time_first_resp") Long startTimeFirstResp, + @SerializedName("latency_first_resp") Long latencyFirstResp, + @SerializedName("think_messages") List thinkMessages, + @SerializedName("tool_messages") List toolMessages +){ +} diff --git a/src/main/java/xyz/zcraft/seira/ai/provider/ChatProvider.java b/src/main/java/xyz/zcraft/seira/ai/provider/ChatProvider.java new file mode 100644 index 00000000..370c57bf --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/ai/provider/ChatProvider.java @@ -0,0 +1,41 @@ +package xyz.zcraft.seira.ai.provider; + +import xyz.zcraft.seira.ai.StreamHandler; +import xyz.zcraft.seira.ai.data.AgentFile; + +import java.util.List; +import java.util.Set; +import java.util.function.Function; + +public interface ChatProvider { + int CONTEXT_SIZE = 30; + + void recordHistory(String groupId, String sender, String message); + + String input( + String groupId, String openId, String rawContent, + Function contextFunc, StreamHandler handler, + List attachments, String refContent + ); + + StopStatus requireStop(String groupId, String openId); + + boolean isRunning(String groupId, String openId); + + int runningCount(String groupId); + + boolean clearState(String groupId, String openId); + + int clearStateOfGroup(String groupId); + + int clearStateOfUser(String openId); + + Set activeGroupIds(); + + enum StopStatus { + SUCCESS, + NO_CONVERSATION, + NOT_SUPPORTED, + FAILED + } +} diff --git a/src/main/java/xyz/zcraft/seira/ai/provider/ChatProviders.java b/src/main/java/xyz/zcraft/seira/ai/provider/ChatProviders.java new file mode 100644 index 00000000..700e7189 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/ai/provider/ChatProviders.java @@ -0,0 +1,9 @@ +package xyz.zcraft.seira.ai.provider; + +import xyz.zcraft.seira.config.LLMConfig; + +public class ChatProviders { + public static ChatProvider newHiAgentChatProvider(LLMConfig config) { + return new HiAgentProvider(config); + } +} diff --git a/src/main/java/xyz/zcraft/seira/ai/provider/HiAgentProvider.java b/src/main/java/xyz/zcraft/seira/ai/provider/HiAgentProvider.java new file mode 100644 index 00000000..6c178a56 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/ai/provider/HiAgentProvider.java @@ -0,0 +1,729 @@ +package xyz.zcraft.seira.ai.provider; + +import com.google.gson.Gson; +import com.google.gson.JsonObject; +import com.google.gson.JsonParser; +import lombok.Getter; +import lombok.Setter; +import org.apache.logging.log4j.LogManager; +import org.apache.logging.log4j.Logger; +import org.jetbrains.annotations.NotNull; +import xyz.zcraft.seira.ai.StreamHandler; +import xyz.zcraft.seira.ai.data.AgentFile; +import xyz.zcraft.seira.ai.data.AppConversationBrief; +import xyz.zcraft.seira.ai.data.ChatQueryResponse; +import xyz.zcraft.seira.command.Context; +import xyz.zcraft.seira.config.LLMConfig; +import xyz.zcraft.seira.services.AiPermission; + +import java.io.BufferedReader; +import java.io.InputStreamReader; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.nio.charset.StandardCharsets; +import java.util.*; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.function.Consumer; +import java.util.function.Function; +import java.util.regex.Matcher; +import java.util.regex.Pattern; +import java.util.stream.Collectors; + +class HiAgentProvider implements ChatProvider { + // + // + private static final Pattern QQ_FACE = Pattern.compile( + "", Pattern.CASE_INSENSITIVE + ); + // ![ECA4B87395655D69E285314BAB3A2105.jpg](https://multimedia.nt.qq.com.cn/download?appid=14....9_U&spec=0) + private static final Pattern QQ_MEME = Pattern.compile( + "(?: !\\[.*]\\(.*\\))?", Pattern.CASE_INSENSITIVE + ); + // + private static final Pattern QQ_MEME_ALT = Pattern.compile( + "", Pattern.CASE_INSENSITIVE + ); + private final Api api; + private final Map states = new ConcurrentHashMap<>(); + private final Map> chatLog = new ConcurrentHashMap<>(); + private final Map stateCreationLocks = new ConcurrentHashMap<>(); + + protected HiAgentProvider(LLMConfig config) { + this.api = new Api(config); + } + + private static void restoreIncomingMessages( + State state, + Collection pendingMessages + ) { + if (pendingMessages.isEmpty()) { + return; + } + + synchronized (state) { + final Deque merged = new ArrayDeque<>( + pendingMessages.size() + + state.incomingMessages.size() + ); + + merged.addAll(pendingMessages); + merged.addAll(state.incomingMessages); + + state.incomingMessages.clear(); + state.incomingMessages.addAll(merged); + + trimToContextSize(state.incomingMessages); + } + } + + private static void trimToContextSize(Deque messages) { + while (messages.size() > CONTEXT_SIZE) { + messages.removeFirst(); + } + } + + private static String shorten(String input) { + if (input == null || input.isEmpty()) { + return input; + } + + if (input.length() <= 300) { + return input; + } + + return input.substring(0, 150) + "... 已省略 ..." + input.substring(input.length() - 150); + } + + private static String parseQqMeme(String original) { + final Matcher qqFaceMatcher = QQ_FACE.matcher(original); + + if (qqFaceMatcher.matches()) { + original = qqFaceMatcher.replaceAll(r -> { + final String extBase64 = r.group(2); + final String ext = new String(Base64.getDecoder().decode(extBase64)); + final String text = JsonParser.parseString(ext).getAsJsonObject().get("text").getAsString(); + return "[表情:" + text + "]"; + }); + } + + final Matcher qqMemeMatcher = QQ_MEME.matcher(original); + + if (qqMemeMatcher.matches()) { + original = qqMemeMatcher.replaceAll("[表情]"); + } + + final Matcher qqMemeAltMatcher = QQ_MEME_ALT.matcher(original); + + if (qqMemeAltMatcher.matches()) { + original = qqMemeAltMatcher.replaceAll(r -> { + final String extBase64 = r.group(1); + final String ext = new String(Base64.getDecoder().decode(extBase64)); + final String text = JsonParser.parseString(ext).getAsJsonObject().get("text").getAsString(); + return "[表情:" + text + "]"; + }); + } + + return original; + } + + private static String processMessage(String original) { + try { + return parseQqMeme(original); + } catch (Exception e) { + return original; + } + } + + @Override + public void recordHistory(String groupId, String sender, String message) { + if (groupId == null || groupId.isEmpty() + || sender == null || sender.isEmpty() + || message == null || message.isEmpty()) { + return; + } + + final String historyMessage = "**<@" + sender + ">**" + ": " + processMessage(message); + + final Deque log = chatLog.computeIfAbsent(groupId, _ -> new ArrayDeque<>()); + + synchronized (log) { + log.addLast(historyMessage); + trimToContextSize(log); + + states.forEach((owner, state) -> { + if (!owner.groupId().equals(groupId)) { + return; + } + + synchronized (state) { + state.incomingMessages.addLast(historyMessage); + trimToContextSize(state.incomingMessages); + } + }); + } + } + + @Override + public String input( + String groupId, String openId, String rawContent, + Function contextFunc, StreamHandler handler, + List attachments, String refContent + ) { + final StateOwner owner = StateOwner.of(groupId, openId); + final State state = getOrCreateState(owner); + + if (!state.running.compareAndSet(false, true)) { + throw new IllegalStateException("已有请求正在运行"); + } + + try { + final Deque pendingMessages; + + synchronized (state) { + pendingMessages = new ArrayDeque<>(state.incomingMessages); + state.incomingMessages.clear(); + } + + var query = new StringBuilder(); + + final String message = processMessage(rawContent); + + if (pendingMessages.isEmpty()) { + query.append("**<@").append(openId).append(">**").append(": ").append(message); + } else { + if (pendingMessages.size() == CONTEXT_SIZE) { + query.append("====== ...历史消息较多已省略 ======"); + } + query.append(String.join("\n", pendingMessages)) + .append("\n") + .append("====== 以上是最近的所有消息 ======\n"); + + if (refContent != null && !refContent.isBlank()) { + query.append("====== 以下本次询问引用的消息 ======\n") + .append(refContent) + .append("\n"); + } + + query.append("====== 以下是本次询问的内容 ======\n") + .append("\n") + .append("**<@").append(openId).append(">**").append(": ").append(message); + } + + recordHistory(groupId, openId, message); + + state.resetIfNeeded(api, groupId); + + if (contextFunc != null) { + state.getVars().put("CONTEXT", contextFunc.apply(query.toString())); + } + + api.updateConversation(groupId, state.conv.appConversationID(), state.vars); + + final String answer; + try { + if (handler != null) { + answer = api.chatQueryStreaming( + groupId, + state.conv.appConversationID(), + query.toString(), + handler, + attachments, + state::setRunningMessageId + ); + } else { + var response = api.chatQuery(groupId, state.conv.appConversationID(), query.toString(), attachments); + answer = response.answer(); + } + } catch (RuntimeException | Error e) { + restoreIncomingMessages(state, pendingMessages); + + throw e; + } + + recordHistory(groupId, "Seira(你,回复" + openId + "的消息)", shorten(answer)); + return answer; + } finally { + state.running.set(false); + } + } + + @Override + public StopStatus requireStop(String groupId, String openId) { + final StateOwner owner = StateOwner.of(groupId, openId); + final State state = states.get(owner); + + if (state == null || !state.running.get()) { + return StopStatus.NO_CONVERSATION; + } + + if (state.getRunningMessageId() == null) { + return StopStatus.NOT_SUPPORTED; + } + + try { + api.stopConversation(state.getRunningMessageId(), groupId); + return StopStatus.SUCCESS; + } catch (Exception e) { + return StopStatus.FAILED; + } + } + + @Override + public boolean isRunning(String groupId, String openId) { + final State state = states.get( + StateOwner.of(groupId, openId) + ); + + return state != null && state.running.get(); + } + + @Override + public int runningCount(String groupId) { + int count = 0; + + for (Map.Entry entry : states.entrySet()) { + if (!entry.getKey().groupId().equals(groupId)) { + continue; + } + + count++; + } + + return count; + } + + @Override + public boolean clearState(String groupId, String openId) { + final State state = states.get( + StateOwner.of(groupId, openId) + ); + + if (state == null) { + return false; + } + + state.resetting.set(true); + return true; + } + + @Override + public int clearStateOfGroup(String groupId) { + int count = 0; + + for (Map.Entry entry : states.entrySet()) { + if (!entry.getKey().groupId().equals(groupId)) { + continue; + } + + entry.getValue().resetting.set(true); + count++; + } + + return count; + } + + @Override + public int clearStateOfUser(String openId) { + int count = 0; + + for (Map.Entry entry : states.entrySet()) { + if (!entry.getKey().openId().equals(openId)) { + continue; + } + + entry.getValue().resetting.set(true); + count++; + } + + return count; + } + + @Override + public Set activeGroupIds() { + return states.entrySet() + .stream() + .filter(entry -> entry.getValue().running.get()) + .map(entry -> entry.getKey().groupId()) + .collect(Collectors.toSet()); + } + + private State getOrCreateState(StateOwner owner) { + State state = states.get(owner); + + if (state != null) { + return state; + } + + final Object creationLock = stateCreationLocks.computeIfAbsent(owner, _ -> new Object()); + + try { + synchronized (creationLock) { + state = states.get(owner); + + if (state != null) { + return state; + } + + final String groupId = owner.groupId(); + final Deque log = chatLog.computeIfAbsent(groupId, _ -> new ArrayDeque<>()); + final AppConversationBrief conv = api.createConversation(groupId); + + synchronized (log) { + final State created = State.create(conv, new ArrayList<>(log)); + + states.put(owner, created); + + return created; + } + } + } finally { + stateCreationLocks.remove(owner, creationLock); + } + } + + @Getter + static final class State { + private final ConcurrentHashMap vars; + private final AtomicBoolean running; + private final AtomicBoolean resetting; + private final Deque incomingMessages; + private AppConversationBrief conv; + @Setter + private volatile String runningMessageId; + + private State( + AppConversationBrief conv, + ConcurrentHashMap vars, + AtomicBoolean running, + AtomicBoolean resetting, + Deque incomingMessages + ) { + this.conv = conv; + this.vars = vars; + this.running = running; + this.resetting = resetting; + this.incomingMessages = incomingMessages; + } + + + public static State create( + AppConversationBrief conv, + Collection incomingMessages + ) { + return new State( + conv, + new ConcurrentHashMap<>(), + new AtomicBoolean(false), + new AtomicBoolean(false), + new ArrayDeque<>(incomingMessages) + ); + } + + public void resetIfNeeded(Api api, String groupId) { + if (!resetting.compareAndSet(true, false)) { + return; + } + + try { + api.clearConversation(groupId, conv.appConversationID()); + conv = api.createConversation(groupId); + vars.clear(); + incomingMessages.clear(); + } catch (RuntimeException | Error e) { + resetting.set(true); + throw e; + } + } + } + + record StateOwner(String groupId, String openId) { + public static StateOwner of( + String groupId, + String openId + ) { + return new StateOwner(groupId, openId); + } + + public static StateOwner of(Context ctx) { + return new StateOwner( + ctx.groupId(), + ctx.senderUserId() + ); + } + } +} + +class Api { + private static final Logger LOG = LogManager.getLogger(Api.class); + public final HttpClient CLIENT = HttpClient.newHttpClient(); + public final Gson GSON = new Gson(); + public final String apiKey; + public final String endpoint; + + public Api(LLMConfig config) { + this.endpoint = config.baseUrl(); + this.apiKey = config.apiKey(); + } + + public AppConversationBrief createConversation(String groupId) { + LOG.info("Creating conversation for id {}", groupId); + JsonObject body = new JsonObject(); + body.addProperty("UserID", groupId); + try { + var request = newRequest("/api/proxy/api/v1/create_conversation") + .POST(HttpRequest.BodyPublishers.ofString(body.toString())) + .build(); + + final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); + + if (send.statusCode() != 200) { + throw new RuntimeException("Failed to create conversation: " + send.statusCode()); + } + + final var response = JsonParser.parseString(send.body()).getAsJsonObject(); + + final AppConversationBrief conversation = GSON.fromJson( + response.getAsJsonObject("Conversation"), + AppConversationBrief.class + ); + + LOG.info("Conversation for id {} created, conv id {}", groupId, conversation.appConversationID()); + return conversation; + } catch (Exception e) { + throw new RuntimeException("Error creating conversation", e); + } + } + + public void updateConversation(String groupId, String appConvId, Map variables) { + LOG.info("Updating conversation for id {}", groupId); + + JsonObject body = new JsonObject(); + body.addProperty("UserID", groupId); + body.addProperty("AppConversationID", appConvId); + body.add("Inputs", GSON.toJsonTree(variables)); + + try { + var request = newRequest("/api/proxy/api/v1/update_conversation") + .POST(HttpRequest.BodyPublishers.ofString(body.toString())) + .build(); + + final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); + + if (send.statusCode() != 200) { + throw new RuntimeException("Failed to update conversation: " + send.statusCode()); + } + + LOG.info("Updated conversation for id {}", groupId); + } catch (Exception e) { + throw new RuntimeException("Error updating conversation", e); + } + } + + public void clearConversation(String groupId, String appConvId) { + LOG.info("Clearing conversation for id {}", groupId); + + JsonObject body = new JsonObject(); + body.addProperty("UserID", groupId); + body.addProperty("AppConversationID", appConvId); + + try { + var request = newRequest("/api/proxy/api/v1/clear_message") + .POST(HttpRequest.BodyPublishers.ofString(body.toString())) + .build(); + + final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); + + if (send.statusCode() != 200) { + throw new RuntimeException("Failed to clear conversation: " + send.statusCode()); + } + + LOG.info("Cleared conversation for id {}", groupId); + } catch (Exception e) { + throw new RuntimeException("Error clearing conversation", e); + } + } + + public void stopConversation(String messageId, String groupId) { + LOG.info("Stopping conversation for id {}", groupId); + + JsonObject body = new JsonObject(); + body.addProperty("UserID", groupId); + body.addProperty("MessageID", messageId); + + try { + var request = newRequest("/api/proxy/api/v1/stop_message") + .POST(HttpRequest.BodyPublishers.ofString(body.toString())) + .build(); + + final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); + + if (send.statusCode() != 200) { + throw new RuntimeException("Failed to stop conversation: " + send.statusCode()); + } + + LOG.info("Stopped conversation for id {}", groupId); + } catch (Exception e) { + throw new RuntimeException("Error clearing conversation", e); + } + } + + public ChatQueryResponse chatQuery( + String groupId, String appConvId, String query, List attachments + ) { + LOG.info("Running chat query for id {}", groupId); + + JsonObject body = new JsonObject(); + body.addProperty("UserID", groupId); + body.addProperty("AppConversationID", appConvId); + body.addProperty("Query", query); + body.addProperty("ResponseMode", "blocking"); + + if (attachments != null && !attachments.isEmpty()) { + body.add("QueryExtends", GSON.toJsonTree(Map.of("Files", attachments))); + } + + try { + var request = newRequest("/api/proxy/api/v1/chat_query_v2") + .POST(HttpRequest.BodyPublishers.ofString(body.toString())) + .build(); + + final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); + + if (send.statusCode() != 200) { + throw new RuntimeException("Failed to query conversation: " + send.statusCode()); + } + + final ChatQueryResponse chatQueryResponse = GSON.fromJson(send.body(), ChatQueryResponse.class); + + LOG.info( + "Chat query success for id {}, Token input:{}, output:{}", + groupId, + chatQueryResponse.inputTokens(), + chatQueryResponse.outputTokens() + ); + + return chatQueryResponse; + } catch (Exception e) { + throw new RuntimeException("Error querying conversation", e); + } + } + + public String chatQueryStreaming( + String groupId, String appConvId, String query, + StreamHandler handler, List attachments, + Consumer taskIdSetter + ) { + LOG.info("Running chat query for id {}", groupId); + + JsonObject body = new JsonObject(); + body.addProperty("UserID", groupId); + body.addProperty("AppConversationID", appConvId); + body.addProperty("Query", query); + body.addProperty("ResponseMode", "streaming"); + + if (attachments != null && !attachments.isEmpty()) { + body.add("QueryExtends", GSON.toJsonTree(Map.of("Files", attachments))); + } + + try { + var request = newRequest("/api/proxy/api/v1/chat_query_v2") + .POST(HttpRequest.BodyPublishers.ofString(body.toString())) + .build(); + + final var response = CLIENT.send(request, HttpResponse.BodyHandlers.ofInputStream()); + + if (response.statusCode() != 200) { + throw new RuntimeException("Failed to query conversation: " + response.statusCode()); + } + + StringBuilder fullAnswer = new StringBuilder(); + StringBuilder currentMessage = new StringBuilder(); + boolean ended = false; + + try (BufferedReader reader = new BufferedReader( + new InputStreamReader(response.body(), StandardCharsets.UTF_8) + )) { + String line; + + while ((line = reader.readLine()) != null) { + if (!line.startsWith("data:")) { + continue; + } + + String data = line.substring(5).trim(); + + if (data.isEmpty()) { + continue; + } + + JsonObject event = JsonParser.parseString(data).getAsJsonObject(); + String eventType = event.get("event").getAsString(); + String taskId = event.get("task_id").getAsString(); + + if (taskIdSetter != null && taskId != null && !taskId.isEmpty()) { + taskIdSetter.accept(taskId); + } + + switch (eventType) { + case "message" -> { + String delta = event.get("answer").getAsString(); + + currentMessage.append(delta); + fullAnswer.append(delta); + } + + case "agent_thought", "message_output_end" -> flushMessage(currentMessage, handler); + + case "agent_error" -> { + final String errorMsg = event.get("error_msg").getAsString(); + final String errorCode = event.get("error_code").getAsString(); + handler.onError(errorCode, errorMsg); + throw new RuntimeException("Error when querying stream conversation: " + errorCode + ": " + errorMsg); + } + + case "message_cost" -> { + // token statistics + } + + case "message_end" -> ended = true; + } + + if (ended) { + break; + } + } + } + + final String result = fullAnswer.toString(); + + handler.onComplete(result); + + LOG.info("Chat query success for id {}", groupId); + + return result; + } catch (Exception e) { + throw new RuntimeException("Error querying conversation", e); + } + } + + private void flushMessage(@NotNull StringBuilder currentMessage, @NotNull StreamHandler handler) { + final String message = currentMessage.toString().trim(); + + if (message.isBlank() || "大模型接口调用出错,请联系管理员".equals(message)) { + return; + } + + currentMessage.setLength(0); + + handler.onText(message); + } + + private HttpRequest.Builder newRequest(String path) { + return HttpRequest.newBuilder() + .uri(java.net.URI.create(this.endpoint + path)) + .header("Apikey", this.apiKey) + .header("Content-Type", "application/json"); + } +} diff --git a/src/main/java/xyz/zcraft/seira/api/ApiUtil.java b/src/main/java/xyz/zcraft/seira/api/ApiUtil.java new file mode 100644 index 00000000..70ca2bc3 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/api/ApiUtil.java @@ -0,0 +1,99 @@ +package xyz.zcraft.seira.api; + +import com.google.gson.Gson; +import com.google.gson.JsonArray; +import com.google.gson.JsonElement; +import com.google.gson.JsonObject; +import xyz.zcraft.seira.api.data.RawResponse; + +import java.nio.charset.StandardCharsets; + +public class ApiUtil { + private static final Gson GSON = new Gson(); + + public static RuntimeException parseHttpError(String responseBody, int statusCode, String fallbackMessage) { + Integer errorCode = null; + String message = fallbackMessage; + try { + JsonObject root = GSON.fromJson(responseBody, JsonObject.class); + if (root != null) { + if (root.has("data") && root.get("data").isJsonObject()) { + JsonObject data = root.getAsJsonObject("data"); + errorCode = readCodeFromJsonObject(data); + } + if (errorCode == null) { + errorCode = readCodeFromJsonObject(root); + } + } + } catch (Exception ignored) { + } + + if (statusCode == 500) { + message += "(" + (errorCode == null ? "未知错误码" : errorCode) + " / HTTP " + statusCode + " / 发生了一个内部错误)"; + } + + return new ApiRequestException(errorCode, message); + } + + static Integer readCodeFromJsonObject(JsonObject object) { + if (object == null || !object.has("code") || !object.get("code").isJsonPrimitive()) { + return null; + } + try { + return object.get("code").getAsInt(); + } catch (Exception ignored) { + return null; + } + } + + static void ensureApiSuccess(RawResponse payload, String fallbackMessage) { + if (payload == null) { + throw new RuntimeException(fallbackMessage); + } + if (!payload.isSuccess()) { + Integer errorCode = extractErrorCode(payload); + String message = payload.getMessage() != null ? payload.getMessage() : fallbackMessage; + throw new ApiRequestException(errorCode, message); + } + } + + static RuntimeException parseHttpError(byte[] responseBody, int statusCode, String fallbackMessage) { + String bodyAsText = responseBody == null ? null : new String(responseBody, StandardCharsets.UTF_8); + return parseHttpError(bodyAsText, statusCode, fallbackMessage); + } + + static Integer extractErrorCode(RawResponse payload) { + if (payload.getData() != null && payload.getData().isJsonObject()) { + JsonObject data = payload.getData().getAsJsonObject(); + return readCodeFromJsonObject(data); + } + return null; + } + + static boolean codeNotOk(int statusCode) { + return statusCode < 200 || statusCode >= 300; + } + + static JsonObject requireDataObject(RawResponse payload, String message) { + if (payload.getData() == null || !payload.getData().isJsonObject()) { + throw new RuntimeException(message); + } + return payload.getData().getAsJsonObject(); + } + + static JsonArray requireResultArray(RawResponse payload, String message) { + JsonElement data = payload.getData(); + if (data != null && data.isJsonArray()) { + return data.getAsJsonArray(); + } + if (data == null || !data.isJsonObject()) { + throw new RuntimeException(message); + } + + JsonElement result = data.getAsJsonObject().get("result"); + if (result == null || !result.isJsonArray()) { + throw new RuntimeException(message); + } + return result.getAsJsonArray(); + } +} diff --git a/src/main/java/xyz/zcraft/seira/api/AsteroidApi.java b/src/main/java/xyz/zcraft/seira/api/AsteroidApi.java new file mode 100644 index 00000000..0293d97f --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/api/AsteroidApi.java @@ -0,0 +1,92 @@ +package xyz.zcraft.seira.api; + +import com.google.gson.Gson; +import com.google.gson.JsonElement; +import com.google.gson.JsonObject; +import org.apache.logging.log4j.LogManager; +import org.apache.logging.log4j.Logger; +import xyz.zcraft.seira.Seira; +import xyz.zcraft.seira.api.data.MinecraftServerStatus; +import xyz.zcraft.seira.api.data.RawResponse; + +import java.net.URI; +import java.net.URLEncoder; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.nio.charset.StandardCharsets; +import java.time.Duration; + +public class AsteroidApi { + private static final String ENDPOINT; + private static final String TOKEN; + private static final HttpClient CLIENT = HttpClient.newBuilder().connectTimeout(Duration.ofMinutes(5)).build(); + private static final Gson GSON = new Gson(); + + static { + ENDPOINT = Seira.getConfig().asteroid().endpoint(); + TOKEN = Seira.getConfig().asteroid().token(); + } + + private static HttpRequest.Builder requestBuilder(String path) { + HttpRequest.Builder builder = HttpRequest.newBuilder(); + if (TOKEN != null && !TOKEN.isBlank()) { + builder.header("Authorization", "Bearer " + TOKEN); + } + builder.uri(URI.create(ENDPOINT + path)); + return builder; + } + + public static MinecraftServerStatus getMinecraftServerStatus(String addr) { + try { + var request = requestBuilder("/minecraft/servers/" + URLEncoder.encode(addr, StandardCharsets.UTF_8) + "/status") + .GET() + .build(); + + final HttpResponse response = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); + + if (response.statusCode() != 200) { + throw ApiUtil.parseHttpError(response.body(), response.statusCode(), "获取 MC 服务器状态失败"); + } + + final RawResponse r = GSON.fromJson(response.body(), RawResponse.class); + ApiUtil.ensureApiSuccess(r, "获取 MC 服务器状态失败"); + + return GSON.fromJson(r.getData(), MinecraftServerStatus.class); + } catch (Exception e) { + throw new RuntimeException("获取 MC 服务器状态失败", e); + } + } + + public record ServerStatus(boolean online, String version){} + + public static ServerStatus getServerStatus() { + try { + var request = requestBuilder("/health") + .GET() + .build(); + + final HttpResponse response = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); + final RawResponse rawResponse = GSON.fromJson(response.body(), RawResponse.class); + + if (response.statusCode() == 200 + && response.body() != null + && rawResponse != null + && rawResponse.isSuccess()) { + final JsonElement rawResponseData = rawResponse.getData(); + if (rawResponseData != null && rawResponseData.isJsonObject()) { + JsonObject data = rawResponseData.getAsJsonObject(); + return new ServerStatus(true, data.get("version").getAsString()); + } + return new ServerStatus(true, "?"); + } + } catch (Exception e) { + LOG.error("Failed to get server status", e); + } + + LOG.warn("Asteroid server is down."); + return new ServerStatus(false, null); + } + + private static final Logger LOG = LogManager.getLogger(AsteroidApi.class); +} diff --git a/src/main/java/xyz/zcraft/seira/api/APIHelper.java b/src/main/java/xyz/zcraft/seira/api/OstellaApi.java similarity index 66% rename from src/main/java/xyz/zcraft/seira/api/APIHelper.java rename to src/main/java/xyz/zcraft/seira/api/OstellaApi.java index 828b2f6f..fbec79ba 100644 --- a/src/main/java/xyz/zcraft/seira/api/APIHelper.java +++ b/src/main/java/xyz/zcraft/seira/api/OstellaApi.java @@ -11,8 +11,6 @@ import xyz.zcraft.seira.api.data.*; import xyz.zcraft.seira.bot.data.FileInfo; import xyz.zcraft.seira.command.ResolutionException; -import xyz.zcraft.seira.command.parse.ShortcutTarget; -import xyz.zcraft.seira.data.UserRef; import xyz.zcraft.seira.util.TimeDurationParser; import java.io.IOException; @@ -27,8 +25,10 @@ import java.util.List; import java.util.Map; -public class APIHelper { +public class OstellaApi { + private static final String OSU_AUTHORIZATION_HEADER = "X-Osu-Authorization"; private static final String ENDPOINT; + private static final String TOKEN; private static final HttpClient CLIENT = HttpClient.newBuilder().connectTimeout(Duration.ofMinutes(5)).build(); private static final Gson GSON = new Gson(); private static final int REPLAY_POLL_INTERVAL_MS = 5000; @@ -36,25 +36,37 @@ public class APIHelper { static { ENDPOINT = Seira.getConfig().ostella().endpoint(); + TOKEN = Seira.getConfig().ostella().token(); + } + + private static HttpRequest.Builder requestBuilder() { + HttpRequest.Builder builder = HttpRequest.newBuilder(); + if (TOKEN != null && !TOKEN.isBlank()) { + builder.header("Authorization", "Bearer " + TOKEN); + } + return builder; + } + + private static HttpRequest.Builder withOsuAuthorization(HttpRequest.Builder builder, String accessToken) { + return builder.header(OSU_AUTHORIZATION_HEADER, "Bearer " + accessToken); } public static Response> getFollowed(String accessToken) { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = withOsuAuthorization(requestBuilder(), accessToken) .uri(URI.create(ENDPOINT + "/users/me/friends")) - .header("Authorization", "Bearer " + accessToken) .GET() .build(); final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "获取多人房间失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "获取多人房间失败"); } final RawResponse r = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(r, "获取多人房间失败"); - final JsonArray data = r.getData().getAsJsonArray(); + ApiUtil.ensureApiSuccess(r, "获取多人房间失败"); + final JsonArray data = ApiUtil.requireResultArray(r, "获取好友响应缺少结果数组"); LinkedList followed = new LinkedList<>(); @@ -72,20 +84,19 @@ public static Response> getFollowed(String accessToken) { public static Response getSelf(String accessToken) { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = withOsuAuthorization(requestBuilder(), accessToken) .uri(URI.create(ENDPOINT + "/users/me")) - .header("Authorization", "Bearer " + accessToken) .GET() .build(); final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "获取用户信息失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "获取用户信息失败"); } final RawResponse r = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(r, "获取用户信息失败"); + ApiUtil.ensureApiSuccess(r, "获取用户信息失败"); final var data = r.getData().getAsJsonObject(); return Response.fromHeaders(send.headers()) @@ -97,16 +108,15 @@ public static Response getSelf(String accessToken) { } @SuppressWarnings("unused") - public static Response getBoNResponse(int n, UserRef userRef) { - return getBoNResponse(n, userRef, List.of()); + public static Response getBoNResponse(int n, long uid) { + return getBoNResponse(n, uid, List.of()); } - public static Response getBoNResponse(int n, UserRef userRef, List filters) { - return getBoNResponse(n, 1, userRef, filters); + public static Response getBoNResponse(int n, long uid, List filters) { + return getBoNResponse(n, 1, uid, filters); } - public static Response getBoNResponse(int n, int start, UserRef userRef, List filters) { - long uid = resolveUid(userRef); + public static Response getBoNResponse(int n, int start, long uid, List filters) { return getBase64BytesResponse( "/users/" + uid + "/scores/bestof?n=" + n + encodeScoreRangeStart(start) + encodeScoreFilters(filters), "获取最好成绩失败", @@ -114,8 +124,7 @@ public static Response getBoNResponse(int n, int start, UserRef use ); } - public static Response getUserInfoResponse(UserRef userRef) { - long uid = resolveUid(userRef); + public static Response getUserInfoResponse(long uid) { return getBase64BytesResponse( "/users/" + uid, "获取玩家资料失败", @@ -123,10 +132,9 @@ public static Response getUserInfoResponse(UserRef userRef) { ); } - public static UserExtended getUserRaw(UserRef userRef) { - long uid = resolveUid(userRef); + public static UserExtended getUserRaw(long uid) { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/users/" + uid)) .header("Accept", "application/json") .GET() @@ -135,11 +143,11 @@ public static UserExtended getUserRaw(UserRef userRef) { final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "获取用户信息失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "获取用户信息失败"); } final RawResponse r = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(r, "获取用户信息失败"); + ApiUtil.ensureApiSuccess(r, "获取用户信息失败"); final var data = r.getData().getAsJsonObject(); return GSON.fromJson(data, UserExtended.class); @@ -149,12 +157,11 @@ public static UserExtended getUserRaw(UserRef userRef) { } @SuppressWarnings("unused") - public static Response getTodayBestResponse(UserRef userRef) { - return getTodayBestResponse(userRef, 1); + public static Response getTodayBestResponse(long uid) { + return getTodayBestResponse(uid, 1); } - public static Response getTodayBestResponse(UserRef userRef, int days) { - long uid = resolveUid(userRef); + public static Response getTodayBestResponse(long uid, int days) { return getBase64BytesResponse( "/users/" + uid + "/scores/today-best?days=" + days, "获取近期BP失败", @@ -180,7 +187,7 @@ public static Response getLeaderboardResponse(List uids) { public static String getDaily() { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/daily")) .GET() .build(); @@ -188,11 +195,11 @@ public static String getDaily() { final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "获取每日挑战失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "获取每日挑战失败"); } final RawResponse r = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(r, "获取每日挑战失败"); + ApiUtil.ensureApiSuccess(r, "获取每日挑战失败"); final JsonObject data = r.getData().getAsJsonObject(); String mods = null; @@ -220,20 +227,19 @@ public static String getDaily() { public static Response getMultiplayerRoom(String accessToken) { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = withOsuAuthorization(requestBuilder(), accessToken) .uri(URI.create(ENDPOINT + "/multiplayer/rooms/current")) - .header("Authorization", "Bearer " + accessToken) .GET() .build(); final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "获取多人房间失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "获取多人房间失败"); } final RawResponse r = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(r, "获取多人房间失败"); + ApiUtil.ensureApiSuccess(r, "获取多人房间失败"); final JsonObject data = r.getData().getAsJsonObject(); return Response.fromHeaders(send.headers()) @@ -245,22 +251,21 @@ public static Response getMultiplayerRoom(String accessToken) { } @SuppressWarnings("unused") - public static Response getRecentResponse(int n, UserRef userRef, boolean includeFail) { - return getRecentResponse(n, userRef, includeFail, List.of()); + public static Response getRecentResponse(int n, long uid, boolean includeFail) { + return getRecentResponse(n, uid, includeFail, List.of()); } - public static Response getRecentResponse(int n, UserRef userRef, boolean includeFail, List filters) { - return getRecentResponse(n, 1, userRef, includeFail, filters); + public static Response getRecentResponse(int n, long uid, boolean includeFail, List filters) { + return getRecentResponse(n, 1, uid, includeFail, filters); } public static Response getRecentResponse( int n, int start, - UserRef userRef, + long uid, boolean includeFail, List filters ) { - long uid = resolveUid(userRef); return getBase64BytesResponse( "/users/" + uid + "/scores/recent?n=" + n + "&fail=" + includeFail + encodeScoreRangeStart(start) + encodeScoreFilters(filters), @@ -301,81 +306,24 @@ public static Response getBeatmapBgResponse(long beatmapId) { return getBase64BytesResponse("/beatmaps/" + beatmapId + "/background", "获取谱面失败", null); } - public static long lookupBeatmap(ShortcutTarget target, String auth) { - long beatmapId; - if (target.isLocalScore() || "s".equals(target.macroType())) { - beatmapId = lookupScoreData(lookupScoreId(target, List.of(), null)).get("beatmap_id").getAsLong(); - } else if ("m".equals(target.macroType())) { - beatmapId = target.explicitId(); - } else if (!target.isMacro()) { - beatmapId = target.explicitId(); - } else { - try { - final String query = getBeatmapQuery(target); - - HttpRequest localRequest = HttpRequest.newBuilder() - .uri(URI.create(ENDPOINT + query)) - .header("Authorization", "Bearer " + auth) - .GET() - .build(); - - final HttpResponse send = CLIENT.send(localRequest, HttpResponse.BodyHandlers.ofString()); - - if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "查找谱面失败"); - } - - final RawResponse rawResponse = GSON.fromJson(send.body(), RawResponse.class); - - ensureApiSuccess(rawResponse, "查找谱面失败"); - - beatmapId = rawResponse.getData().getAsJsonObject().get("beatmap_id").getAsLong(); - } catch (IOException | InterruptedException e) { - throw requestFailure(e); - } - } - return beatmapId; - } - - private static String getBeatmapQuery(ShortcutTarget target) { - String query = "/beatmaps/lookup?"; - if (target.isMacro()) { - switch (target.macroType().toLowerCase()) { - case "rs", "bp", "rp" -> { - query += "&of=" + target.macroType() + "&u=" + resolveUid(target.userRef()); - query += "&i=" + target.macroIndex(); - } - case "ms" -> { - query += "&ms=" + target.explicitId(); - query += "&i=" + target.macroIndex(); - } - case "mp" -> query += "&of=mp"; - } - } else { - query = "/beatmap/lookup?m=" + target.explicitId(); - } - - return query; - } - public static Response getBeatmapsetResponse(long beatmapsetId) { return getBase64BytesResponse("/beatmapsets/" + beatmapsetId, "获取谱面集失败", null); } public static Beatmapset getBeatmapsetRaw(long id) { try { - var builder = HttpRequest.newBuilder() + var builder = requestBuilder() .uri(URI.create(ENDPOINT + "/beatmapsets/" + id)) .header("Accept", "application/json"); final HttpResponse send = CLIENT.send(builder.build(), HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "获取谱面集失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "获取谱面集失败"); } final RawResponse r = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(r, "获取谱面集失败"); + ApiUtil.ensureApiSuccess(r, "获取谱面集失败"); final JsonObject data = r.getData().getAsJsonObject(); return GSON.fromJson(data, Beatmapset.class); @@ -384,53 +332,6 @@ public static Beatmapset getBeatmapsetRaw(long id) { } } - public static long lookupBeatmapset(ShortcutTarget target, String auth) { - long beatmapsetId; - if (target.isLocalScore() || "s".equals(target.macroType())) { - return lookupBeatmapset(new ShortcutTarget(lookupBeatmap(target, auth), null, "m", null, null), auth); - } else if (!target.isMacro() || "ms".equals(target.macroType())) { - beatmapsetId = target.explicitId(); - } else { - try { - final String query = getBeatmapsetQuery(target); - - HttpRequest localRequest = HttpRequest.newBuilder() - .uri(URI.create(ENDPOINT + query)) - .header("Authorization", "Bearer " + auth) - .GET() - .build(); - - final HttpResponse send = CLIENT.send(localRequest, HttpResponse.BodyHandlers.ofString()); - - if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "查找谱面集失败"); - } - - final RawResponse rawResponse = GSON.fromJson(send.body(), RawResponse.class); - - ensureApiSuccess(rawResponse, "查找谱面集失败"); - - beatmapsetId = rawResponse.getData().getAsJsonObject().get("beatmapset_id").getAsLong(); - } catch (IOException | InterruptedException e) { - throw requestFailure(e); - } - } - return beatmapsetId; - } - - private static String getBeatmapsetQuery(ShortcutTarget target) { - String query = "/beatmapsets/lookup"; - - return switch (target.macroType().toLowerCase()) { - case "m" -> query + "?m=" + target.explicitId(); - case "ms" -> query + "?ms=" + target.explicitId(); - case "rs", "bp", "rp" -> - query + "?of=" + target.macroType() + "&i=" + target.macroIndex() + "&u=" + resolveUid(target.userRef()); - case "mp" -> query + "?of=mp"; - case null, default -> throw new ResolutionException("快捷查询格式错误。"); - }; - } - public static Response getScoreResponse(String scoreId) { return getBase64BytesResponse("/scores/" + scoreId, "获取成绩失败", null); } @@ -445,7 +346,7 @@ public static Response getMissVisualizeResponse(String scoreId, int private static Response getBase64BytesResponse(String query, String failMessage, @Nullable String postBody) { try { - var builder = HttpRequest.newBuilder() + var builder = requestBuilder() .uri(URI.create(ENDPOINT + query)); if (postBody != null) { @@ -457,7 +358,7 @@ private static Response getBase64BytesResponse(String query, String final HttpResponse send = CLIENT.send(builder.build(), HttpResponse.BodyHandlers.ofByteArray()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), failMessage); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), failMessage); } byte[] imageBytes = send.body(); @@ -470,35 +371,23 @@ private static Response getBase64BytesResponse(String query, String } } - private static String getScoreQuery(ShortcutTarget target) { - return switch (target.macroType().toLowerCase()) { - case "rs", "bp", "rp" -> - "/scores/lookup?of=" + target.macroType() + "&i=" + target.macroIndex() + "&u=" + resolveUid(target.userRef()); - case "m" -> "/scores/lookup?m=" + target.explicitId() + "&u=" + resolveUid(target.userRef()); - case "ms" -> - "/scores/lookup?ms=" + target.explicitId() + "&i=" + target.macroIndex() + "&u=" + resolveUid(target.userRef()); - case null, default -> throw new IllegalArgumentException("Invalid macro type"); - }; - } - public static Response getLookupBeatmapsetResponse(long beatmapsetId, String auth) { try { final String query = "/beatmapsets/lookup?ms=" + beatmapsetId; - HttpRequest localRequest = HttpRequest.newBuilder() + HttpRequest localRequest = withOsuAuthorization(requestBuilder(), auth) .uri(URI.create(ENDPOINT + query)) - .header("Authorization", "Bearer " + auth) .GET() .build(); final HttpResponse send = CLIENT.send(localRequest, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "获取谱面集失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "获取谱面集失败"); } final RawResponse rawResponse = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(rawResponse, "查找谱面集失败"); + ApiUtil.ensureApiSuccess(rawResponse, "查找谱面集失败"); final JsonObject data = rawResponse.getData().getAsJsonObject(); return Response.fromHeaders(send.headers()) @@ -511,7 +400,7 @@ public static Response getLookupBeatmapsetResponse(long beatmapsetId, String public static Response> searchBeatmapSetResponse(SearchQuery query) { try { - HttpRequest localRequest = HttpRequest.newBuilder() + HttpRequest localRequest = requestBuilder() .uri(URI.create(ENDPOINT + "/beatmapsets/search?" + "q=" + URLEncoder.encode(query.query(), StandardCharsets.UTF_8))) .GET() .build(); @@ -519,12 +408,12 @@ public static Response> searchBeatmapSetResponse(SearchQu final var send = CLIENT.send(localRequest, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "搜索谱面集失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "搜索谱面集失败"); } final RawResponse rawResponse = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(rawResponse, "搜索谱面集失败"); - final JsonArray data = rawResponse.getData().getAsJsonArray(); + ApiUtil.ensureApiSuccess(rawResponse, "搜索谱面集失败"); + final JsonArray data = ApiUtil.requireResultArray(rawResponse, "搜索谱面集响应缺少结果数组"); final LinkedList items = new LinkedList<>(); @@ -549,7 +438,7 @@ public static ReplayTaskInfo createObscuredReplayRenderTask(long scoreId, QqUplo } TimeDurationParser.TimeRange timeRange = getScoreHighlight(scoreId, 10); - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/replays/renders/score/" + scoreId + "?obscured=true" + timeRange.toQueryString())) .header("Content-Type", "application/json") @@ -561,19 +450,19 @@ public static ReplayTaskInfo createObscuredReplayRenderTask(long scoreId, QqUplo public static RandomScore getRandomScore() { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/scores/random?min_rank=500000")) .GET() .build(); HttpResponse response = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); - if (codeNotOk(response.statusCode())) { - throw parseHttpError(response.body(), response.statusCode(), "获取随机成绩失败"); + if (ApiUtil.codeNotOk(response.statusCode())) { + throw ApiUtil.parseHttpError(response.body(), response.statusCode(), "获取随机成绩失败"); } RawResponse payload = GSON.fromJson(response.body(), RawResponse.class); - ensureApiSuccess(payload, "获取随机成绩失败"); - JsonObject data = requireDataObject(payload, "随机成绩响应缺少data"); + ApiUtil.ensureApiSuccess(payload, "获取随机成绩失败"); + JsonObject data = ApiUtil.requireDataObject(payload, "随机成绩响应缺少data"); if (!data.has("user") || !data.get("user").isJsonObject() || !data.has("score") || !data.get("score").isJsonObject()) { throw new RuntimeException("随机成绩响应缺少用户或成绩数据"); @@ -599,18 +488,18 @@ public static String getRandomScoreWeight(Long userId, JsonObject weights, boole body.add("weight_factor", weights); - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/scores/random/users/" + userId + "/weights?all=" + all)) .POST(HttpRequest.BodyPublishers.ofString(body.toString())) .build(); HttpResponse response = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); - if (codeNotOk(response.statusCode())) { - throw parseHttpError(response.body(), response.statusCode(), "获取成绩权重失败"); + if (ApiUtil.codeNotOk(response.statusCode())) { + throw ApiUtil.parseHttpError(response.body(), response.statusCode(), "获取成绩权重失败"); } RawResponse payload = GSON.fromJson(response.body(), RawResponse.class); - ensureApiSuccess(payload, "获取成绩权重失败"); + ApiUtil.ensureApiSuccess(payload, "获取成绩权重失败"); return payload.getData().getAsString(); } catch (IOException e) { @@ -632,19 +521,19 @@ public static RandomScore getRandomScoreFromUsers(List uids, JsonObject we body.add("uids", uidsArray); body.add("weight_factor", weights); - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/scores/random/users")) .POST(HttpRequest.BodyPublishers.ofString(body.toString())) .build(); HttpResponse response = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); - if (codeNotOk(response.statusCode())) { - throw parseHttpError(response.body(), response.statusCode(), "获取随机成绩失败"); + if (ApiUtil.codeNotOk(response.statusCode())) { + throw ApiUtil.parseHttpError(response.body(), response.statusCode(), "获取随机成绩失败"); } RawResponse payload = GSON.fromJson(response.body(), RawResponse.class); - ensureApiSuccess(payload, "获取随机成绩失败"); - JsonObject data = requireDataObject(payload, "随机成绩响应缺少data"); + ApiUtil.ensureApiSuccess(payload, "获取随机成绩失败"); + JsonObject data = ApiUtil.requireDataObject(payload, "随机成绩响应缺少data"); if (!data.has("user") || !data.get("user").isJsonObject() || !data.has("score") || !data.get("score").isJsonObject()) { throw new RuntimeException("随机成绩响应缺少用户或成绩数据"); @@ -664,8 +553,7 @@ public static RandomScore getRandomScoreFromUsers(List uids, JsonObject we } } - - public static ReplayTaskInfo createBeatmapPreviewTask(long beatmapId, String mods, + public static ReplayTaskInfo createBeatmapPreviewTask(long beatmapId, String mods, TimeDurationParser.TimeRange range, QqUploadRequest qqUpload) { JsonObject body = new JsonObject(); if (mods != null && !mods.isBlank()) { @@ -675,8 +563,14 @@ public static ReplayTaskInfo createBeatmapPreviewTask(long beatmapId, String mod body.add("qqUpload", GSON.toJsonTree(qqUpload)); } - HttpRequest request = HttpRequest.newBuilder() - .uri(URI.create(ENDPOINT + "/replays/renders/preview/" + beatmapId)) + String rangeQuery = "?"; + + if (range != null) { + rangeQuery += range.toQueryString(); + } + + HttpRequest request = requestBuilder() + .uri(URI.create(ENDPOINT + "/replays/renders/preview/" + beatmapId + rangeQuery)) .header("Content-Type", "application/json") .POST(HttpRequest.BodyPublishers.ofString(body.toString())) .build(); @@ -695,7 +589,7 @@ public static ReplayTaskInfo createReplayShowcaseTask(long beatmapId, String[] s if (qqUpload != null) { body.add("qqUpload", GSON.toJsonTree(qqUpload)); } - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/replays/renders/showcase/" + beatmapId)) .header("Content-Type", "application/json") .POST(HttpRequest.BodyPublishers.ofString(body.toString())) @@ -717,7 +611,7 @@ public static ReplayTaskInfo createReplayRenderTask(String scoreId, timeRange = getScoreHighlight(scoreId, 5); } - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/replays/renders/score/" + scoreId + "?" + timeRange.toQueryString())) .header("Content-Type", "application/json") .POST(renderRequestBody(qqUpload)) @@ -733,7 +627,7 @@ private static TimeDurationParser.TimeRange getScoreHighlight(long scoreId, int private static TimeDurationParser.TimeRange getScoreHighlight(String scoreId, int extend) { try { - HttpRequest localRequest = HttpRequest.newBuilder() + HttpRequest localRequest = requestBuilder() .uri(URI.create(ENDPOINT + "/scores/" + scoreId + "/highlight")) .GET() .build(); @@ -741,11 +635,11 @@ private static TimeDurationParser.TimeRange getScoreHighlight(String scoreId, in final var send = CLIENT.send(localRequest, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "高光获取失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "高光获取失败"); } final RawResponse rawResponse = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(rawResponse, "高光获取失败"); + ApiUtil.ensureApiSuccess(rawResponse, "高光获取失败"); final JsonObject data = rawResponse.getData().getAsJsonObject(); return new TimeDurationParser.TimeRange( @@ -757,38 +651,65 @@ private static TimeDurationParser.TimeRange getScoreHighlight(String scoreId, in } } - public static String lookupScoreId(ShortcutTarget target, List filters, String mod) { - String scoreId; - if (target.isLocalScore()) { - scoreId = target.localScoreId(); - } else if (!target.isMacro() || "s".equals(target.macroType())) { - scoreId = String.valueOf(target.explicitId()); - } else { - try { - final String query = getScoreQuery(target) + encodeScoreFilters(filters) - + (mod == null ? "" : "&mod=" + URLEncoder.encode(mod, StandardCharsets.UTF_8)); + /** + * 每个查找方法只请求一个接口;目标类型转换和记忆由指令处理方法决定。 + */ + public static long lookupBeatmapInSet(long setId, long index, String auth) { + return lookupTargetData("/beatmaps/lookup?ms=" + setId + "&i=" + index, auth, "查找谱面失败") + .get("beatmap_id").getAsLong(); + } - HttpRequest localRequest = HttpRequest.newBuilder() - .uri(URI.create(ENDPOINT + query)) - .GET() - .build(); + public static long lookupMultiplayerBeatmap(String auth) { + return lookupTargetData("/beatmaps/lookup?of=mp", auth, "查找谱面失败").get("beatmap_id").getAsLong(); + } - final HttpResponse send = CLIENT.send(localRequest, HttpResponse.BodyHandlers.ofString()); + public static long lookupPlayerScoreBeatmap(long uid, String list, long index, String auth) { + if (!List.of("rs", "rp", "bp").contains(list)) throw new IllegalArgumentException("Invalid score list"); + return lookupTargetData("/beatmaps/lookup?of=" + list + "&u=" + uid + "&i=" + index, + auth, "查找谱面失败").get("beatmap_id").getAsLong(); + } - if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "获取成绩失败"); - } + public static long lookupMultiplayerBeatmapset(String auth) { + return lookupTargetData("/beatmapsets/lookup?of=mp", auth, "查找谱面集失败").get("beatmapset_id").getAsLong(); + } - final RawResponse rawResponse = GSON.fromJson(send.body(), RawResponse.class); + public static long lookupBeatmapsetForBeatmap(long beatmapId, String auth) { + return lookupTargetData("/beatmapsets/lookup?m=" + beatmapId, auth, "查找谱面集失败") + .get("beatmapset_id").getAsLong(); + } - ensureApiSuccess(rawResponse, "获取成绩失败"); + public static String lookupPlayerScore(long uid, String list, long index, List filters, String mod) { + if (!List.of("rs", "rp", "bp").contains(list)) throw new IllegalArgumentException("Invalid score list"); + return lookupScore("/scores/lookup?of=" + list + "&i=" + index + "&u=" + uid, filters, mod); + } - scoreId = rawResponse.getData().getAsJsonObject().get("score_id").getAsString(); - } catch (IOException | InterruptedException e) { - throw requestFailure(e); - } + public static String lookupBeatmapScore(long beatmapId, long uid, List filters, String mod) { + return lookupScore("/scores/lookup?m=" + beatmapId + "&u=" + uid, filters, mod); + } + + public static String lookupBeatmapsetScore(long setId, long index, long uid, List filters, String mod) { + return lookupScore("/scores/lookup?ms=" + setId + "&i=" + index + "&u=" + uid, filters, mod); + } + + private static String lookupScore(String query, List filters, String mod) { + return lookupTargetData(query + encodeScoreFilters(filters) + + (mod == null ? "" : "&mod=" + URLEncoder.encode(mod, StandardCharsets.UTF_8)), + null, "获取成绩失败").get("score_id").getAsString(); + } + + private static JsonObject lookupTargetData(String query, String auth, String error) { + try { + var request = requestBuilder().uri(URI.create(ENDPOINT + query)).GET(); + if (auth != null) withOsuAuthorization(request, auth); + var response = CLIENT.send(request.build(), HttpResponse.BodyHandlers.ofString()); + if (response.statusCode() != 200) + throw ApiUtil.parseHttpError(response.body(), response.statusCode(), error); + RawResponse payload = GSON.fromJson(response.body(), RawResponse.class); + ApiUtil.ensureApiSuccess(payload, error); + return ApiUtil.requireDataObject(payload, error); + } catch (IOException | InterruptedException e) { + throw requestFailure(e); } - return scoreId; } public static long getScoreBeatmapId(String scoreId) { @@ -797,17 +718,17 @@ public static long getScoreBeatmapId(String scoreId) { private static JsonObject lookupScoreData(String scoreId) { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/scores/lookup?s=" + scoreId)) .GET() .build(); HttpResponse response = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); if (response.statusCode() != 200) { - throw parseHttpError(response.body(), response.statusCode(), "获取本地成绩信息失败"); + throw ApiUtil.parseHttpError(response.body(), response.statusCode(), "获取本地成绩信息失败"); } RawResponse payload = GSON.fromJson(response.body(), RawResponse.class); - ensureApiSuccess(payload, "获取本地成绩信息失败"); - return requireDataObject(payload, "本地成绩响应缺少data"); + ApiUtil.ensureApiSuccess(payload, "获取本地成绩信息失败"); + return ApiUtil.requireDataObject(payload, "本地成绩响应缺少data"); } catch (IOException | InterruptedException e) { throw requestFailure(e); } @@ -818,13 +739,13 @@ private static ReplayTaskInfo getReplayTaskInfo(HttpRequest request) { try { HttpResponse response = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); RawResponse payload = GSON.fromJson(response.body(), RawResponse.class); - if (codeNotOk(response.statusCode())) { - throw parseHttpError(response.body(), response.statusCode(), "回放渲染请求失败"); + if (ApiUtil.codeNotOk(response.statusCode())) { + throw ApiUtil.parseHttpError(response.body(), response.statusCode(), "回放渲染请求失败"); } - ensureApiSuccess(payload, "回放渲染请求失败"); + ApiUtil.ensureApiSuccess(payload, "回放渲染请求失败"); - JsonObject data = requireDataObject(payload, "回放渲染请求缺少任务信息"); + JsonObject data = ApiUtil.requireDataObject(payload, "回放渲染请求缺少任务信息"); if (!data.has("id") || data.get("id").isJsonNull()) { throw new RuntimeException("回放渲染请求缺少任务ID"); @@ -891,17 +812,17 @@ private static FileInfo waitReplayDone(String taskId, long timeout) { private static JsonObject getReplayStatus(String taskId) { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/replays/" + taskId + "/status")) .GET() .build(); HttpResponse response = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); RawResponse payload = GSON.fromJson(response.body(), RawResponse.class); - if (codeNotOk(response.statusCode())) { - throw parseHttpError(response.body(), response.statusCode(), "查询回放渲染状态失败"); + if (ApiUtil.codeNotOk(response.statusCode())) { + throw ApiUtil.parseHttpError(response.body(), response.statusCode(), "查询回放渲染状态失败"); } - ensureApiSuccess(payload, "查询回放渲染状态失败"); - JsonObject data = requireDataObject(payload, "回放渲染状态响应缺少data"); + ApiUtil.ensureApiSuccess(payload, "查询回放渲染状态失败"); + JsonObject data = ApiUtil.requireDataObject(payload, "回放渲染状态响应缺少data"); if (!data.has("status") || data.get("status").isJsonNull()) { throw new RuntimeException("回放渲染状态响应缺少status"); } @@ -923,79 +844,9 @@ private static HttpRequest.BodyPublisher renderRequestBody(QqUploadRequest qqUpl return HttpRequest.BodyPublishers.ofString(body.toString(), StandardCharsets.UTF_8); } - private static void ensureApiSuccess(RawResponse payload, String fallbackMessage) { - if (payload == null) { - throw new RuntimeException(fallbackMessage); - } - if (!payload.isSuccess()) { - Integer errorCode = extractErrorCode(payload); - String message = payload.getMessage() != null ? payload.getMessage() : fallbackMessage; - throw new ApiRequestException(errorCode, message); - } - } - - private static RuntimeException parseHttpError(String responseBody, int statusCode, String fallbackMessage) { - Integer errorCode = null; - String message = fallbackMessage; - try { - JsonObject root = GSON.fromJson(responseBody, JsonObject.class); - if (root != null) { - if (root.has("data") && root.get("data").isJsonObject()) { - JsonObject data = root.getAsJsonObject("data"); - errorCode = readCodeFromJsonObject(data); - } - if (errorCode == null) { - errorCode = readCodeFromJsonObject(root); - } - } - } catch (Exception ignored) { - } - - if (statusCode == 500) { - message += "(" + (errorCode == null ? "未知错误码" : errorCode) + " / HTTP " + statusCode + " / 发生了一个内部错误)"; - } - - return new ApiRequestException(errorCode, message); - } - - private static RuntimeException parseHttpError(byte[] responseBody, int statusCode, String fallbackMessage) { - String bodyAsText = responseBody == null ? null : new String(responseBody, StandardCharsets.UTF_8); - return parseHttpError(bodyAsText, statusCode, fallbackMessage); - } - - private static Integer extractErrorCode(RawResponse payload) { - if (payload.getData() != null && payload.getData().isJsonObject()) { - JsonObject data = payload.getData().getAsJsonObject(); - return readCodeFromJsonObject(data); - } - return null; - } - - private static boolean codeNotOk(int statusCode) { - return statusCode < 200 || statusCode >= 300; - } - - private static Integer readCodeFromJsonObject(JsonObject object) { - if (object == null || !object.has("code") || !object.get("code").isJsonPrimitive()) { - return null; - } - try { - return object.get("code").getAsInt(); - } catch (Exception ignored) { - return null; - } - } - - private static JsonObject requireDataObject(RawResponse payload, String message) { - if (payload.getData() == null || !payload.getData().isJsonObject()) { - throw new RuntimeException(message); - } - return payload.getData().getAsJsonObject(); - } - public static RenderStat getRenderStat(String jobId) { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/replays/" + jobId + "/status")) .GET() .build(); @@ -1003,11 +854,11 @@ public static RenderStat getRenderStat(String jobId) { final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "获取渲染进度失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "获取渲染进度失败"); } final RawResponse r = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(r, "获取渲染进度失败"); + ApiUtil.ensureApiSuccess(r, "获取渲染进度失败"); final JsonObject data = r.getData().getAsJsonObject(); return GSON.fromJson(data, RenderStat.class); @@ -1018,18 +869,18 @@ public static RenderStat getRenderStat(String jobId) { public static RenderStat cancelReplayRender(String jobId) { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/replays/" + jobId + "/cancel")) .POST(HttpRequest.BodyPublishers.noBody()) .build(); HttpResponse response = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); - if (codeNotOk(response.statusCode())) { - throw parseHttpError(response.body(), response.statusCode(), "取消回放渲染失败"); + if (ApiUtil.codeNotOk(response.statusCode())) { + throw ApiUtil.parseHttpError(response.body(), response.statusCode(), "取消回放渲染失败"); } RawResponse payload = GSON.fromJson(response.body(), RawResponse.class); - ensureApiSuccess(payload, "取消回放渲染失败"); - return GSON.fromJson(requireDataObject(payload, "取消回放渲染响应缺少data"), RenderStat.class); + ApiUtil.ensureApiSuccess(payload, "取消回放渲染失败"); + return GSON.fromJson(ApiUtil.requireDataObject(payload, "取消回放渲染响应缺少data"), RenderStat.class); } catch (IOException e) { throw requestFailure(e); } catch (InterruptedException e) { @@ -1042,8 +893,10 @@ public static ServerStatus getServerStatus() { boolean oStella = false; boolean osu = false; String oStellaVersion = null; + int allWorkers = 0; + int onlineWorkers = 0; try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/health")) .GET() .build(); @@ -1064,17 +917,23 @@ public static ServerStatus getServerStatus() { if (data.has("osu_api") && !data.get("osu_api").isJsonNull()) { osu = data.get("osu_api").getAsBoolean(); } + if (data.has("all_render_workers") && !data.get("all_render_workers").isJsonNull()) { + allWorkers = data.get("all_render_workers").getAsInt(); + } + if (data.has("online_render_workers") && !data.get("online_render_workers").isJsonNull()) { + onlineWorkers = data.get("online_render_workers").getAsInt(); + } } } } catch (Exception _) { } - return new ServerStatus(true, oStella, oStellaVersion, osu); + return new ServerStatus(true, oStella, oStellaVersion, allWorkers, onlineWorkers, osu); } public static Response> getScoreMissesResponse(String scoreId) { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/scores/" + scoreId + "/misses")) .GET() .build(); @@ -1082,12 +941,12 @@ public static Response> getScoreMissesResponse(String scoreId) { final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "获取 Miss 数据失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "获取 Miss 数据失败"); } final RawResponse r = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(r, "获取 Miss 数据失败"); - final JsonArray data = r.getData().getAsJsonArray(); + ApiUtil.ensureApiSuccess(r, "获取 Miss 数据失败"); + final JsonArray data = ApiUtil.requireResultArray(r, "获取 Miss 数据响应缺少结果数组"); List misses = new LinkedList<>(); for (JsonElement datum : data) { @@ -1107,20 +966,19 @@ public static Response> getScoreMissesResponse(String scoreId) { } } - public static long resolveUid(UserRef userRef) { - if (userRef instanceof UserRef.ByUid byUid) { - return byUid.getUid(); - } - if (userRef instanceof UserRef.ByUsername byUsername) { - return lookupUser(byUsername.getUsername()).getContent().getId(); + public static long resolveUid(String player) { + if (player == null || player.isBlank()) throw new ResolutionException("无法识别指定的玩家"); + try { + long uid = Long.parseLong(player); + if (uid > 0) return uid; + } catch (NumberFormatException ignored) { } - throw new ResolutionException("无法识别指定的玩家"); + return lookupUser(player).getContent().getId(); } - public static long getUserRank(UserRef userRef) { - long uid = resolveUid(userRef); + public static long getUserRank(long uid) { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/users/" + uid + "/rank")) .header("Content-Type", "application/json") .GET() @@ -1129,12 +987,12 @@ public static long getUserRank(UserRef userRef) { final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "获取玩家Rank失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "获取玩家Rank失败"); } final RawResponse response = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(response, "获取玩家Rank失败"); - final JsonObject data = requireDataObject(response, "获取玩家Rank响应缺少用户数据"); + ApiUtil.ensureApiSuccess(response, "获取玩家Rank失败"); + final JsonObject data = ApiUtil.requireDataObject(response, "获取玩家Rank响应缺少用户数据"); return data.get("global_rank").getAsLong(); } catch (IOException e) { throw requestFailure(e); @@ -1146,7 +1004,7 @@ public static long getUserRank(UserRef userRef) { public static Response lookupUser(String username) { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/users/lookup")) .header("Content-Type", "application/json") .POST(HttpRequest.BodyPublishers.ofString( @@ -1157,12 +1015,12 @@ public static Response lookupUser(String username) { final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "查找玩家失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "查找玩家失败"); } final RawResponse response = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(response, "查找玩家失败"); - final JsonObject data = requireDataObject(response, "查找玩家响应缺少用户数据"); + ApiUtil.ensureApiSuccess(response, "查找玩家失败"); + final JsonObject data = ApiUtil.requireDataObject(response, "查找玩家响应缺少用户数据"); return Response.fromHeaders(send.headers()) .content(GSON.fromJson(data, User.class)) @@ -1177,7 +1035,7 @@ public static Response lookupUser(String username) { public static List getUsers(List u) { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/users")) .POST(HttpRequest.BodyPublishers.ofString(GSON.toJsonTree(Map.of("ids", u)).toString())) .build(); @@ -1185,12 +1043,12 @@ public static List getUsers(List u) { final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200) { - throw parseHttpError(send.body(), send.statusCode(), "获取用户信息失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "获取用户信息失败"); } final RawResponse r = GSON.fromJson(send.body(), RawResponse.class); - ensureApiSuccess(r, "获取用户信息失败"); - final JsonArray data = r.getData().getAsJsonArray(); + ApiUtil.ensureApiSuccess(r, "获取用户信息失败"); + final JsonArray data = ApiUtil.requireResultArray(r, "获取用户信息响应缺少结果数组"); List users = new LinkedList<>(); for (JsonElement datum : data) { @@ -1205,7 +1063,7 @@ public static List getUsers(List u) { public static ReplayUploadInfo uploadReplay(byte[] replayBytes) { try { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest request = requestBuilder() .uri(URI.create(ENDPOINT + "/replays/upload")) .POST(HttpRequest.BodyPublishers.ofByteArray(replayBytes)) .build(); @@ -1213,19 +1071,19 @@ public static ReplayUploadInfo uploadReplay(byte[] replayBytes) { final HttpResponse send = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); if (send.statusCode() != 200 && send.statusCode() != 404) { - throw parseHttpError(send.body(), send.statusCode(), "回放上传失败"); + throw ApiUtil.parseHttpError(send.body(), send.statusCode(), "回放上传失败"); } final RawResponse r = GSON.fromJson(send.body(), RawResponse.class); if (send.statusCode() == 404) { - final Integer errCode = extractErrorCode(r); + final Integer errCode = ApiUtil.extractErrorCode(r); if (errCode != null && errCode == ErrorCode.NO_SCORE_FOUND.getCode()) { throw new ApiRequestException(ErrorCode.NO_SCORE_FOUND.getCode(), "回放上传失败:无法获取对应的成绩"); } } - ensureApiSuccess(r, "回放上传失败"); + ApiUtil.ensureApiSuccess(r, "回放上传失败"); return GSON.fromJson(r.getData().getAsJsonObject(), ReplayUploadInfo.class); } catch (IOException | InterruptedException e) { @@ -1240,7 +1098,8 @@ private static RuntimeException requestFailure(Exception exception) { return new RuntimeException(exception); } - public record ServerStatus(boolean gateway, boolean oStella, String oStellaVersion, boolean osu) { + public record ServerStatus(boolean gateway, boolean oStella, String oStellaVersion, + int allWorkers, int onlineWorkers, boolean osu) { } public record ReplayRenderResult(String videoUrl, String taskId, FileInfo qqFile) { diff --git a/src/main/java/xyz/zcraft/seira/api/RomAIApi.java b/src/main/java/xyz/zcraft/seira/api/RomAIApi.java new file mode 100644 index 00000000..e83096fb --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/api/RomAIApi.java @@ -0,0 +1,70 @@ +package xyz.zcraft.seira.api; + +import com.google.gson.Gson; +import com.google.gson.JsonArray; +import com.google.gson.JsonElement; +import com.google.gson.JsonParser; +import org.apache.logging.log4j.LogManager; +import org.apache.logging.log4j.Logger; +import xyz.zcraft.seira.Seira; +import xyz.zcraft.seira.api.data.RomAIMatch; + +import java.net.InetSocketAddress; +import java.net.ProxySelector; +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.time.Duration; +import java.util.LinkedList; +import java.util.List; + +public class RomAIApi { + private static final Gson GSON = new Gson(); + private static final Logger LOG = LogManager.getLogger(RomAIApi.class); + + private static final HttpClient CLIENT; + + static { + var proxy = Seira.getConfig().discord().proxy(); + CLIENT = HttpClient.newBuilder() + .proxy(ProxySelector.of( + new InetSocketAddress(proxy.host(), proxy.port()) + )) + .connectTimeout(Duration.ofSeconds(20)) + .build(); + } + + public static List getActiveMatches() { + try { + final HttpRequest request = HttpRequest.newBuilder() + .uri(URI.create("https://rom-ai-site.vercel.app/api/active-matches")) + .GET() + .build(); + + final HttpResponse response = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); + + final JsonArray arr = JsonParser.parseString(response.body()).getAsJsonArray(); + + List result = new LinkedList<>(); + + for (JsonElement jsonElement : arr) { + result.add(GSON.fromJson(jsonElement, RomAIMatch.class)); + } + + return result; + } catch (Exception e) { + LOG.error("Error occurred while fetching active matches", e); + throw new RuntimeException("Error occurred while fetching active matches", e); + } + } + + public static RomAIMatch getMatchFor(String username) { + final List activeMatches = getActiveMatches(); + + return activeMatches.stream() + .filter(m -> m.players().contains(username)) + .findFirst() + .orElse(null); + } +} diff --git a/src/main/java/xyz/zcraft/seira/api/data/MinecraftServerStatus.java b/src/main/java/xyz/zcraft/seira/api/data/MinecraftServerStatus.java new file mode 100644 index 00000000..39fb8582 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/api/data/MinecraftServerStatus.java @@ -0,0 +1,38 @@ +package xyz.zcraft.seira.api.data; + +import com.google.gson.JsonElement; + +import java.util.List; + +public record MinecraftServerStatus( + String host, + Integer port, + Long latency, + Status status +) { + public record Status( + Version version, + Players players, + String description, + JsonElement descriptionRaw, + String favicon + ){} + + public record Version( + String name, + Integer protocol + ) { + } + + public record Players( + Long max, + Long online, + List samples + ) { + public record Sample( + String id, + String name + ) { + } + } +} diff --git a/src/main/java/xyz/zcraft/seira/api/data/RomAIMatch.java b/src/main/java/xyz/zcraft/seira/api/data/RomAIMatch.java new file mode 100644 index 00000000..a4c1b8fe --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/api/data/RomAIMatch.java @@ -0,0 +1,94 @@ +package xyz.zcraft.seira.api.data; + +import java.util.List; +import java.util.Map; + +/* +{ + "_id": "example", + "osuUser1": null, + "osuUser2": null, + "players": ["example","example","example","example","example","example"], + "teams": { + "teamA": ["example","example","example"], + "teamB": ["example","example","example"] + }, + "tournament": null, + "selectedPool": null, + "customBO": 9, + "customELO": 1335, + "interactionIds": ["1551915321242239126"], + "discordChannelId": "1295392984787811380", + "isMatchmaking": false, + "lobbyId": "121100588", + "matchState": "Picking", + "score": [2,3], + "bans": ["dt2","hd2","dt1","nm6"], + "picks": [ + { + "map": 4722846, + "mod": "nm2", + "scores": [733676,342836] + } + ], + "poolOptions": [ + "nm4","nm5","hd1","hd3","hr1","hr2","hr4","dt3","dt4" + ], + "availablePools": [ + { + "maps": { + "noMod": [4168602,3596841,5019445,5375902,4706834], + "hidden": [4836596,3505413,5074450], + "hardRock": [2676188,4641931,5445369], + "doubleTime": [5014531,2246430,3632103], + "freeMod": [], + "tieBreaker": 5384525 + }, + "_id": "696d20ff8f8d4a2449abe8d1", + "name": "example", + "elo": 1334, + "__v": 0 + } + ], + "poolInfoName": "Example Tournament", + "firstPick": "example", + "secondPick": "example", + "currentMapId": 5341088, + "currentMapMod": "hr", + "timerEndsAt": "2026-09-22T15:09:19.111Z", + "startedAt": "2026-09-22T14:36:47.445Z", + "__v": 40, + "playerData": [ + { + "osuUserName": "example", + "osuUserId": 12345678, + "elo": {"3v3": 1300,"1v1": 1300,"2v2": 1300}, + "country": "CN" + } + ], + "mode": "3v3", + "avgElo": 1251, + "isLeagueMatch": false + } + */ +public record RomAIMatch( + List players, + Teams teams, + String lobbyId, + List score, + Long currentMapId, + String mode, + List playerData, + Integer customBO, + Integer customELO +) { + + public record Teams(List teamA, List teamB){} + + public record PlayerData( + String osuUserName, + Long osuUserId, + Map elo, + String country + ){} +} diff --git a/src/main/java/xyz/zcraft/seira/api/data/VideoRenderRecord.java b/src/main/java/xyz/zcraft/seira/api/data/VideoRenderRecord.java index 59c51ad4..0d3bb7bd 100644 --- a/src/main/java/xyz/zcraft/seira/api/data/VideoRenderRecord.java +++ b/src/main/java/xyz/zcraft/seira/api/data/VideoRenderRecord.java @@ -4,9 +4,11 @@ public class VideoRenderRecord { private final ConcurrentHashMap renderRecord = new ConcurrentHashMap<>(); + private final ConcurrentHashMap taskOwner = new ConcurrentHashMap<>(); public void updateRenderTask(String uid, String jobId) { renderRecord.put(uid, jobId); + taskOwner.put(jobId, uid); } public boolean hasRenderTask(String uid) { @@ -17,6 +19,10 @@ public String getRenderTask(String uid) { return renderRecord.get(uid); } + public String getTaskOwner(String jobId) { + return taskOwner.get(jobId); + } + @SuppressWarnings("unused") public void removeRenderTask(String uid) { renderRecord.remove(uid); @@ -24,5 +30,6 @@ public void removeRenderTask(String uid) { public void removeRenderTask(String uid, String jobId) { renderRecord.remove(uid, jobId); + taskOwner.remove(jobId, uid); } } diff --git a/src/main/java/xyz/zcraft/seira/bot/MessageSender.java b/src/main/java/xyz/zcraft/seira/bot/MessageSender.java index 11f491a3..3579d6cd 100644 --- a/src/main/java/xyz/zcraft/seira/bot/MessageSender.java +++ b/src/main/java/xyz/zcraft/seira/bot/MessageSender.java @@ -1,6 +1,5 @@ package xyz.zcraft.seira.bot; -import com.google.gson.Gson; import org.apache.logging.log4j.LogManager; import org.apache.logging.log4j.Logger; import xyz.zcraft.seira.api.data.QqUploadRequest; @@ -12,7 +11,6 @@ import xyz.zcraft.seira.services.CosService; import xyz.zcraft.seira.util.TokenManager; -import java.util.Map; import java.util.function.Supplier; public class MessageSender { @@ -149,10 +147,17 @@ public SentMessage sendGroupText(String groupId, String content) { public SentMessage sendGroupMarkdown(String groupId, String content) { Message message = new Message(); message.setMsgType(PendingMessage.MSG_TYPE_MARKDOWN); - message.setMarkdown(new Gson().toJsonTree(Map.of("content", content)).getAsJsonObject()); + message.setMarkdown(Message.MessageMarkdown.of(content)); return sendGroupMessage(groupId, message); } + public SentMessage sendPrivateMarkdown(String userId, String content) { + Message message = new Message(); + message.setMsgType(PendingMessage.MSG_TYPE_MARKDOWN); + message.setMarkdown(Message.MessageMarkdown.of(content)); + return sendPrivateMessage(userId, message); + } + private FileInfo retryUpload(Supplier operation, long baseDelayMillis, String description) { for (int attempt = 1; attempt <= MAX_UPLOAD_ATTEMPTS; attempt++) { try { diff --git a/src/main/java/xyz/zcraft/seira/bot/QQApi.java b/src/main/java/xyz/zcraft/seira/bot/QQApi.java index 9dde38f6..ab3df75f 100644 --- a/src/main/java/xyz/zcraft/seira/bot/QQApi.java +++ b/src/main/java/xyz/zcraft/seira/bot/QQApi.java @@ -644,6 +644,50 @@ public static int editPanel(AccessToken accessToken, String panelId, Panel newPa } } + public static String getAvatarUrl(String appId, String openId) { + return "https://thirdqq.qlogo.cn/qqapp/" + appId + "/" + openId + "/100"; + } + + public static GroupInfo getGroupInfo(AccessToken accessToken, String groupId) { + try { + final var request = newRequestBuilder(accessToken) + .uri(URI.create(ENDPOINT + "/v2/groups/" + groupId + "/info")) + .GET() + .build(); + + final HttpResponse response = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); + + if (response.statusCode() != 200) { + LOG.error("Failed to get group info, status code: {} body={}", response.statusCode(), response.body()); + throw new RuntimeException("Failed to get group info, status code: " + response.statusCode() + " body=" + response.body()); + } + + return GSON.fromJson(response.body(), GroupInfo.class); + } catch (IOException | InterruptedException e) { + throw requestFailure(e); + } + } + + public static GroupBotState getGroupBotState(AccessToken accessToken, String groupId) { + try { + final var request = newRequestBuilder(accessToken) + .uri(URI.create(ENDPOINT + "/v2/groups/" + groupId + "/bot_state")) + .GET() + .build(); + + final HttpResponse response = CLIENT.send(request, HttpResponse.BodyHandlers.ofString()); + + if (response.statusCode() != 200) { + LOG.error("Failed to get group bot state, status code: {} body={}", response.statusCode(), response.body()); + throw new RuntimeException("Failed to get group bot state, status code: " + response.statusCode() + " body=" + response.body()); + } + + return GSON.fromJson(response.body(), GroupBotState.class); + } catch (IOException | InterruptedException e) { + throw requestFailure(e); + } + } + private record MediaDigests(String md5, String sha1, String md5First10m) { } diff --git a/src/main/java/xyz/zcraft/seira/bot/QQBot.java b/src/main/java/xyz/zcraft/seira/bot/QQBot.java index 7db86c07..d9b01304 100644 --- a/src/main/java/xyz/zcraft/seira/bot/QQBot.java +++ b/src/main/java/xyz/zcraft/seira/bot/QQBot.java @@ -4,10 +4,9 @@ import lombok.Getter; import org.apache.logging.log4j.LogManager; import org.apache.logging.log4j.Logger; -import xyz.zcraft.seira.bot.data.Panel; -import xyz.zcraft.seira.bot.data.PanelItem; -import xyz.zcraft.seira.bot.data.PanelRecord; -import xyz.zcraft.seira.bot.data.QQUser; +import xyz.zcraft.seira.ai.provider.ChatProvider; +import xyz.zcraft.seira.ai.provider.ChatProviders; +import xyz.zcraft.seira.bot.data.*; import xyz.zcraft.seira.command.AttachmentHandler; import xyz.zcraft.seira.command.route.Router; import xyz.zcraft.seira.config.AppConfig; @@ -37,7 +36,7 @@ public class QQBot implements AutoCloseable, ConsoleRuntimeControl { private static final Logger LOG = LogManager.getLogger(QQBot.class); - + final AtomicReference self = new AtomicReference<>(); @Getter private final TokenManager tokenManager; @Getter @@ -45,8 +44,9 @@ public class QQBot implements AutoCloseable, ConsoleRuntimeControl { @Getter private final MessageSender sender; private final ScoreWatchService watchService; - private final MultiplayerRoomWatchService multiplayerRoomWatchService; + private final MPWatchService mpWatchService; private final RankGuessGameService rankGuessGameService; + private final ChatProvider chatProvider; private final RealtimeServiceInterruptionNotifier interruptionNotifier; private final DiscordBridgeService discordBridgeService; private final AppConfig startupConfig; @@ -69,7 +69,7 @@ public QQBot( AppConfig config = runtimeConfig.current(); this.startupConfig = config; this.executors = executors; - this.cacheControlClient = new OstellaCacheControlClient(config.ostella().endpoint()); + this.cacheControlClient = new OstellaCacheControlClient(config.ostella().endpoint(), config.ostella().token()); LOG.info("Authorizing QQ API"); this.tokenManager = new TokenManager(config.qq().appId(), config.qq().appSecret()); @@ -85,7 +85,7 @@ public QQBot( LOG.info("Initializing score watch service"); this.watchService = new ScoreWatchService( - new OstellaWatchApi(config.ostella().endpoint()), + new ScoreWatchApi(config.ostella().endpoint(), config.ostella().token()), new WatchScoreNotifier(sender), new SpecificScoreNotifier(sender), new SqliteSpecificScoreWatchStore(), @@ -93,14 +93,19 @@ public QQBot( ); LOG.info("Initializing multiplayer room watch service"); - this.multiplayerRoomWatchService = new MultiplayerRoomWatchService( - new OstellaMultiplayerRoomWatchApi(config.ostella().endpoint()), - new QqMultiplayerRoomNotifier(sender), + this.mpWatchService = new MPWatchService( + new MPWatchApi(config.ostella().endpoint(), config.ostella().token()), + new MPNotifier(sender), Duration.ofSeconds(config.seira().effectiveMultiplayerWatchIntervalSeconds()) ); + LOG.info("Initializing rank guess service"); this.rankGuessGameService = new RankGuessGameService(); + + LOG.info("Initializing agents service"); + this.chatProvider = ChatProviders.newHiAgentChatProvider(config.llm()); + this.attachmentHandler = new AttachmentHandler(executors.attachmentDownloads()); this.router = new Router( sender, @@ -108,11 +113,23 @@ public QQBot( admins, bindingService, watchService, - multiplayerRoomWatchService, + mpWatchService, discordBridgeService, rankGuessGameService, executors.commandTasks(), - BotStat::incrementCommands + BotStat::incrementCommands, + bytes -> { + try { + return cos.uploadImage(bytes); + } catch (Exception e) { + LOG.error("Error uploading image", e); + return null; + } + }, + self::get, + chatProvider, + s -> QQApi.getGroupBotState(tokenManager.getToken(), s) + ); } @@ -126,7 +143,7 @@ public void start() { runnerThread = Thread.currentThread(); tokenManager.start(); watchService.start(); - multiplayerRoomWatchService.start(); + mpWatchService.start(); discordBridgeService.start(); LOG.info("Starting bot connection loop..."); @@ -141,9 +158,9 @@ public void start() { String wssEndpoint = QQApi.getWSSEndpoint(tokenManager.getToken()); LOG.info("Endpoint: {}", wssEndpoint); - final QQUser self = QQApi.getSelf(tokenManager.getToken()); + self.set(QQApi.getSelf(tokenManager.getToken())); - LOG.info("Self info: id={}, nickname={}", self.id(), self.username()); + LOG.info("Self info: id={}, nickname={}", self.get().id(), self.get().username()); WSClient client = new WSClient( URI.create(wssEndpoint), @@ -223,7 +240,7 @@ public void stop() { thread.interrupt(); } watchService.close(); - multiplayerRoomWatchService.close(); + mpWatchService.close(); discordBridgeService.close(); tokenManager.close(); } @@ -296,7 +313,8 @@ public void requestStop() { RealtimeServiceInterruptionNotifier.NotificationResult result = interruptionNotifier.notifyGroups( watchService.activeTransientGroupIds(), rankGuessGameService.activeGroupIds(), - multiplayerRoomWatchService.activeGroupIds() + mpWatchService.activeGroupIds(), + chatProvider.activeGroupIds() ); if (result.failedGroups() == 0) { LOG.info("Sent restart interruption notices to {} affected groups", result.sentGroups()); @@ -403,4 +421,44 @@ public ConsoleCommandProcessor.ConsoleResult editPanel(String panelId, String js return ConsoleCommandProcessor.ConsoleResult.failure("Error editing panel: " + e.getMessage()); } } + + @Override + public ConsoleCommandProcessor.ConsoleResult getGroupInfo(String groupId) { + try { + final GroupInfo groupInfo = QQApi.getGroupInfo(tokenManager.getToken(), groupId); + + String sb = "=== Group info ===\n" + + "group_id: " + groupId + "\n" + + "group_name: " + groupInfo.groupName() + "\n" + + "group_finger_memo: " + groupInfo.groupFingerMemo() + "\n" + + "group_class_text: " + groupInfo.groupClassText() + "\n" + + "group_tags: " + String.join(", ", groupInfo.groupTags()) + "\n" + + "group_member_num: " + groupInfo.groupMemberNum() + "\n" + + "=================="; + + return ConsoleCommandProcessor.ConsoleResult.success(sb); + } catch (Exception e) { + return ConsoleCommandProcessor.ConsoleResult.failure("Error editing panel: " + e.getMessage()); + } + } + + @Override + public ConsoleCommandProcessor.ConsoleResult getGroupBotState(String groupId) { + try { + final GroupBotState groupBotState = QQApi.getGroupBotState(tokenManager.getToken(), groupId); + + String sb = "=== Group bot state ===\n" + + "group_id: " + groupId + "\n" + + "member_openid: " + groupBotState.memberOpenId() + "\n" + + "joined_at: " + groupBotState.joinedAt() + "\n" + + "allow_proactive_msg: " + groupBotState.allowProactiveMsg() + "\n" + + "recv_msg_setting: " + groupBotState.receiveMsgSetting() + "\n" + + "member_role: " + groupBotState.memberRole() + "\n" + + "======================="; + + return ConsoleCommandProcessor.ConsoleResult.success(sb); + } catch (Exception e) { + return ConsoleCommandProcessor.ConsoleResult.failure("Error getting group bot state: " + e.getMessage()); + } + } } diff --git a/src/main/java/xyz/zcraft/seira/bot/RealtimeServiceInterruptionNotifier.java b/src/main/java/xyz/zcraft/seira/bot/RealtimeServiceInterruptionNotifier.java index 0a75db7b..70362d3b 100644 --- a/src/main/java/xyz/zcraft/seira/bot/RealtimeServiceInterruptionNotifier.java +++ b/src/main/java/xyz/zcraft/seira/bot/RealtimeServiceInterruptionNotifier.java @@ -6,6 +6,7 @@ public final class RealtimeServiceInterruptionNotifier { private static final String SCORE_WATCH = "成绩监视"; private static final String RANK_GUESS = "猜 Rank"; private static final String MULTIPLAYER_WATCH = "MP 监视"; + private static final String AI_CHAT = "AI 对话"; private final MessageSender sender; @@ -14,9 +15,7 @@ public RealtimeServiceInterruptionNotifier(MessageSender sender) { } private static void addService( - Map> servicesByGroup, - Set groupIds, - String service + Map> servicesByGroup, Set groupIds, String service ) { Objects.requireNonNull(groupIds); groupIds.stream().sorted().forEach(groupId -> servicesByGroup @@ -33,12 +32,14 @@ private static String message(Set services) { public NotificationResult notifyGroups( Set scoreWatchGroups, Set rankGuessGroups, - Set multiplayerWatchGroups + Set multiplayerWatchGroups, + Set agentGroups ) { Map> servicesByGroup = new LinkedHashMap<>(); addService(servicesByGroup, scoreWatchGroups, SCORE_WATCH); addService(servicesByGroup, rankGuessGroups, RANK_GUESS); addService(servicesByGroup, multiplayerWatchGroups, MULTIPLAYER_WATCH); + addService(servicesByGroup, agentGroups, AI_CHAT); int sent = 0; for (Map.Entry> entry : servicesByGroup.entrySet()) { diff --git a/src/main/java/xyz/zcraft/seira/bot/WSClient.java b/src/main/java/xyz/zcraft/seira/bot/WSClient.java index 54a4db30..380ca9ec 100644 --- a/src/main/java/xyz/zcraft/seira/bot/WSClient.java +++ b/src/main/java/xyz/zcraft/seira/bot/WSClient.java @@ -9,6 +9,7 @@ import org.java_websocket.handshake.ServerHandshake; import xyz.zcraft.seira.bot.data.AccessToken; import xyz.zcraft.seira.bot.data.Attachment; +import xyz.zcraft.seira.bot.data.MsgElem; import xyz.zcraft.seira.command.AttachmentHandler; import xyz.zcraft.seira.command.route.Router; import xyz.zcraft.seira.config.AppConfig; @@ -147,7 +148,38 @@ private void onC2CMsg(JsonObject payload) { String content = data.get("content").getAsString(); String msgId = data.get("id").getAsString(); String openId = data.get("author").getAsJsonObject().get("user_openid").getAsString(); - router.onPrivateMessageReceived(openId, msgId, content); + + final JsonArray attachments = data.getAsJsonArray("attachments"); + List attachmentList = new ArrayList<>(); + + if (attachments != null && !attachments.isJsonNull()) { + for (JsonElement attachmentElem : attachments) { + JsonObject attachmentObj = attachmentElem.getAsJsonObject(); + Attachment attachment = gson.fromJson(attachmentObj, Attachment.class); + attachmentList.add(attachment); + } + } + + String msgIdx = null; + + final JsonArray extArr = data.get("message_scene").getAsJsonObject().get("ext").getAsJsonArray(); + + for (JsonElement elem : extArr) { + final String str = elem.getAsString(); + + if (str.startsWith("msg_idx=")) { + msgIdx = str.substring("msg_idx=".length()); + } + } + + List msgElemList = new ArrayList<>(); + + if (data.has("msg_elements")) { + data.get("msg_elements").getAsJsonArray() + .forEach(elem -> msgElemList.add(gson.fromJson(elem, MsgElem.class))); + } + + router.onPrivateMessageReceived(openId, msgId, content, msgIdx, attachmentList, msgElemList); } private void onC2CFile(JsonObject payload) { @@ -177,6 +209,26 @@ private void onGroupMsg(JsonObject payload) { JsonObject author = data.get("author").getAsJsonObject(); String openId = author.get("member_openid").getAsString(); String groupId = data.get("group_openid").getAsString(); + + String msgIdx = null; + + final JsonArray extArr = data.get("message_scene").getAsJsonObject().get("ext").getAsJsonArray(); + + for (JsonElement elem : extArr) { + final String str = elem.getAsString(); + + if (str.startsWith("msg_idx=")) { + msgIdx = str.substring("msg_idx=".length()); + } + } + + List msgElemList = new ArrayList<>(); + + if (data.has("msg_elements")) { + data.get("msg_elements").getAsJsonArray() + .forEach(elem -> msgElemList.add(gson.fromJson(elem, MsgElem.class))); + } + List attachments = parseAttachments(data); Map mentions = parseMentions(data); @@ -191,7 +243,8 @@ private void onGroupMsg(JsonObject payload) { mentions )); } - router.onGroupMessageReceived(groupId, openId, msgId, content); + + router.onGroupMessageReceived(groupId, openId, msgId, content, msgIdx, attachments, msgElemList); } private Map parseMentions(JsonObject data) { @@ -242,7 +295,20 @@ private String stripSelfMention(String content) { private void sendIdentify() { JsonObject data = new JsonObject(); data.addProperty("token", "QQBot " + tokenSupplier.get().token()); - data.addProperty("intents", 1 << 25); + /* + GUILDS (1 << 0) + GUILD_MEMBERS (1 << 1) + GUILD_MESSAGES (1 << 9) // 消息事件,仅 *私域* 机器人能够设置此 intents。 + GUILD_MESSAGE_REACTIONS (1 << 10) + DIRECT_MESSAGE (1 << 12) + GROUP_AND_C2C_EVENT (1 << 25) + INTERACTION (1 << 26) + MESSAGE_AUDIT (1 << 27) + FORUMS_EVENT (1 << 28) // 论坛事件,仅 *私域* 机器人能够设置此 intents。 + AUDIO_ACTION (1 << 29) + PUBLIC_GUILD_MESSAGES (1 << 30) // 消息事件,此为公域的消息事件 + */ + data.addProperty("intents", 1 << 25 | 1 << 26 | 1 << 28 | 1 << 30); JsonObject payload = new JsonObject(); payload.addProperty("op", 2); diff --git a/src/main/java/xyz/zcraft/seira/bot/data/Button.java b/src/main/java/xyz/zcraft/seira/bot/data/Button.java index a0f75ea3..e9545b0d 100644 --- a/src/main/java/xyz/zcraft/seira/bot/data/Button.java +++ b/src/main/java/xyz/zcraft/seira/bot/data/Button.java @@ -97,6 +97,16 @@ public Button permit(String userId) { return this; } + public Button modal(String content) { + if (content != null && !content.isBlank()) { + if (this.getAction() != null) { + this.getAction().setModal(Action.Modal.of(content)); + } + } + + return this; + } + public Button disable() { this.renderData.setStyle(0); @@ -114,7 +124,7 @@ public Button disable() { @Data @NoArgsConstructor @AllArgsConstructor - private static class RenderData { + public static class RenderData { private String label; @SerializedName("visited_label") private String visitedLabel; @@ -124,20 +134,31 @@ private static class RenderData { @Data @NoArgsConstructor @AllArgsConstructor - private static class Action { + public static class Action { private int type; private Permission permission; private String data; private boolean enter; private int anchor; @SerializedName("unsupport_tips") - private String unsupportTips; + private String unsupportedTips; + private Modal modal; @Data - private static class Permission { + public static class Permission { private int type; @SerializedName("specify_user_ids") private List specifyUserIds; } + + public record Modal( + String content, + @SerializedName("confirm_text") String confirmText, + @SerializedName("cancel_text") String cancelText + ) { + public static Modal of(String content) { + return new Modal(content, null, null); + } + } } } diff --git a/src/main/java/xyz/zcraft/seira/bot/data/GroupBotState.java b/src/main/java/xyz/zcraft/seira/bot/data/GroupBotState.java new file mode 100644 index 00000000..862a2ad3 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/bot/data/GroupBotState.java @@ -0,0 +1,23 @@ +package xyz.zcraft.seira.bot.data; + +import com.google.gson.annotations.SerializedName; + +public record GroupBotState( + @SerializedName("member_openid") String memberOpenId, + @SerializedName("joined_at") String joinedAt, + @SerializedName("allow_proactive_msg") Boolean allowProactiveMsg, + @SerializedName("recv_msg_setting") ReceiveMsgSetting receiveMsgSetting, + @SerializedName("member_role") MemberRole memberRole +) { + public enum ReceiveMsgSetting { + @SerializedName("all") ALL, + @SerializedName("only_mention") ONLY_MENTION, + @SerializedName("mention_and_context") MENTION_AND_CONTEXT, + } + + public enum MemberRole { + @SerializedName("member") MEMBER, + @SerializedName("owner") OWNER, + @SerializedName("admin") ADMIN + } +} diff --git a/src/main/java/xyz/zcraft/seira/bot/data/GroupInfo.java b/src/main/java/xyz/zcraft/seira/bot/data/GroupInfo.java new file mode 100644 index 00000000..b1db711f --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/bot/data/GroupInfo.java @@ -0,0 +1,15 @@ +package xyz.zcraft.seira.bot.data; + +import com.google.gson.annotations.SerializedName; + +import java.util.List; + +public record GroupInfo( + @SerializedName("group_openid") String groupOpenId, + @SerializedName("group_name") String groupName, + @SerializedName("group_finger_memo") String groupFingerMemo, + @SerializedName("group_class_text") String groupClassText, + @SerializedName("group_tags") List groupTags, + @SerializedName("group_member_num") Integer groupMemberNum +) { +} diff --git a/src/main/java/xyz/zcraft/seira/bot/data/Message.java b/src/main/java/xyz/zcraft/seira/bot/data/Message.java index 53c2c0dd..c25b59cd 100644 --- a/src/main/java/xyz/zcraft/seira/bot/data/Message.java +++ b/src/main/java/xyz/zcraft/seira/bot/data/Message.java @@ -17,7 +17,7 @@ public class Message { @SerializedName("msg_type") private int msgType; - private JsonObject markdown; + private MessageMarkdown markdown; private JsonObject keyboard; @@ -39,4 +39,13 @@ public class Message { @SerializedName("message_reference") private MessageReference messageReference; + + public record MessageMarkdown( + String content, + @SerializedName("force_verify_image_resource") Boolean forceVerifyImageResource + ) { + public static MessageMarkdown of(String content) { + return new MessageMarkdown(content, null); + } + } } diff --git a/src/main/java/xyz/zcraft/seira/bot/data/MsgElem.java b/src/main/java/xyz/zcraft/seira/bot/data/MsgElem.java new file mode 100644 index 00000000..fa6f884e --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/bot/data/MsgElem.java @@ -0,0 +1,14 @@ +package xyz.zcraft.seira.bot.data; + + +import com.google.gson.annotations.SerializedName; + +import java.util.List; + +public record MsgElem( + List attachments, + String content, + @SerializedName("message_type") Integer messageType, + @SerializedName("msg_idx") String msgIdx +) { +} diff --git a/src/main/java/xyz/zcraft/seira/bot/data/PendingMessage.java b/src/main/java/xyz/zcraft/seira/bot/data/PendingMessage.java index d9da7f41..4f87c143 100644 --- a/src/main/java/xyz/zcraft/seira/bot/data/PendingMessage.java +++ b/src/main/java/xyz/zcraft/seira/bot/data/PendingMessage.java @@ -95,4 +95,12 @@ public PendingMessage ref(MessageReference reference) { this.messageReference = reference; return this; } + + public String getRealContent() { + if (this instanceof MDMessage md) { + return md.getMarkdown(); + } else { + return getContent(); + } + } } diff --git a/src/main/java/xyz/zcraft/seira/command/AttachmentHandler.java b/src/main/java/xyz/zcraft/seira/command/AttachmentHandler.java index 932a6014..8182e7c0 100644 --- a/src/main/java/xyz/zcraft/seira/command/AttachmentHandler.java +++ b/src/main/java/xyz/zcraft/seira/command/AttachmentHandler.java @@ -2,7 +2,7 @@ import org.apache.logging.log4j.LogManager; import org.apache.logging.log4j.Logger; -import xyz.zcraft.seira.api.APIHelper; +import xyz.zcraft.seira.api.OstellaApi; import xyz.zcraft.seira.api.data.ReplayUploadInfo; import xyz.zcraft.seira.bot.data.Attachment; import xyz.zcraft.seira.bot.data.PendingMessage; @@ -74,7 +74,7 @@ public void handleAttachment(Attachment attachment, Consumer msg throw new IllegalArgumentException("Replay file exceeds the 512 KiB limit"); } - final ReplayUploadInfo replayUploadInfo = APIHelper.uploadReplay(bytes); + final ReplayUploadInfo replayUploadInfo = OstellaApi.uploadReplay(bytes); msgSender.accept(ReplyFactory.replayUploadMessage(replayUploadInfo)); } catch (InterruptedException e) { diff --git a/src/main/java/xyz/zcraft/seira/command/CommandReplyChannel.java b/src/main/java/xyz/zcraft/seira/command/CommandReplyChannel.java deleted file mode 100644 index 5b9228fd..00000000 --- a/src/main/java/xyz/zcraft/seira/command/CommandReplyChannel.java +++ /dev/null @@ -1,19 +0,0 @@ -package xyz.zcraft.seira.command; - -import xyz.zcraft.seira.bot.data.PendingMessage; -import xyz.zcraft.seira.data.SendResult; - -/** - * The outbound side of one command invocation. - * - *

A reply is associated with the inbound QQ message. A proactive message is - * sent to the same conversation without that association and therefore does - * not consume the passive reply sequence.

- */ -public interface CommandReplyChannel { - SendResult sendReply(PendingMessage message); - - SendResult sendProactive(PendingMessage message); - - SendResult sendQueueNotice(PendingMessage message); -} diff --git a/src/main/java/xyz/zcraft/seira/command/Context.java b/src/main/java/xyz/zcraft/seira/command/Context.java index ccc13d70..9ff9febb 100644 --- a/src/main/java/xyz/zcraft/seira/command/Context.java +++ b/src/main/java/xyz/zcraft/seira/command/Context.java @@ -4,30 +4,20 @@ import xyz.zcraft.seira.data.SendResult; import java.util.Objects; +import java.util.function.Consumer; public record Context( - String senderUserId, - String groupId, - String messageId, - String command, - String[] args, - String query, - String rawContent, - CommandReplyChannel replies) { + String senderUserId, String groupId, String messageId, String command, + String[] args, String query, String rawContent, ReplyChannel replies, Consumer recorder +) { public Context( - String senderUserId, - String groupId, - String messageId, - String command, - String[] args, - String rawContent, - String query + String senderUserId, String groupId, String messageId, String command, + String[] args, String rawContent, String query ) { - this(senderUserId, groupId, messageId, command, args, query, rawContent, null); + this(senderUserId, groupId, messageId, command, args, query, rawContent, null, null); } public Context { - command = Objects.requireNonNull(command, "command"); args = args == null ? new String[0] : args.clone(); query = query == null ? "" : query; } @@ -49,16 +39,23 @@ public boolean inGroup() { return groupId != null && !groupId.isBlank(); } - public Context withReplies(CommandReplyChannel replyChannel) { + public Context withReplies(ReplyChannel replyChannel) { return new Context( senderUserId, groupId, messageId, command, args, query, rawContent, - Objects.requireNonNull(replyChannel, "replyChannel") + Objects.requireNonNull(replyChannel, "replyChannel"), recorder ); } public Context asCommand(String nextCommand, String[] nextArgs, String nextQuery) { return new Context( - senderUserId, groupId, messageId, nextCommand, nextArgs, nextQuery, rawContent, replies + senderUserId, groupId, messageId, nextCommand, nextArgs, nextQuery, rawContent, replies, recorder + ); + } + + public Context withRecorder(Consumer recorder) { + return new Context( + senderUserId, groupId, messageId, command, args, query, rawContent, + replies, Objects.requireNonNull(recorder, "recorder") ); } @@ -66,28 +63,67 @@ public Context asCommand(String nextCommand, String[] nextArgs, String nextQuery * Sends a passive reply associated with the message that invoked this command. */ public SendResult sendReply(PendingMessage message) { - return requireReplies().sendReply(Objects.requireNonNull(message, "message")); + return sendReply(message, false); + } + + public SendResult sendReply(PendingMessage message, boolean ref) { + record(message); + return requireReplies().sendReply(Objects.requireNonNull(message, "message"), ref); } public SendResult sendReply(String message) { - return requireReplies().sendReply(PendingMessage.ofString(message)); + final PendingMessage msg = PendingMessage.ofMarkdownRaw(message); + record(msg); + return requireReplies().sendReply(msg); + } + + public SendResult send(boolean replyFirst, PendingMessage message) { + return send(replyFirst, message, false); + } + + public SendResult send(boolean replyFirst, PendingMessage message, boolean ref) { + SendResult sendResult; + if (replyFirst) { + sendResult = sendReply(message, ref); + if (!sendResult.success()) { + sendResult = sendMessage(message, ref); + } + } else { + sendResult = sendMessage(message, ref); + if (!sendResult.success()) { + sendResult = sendReply(message, ref); + } + } + return sendResult; } /** * Sends an active message to the same user or group, without an inbound message reference. */ + public SendResult sendMessage(PendingMessage message, boolean ref) { + record(message); + return requireReplies().sendProactive(Objects.requireNonNull(message, "message"), ref); + } + public SendResult sendMessage(PendingMessage message) { + record(message); return requireReplies().sendProactive(Objects.requireNonNull(message, "message")); } public SendResult sendQueueNotice(PendingMessage message) { + record(message); return requireReplies().sendQueueNotice(Objects.requireNonNull(message, "message")); } - private CommandReplyChannel requireReplies() { + private ReplyChannel requireReplies() { if (replies == null) { throw new IllegalStateException("This command context is not bound to a reply channel"); } return replies; } + + private void record(PendingMessage message) { + if (message == null || recorder == null) return; + recorder.accept(message.getRealContent()); + } } diff --git a/src/main/java/xyz/zcraft/seira/command/ReplayResultStore.java b/src/main/java/xyz/zcraft/seira/command/ReplayResultStore.java index de5e26c3..b4d7182e 100644 --- a/src/main/java/xyz/zcraft/seira/command/ReplayResultStore.java +++ b/src/main/java/xyz/zcraft/seira/command/ReplayResultStore.java @@ -1,13 +1,13 @@ package xyz.zcraft.seira.command; -import xyz.zcraft.seira.api.APIHelper; +import xyz.zcraft.seira.api.OstellaApi; import java.util.Objects; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.ConcurrentMap; public final class ReplayResultStore { - private final ConcurrentMap results = new ConcurrentHashMap<>(); + private final ConcurrentMap results = new ConcurrentHashMap<>(); private static String requireTaskId(String taskId) { if (taskId == null || taskId.isBlank()) { @@ -16,11 +16,11 @@ private static String requireTaskId(String taskId) { return taskId; } - void put(String taskId, APIHelper.ReplayRenderResult result) { + void put(String taskId, OstellaApi.ReplayRenderResult result) { results.put(requireTaskId(taskId), Objects.requireNonNull(result)); } - public APIHelper.ReplayRenderResult get(String taskId) { + public OstellaApi.ReplayRenderResult get(String taskId) { return results.get(requireTaskId(taskId)); } diff --git a/src/main/java/xyz/zcraft/seira/command/ReplyChannel.java b/src/main/java/xyz/zcraft/seira/command/ReplyChannel.java new file mode 100644 index 00000000..761d1ea6 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/command/ReplyChannel.java @@ -0,0 +1,58 @@ +package xyz.zcraft.seira.command; + +import xyz.zcraft.seira.bot.data.MessageReference; +import xyz.zcraft.seira.bot.data.PendingMessage; +import xyz.zcraft.seira.data.SendResult; + +import java.util.concurrent.atomic.AtomicInteger; + +public class ReplyChannel { + private final TaskCoordinator taskCoordinator; + private final String targetId; + private final String messageId; + private final String refMsgIdx; + private final boolean groupMessage; + private final boolean queueMessageInGroup; + private final AtomicInteger passiveSequence = new AtomicInteger(1); + + ReplyChannel(TaskCoordinator taskCoordinator, + String targetId, + String messageId, + boolean groupMessage, + boolean queueMessageInGroup, + String refMsgIdx + ) { + this.taskCoordinator = taskCoordinator; + this.targetId = targetId; + this.messageId = messageId; + this.groupMessage = groupMessage; + this.queueMessageInGroup = queueMessageInGroup; + this.refMsgIdx = refMsgIdx; + } + + public synchronized SendResult sendReply(PendingMessage message) { + return taskCoordinator.sendOutboundMessage(targetId, messageId, groupMessage, message, passiveSequence); + } + + public synchronized SendResult sendReply(PendingMessage message, boolean ref) { + if (ref) message.ref(new MessageReference(refMsgIdx)); + return taskCoordinator.sendOutboundMessage(targetId, messageId, groupMessage, message, passiveSequence); + } + + public synchronized SendResult sendProactive(PendingMessage message) { + return sendProactive(message, false); + } + + public synchronized SendResult sendProactive(PendingMessage message, boolean ref) { + if (ref) message.ref(new MessageReference(refMsgIdx)); + return taskCoordinator.sendOutboundMessage(targetId, null, groupMessage, message, null); + } + + public synchronized SendResult sendQueueNotice(PendingMessage message) { + if (groupMessage && !queueMessageInGroup) { + return new SendResult(true, null); + } + return sendReply(message); + } + +} diff --git a/src/main/java/xyz/zcraft/seira/command/TargetHistory.java b/src/main/java/xyz/zcraft/seira/command/TargetHistory.java index eff75bab..3aff3f6e 100644 --- a/src/main/java/xyz/zcraft/seira/command/TargetHistory.java +++ b/src/main/java/xyz/zcraft/seira/command/TargetHistory.java @@ -1,285 +1,22 @@ package xyz.zcraft.seira.command; -import xyz.zcraft.seira.api.APIHelper; -import xyz.zcraft.seira.bot.data.PendingMessage; -import xyz.zcraft.seira.command.parse.Resolver; -import xyz.zcraft.seira.command.parse.ShortcutTarget; -import xyz.zcraft.seira.command.parse.TargetResolution; -import xyz.zcraft.seira.command.parse.UserRefResolution; -import xyz.zcraft.seira.data.UserRef; - -import java.util.List; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.ConcurrentMap; -import java.util.function.Function; -import java.util.function.Predicate; - -import static xyz.zcraft.seira.command.reply.ReplyFactory.at; public final class TargetHistory { private final ConcurrentMap users = new ConcurrentHashMap<>(); - private final Resolver resolver; - private final Function accessToken; - public TargetHistory(Resolver resolver, Function accessToken) { - this.resolver = resolver; - this.accessToken = accessToken; - } - - private static boolean isLocalId(String id) { - return !id.chars().allMatch(Character::isDigit); - } - /** - * 显式记忆一个新目标,清除旧目标的关联 ID。 - */ - public void remember(Context ctx, long id, Type type) { - remember(ctx, Long.toString(id), type); + public Ids get(Context ctx) { + return users.get(ctx.senderUserId()); } - public void remember(Context ctx, String id, Type type) { - Ids ids = new Ids(); - switch (type) { - case BEATMAPSET -> ids.beatmapsetId = Long.parseLong(id); - case BEATMAP -> ids.beatmapId = Long.parseLong(id); - case SCORE -> ids.scoreId = id; - } - users.put(ctx.senderUserId(), ids); - } - - /** - * 只取指定类型已经记住的 ID,不进行查找。 - */ - public ShortcutTarget get(Context ctx, Type type) { - Ids ids = users.get(ctx.senderUserId()); - if (ids == null) return null; - String id = switch (type) { - case BEATMAPSET -> ids.beatmapsetId == null ? null : ids.beatmapsetId.toString(); - case BEATMAP -> ids.beatmapId == null ? null : ids.beatmapId.toString(); - case SCORE -> ids.scoreId; - }; - if (id == null) return null; - if (type == Type.SCORE && isLocalId(id)) return ShortcutTarget.localScore(id); - return new ShortcutTarget(Long.parseLong(id), null, null, null, null); + public void remember(Context ctx, Long beatmapsetId, Long beatmapId, String scoreId) { + remember(ctx, new Ids(beatmapsetId, beatmapId, scoreId)); } public void remember(Context ctx, Ids ids) { - users.put(ctx.senderUserId(), new Ids(ids)); - } - - public Ids resolve(Context ctx, Type type, TargetResolution args) { - return resolve(ctx, type, args, List.of(), null); - } - - /** - * 查找只修改本次结果;调用者显式 remember 后才更新历史。 - */ - public Ids resolve(Context ctx, Type type, TargetResolution args, List filters, String mod) { - ShortcutTarget target = args.target(); - if (target != null && target.isError()) throw new ResolutionException(target.errorMessage()); - - Ids previous = users.get(ctx.senderUserId()); - // 省略目标时沿用三个 ID;显式输入目标时从空记忆开始。 - Ids ids = target == null ? new Ids(previous) : new Ids(); - UserRef player = args.userOverride() != null ? args.userOverride() - : target == null ? null : target.userRef(); - boolean selectedPlayerScore = false; - - if (target != null) { - if (target.isLocalScore()) { - ids.scoreId = target.localScoreId(); - } else if (!target.isMacro()) { - switch (type) { - case BEATMAPSET -> ids.beatmapsetId = target.explicitId(); - case BEATMAP -> ids.beatmapId = target.explicitId(); - case SCORE -> ids.scoreId = target.explicitId().toString(); - } - } else { - switch (target.macroType()) { - case "m" -> ids.beatmapId = target.explicitId(); - case "s" -> ids.scoreId = target.explicitId().toString(); - case "ms" -> { - ids.beatmapsetId = target.explicitId(); - if (type != Type.BEATMAPSET) { - if (target.macroIndex() == null) throw new ResolutionException("请指定指令目标谱面喵"); - if (type == Type.BEATMAP) { - ids.beatmapId = APIHelper.lookupBeatmap(target, accessToken.apply(ctx.senderUserId())); - } else { - player = requirePlayer(ctx, player); - ids.scoreId = APIHelper.lookupScoreId(new ShortcutTarget( - target.explicitId(), player, "ms", target.macroIndex(), null), filters, mod); - selectedPlayerScore = true; - } - } - } - case "rs", "rp", "bp" -> { - // 先把列表位置固定为实际成绩 ID,后续指令不再重新查询列表。 - player = requirePlayer(ctx, player); - ids.scoreId = APIHelper.lookupScoreId(new ShortcutTarget( - null, player, target.macroType(), target.macroIndex(), null), filters, mod); - selectedPlayerScore = true; - } - case "mp" -> { - if (type == Type.BEATMAPSET) { - ids.beatmapsetId = APIHelper.lookupBeatmapset(target, accessToken.apply(ctx.senderUserId())); - } else { - ids.beatmapId = APIHelper.lookupBeatmap(target, accessToken.apply(ctx.senderUserId())); - } - } - default -> throw new ResolutionException("未知的快捷查询"); - } - } - } - - // 指定用户或 Mods 时,按同一谱面重新查成绩,不能直接沿用旧成绩 ID。 - // /s rs2 @用户 已经选好了该用户的 rs2,不再改查谱面最佳成绩。 - if (type == Type.SCORE && (args.userOverride() != null || mod != null) - && ids.scoreId != null && !selectedPlayerScore) { - if (ids.beatmapId == null) { - ids.beatmapId = APIHelper.getScoreBeatmapId(ids.scoreId); - } - ids.scoreId = null; - } - - switch (type) { - case BEATMAP -> { - if (ids.beatmapId == null) { - if (ids.scoreId == null) throw new ResolutionException("请指定指令目标谱面喵"); - ids.beatmapId = APIHelper.getScoreBeatmapId(ids.scoreId); - } - } - case BEATMAPSET -> { - if (ids.beatmapsetId == null) { - if (ids.beatmapId == null && ids.scoreId != null) { - ids.beatmapId = APIHelper.getScoreBeatmapId(ids.scoreId); - } - if (ids.beatmapId == null) throw new ResolutionException("请指定指令目标喵"); - ids.beatmapsetId = APIHelper.lookupBeatmapset(new ShortcutTarget( - ids.beatmapId, null, "m", null, null), accessToken.apply(ctx.senderUserId())); - } - } - case SCORE -> { - if (ids.scoreId == null) { - if (ids.beatmapId == null) throw new ResolutionException("请指定指令目标谱面喵"); - player = requirePlayer(ctx, player); - ids.scoreId = APIHelper.lookupScoreId(new ShortcutTarget( - ids.beatmapId, player, "m", null, null), filters, mod); - } - } - } - return ids; - } - - private UserRef requirePlayer(Context ctx, UserRef player) { - if (player != null) return player; - Long uid = resolver.resolveBoundUid(ctx.senderUserId()); - if (uid == null) throw new ResolutionException("请先绑定 osu! 账号,再查找记忆谱面上的成绩喵"); - return new UserRef.ByUid(uid); - } - - /** - * 同屏回放需要保留本地成绩本身,以便把它加入回放列表。 - */ - public boolean isLocalScore(Context ctx, TargetResolution args) { - if (args.target() != null) return args.target().isLocalScore(); - Ids ids = users.get(ctx.senderUserId()); - return ids != null && ids.scoreId != null && isLocalId(ids.scoreId); - } - - public TargetResolution parseArguments(Context ctx, String usage, int maxOptions) { - return parseArguments(ctx, usage, maxOptions, _ -> false); - } - - /** - * optional 判断首个参数是否是省略目标后的选项;返回的 consumedArgs 标记选项起点。 - */ - public TargetResolution parseArguments(Context ctx, String usage, int maxOptions, Predicate optional) { - try { - TargetResolution args; - if (ctx.argumentCount() == 0 || optional.test(ctx.argument(0))) { - if (!users.containsKey(ctx.senderUserId())) throw new ResolutionException(usage); - args = new TargetResolution(null, 0); - } else { - args = resolver.resolveTargetWithOptionalMention(ctx.args(), ctx.senderUserId()); - if (args.target().isError()) throw new ResolutionException(args.target().errorMessage()); - } - if (ctx.argumentCount() - args.consumedArgs() > maxOptions) throw new ResolutionException(usage); - return args; - } catch (ResolutionException e) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + e.getMessage())); - return null; - } - } - - public TargetResolution parseScoreArguments(Context ctx, String usage) { - return parseScoreArguments(ctx, usage, 0, _ -> false); - } - - /** - * 目标和可选用户在前,其余参数交给指令;这里不认识 +mod 等具体选项。 - */ - public TargetResolution parseScoreArguments(Context ctx, String usage, int maxOptions, Predicate optional) { - try { - TargetResolution args; - if (ctx.argumentCount() > 0 && resolver.looksLikeMention(ctx.argument(0)) - && (ctx.argumentCount() == 1 || optional.test(ctx.argument(1)))) { - UserRef player = parsePlayer(ctx.argument(0), usage); - if (!users.containsKey(ctx.senderUserId())) throw new ResolutionException(usage); - args = new TargetResolution(null, 1, player); - } else { - args = parseArguments(ctx, usage, Integer.MAX_VALUE, optional); - if (args == null) return null; - String next = args.nextArgument(ctx); - if (next != null && !optional.test(next)) { - args = new TargetResolution(args.target(), args.consumedArgs() + 1, parsePlayer(next, usage)); - } - } - if (ctx.argumentCount() - args.consumedArgs() > maxOptions) throw new ResolutionException(usage); - for (int i = args.consumedArgs(); i < ctx.argumentCount(); i++) { - if (!optional.test(ctx.argument(i))) throw new ResolutionException(usage); - } - return args; - } catch (ResolutionException e) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + e.getMessage())); - return null; - } - } - - private UserRef parsePlayer(String argument, String usage) { - UserRefResolution result = resolver.resolveUserRefArgument(argument); - if (result.errorMessage() != null) throw new ResolutionException(result.errorMessage()); - if (result.userRef() == null) throw new ResolutionException(usage); - return result.userRef(); + users.put(ctx.senderUserId(), ids); } - public enum Type {BEATMAPSET, BEATMAP, SCORE} - - // 每个调用者只保存三个 ID。成绩 ID 使用字符串以兼容 loc... 本地成绩。 - public static final class Ids { - private Long beatmapsetId; - private Long beatmapId; - private String scoreId; - - Ids() { - } - - Ids(Ids previous) { - if (previous != null) { - beatmapsetId = previous.beatmapsetId; - beatmapId = previous.beatmapId; - scoreId = previous.scoreId; - } - } - - public Long beatmapsetId() { - return beatmapsetId; - } - - public Long beatmapId() { - return beatmapId; - } - - public String scoreId() { - return scoreId; - } - } + public record Ids(Long beatmapsetId, Long beatmapId, String scoreId) {} } diff --git a/src/main/java/xyz/zcraft/seira/command/TaskCoordinator.java b/src/main/java/xyz/zcraft/seira/command/TaskCoordinator.java index fe6fb070..22006e83 100644 --- a/src/main/java/xyz/zcraft/seira/command/TaskCoordinator.java +++ b/src/main/java/xyz/zcraft/seira/command/TaskCoordinator.java @@ -1,11 +1,10 @@ package xyz.zcraft.seira.command; -import com.google.gson.Gson; import org.apache.logging.log4j.LogManager; import org.apache.logging.log4j.Logger; import org.jetbrains.annotations.NotNull; -import xyz.zcraft.seira.api.APIHelper; import xyz.zcraft.seira.api.ApiRequestException; +import xyz.zcraft.seira.api.OstellaApi; import xyz.zcraft.seira.api.ReplayRenderException; import xyz.zcraft.seira.api.data.Base64Bytes; import xyz.zcraft.seira.api.data.QqUploadRequest; @@ -19,14 +18,17 @@ import xyz.zcraft.seira.services.BotStat; import java.nio.channels.ClosedChannelException; -import java.util.Map; +import java.util.concurrent.Executors; +import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.ScheduledFuture; +import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicInteger; import static xyz.zcraft.seira.command.reply.ReplyFactory.at; public final class TaskCoordinator { private static final Logger LOG = LogManager.getLogger(TaskCoordinator.class); - + private static final ScheduledExecutorService TIMEOUT_SCHEDULER = Executors.newSingleThreadScheduledExecutor(); private final MessageSender messageSender; private final DiscordBridgeService discordBridgeService; private final ApiRequestStats apiRequestStats = new ApiRequestStats(); @@ -50,7 +52,7 @@ public static String resolveErrorMessage(Exception exception) { return ApiRequestException.getDefaultMessage(e.getErrorCode()); } case ClosedChannelException _ -> { - return "oStella API 无法连接,请稍后再试。"; + return "oStella API 无法连接,请稍后再试喵"; } case ResolutionException e -> { return e.getMessage(); @@ -63,28 +65,34 @@ public static String resolveErrorMessage(Exception exception) { } cursor = cursor.getCause(); } - return "请求处理失败,请稍后再试。"; + return "请求处理失败,请稍后再试喵"; } - public CommandReplyChannel openReplyChannel( - String targetId, - String messageId, - boolean groupMessage, - boolean queueMessageInGroup + public ReplyChannel openReplyChannel( + String targetId, String messageId, boolean groupMessage, boolean queueMessageInGroup, String refMsgIdx ) { - return new OutboundReplyChannel(targetId, messageId, groupMessage, queueMessageInGroup); + return new ReplyChannel(this, targetId, messageId, groupMessage, queueMessageInGroup, refMsgIdx); } + public RequestTiming beginRequest(Context ctx, String requestType, boolean timeoutNotify) { + return beginRequest(ctx, requestType, 60, timeoutNotify ? "请求处理时间超过预期,这可能是由于相关数据缺少缓存,请耐心等待喵。" : null); + } - /** - * Tracks queue estimates and elapsed time; the caller executes the request directly. - */ - public RequestTiming beginRequest(Context ctx, String requestType) { + public RequestTiming beginRequest(Context ctx, String requestType, int timeout, String timeoutNotify) { long estimatedSeconds = apiRequestStats.estimateAndEnqueue(requestType); - RequestTiming timing = new RequestTiming(requestType); + ScheduledFuture schedule = null; + + if (timeoutNotify != null && !timeoutNotify.isBlank()) { + schedule = TIMEOUT_SCHEDULER.schedule( + () -> ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + timeoutNotify)), + timeout, TimeUnit.SECONDS + ); + } + + RequestTiming timing = new RequestTiming(requestType, schedule); + try { - ctx.sendQueueNotice(PendingMessage.ofMarkdownRaw( - at(ctx) + "请求已加入队列,预计等待时间" + estimatedSeconds + "秒。")); + ctx.sendQueueNotice(PendingMessage.ofMarkdownRaw(at(ctx) + "请求已加入队列,预计等待时间" + estimatedSeconds + "秒。")); return timing; } catch (RuntimeException e) { timing.close(); @@ -92,6 +100,10 @@ public RequestTiming beginRequest(Context ctx, String requestType) { } } + public RequestTiming beginRequest(Context ctx, String requestType) { + return beginRequest(ctx, requestType, true); + } + public PendingMessage imageMessage(Response response, PendingMessage completion) { UploadedImage image = messageSender.uploadImageToCos(response.getContent().bytes()); return combineImageAndCompletion(image, completion); @@ -102,22 +114,22 @@ public QqUploadRequest createVideoUploadRequest(Context ctx) { return messageSender.createVideoUploadRequest(targetId, ctx.inGroup()); } - public APIHelper.ReplayRenderResult waitForReplay(APIHelper.ReplayTaskInfo taskInfo) { + public OstellaApi.ReplayRenderResult waitForReplay(OstellaApi.ReplayTaskInfo taskInfo) { return waitForReplay(taskInfo, -1); } - public APIHelper.ReplayRenderResult waitForReplay(APIHelper.ReplayTaskInfo taskInfo, long timeout) { + public OstellaApi.ReplayRenderResult waitForReplay(OstellaApi.ReplayTaskInfo taskInfo, long timeout) { if (taskInfo == null || taskInfo.taskId() == null || taskInfo.taskId().isBlank()) { throw new IllegalArgumentException("回放任务未返回有效请求ID,无法获取视频结果。"); } - APIHelper.ReplayRenderResult result = APIHelper.waitReplayVideo(taskInfo.taskId(), timeout); + OstellaApi.ReplayRenderResult result = OstellaApi.waitReplayVideo(taskInfo.taskId(), timeout); replayResults.put(taskInfo.taskId(), result); BotStat.incrementReplays(); return result; } - public PendingMessage replayVideoMessage(APIHelper.ReplayRenderResult result) { + public PendingMessage replayVideoMessage(OstellaApi.ReplayRenderResult result) { if (result == null) { return PendingMessage.ofString("回放视频生成失败,请稍后重试。"); } @@ -164,12 +176,12 @@ public SendResult sendOutboundMessage(String targetId, String messageId, boolean if (pendingMsg instanceof MDMessage md) { message.setMsgType(PendingMessage.MSG_TYPE_MARKDOWN); - message.setMarkdown(new Gson().toJsonTree(Map.of("content", md.getMarkdown())).getAsJsonObject()); + message.setMarkdown(Message.MessageMarkdown.of(md.getMarkdown())); if (md.hasKeyboard()) { message.setKeyboard(md.getKeyboard()); } } else if (pendingMsg.getMsgType() == PendingMessage.MSG_TYPE_MARKDOWN) { - message.setMarkdown(new Gson().toJsonTree(Map.of("content", pendingMsg.getContent())).getAsJsonObject()); + message.setMarkdown(Message.MessageMarkdown.of(pendingMsg.getContent())); } else { message.setContent(pendingMsg.getContent()); } @@ -220,63 +232,32 @@ public SendResult sendOutboundMessage(String targetId, String messageId, boolean discordBridgeService.acceptQqCommandReply(targetId, portableResult); } - return new SendResult(uploadResult && sentMessage != null, sentMessage); + final boolean success = uploadResult && sentMessage != null; + + return new SendResult(success, sentMessage); } public final class RequestTiming implements AutoCloseable { private final String requestType; private final long startedAt = System.nanoTime(); + private final ScheduledFuture scheduledFuture; private boolean closed; - private RequestTiming(String requestType) { + private RequestTiming(String requestType, ScheduledFuture schedule) { this.requestType = requestType; + this.scheduledFuture = schedule; } @Override public void close() { + if (scheduledFuture != null) { + scheduledFuture.cancel(true); + } + if (closed) return; closed = true; apiRequestStats.complete(requestType, Math.max(1L, (System.nanoTime() - startedAt) / 1_000_000L)); } } - - private final class OutboundReplyChannel implements CommandReplyChannel { - private final String targetId; - private final String messageId; - private final boolean groupMessage; - private final boolean queueMessageInGroup; - private final AtomicInteger passiveSequence = new AtomicInteger(1); - - private OutboundReplyChannel( - String targetId, - String messageId, - boolean groupMessage, - boolean queueMessageInGroup - ) { - this.targetId = targetId; - this.messageId = messageId; - this.groupMessage = groupMessage; - this.queueMessageInGroup = queueMessageInGroup; - } - - @Override - public synchronized SendResult sendReply(PendingMessage message) { - return sendOutboundMessage(targetId, messageId, groupMessage, message, passiveSequence); - } - - @Override - public synchronized SendResult sendProactive(PendingMessage message) { - return sendOutboundMessage(targetId, null, groupMessage, message, null); - } - - @Override - public synchronized SendResult sendQueueNotice(PendingMessage message) { - if (groupMessage && !queueMessageInGroup) { - return new SendResult(true, null); - } - return sendReply(message); - } - - } } diff --git a/src/main/java/xyz/zcraft/seira/command/handler/BeatmapCommandHandler.java b/src/main/java/xyz/zcraft/seira/command/handler/BeatmapCommandHandler.java index 25251e58..ebee30e0 100644 --- a/src/main/java/xyz/zcraft/seira/command/handler/BeatmapCommandHandler.java +++ b/src/main/java/xyz/zcraft/seira/command/handler/BeatmapCommandHandler.java @@ -1,25 +1,26 @@ package xyz.zcraft.seira.command.handler; import org.jline.utils.Log; -import xyz.zcraft.seira.api.APIHelper; +import xyz.zcraft.seira.api.OstellaApi; import xyz.zcraft.seira.api.data.Response; import xyz.zcraft.seira.api.data.SearchQuery; import xyz.zcraft.seira.api.data.SearchResultItem; import xyz.zcraft.seira.api.data.VideoRenderRecord; import xyz.zcraft.seira.bot.data.PendingMessage; import xyz.zcraft.seira.command.Context; +import xyz.zcraft.seira.command.ResolutionException; import xyz.zcraft.seira.command.TargetHistory; import xyz.zcraft.seira.command.TaskCoordinator; import xyz.zcraft.seira.command.parse.Resolver; +import xyz.zcraft.seira.command.parse.TargetInput; import xyz.zcraft.seira.command.reply.CommandUsage; import xyz.zcraft.seira.command.reply.ReplyFactory; import xyz.zcraft.seira.data.SendResult; +import xyz.zcraft.seira.util.TimeDurationParser; import java.util.List; import java.util.function.Function; -import static xyz.zcraft.seira.command.TargetHistory.Type.BEATMAP; -import static xyz.zcraft.seira.command.TargetHistory.Type.BEATMAPSET; import static xyz.zcraft.seira.command.reply.ReplyFactory.at; public final class BeatmapCommandHandler { @@ -47,58 +48,181 @@ public BeatmapCommandHandler( } public void handleDaily(Context ctx) { - try (var timing = taskCoordinator.beginRequest(ctx, "Daily Challenge")) { - var daily = APIHelper.getDaily(); + try (var _ = taskCoordinator.beginRequest(ctx, "Daily Challenge")) { + var daily = OstellaApi.getDaily(); ctx.sendReply(PendingMessage.ofMarkdownRaw(daily)); } } public void handleM(Context ctx) { - var target = history.parseArguments(ctx, CommandUsage.M, 1); - if (target == null) return; - try (var timing = taskCoordinator.beginRequest(ctx, "Beatmap")) { - var ids = history.resolve(ctx, BEATMAP, target); - history.remember(ctx, ids); - var response = APIHelper.getBeatmapResponse(ids.beatmapId(), target.nextArgument(ctx)); + var target = ctx.argumentCount() == 0 + ? TargetInput.memory() : TargetInput.read(ctx.args()); + var remembered = history.get(ctx); + if ((target.kind() == TargetInput.Kind.MEMORY && remembered == null) + || ctx.argumentCount() - target.consumedArgs() > 1) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.M)); + return; + } + try (var _ = taskCoordinator.beginRequest(ctx, "Beatmap")) { + var previous = target.kind() == TargetInput.Kind.MEMORY ? remembered : null; + Long beatmapId = previous == null ? null : previous.beatmapId(); + Long beatmapsetId = previous == null ? null : previous.beatmapsetId(); + String scoreId = previous == null ? null : previous.scoreId(); + switch (target.kind()) { + case ID, MAP -> beatmapId = Long.parseLong(target.id()); + case SCORE -> scoreId = target.id(); + case SET -> { + beatmapsetId = Long.parseLong(target.id()); + beatmapId = OstellaApi.lookupBeatmapInSet(beatmapsetId, target.index(), accessTokenProvider.apply(ctx.senderUserId())); + } + case RS, RP, BP -> { + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupPlayerScore(uid, target.scoreList(), target.index(), List.of(), null); + } + case MP -> beatmapId = OstellaApi.lookupMultiplayerBeatmap(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> {} + } + if (beatmapId == null && scoreId != null) beatmapId = OstellaApi.getScoreBeatmapId(scoreId); + if (beatmapId == null) throw new ResolutionException("请指定指令目标谱面喵"); + var response = OstellaApi.getBeatmapResponse(beatmapId, (ctx.argumentCount() > target.consumedArgs() ? ctx.argument(target.consumedArgs()) : null)); + history.remember(ctx, beatmapsetId, beatmapId, scoreId); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.beatmapMessage(ctx, response))); } } public void handleBma(Context ctx) { - var target = history.parseArguments(ctx, CommandUsage.BMA, 1); - if (target == null) return; - try (var timing = taskCoordinator.beginRequest(ctx, "Beatmap Analysis")) { - var ids = history.resolve(ctx, BEATMAP, target); - history.remember(ctx, ids); - var response = APIHelper.getBeatmapAnalysisResponse(ids.beatmapId(), target.nextArgument(ctx)); + var target = ctx.argumentCount() == 0 + ? TargetInput.memory() : TargetInput.read(ctx.args()); + var remembered = history.get(ctx); + if ((target.kind() == TargetInput.Kind.MEMORY && remembered == null) + || ctx.argumentCount() - target.consumedArgs() > 1) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.BMA)); + return; + } + try (var _ = taskCoordinator.beginRequest(ctx, "Beatmap Analysis")) { + var previous = target.kind() == TargetInput.Kind.MEMORY ? remembered : null; + Long beatmapId = previous == null ? null : previous.beatmapId(); + Long beatmapsetId = previous == null ? null : previous.beatmapsetId(); + String scoreId = previous == null ? null : previous.scoreId(); + switch (target.kind()) { + case ID, MAP -> beatmapId = Long.parseLong(target.id()); + case SCORE -> scoreId = target.id(); + case SET -> { + beatmapsetId = Long.parseLong(target.id()); + beatmapId = OstellaApi.lookupBeatmapInSet(beatmapsetId, target.index(), accessTokenProvider.apply(ctx.senderUserId())); + } + case RS, RP, BP -> { + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupPlayerScore(uid, target.scoreList(), target.index(), List.of(), null); + } + case MP -> beatmapId = OstellaApi.lookupMultiplayerBeatmap(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> {} + } + if (beatmapId == null && scoreId != null) beatmapId = OstellaApi.getScoreBeatmapId(scoreId); + if (beatmapId == null) throw new ResolutionException("请指定指令目标谱面喵"); + var response = OstellaApi.getBeatmapAnalysisResponse(beatmapId, (ctx.argumentCount() > target.consumedArgs() ? ctx.argument(target.consumedArgs()) : null)); + history.remember(ctx, beatmapsetId, beatmapId, scoreId); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.beatmapMessage(ctx, response))); } } public void handleAp(Context ctx) { - var target = history.parseArguments(ctx, CommandUsage.AP, Integer.MAX_VALUE); - if (target == null) return; - try (var timing = taskCoordinator.beginRequest(ctx, "Audio Preview")) { - var ids = history.resolve(ctx, BEATMAPSET, target); - history.remember(ctx, ids); - long id = ids.beatmapsetId(); - ctx.sendReply(PendingMessage.ofVoiceUrl("https://b.ppy.sh/preview/" + id + ".mp3").doUpload(false)); + var target = ctx.argumentCount() == 0 + ? TargetInput.memory() : TargetInput.read(ctx.args()); + var remembered = history.get(ctx); + if (target.kind() == TargetInput.Kind.MEMORY && remembered == null) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.AP)); + return; + } + try (var _ = taskCoordinator.beginRequest(ctx, "Audio Preview")) { + var previous = target.kind() == TargetInput.Kind.MEMORY ? remembered : null; + Long beatmapId = previous == null ? null : previous.beatmapId(); + Long beatmapsetId = previous == null ? null : previous.beatmapsetId(); + String scoreId = previous == null ? null : previous.scoreId(); + switch (target.kind()) { + case ID -> beatmapsetId = Long.parseLong(target.id()); + case MAP -> beatmapId = Long.parseLong(target.id()); + case SCORE -> scoreId = target.id(); + case SET -> beatmapsetId = Long.parseLong(target.id()); + case RS, RP, BP -> { + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupPlayerScore(uid, target.scoreList(), target.index(), List.of(), null); + } + case MP -> beatmapsetId = OstellaApi.lookupMultiplayerBeatmapset(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> {} + } + if (beatmapsetId == null) { + if (beatmapId == null && scoreId != null) beatmapId = OstellaApi.getScoreBeatmapId(scoreId); + if (beatmapId == null) throw new ResolutionException("请指定指令目标喵"); + beatmapsetId = OstellaApi.lookupBeatmapsetForBeatmap(beatmapId, accessTokenProvider.apply(ctx.senderUserId())); + } + history.remember(ctx, beatmapsetId, beatmapId, scoreId); + ctx.sendReply(PendingMessage.ofVoiceUrl("https://b.ppy.sh/preview/" + beatmapsetId + ".mp3").doUpload(false)); } } public void handleBpv(Context ctx) { - var target = history.parseArguments(ctx, CommandUsage.BPV, 1, arg -> arg.startsWith("+")); - if (target == null) return; - try (var timing = taskCoordinator.beginRequest(ctx, "Beatmap Preview Render")) { + var target = ctx.argumentCount() == 0 || ctx.argument(0).startsWith("+") || TimeDurationParser.isTimeRange(ctx.argument(0)) + ? TargetInput.memory() : TargetInput.read(ctx.args()); + var remembered = history.get(ctx); + if ((target.kind() == TargetInput.Kind.MEMORY && remembered == null) + || ctx.argumentCount() - target.consumedArgs() > 2) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.BPV)); + return; + } + + TimeDurationParser.TimeRange range = null; + String mods = null; + for (int optionIndex = target.consumedArgs(); optionIndex < ctx.argumentCount(); optionIndex++) { + String option = ctx.argument(optionIndex); + if (TimeDurationParser.isTimeRange(option) && range == null) { + try { + range = TimeDurationParser.parseRange(option); + } catch (IllegalArgumentException e) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "无法解析时间范围")); + return; + } + } else if (mods == null && !TimeDurationParser.isTimeRange(option)) { + mods = option; + } else { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.BPV)); + return; + } + } + + try (var _ = taskCoordinator.beginRequest(ctx, "Beatmap Preview Render")) { ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "正在获取谱面以及回放文件,请稍作等待喵...")); var qqUpload = taskCoordinator.createVideoUploadRequest(ctx); - var ids = history.resolve(ctx, BEATMAP, target); - history.remember(ctx, ids); - var task = APIHelper.createBeatmapPreviewTask(ids.beatmapId(), target.nextArgument(ctx), qqUpload); + var previous = target.kind() == TargetInput.Kind.MEMORY ? remembered : null; + Long beatmapId = previous == null ? null : previous.beatmapId(); + Long beatmapsetId = previous == null ? null : previous.beatmapsetId(); + String scoreId = previous == null ? null : previous.scoreId(); + switch (target.kind()) { + case ID, MAP -> beatmapId = Long.parseLong(target.id()); + case SCORE -> scoreId = target.id(); + case SET -> { + beatmapsetId = Long.parseLong(target.id()); + beatmapId = OstellaApi.lookupBeatmapInSet(beatmapsetId, target.index(), accessTokenProvider.apply(ctx.senderUserId())); + } + case RS, RP, BP -> { + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupPlayerScore(uid, target.scoreList(), target.index(), List.of(), null); + } + case MP -> beatmapId = OstellaApi.lookupMultiplayerBeatmap(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> {} + } + if (beatmapId == null && scoreId != null) beatmapId = OstellaApi.getScoreBeatmapId(scoreId); + if (beatmapId == null) throw new ResolutionException("请指定指令目标谱面喵"); + var task = OstellaApi.createBeatmapPreviewTask(beatmapId, mods, range, qqUpload); + history.remember(ctx, beatmapsetId, beatmapId, scoreId); videoRenderRecord.updateRenderTask(ctx.senderUserId(), task.taskId()); ctx.sendReply(replyFactory.replayMessage(ctx, task)); - APIHelper.ReplayRenderResult result; + OstellaApi.ReplayRenderResult result; try { result = taskCoordinator.waitForReplay(task); @@ -121,34 +245,113 @@ public void handleBpv(Context ctx) { } public void handleBgp(Context ctx) { - var target = history.parseArguments(ctx, CommandUsage.BGP, Integer.MAX_VALUE); - if (target == null) return; - try (var timing = taskCoordinator.beginRequest(ctx, "Background Preview")) { - var ids = history.resolve(ctx, BEATMAP, target); - history.remember(ctx, ids); - var response = APIHelper.getBeatmapBgResponse(ids.beatmapId()); + var target = ctx.argumentCount() == 0 + ? TargetInput.memory() : TargetInput.read(ctx.args()); + var remembered = history.get(ctx); + if (target.kind() == TargetInput.Kind.MEMORY && remembered == null) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.BGP)); + return; + } + try (var _ = taskCoordinator.beginRequest(ctx, "Background Preview")) { + var previous = target.kind() == TargetInput.Kind.MEMORY ? remembered : null; + Long beatmapId = previous == null ? null : previous.beatmapId(); + Long beatmapsetId = previous == null ? null : previous.beatmapsetId(); + String scoreId = previous == null ? null : previous.scoreId(); + switch (target.kind()) { + case ID, MAP -> beatmapId = Long.parseLong(target.id()); + case SCORE -> scoreId = target.id(); + case SET -> { + beatmapsetId = Long.parseLong(target.id()); + beatmapId = OstellaApi.lookupBeatmapInSet(beatmapsetId, target.index(), accessTokenProvider.apply(ctx.senderUserId())); + } + case RS, RP, BP -> { + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupPlayerScore(uid, target.scoreList(), target.index(), List.of(), null); + } + case MP -> beatmapId = OstellaApi.lookupMultiplayerBeatmap(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> {} + } + if (beatmapId == null && scoreId != null) beatmapId = OstellaApi.getScoreBeatmapId(scoreId); + if (beatmapId == null) throw new ResolutionException("请指定指令目标谱面喵"); + var response = OstellaApi.getBeatmapBgResponse(beatmapId); + history.remember(ctx, beatmapsetId, beatmapId, scoreId); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.bgpMessage(ctx, response))); } } public void handleDl(Context ctx) { - var target = history.parseArguments(ctx, CommandUsage.DL, 0); - if (target == null) return; - try (var timing = taskCoordinator.beginRequest(ctx, "Download Beatmap")) { - var ids = history.resolve(ctx, BEATMAPSET, target); - history.remember(ctx, ids); - var response = APIHelper.getLookupBeatmapsetResponse(ids.beatmapsetId(), accessTokenProvider.apply(ctx.senderUserId())); + var target = ctx.argumentCount() == 0 + ? TargetInput.memory() : TargetInput.read(ctx.args()); + var remembered = history.get(ctx); + if ((target.kind() == TargetInput.Kind.MEMORY && remembered == null) + || ctx.argumentCount() - target.consumedArgs() > 0) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.DL)); + return; + } + try (var _ = taskCoordinator.beginRequest(ctx, "Download Beatmap")) { + var previous = target.kind() == TargetInput.Kind.MEMORY ? remembered : null; + Long beatmapId = previous == null ? null : previous.beatmapId(); + Long beatmapsetId = previous == null ? null : previous.beatmapsetId(); + String scoreId = previous == null ? null : previous.scoreId(); + switch (target.kind()) { + case ID -> beatmapsetId = Long.parseLong(target.id()); + case MAP -> beatmapId = Long.parseLong(target.id()); + case SCORE -> scoreId = target.id(); + case SET -> beatmapsetId = Long.parseLong(target.id()); + case RS, RP, BP -> { + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupPlayerScore(uid, target.scoreList(), target.index(), List.of(), null); + } + case MP -> beatmapsetId = OstellaApi.lookupMultiplayerBeatmapset(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> {} + } + if (beatmapsetId == null) { + if (beatmapId == null && scoreId != null) beatmapId = OstellaApi.getScoreBeatmapId(scoreId); + if (beatmapId == null) throw new ResolutionException("请指定指令目标喵"); + beatmapsetId = OstellaApi.lookupBeatmapsetForBeatmap(beatmapId, accessTokenProvider.apply(ctx.senderUserId())); + } + var response = OstellaApi.getLookupBeatmapsetResponse(beatmapsetId, accessTokenProvider.apply(ctx.senderUserId())); + history.remember(ctx, beatmapsetId, beatmapId, scoreId); ctx.sendReply(replyFactory.dlMessage(ctx, response)); } } public void handleMs(Context ctx) { - var target = history.parseArguments(ctx, "用法:/ms <谱面集ID 或 快捷查询>", 0); - if (target == null) return; - try (var timing = taskCoordinator.beginRequest(ctx, "Beatmapset")) { - var ids = history.resolve(ctx, BEATMAPSET, target); - history.remember(ctx, ids); - var response = APIHelper.getBeatmapsetResponse(ids.beatmapsetId()); + var target = ctx.argumentCount() == 0 + ? TargetInput.memory() : TargetInput.read(ctx.args()); + var remembered = history.get(ctx); + if ((target.kind() == TargetInput.Kind.MEMORY && remembered == null) + || ctx.argumentCount() - target.consumedArgs() > 0) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "用法:/ms <谱面集ID 或 快捷查询>")); + return; + } + try (var _ = taskCoordinator.beginRequest(ctx, "Beatmapset")) { + var previous = target.kind() == TargetInput.Kind.MEMORY ? remembered : null; + Long beatmapId = previous == null ? null : previous.beatmapId(); + Long beatmapsetId = previous == null ? null : previous.beatmapsetId(); + String scoreId = previous == null ? null : previous.scoreId(); + switch (target.kind()) { + case ID -> beatmapsetId = Long.parseLong(target.id()); + case MAP -> beatmapId = Long.parseLong(target.id()); + case SCORE -> scoreId = target.id(); + case SET -> beatmapsetId = Long.parseLong(target.id()); + case RS, RP, BP -> { + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupPlayerScore(uid, target.scoreList(), target.index(), List.of(), null); + } + case MP -> beatmapsetId = OstellaApi.lookupMultiplayerBeatmapset(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> {} + } + if (beatmapsetId == null) { + if (beatmapId == null && scoreId != null) beatmapId = OstellaApi.getScoreBeatmapId(scoreId); + if (beatmapId == null) throw new ResolutionException("请指定指令目标喵"); + beatmapsetId = OstellaApi.lookupBeatmapsetForBeatmap(beatmapId, accessTokenProvider.apply(ctx.senderUserId())); + } + var response = OstellaApi.getBeatmapsetResponse(beatmapsetId); + history.remember(ctx, beatmapsetId, beatmapId, scoreId); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.beatmapsetMessage(ctx, response))); } } @@ -159,8 +362,8 @@ public void handleSms(Context ctx) { ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "用法:/sms [#页数] <搜索关键字>")); return; } - try (var timing = taskCoordinator.beginRequest(ctx, "Search Beatmapset")) { - Response> searchResponse = APIHelper.searchBeatmapSetResponse(searchQuery); + try (var _ = taskCoordinator.beginRequest(ctx, "Search Beatmapset")) { + Response> searchResponse = OstellaApi.searchBeatmapSetResponse(searchQuery); ctx.sendReply(replyFactory.searchMessage(ctx, searchResponse, searchQuery)); } } diff --git a/src/main/java/xyz/zcraft/seira/command/handler/DcsCommandHandler.java b/src/main/java/xyz/zcraft/seira/command/handler/DcsCommandHandler.java index e598e5e7..3d8990ed 100644 --- a/src/main/java/xyz/zcraft/seira/command/handler/DcsCommandHandler.java +++ b/src/main/java/xyz/zcraft/seira/command/handler/DcsCommandHandler.java @@ -8,6 +8,7 @@ import java.util.Locale; import java.util.Objects; +import static xyz.zcraft.seira.command.reply.ReplyFactory.ExternalUrls.PERMISSION; import static xyz.zcraft.seira.command.reply.ReplyFactory.at; public final class DcsCommandHandler { @@ -52,7 +53,7 @@ private void handleStart(Context ctx) { final boolean b = ctx.sendMessage(PendingMessage.ofString("正在尝试开启 Discord 消息同步,请稍候...")).success(); if (!b) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "由于缺少主动消息权限,无法添加消息同步!权限配置请见[这里](https://docs.seira.top/overview/use.html#extra-permission)~")); + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "由于缺少主动消息权限,无法添加消息同步!权限配置请见[这里](" + PERMISSION + ")~")); return; } diff --git a/src/main/java/xyz/zcraft/seira/command/handler/GeneralCommandHandler.java b/src/main/java/xyz/zcraft/seira/command/handler/GeneralCommandHandler.java index e20a8b10..329ee9a4 100644 --- a/src/main/java/xyz/zcraft/seira/command/handler/GeneralCommandHandler.java +++ b/src/main/java/xyz/zcraft/seira/command/handler/GeneralCommandHandler.java @@ -1,21 +1,26 @@ package xyz.zcraft.seira.command.handler; import xyz.zcraft.osu.model.Beatmapset; -import xyz.zcraft.seira.api.APIHelper; +import xyz.zcraft.seira.api.AsteroidApi; +import xyz.zcraft.seira.api.OstellaApi; +import xyz.zcraft.seira.api.data.MinecraftServerStatus; import xyz.zcraft.seira.bot.MessageSender; import xyz.zcraft.seira.bot.data.PendingMessage; import xyz.zcraft.seira.command.Context; import xyz.zcraft.seira.command.TaskCoordinator; import xyz.zcraft.seira.command.parse.Resolver; -import xyz.zcraft.seira.command.parse.UserRefResolution; import xyz.zcraft.seira.command.reply.ReplyFactory; import xyz.zcraft.seira.data.Notice; import xyz.zcraft.seira.data.UploadedImage; -import xyz.zcraft.seira.data.UserRef; import xyz.zcraft.seira.services.DailyLuck; import xyz.zcraft.seira.services.NoticeStore; +import xyz.zcraft.seira.util.dice.Dice; +import xyz.zcraft.seira.util.dice.expr.DiceExpr; +import xyz.zcraft.seira.util.dice.result.DiceResult; +import java.util.Objects; import java.util.function.Predicate; +import java.util.regex.Pattern; import static xyz.zcraft.seira.command.reply.ReplyFactory.at; @@ -41,29 +46,38 @@ public GeneralCommandHandler( } public void handleU(Context context) { - UserRef userRef; - if (context.argumentCount() == 0) { - Long boundUid = resolver.resolveBoundUid(context.senderUserId()); - userRef = boundUid == null ? null : new UserRef.ByUid(boundUid); + String player = resolver.player(context.argumentCount() == 0 ? null : context.argument(0), context.senderUserId()); + + try (var _ = taskCoordinator.beginRequest(context, "User Info")) { + long uid = OstellaApi.resolveUid(player); + var response = OstellaApi.getUserInfoResponse(uid); + var completion = replyFactory.userInfoMessage(context, response); + context.sendReply(taskCoordinator.imageMessage(response, completion)); + } + } + + public void handleRoll(Context ctx) { + DiceExpr diceExpr; + + if (ctx.argumentCount() == 0) { + diceExpr = DiceExpr.HUNDRED; } else { - UserRefResolution target = resolver.resolveUserRefArgument(context.argument(0)); - if (target.errorMessage() != null) { - context.sendReply(PendingMessage.ofMarkdownRaw(at(context) + target.errorMessage())); + try { + diceExpr = DiceExpr.parse(ctx.query()); + } catch (Exception e) { + ctx.sendReply(at(ctx) + "无法解析骰子表达式喵。"); return; } - userRef = target.userRef(); } - if (userRef == null) { - context.sendReply(PendingMessage.ofMarkdownRaw(at(context) + "用法:/u [玩家ID/用户名/@用户]")); - return; - } + final Dice dice = new Dice(); + final DiceResult rollResult = dice.roll(diceExpr); - try (var _ = taskCoordinator.beginRequest(context, "User Info")) { - var response = APIHelper.getUserInfoResponse(userRef); - var completion = replyFactory.userInfoMessage(context, response); - context.sendReply(taskCoordinator.imageMessage(response, completion)); - } + final String rollResultStr = rollResult.toString(); + final String totalStr = String.valueOf(rollResult.total()); + + final String message = at(ctx) + diceExpr + (Objects.equals(rollResultStr, totalStr) ? "" : " = " + rollResultStr) + " = __" + totalStr + "__"; + ctx.sendReply(message.replace("*", "\\*")); } public void handleLuck(Context context) { @@ -74,7 +88,7 @@ public void handleLuck(Context context) { try (var _ = taskCoordinator.beginRequest(context, "Luck")) { DailyLuck.Luck luck = DailyLuck.getLuck(context.senderUserId()); - Beatmapset mapset = APIHelper.getBeatmapsetRaw(luck.dailyMapset()); + Beatmapset mapset = OstellaApi.getBeatmapsetRaw(luck.dailyMapset()); UploadedImage cover = messageSender.uploadImageToCos(mapset.getCovers().getCover()); context.sendReply(replyFactory.luckMessage(context, luck, mapset, cover)); } @@ -91,12 +105,16 @@ public void handleHelp(Context context) { context.sendReply(replyFactory.helpMessage(context)); } + public void handleUsages(Context context) { + context.sendReply(replyFactory.usagesMessage(context)); + } + public void handleFaq(Context context) { context.sendReply(replyFactory.faqMessage(context)); } public void handleStat(Context context) { - context.sendReply(replyFactory.statusMessage(context, APIHelper.getServerStatus())); + context.sendReply(replyFactory.statusMessage(context, OstellaApi.getServerStatus(), AsteroidApi.getServerStatus())); } public void handleUnknown(Context context) { @@ -125,10 +143,70 @@ public void handleNotice(Context context) { .filter(Notice::isActive) .findFirst() .ifPresentOrElse(notice -> { - sb.append(at(context)).append("公告#").append(notice.id()).append(" ").append(notice.title()).append("\n"); + sb.append(at(context)).append("公告 `#").append(notice.id()).append("` - `").append(notice.title()).append("`\n"); sb.append(NoticeStore.getContentFor(notice)); }, () -> sb.append(at(context)).append("未找到公告#").append(noticeId)); context.sendReply(PendingMessage.ofMarkdownRaw(sb.toString())); } + + private static final Pattern SERVER_PATTERN = Pattern.compile("^(?:https?://)?([a-zA-Z0-9-]+(?:\\.[a-zA-Z0-9-]+)+)(?::[0-9]+)?$"); + + public void handleMc(Context ctx) { + if (ctx.argumentCount() != 1) { + ctx.sendReply(at(ctx) + "用法:/mc <服务器地址>"); + return; + } + + try { + final String address = ctx.argument(0); + + if (!SERVER_PATTERN.matcher(address).matches()) { + ctx.sendReply(at(ctx) + "服务器地址无效喵。"); + return; + } + + final var probe = AsteroidApi.getMinecraftServerStatus(address); + + final var status = probe.status(); + final var players = status.players(); + final var samples = players.samples(); + + String playersSample = ""; + + if (samples != null) { + playersSample = String.join(", ", samples.stream().map(MinecraftServerStatus.Players.Sample::name).toList()); + } + + ctx.sendReply(at(ctx) + """ + `%s` 的服务器状态: + - 版本: `%s` + - 描述: `%s` + - 延迟: `%d` ms + - 在线人数: `%d` / `%d` + - 在线玩家: [%s] + """.formatted( + address, + status.version().name(), + status.description(), + probe.latency(), + players.online(), + players.max(), + playersSample + ) + ); + } catch (Exception e) { + ctx.sendReply(at(ctx) + "无法获取目标服务器状态喵,这可能是因为目标服务器未开启或者存在网络问题。"); + } + } + + public void handleUx(Context ctx) { + String player = resolver.player(ctx.argumentCount() == 0 ? null : ctx.argument(0), ctx.senderUserId()); + + try (var _ = taskCoordinator.beginRequest(ctx, "User Info Short")) { + long uid = OstellaApi.resolveUid(player); + var user = OstellaApi.getUserRaw(uid); + ctx.sendReply(replyFactory.userInfoShortMessage(ctx, user)); + } + } } diff --git a/src/main/java/xyz/zcraft/seira/command/handler/MultiplayerRoomWatchCommandHandler.java b/src/main/java/xyz/zcraft/seira/command/handler/MPWatchCommandHandler.java similarity index 50% rename from src/main/java/xyz/zcraft/seira/command/handler/MultiplayerRoomWatchCommandHandler.java rename to src/main/java/xyz/zcraft/seira/command/handler/MPWatchCommandHandler.java index cf161988..fcd97ce4 100644 --- a/src/main/java/xyz/zcraft/seira/command/handler/MultiplayerRoomWatchCommandHandler.java +++ b/src/main/java/xyz/zcraft/seira/command/handler/MPWatchCommandHandler.java @@ -1,16 +1,21 @@ package xyz.zcraft.seira.command.handler; import xyz.zcraft.osu.model.MultiplayerRoom; -import xyz.zcraft.seira.api.APIHelper; +import xyz.zcraft.osu.model.User; +import xyz.zcraft.osu.model.UserExtended; +import xyz.zcraft.seira.api.OstellaApi; +import xyz.zcraft.seira.api.RomAIApi; import xyz.zcraft.seira.api.data.OsuToken; import xyz.zcraft.seira.api.data.Response; +import xyz.zcraft.seira.api.data.RomAIMatch; import xyz.zcraft.seira.bot.data.PendingMessage; import xyz.zcraft.seira.command.Context; import xyz.zcraft.seira.command.ResolutionException; import xyz.zcraft.seira.command.TaskCoordinator; +import xyz.zcraft.seira.command.parse.Resolver; import xyz.zcraft.seira.db.UserDataStore; -import xyz.zcraft.seira.watch.MultiplayerRoomVersion; -import xyz.zcraft.seira.watch.MultiplayerRoomWatchService; +import xyz.zcraft.seira.watch.MPVersion; +import xyz.zcraft.seira.watch.MPWatchService; import xyz.zcraft.seira.watch.RoomWatchView; import java.util.Locale; @@ -18,9 +23,10 @@ import java.util.regex.Matcher; import java.util.regex.Pattern; +import static xyz.zcraft.seira.command.reply.ReplyFactory.ExternalUrls.PERMISSION; import static xyz.zcraft.seira.command.reply.ReplyFactory.at; -public final class MultiplayerRoomWatchCommandHandler { +public final class MPWatchCommandHandler { private static final String USAGE = "用法:/mpwatch [start] <房间ID> [stable|lazer];" + "/mpwatch [start] <房间链接>;/mpwatch stop [all];/mpwatch status"; @@ -32,57 +38,54 @@ public final class MultiplayerRoomWatchCommandHandler { ); private final TaskCoordinator taskCoordinator; - private final MultiplayerRoomWatchService watchService; + private final MPWatchService watchService; + private final Resolver resolver; - public MultiplayerRoomWatchCommandHandler( + public MPWatchCommandHandler( TaskCoordinator taskCoordinator, - MultiplayerRoomWatchService watchService + Resolver resolver, + MPWatchService watchService ) { this.taskCoordinator = Objects.requireNonNull(taskCoordinator); + this.resolver = Objects.requireNonNull(resolver); this.watchService = Objects.requireNonNull(watchService); } - static RoomTarget parseRoomTarget(String value, String explicitVersion) { + static RoomTarget parseRoomTarget(String value, MPVersion requestedVersion, Integer bo) { if (value == null || value.isBlank()) { return null; } String normalized = value.trim(); String numeric = normalized; - MultiplayerRoomVersion inferredVersion = null; + MPVersion inferredVersion = null; Matcher lazerMatcher = LAZER_ROOM_URL.matcher(normalized); Matcher stableMatcher = STABLE_ROOM_URL.matcher(normalized); if (lazerMatcher.matches()) { numeric = lazerMatcher.group(1); - inferredVersion = MultiplayerRoomVersion.LAZER; + inferredVersion = MPVersion.LAZER; } else if (stableMatcher.matches()) { numeric = stableMatcher.group(1); - inferredVersion = MultiplayerRoomVersion.STABLE; + inferredVersion = MPVersion.STABLE; } else if (!normalized.matches("\\d+")) { return null; } - MultiplayerRoomVersion requestedVersion = explicitVersion == null - ? null - : MultiplayerRoomVersion.parse(explicitVersion); - if (explicitVersion != null && requestedVersion == null) { - return null; - } if (inferredVersion != null && requestedVersion != null && inferredVersion != requestedVersion) { return null; } - MultiplayerRoomVersion version = inferredVersion != null - ? inferredVersion - : requestedVersion == null ? MultiplayerRoomVersion.LAZER : requestedVersion; + + MPVersion version = (inferredVersion != null ? inferredVersion : requestedVersion); + try { long roomId = Long.parseLong(numeric); - return roomId > 0 ? new RoomTarget(roomId, version) : null; + return roomId > 0 ? new RoomTarget(roomId, version, bo) : null; } catch (NumberFormatException ignored) { return null; } } private static String formatRoom(RoomWatchView view) { - return view.version().value() + " 多人房间“" + view.roomName() + "” (#" + view.roomId() + ")"; + return view.version().value() + " - " + view.roomId() + " `" + view.roomName() + "`"; } private static void usage(Context ctx) { @@ -108,9 +111,91 @@ public void handleMpWatch(Context ctx) { } } + public void handleRomAI(Context ctx) { + if (!ctx.inGroup()) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "/romai 仅支持群聊使用。")); + return; + } + + String username; + if (ctx.argumentCount() == 0) { + final Long boundUid = resolver.resolveBoundUid(ctx.senderUserId()); + if (boundUid == null) { + ctx.sendReply(at(ctx) + "由于未绑定,无法查找正在进行的 RomAI 比赛喵。"); + return; + } + + final UserExtended userRaw = OstellaApi.getUserRaw(boundUid); + + username = userRaw.getUsername(); + } else if (ctx.argumentCount() == 1) { + final String player = resolver.player(ctx.argument(0), null); + final User content = OstellaApi.getUserRaw(OstellaApi.resolveUid(player)); + username = content.getUsername(); + } else { + ctx.sendReply(at(ctx) + "用法:/romai [@user]"); + return; + } + + final RomAIMatch match = RomAIApi.getMatchFor(username); + + if (match == null) { + ctx.sendReply(at(ctx) + username + "不在打 RomAI 喵。"); + return; + } + + String teamString; + + if (match.teams() != null) { + teamString = """ + > Team A: %s + > Team B: %s + """.formatted( + String.join(", ", match.teams().teamA()), + String.join(", ", match.teams().teamB() + ).trim() + ); + } else { + teamString = "> `%s` vs `%s`".formatted(match.players().getFirst(), match.players().getLast()); + } + + final PendingMessage msg = PendingMessage.ofMarkdownRaw( + at(ctx) + """ + 正在启动 RomAI 监视喵。 + > %s - %s + > ELO %s · BO%s + %s + """.formatted( + match.lobbyId(), match.mode(), + (match.customELO() == null ? "?" : match.customELO().toString()), + (match.customBO() == null ? "?" : match.customBO().toString()), + teamString + ).trim() + ); + + try (var _ = taskCoordinator.beginRequest(ctx, "Start RomAI Watch")) { + if (!ctx.sendMessage(msg).success()) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + + "由于缺少主动消息权限,无法启动监视!权限配置请见[这里](" + PERMISSION + ")~" + )); + return; + } + try { + RoomWatchView view = watchService.watch( + ctx.groupId(), ctx.senderUserId(), MPVersion.STABLE, Long.parseLong(match.lobbyId()), match.customBO() + ); + ctx.sendReply(PendingMessage.ofMarkdownRaw( + at(ctx) + "已开始监视" + formatRoom(view) + "喵。" + )); + } catch (IllegalArgumentException | IllegalStateException e) { + throw new ResolutionException(e.getMessage()); + } + } + } + private void handleStart(Context ctx, int argumentOffset) { int startArgumentCount = ctx.argumentCount() - argumentOffset; - if (startArgumentCount < 0 || startArgumentCount > 2) { + if (startArgumentCount < 0 || startArgumentCount > 3) { usage(ctx); return; } @@ -122,11 +207,26 @@ private void handleStart(Context ctx, int argumentOffset) { ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "由于未绑定账户,无法获取当前房间,请手动提供ID~")); return; } - final Response multiplayerRoom = APIHelper.getMultiplayerRoom(osuToken.accessToken()); - target = new RoomTarget(multiplayerRoom.getContent().getId(), MultiplayerRoomVersion.LAZER); + final Response multiplayerRoom = OstellaApi.getMultiplayerRoom(osuToken.accessToken()); + target = new RoomTarget(multiplayerRoom.getContent().getId(), MPVersion.LAZER, null); } else { - String version = startArgumentCount == 2 ? ctx.argument(argumentOffset + 1) : null; - target = parseRoomTarget(ctx.argument(argumentOffset), version); + MPVersion version = null; + Integer customBo = null; + for (int i = argumentOffset + 1; i < ctx.argumentCount(); i++) { + final String arg = ctx.argument(i); + final Matcher boMatcher = BO.matcher(arg); + if (boMatcher.matches()) { + customBo = Integer.valueOf(boMatcher.group(1)); + } else if ("stable".equalsIgnoreCase(arg) || "stb".equalsIgnoreCase(arg)) { + version = MPVersion.STABLE; + } else if ("lazer".equalsIgnoreCase(arg) || "lzr".equalsIgnoreCase(arg)) { + version = MPVersion.LAZER; + } else { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "未知的参数: " + arg)); + return; + } + } + target = parseRoomTarget(ctx.argument(argumentOffset), version, customBo); } if (target == null) { @@ -134,20 +234,19 @@ private void handleStart(Context ctx, int argumentOffset) { return; } - try (var timing = taskCoordinator.beginRequest(ctx, "Start Multiplayer Room Watch")) { + try (var _ = taskCoordinator.beginRequest(ctx, "Start Multiplayer Room Watch")) { if (!ctx.sendMessage(PendingMessage.ofString("正在尝试启动多人房间监视……")).success()) { ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + - "由于缺少主动消息权限,无法启动监视!权限配置请见:https://docs.seira.top/overview/use.html#extra-permission" + "由于缺少主动消息权限,无法启动监视!权限配置请见[这里](" + PERMISSION + ")~" )); return; } try { RoomWatchView view = watchService.watch( - ctx.groupId(), ctx.senderUserId(), target.version(), target.roomId() + ctx.groupId(), ctx.senderUserId(), target.version(), target.roomId(), null ); ctx.sendReply(PendingMessage.ofMarkdownRaw( - at(ctx) + "已开始监视 `" + formatRoom(view) + "` 。" - + "之后完成的每张图都会自动推送结果。" + at(ctx) + "已开始监视" + formatRoom(view) + "喵。" )); } catch (IllegalArgumentException | IllegalStateException e) { throw new ResolutionException(e.getMessage()); @@ -155,6 +254,9 @@ private void handleStart(Context ctx, int argumentOffset) { } } + private static final Pattern BO = Pattern.compile( + "^bo(\\d+)$" + ); private void handleStop(Context ctx) { if (ctx.argumentCount() == 2 && "all".equalsIgnoreCase(ctx.argument(1))) { int stoppedCount = watchService.stopAll(ctx.groupId()).size(); @@ -184,6 +286,6 @@ private void handleStatus(Context ctx) { : "你当前正在监视" + formatRoom(view) + "。"))); } - record RoomTarget(long roomId, MultiplayerRoomVersion version) { + record RoomTarget(long roomId, MPVersion version, Integer customBo) { } } diff --git a/src/main/java/xyz/zcraft/seira/command/handler/RankGuessCommandHandler.java b/src/main/java/xyz/zcraft/seira/command/handler/RankGuessCommandHandler.java index 2b6cc407..6d9f9bf1 100644 --- a/src/main/java/xyz/zcraft/seira/command/handler/RankGuessCommandHandler.java +++ b/src/main/java/xyz/zcraft/seira/command/handler/RankGuessCommandHandler.java @@ -2,18 +2,16 @@ import org.apache.logging.log4j.LogManager; import org.apache.logging.log4j.Logger; -import xyz.zcraft.seira.api.APIHelper; +import xyz.zcraft.seira.api.OstellaApi; import xyz.zcraft.seira.api.data.RandomScore; import xyz.zcraft.seira.bot.data.MessageReference; import xyz.zcraft.seira.bot.data.PendingMessage; import xyz.zcraft.seira.command.Context; import xyz.zcraft.seira.command.TaskCoordinator; import xyz.zcraft.seira.command.parse.Resolver; -import xyz.zcraft.seira.command.parse.ShortcutTarget; -import xyz.zcraft.seira.command.parse.UserRefResolution; import xyz.zcraft.seira.command.reply.ReplyFactory; import xyz.zcraft.seira.data.SendResult; -import xyz.zcraft.seira.data.UserRef; +import xyz.zcraft.seira.data.UploadedImage; import xyz.zcraft.seira.db.RankGuessRecordStore; import xyz.zcraft.seira.db.UserDataStore; import xyz.zcraft.seira.rankguess.HintUtil; @@ -26,10 +24,12 @@ import java.util.concurrent.Executors; import java.util.concurrent.ScheduledExecutorService; import java.util.concurrent.TimeUnit; +import java.util.function.Function; import java.util.function.Predicate; import java.util.regex.Matcher; import java.util.regex.Pattern; +import static xyz.zcraft.seira.command.reply.ReplyFactory.ExternalUrls.PERMISSION; import static xyz.zcraft.seira.command.reply.ReplyFactory.at; import static xyz.zcraft.seira.command.reply.ReplyFactory.cmd; @@ -44,6 +44,8 @@ public final class RankGuessCommandHandler { private final RankGuessGameService games; private final Resolver resolver; private final Predicate adminAuthorizer; + private final Function avatarUrlGetter; + private final Function imageUploader; private final Pattern BP_PATTERN = Pattern.compile("^bp(\\d+)$"); public RankGuessCommandHandler( @@ -51,13 +53,17 @@ public RankGuessCommandHandler( ReplyFactory replyFactory, RankGuessGameService games, Resolver resolver, - Predicate adminAuthorizer + Predicate adminAuthorizer, + Function avatarUrlGetter, + Function imageUploader ) { this.taskCoordinator = taskCoordinator; this.replyFactory = replyFactory; this.games = games; this.resolver = resolver; this.adminAuthorizer = adminAuthorizer; + this.avatarUrlGetter = avatarUrlGetter; + this.imageUploader = imageUploader; } private static Long parseRank(String argument) { @@ -261,16 +267,9 @@ public void handleRankGuess(Context ctx) { Long rank; if (resolver.looksLikeMention(argument)) { - final UserRefResolution userRefResolution = resolver.resolveUserRefArgument(argument); - - if (userRefResolution.errorMessage() != null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + userRefResolution.errorMessage())); - return; - } - - final UserRef userRef = userRefResolution.userRef(); - - rank = APIHelper.getUserRank(userRef); + String player = resolver.player(argument, ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + rank = OstellaApi.getUserRank(uid); } else { rank = parseRank(argument); if (rank == null) { @@ -331,7 +330,7 @@ private void weight(Context ctx, boolean all, String target) { reply.append(at(ctx)).append("目前%s在本群权重为 `%.2f` (%s)\n".formatted(ref, probability.weight(), factors.isBlank() ? "基础权重" : factors)); reply.append("在本群 `%d` 名玩家中,%s被选中的概率为 `%.3f%%`\n".formatted(totalPlayer, ref, probability.chance() * 100)); - final String randomScoreWeight = APIHelper.getRandomScoreWeight(boundUid, games.generateWeights(ctx.groupId()), all); + final String randomScoreWeight = OstellaApi.getRandomScoreWeight(boundUid, games.generateWeights(ctx.groupId()), all); reply.append("%s的成绩当前抽选概率:\n>".formatted(ref)).append(randomScoreWeight).append("\n"); @@ -408,9 +407,7 @@ private void wishScore(Context ctx, int index) { try { scoreId = Long.parseLong( - APIHelper.lookupScoreId(new ShortcutTarget( - null, new UserRef.ByUid(boundUid), "bp", (long) index, null) - , List.of(), null) + OstellaApi.lookupPlayerScore(boundUid, "bp", index, List.of(), null) ); } catch (Exception e) { LOG.error("Failed to lookup score id", e); @@ -439,7 +436,7 @@ private void start(Context ctx, boolean fromGroup) { } boolean activated = false; - try (var _ = taskCoordinator.beginRequest(ctx, "Rank Guess Render")) { + try (var _ = taskCoordinator.beginRequest(ctx, "Rank Guess Render", false)) { final PendingMessage message = PendingMessage.ofMarkdownRaw(at(ctx) + RandomReply.loading()); final boolean activeMessageEnabled = ctx.sendMessage(message).success(); if (!activeMessageEnabled) { @@ -454,9 +451,9 @@ private void start(Context ctx, boolean fromGroup) { ctx.sendReply(PendingMessage.ofMarkdownRaw("本群没有绑定的用户,无法开始游戏喵")); return; } - randomScore = APIHelper.getRandomScoreFromUsers(uids, games.generateWeights(ctx.groupId())); + randomScore = OstellaApi.getRandomScoreFromUsers(uids, games.generateWeights(ctx.groupId())); } else { - randomScore = APIHelper.getRandomScore(); + randomScore = OstellaApi.getRandomScore(); } Round round = Round.from(randomScore, activeMessageEnabled); @@ -472,12 +469,12 @@ private void start(Context ctx, boolean fromGroup) { content += ",正在渲染回放片段..."; if (!activeMessageEnabled) { - content += "\n\n> 提示: 由于缺少主动消息权限,阶段提示与自动结束已禁用。稍后需要使用 `/rg end` 手动结束。权限配置请见[这里](https://docs.seira.top/overview/use.html#extra-permission)。"; + content += "\n\n> 提示: 由于缺少主动消息权限,阶段提示与自动结束已禁用。稍后需要使用 `/rg end` 手动结束。权限配置请见[这里](" + PERMISSION + ")。"; } ctx.sendReply(PendingMessage.ofMarkdownRaw(content)); - var renderTask = APIHelper.createObscuredReplayRenderTask( + var renderTask = OstellaApi.createObscuredReplayRenderTask( round.scoreId(), taskCoordinator.createVideoUploadRequest(ctx) ); @@ -491,7 +488,7 @@ private void start(Context ctx, boolean fromGroup) { TimeUnit.SECONDS ); - APIHelper.ReplayRenderResult replay = null; + OstellaApi.ReplayRenderResult replay = null; try { replay = taskCoordinator.waitForReplay(renderTask); } catch (Exception e) { @@ -543,7 +540,8 @@ private void start(Context ctx, boolean fromGroup) { hintSource.addAll(round.getNormalHints()); if (fromGroup) { - hintSource.addAll(round.getGroupHints()); + final Optional groupOpenIdByUid = UserDataStore.findGroupOpenIdByUid(ctx.groupId(), round.userId()); + hintSource.addAll(round.getGroupHints(avatarUrlGetter.apply(groupOpenIdByUid.orElse(null)), imageUploader)); } var hints = HintUtil.prepareHints(hintSource, maxHintCount); @@ -683,9 +681,7 @@ private void end(Context ctx, boolean force) { case FINISHED -> replyFactory.rankGuessResultMessage(ctx, result.round(), result.rankType()); }; - if (!ctx.sendReply(message).success()) { - ctx.sendMessage(message); - } + ctx.send(true, message, false); } enum LeaderboardType { diff --git a/src/main/java/xyz/zcraft/seira/command/handler/ReplayCommandHandler.java b/src/main/java/xyz/zcraft/seira/command/handler/ReplayCommandHandler.java index ef781a40..df9877aa 100644 --- a/src/main/java/xyz/zcraft/seira/command/handler/ReplayCommandHandler.java +++ b/src/main/java/xyz/zcraft/seira/command/handler/ReplayCommandHandler.java @@ -1,34 +1,33 @@ package xyz.zcraft.seira.command.handler; import org.jline.utils.Log; -import xyz.zcraft.seira.api.APIHelper; +import xyz.zcraft.seira.api.OstellaApi; import xyz.zcraft.seira.api.data.VideoRenderRecord; import xyz.zcraft.seira.bot.data.PendingMessage; -import xyz.zcraft.seira.command.Context; -import xyz.zcraft.seira.command.ReplayResultStore; -import xyz.zcraft.seira.command.TargetHistory; -import xyz.zcraft.seira.command.TaskCoordinator; +import xyz.zcraft.seira.command.*; import xyz.zcraft.seira.command.parse.Resolver; -import xyz.zcraft.seira.command.parse.RscTarget; +import xyz.zcraft.seira.command.parse.TargetInput; import xyz.zcraft.seira.command.reply.CommandUsage; import xyz.zcraft.seira.command.reply.ReplyFactory; import xyz.zcraft.seira.data.SendResult; import xyz.zcraft.seira.util.TimeDurationParser; +import java.util.List; import java.util.Objects; import java.util.UUID; +import java.util.function.Predicate; -import static xyz.zcraft.seira.command.TargetHistory.Type.BEATMAP; -import static xyz.zcraft.seira.command.TargetHistory.Type.SCORE; import static xyz.zcraft.seira.command.reply.ReplyFactory.at; public final class ReplayCommandHandler { + private final java.util.function.Function accessTokenProvider; private final Resolver resolver; private final TargetHistory history; private final TaskCoordinator taskCoordinator; private final ReplyFactory replyFactory; private final VideoRenderRecord videoRenderRecord; private final ReplayResultStore replayResults; + private final Predicate adminAuthorizer; public ReplayCommandHandler( Resolver resolver, @@ -36,19 +35,29 @@ public ReplayCommandHandler( TaskCoordinator taskCoordinator, ReplyFactory replyFactory, VideoRenderRecord videoRenderRecord, - ReplayResultStore replayResults + ReplayResultStore replayResults, + Predicate adminAuthorizer, + java.util.function.Function accessTokenProvider ) { this.resolver = resolver; + this.accessTokenProvider = accessTokenProvider; this.history = history; this.taskCoordinator = taskCoordinator; this.replyFactory = replyFactory; this.videoRenderRecord = videoRenderRecord; this.replayResults = replayResults; + this.adminAuthorizer = adminAuthorizer; } public void handleR(Context ctx) { - var target = history.parseArguments(ctx, CommandUsage.R, 1, TimeDurationParser::isTimeRange); - if (target == null) return; + var target = ctx.argumentCount() == 0 || TimeDurationParser.isTimeRange(ctx.argument(0)) + ? TargetInput.memory() : TargetInput.read(ctx.args()); + var remembered = history.get(ctx); + if ((target.kind() == TargetInput.Kind.MEMORY && remembered == null) + || ctx.argumentCount() - target.consumedArgs() > 1) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.R)); + return; + } TimeDurationParser.TimeRange range = null; @@ -61,16 +70,42 @@ public void handleR(Context ctx) { } } - try (var _ = taskCoordinator.beginRequest(ctx, "Score Render")) { + try (var _ = taskCoordinator.beginRequest(ctx, "Score Render", false)) { ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "正在获取谱面以及回放文件,请稍作等待喵...")); - var ids = history.resolve(ctx, SCORE, target); - history.remember(ctx, ids); + var previous = target.kind() == TargetInput.Kind.MEMORY ? remembered : null; + Long beatmapId = previous == null ? null : previous.beatmapId(); + Long beatmapsetId = previous == null ? null : previous.beatmapsetId(); + String scoreId = previous == null ? null : previous.scoreId(); + switch (target.kind()) { + case ID, SCORE -> scoreId = target.id(); + case MAP -> beatmapId = Long.parseLong(target.id()); + case SET -> { + beatmapsetId = Long.parseLong(target.id()); + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupBeatmapsetScore(beatmapsetId, target.index(), uid, List.of(), null); + } + case RS, RP, BP -> { + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupPlayerScore(uid, target.scoreList(), target.index(), List.of(), null); + } + case MP -> beatmapId = OstellaApi.lookupMultiplayerBeatmap(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> {} + } + if (scoreId == null) { + if (beatmapId == null) throw new ResolutionException("请指定指令目标谱面喵"); + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupBeatmapScore(beatmapId, uid, List.of(), null); + } var upload = taskCoordinator.createVideoUploadRequest(ctx); - var task = APIHelper.createReplayRenderTask(ids.scoreId(), range, upload); + var task = OstellaApi.createReplayRenderTask(scoreId, range, upload); + history.remember(ctx, beatmapsetId, beatmapId, scoreId); videoRenderRecord.updateRenderTask(ctx.senderUserId(), task.taskId()); ctx.sendReply(replyFactory.replayMessage(ctx, task)); - APIHelper.ReplayRenderResult result; + OstellaApi.ReplayRenderResult result; try { result = taskCoordinator.waitForReplay(task); @@ -98,9 +133,13 @@ public void handleRsc(Context ctx) { return; } - var target = history.parseArguments(ctx, CommandUsage.RSC, Integer.MAX_VALUE, - arg -> arg.startsWith("+") || arg.startsWith("=")); - if (target == null) return; + var target = ctx.argumentCount() == 0 || ctx.argument(0).startsWith("+") || ctx.argument(0).startsWith("=") + ? TargetInput.memory() : TargetInput.read(ctx.args()); + var remembered = history.get(ctx); + if (target.kind() == TargetInput.Kind.MEMORY && remembered == null) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.RSC)); + return; + } String extraUidArg = null; @@ -115,36 +154,79 @@ public void handleRsc(Context ctx) { } } - RscTarget rscTarget = history.isLocalScore(ctx, target) && extraUidArg == null - ? new RscTarget(new String[0], null) - : resolver.resolveRscTarget(ctx.groupId(), extraUidArg); - if (rscTarget.errorMessage() != null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + rscTarget.errorMessage())); - return; + String localScoreId = target.kind() != TargetInput.Kind.MEMORY + ? target.kind() == TargetInput.Kind.SCORE ? target.id() : null + : remembered.scoreId(); + boolean localScore = localScoreId != null && localScoreId.startsWith("loc"); + var participants = new java.util.LinkedHashSet(); + if (!(localScore && extraUidArg == null)) { + if (extraUidArg == null || extraUidArg.trim().startsWith("+")) { + var groupUids = xyz.zcraft.seira.db.UserDataStore.findBoundUidsByGroup(ctx.groupId()); + if (groupUids.isEmpty()) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "本群还没有已绑定的玩家,请先使用 /bind")); + return; + } + groupUids.stream().map(String::valueOf).forEach(participants::add); + } + if (extraUidArg != null) { + String body = extraUidArg.trim().substring(1).trim(); + if (body.isEmpty()) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "追加ID列表不能为空。")); + return; + } + for (String token : body.split(",")) { + if (token.trim().matches("[us]?[0-9]+")) { + participants.add(token.trim()); + } else if (resolver.looksLikeMention(token)) { + String player = resolver.player(token, ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + participants.add("u" + uid); + } else { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "追加ID列表包含非法值。")); + return; + } + } + } } - try (var _ = taskCoordinator.beginRequest(ctx, "Showcase Render")) { + try (var _ = taskCoordinator.beginRequest(ctx, "Showcase Render", false)) { ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "正在获取谱面以及回放文件,请稍作等待喵...")); - var targetType = history.isLocalScore(ctx, target) ? SCORE : BEATMAP; - var resolved = history.resolve(ctx, targetType, target); - history.remember(ctx, resolved); + var previous = target.kind() == TargetInput.Kind.MEMORY ? remembered : null; + Long beatmapId = previous == null ? null : previous.beatmapId(); + Long beatmapsetId = previous == null ? null : previous.beatmapsetId(); + String scoreId = previous == null ? null : previous.scoreId(); + switch (target.kind()) { + case ID, MAP -> beatmapId = Long.parseLong(target.id()); + case SCORE -> scoreId = target.id(); + case SET -> { + beatmapsetId = Long.parseLong(target.id()); + beatmapId = OstellaApi.lookupBeatmapInSet(beatmapsetId, target.index(), accessTokenProvider.apply(ctx.senderUserId())); + } + case RS, RP, BP -> { + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupPlayerScore(uid, target.scoreList(), target.index(), List.of(), null); + } + case MP -> beatmapId = OstellaApi.lookupMultiplayerBeatmap(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> {} + } + if (beatmapId == null && scoreId != null) beatmapId = OstellaApi.getScoreBeatmapId(scoreId); + if (beatmapId == null) throw new ResolutionException("请指定指令目标谱面喵"); var upload = taskCoordinator.createVideoUploadRequest(ctx); - String[] scoreTargets = rscTarget.targets(); - long beatmapId; - if (targetType == SCORE) { + String[] scoreTargets = participants.toArray(String[]::new); + if (localScore) { var ids = new java.util.LinkedHashSet(); - ids.add("s" + resolved.scoreId()); + ids.add("s" + localScoreId); java.util.Collections.addAll(ids, scoreTargets); scoreTargets = ids.toArray(String[]::new); - beatmapId = APIHelper.getScoreBeatmapId(resolved.scoreId()); - } else { - beatmapId = resolved.beatmapId(); + } - var task = APIHelper.createReplayShowcaseTask(beatmapId, scoreTargets, upload); + var task = OstellaApi.createReplayShowcaseTask(beatmapId, scoreTargets, upload); + history.remember(ctx, beatmapsetId, beatmapId, scoreId); videoRenderRecord.updateRenderTask(ctx.senderUserId(), task.taskId()); ctx.sendReply(replyFactory.replayMessage(ctx, task)); - APIHelper.ReplayRenderResult result; + OstellaApi.ReplayRenderResult result; try { result = taskCoordinator.waitForReplay(task); @@ -184,7 +266,7 @@ public void handleRstat(Context ctx) { jobId = ctx.args()[0]; } - APIHelper.ReplayRenderResult replayResult = replayResults.get(jobId); + OstellaApi.ReplayRenderResult replayResult = replayResults.get(jobId); if (replayResult != null) { PendingMessage video = replayResult.qqFile() != null ? PendingMessage.ofUploadedVideo(replayResult.qqFile(), replayResult.videoUrl()) @@ -195,7 +277,7 @@ public void handleRstat(Context ctx) { return; } - ctx.sendReply(replyFactory.replayStatMessage(ctx, jobId, APIHelper.getRenderStat(jobId))); + ctx.sendReply(replyFactory.replayStatMessage(ctx, jobId, OstellaApi.getRenderStat(jobId))); } public void handleRcancel(Context ctx) { @@ -214,7 +296,16 @@ public void handleRcancel(Context ctx) { return; } - var result = APIHelper.cancelReplayRender(jobId); + final String owner = videoRenderRecord.getTaskOwner(jobId); + if (owner != null) { + if (!owner.equalsIgnoreCase(ctx.senderUserId()) + && !adminAuthorizer.test(ctx.senderUserId())) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "你无权取消此任务喵。")); + return; + } + } + + var result = OstellaApi.cancelReplayRender(jobId); String status = Objects.toString(result.getStatus(), "unknown").toLowerCase(); String message = switch (status) { case "canceled" -> "回放渲染已取消。"; @@ -229,5 +320,4 @@ public void handleRcancel(Context ctx) { } ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + message)); } - } diff --git a/src/main/java/xyz/zcraft/seira/command/handler/ScoreCommandHandler.java b/src/main/java/xyz/zcraft/seira/command/handler/ScoreCommandHandler.java index 866b9255..086a34e4 100644 --- a/src/main/java/xyz/zcraft/seira/command/handler/ScoreCommandHandler.java +++ b/src/main/java/xyz/zcraft/seira/command/handler/ScoreCommandHandler.java @@ -1,26 +1,30 @@ package xyz.zcraft.seira.command.handler; -import xyz.zcraft.seira.api.APIHelper; +import xyz.zcraft.seira.api.OstellaApi; +import xyz.zcraft.seira.api.data.MissData; import xyz.zcraft.seira.bot.data.PendingMessage; import xyz.zcraft.seira.command.Context; +import xyz.zcraft.seira.command.ResolutionException; import xyz.zcraft.seira.command.TargetHistory; import xyz.zcraft.seira.command.TaskCoordinator; -import xyz.zcraft.seira.command.parse.*; +import xyz.zcraft.seira.command.parse.Resolver; +import xyz.zcraft.seira.command.parse.ScoreFilterArguments; +import xyz.zcraft.seira.command.parse.TargetInput; import xyz.zcraft.seira.command.reply.CommandUsage; import xyz.zcraft.seira.command.reply.ReplyFactory; -import xyz.zcraft.seira.data.UserRef; import java.util.List; +import java.util.concurrent.ThreadLocalRandom; import java.util.regex.Matcher; import java.util.regex.Pattern; -import static xyz.zcraft.seira.command.TargetHistory.Type.SCORE; import static xyz.zcraft.seira.command.reply.ReplyFactory.at; public final class ScoreCommandHandler { private static final int MAX_SCORE_LIST_COUNT = 200; private static final Pattern SCORE_LIST_RANGE_PATTERN = Pattern.compile("^(\\d+)(?:-(\\d+))?$"); + private final java.util.function.Function accessTokenProvider; private final Resolver resolver; private final TargetHistory history; private final TaskCoordinator taskCoordinator; @@ -30,9 +34,11 @@ public ScoreCommandHandler( Resolver resolver, TargetHistory history, TaskCoordinator taskCoordinator, - ReplyFactory replyFactory + ReplyFactory replyFactory, + java.util.function.Function accessTokenProvider ) { this.resolver = resolver; + this.accessTokenProvider = accessTokenProvider; this.history = history; this.taskCoordinator = taskCoordinator; this.replyFactory = replyFactory; @@ -54,18 +60,29 @@ static TbArguments parseTbArguments(String[] args) { return new TbArguments(days, args.length > targetIndex ? args[targetIndex] : null); } + static ScoreListRange parseScoreListRange(String value) { + Matcher matcher = SCORE_LIST_RANGE_PATTERN.matcher(value); + if (!matcher.matches()) return null; + try { + int first = Integer.parseInt(matcher.group(1)); + String endGroup = matcher.group(2); + int start = endGroup == null ? 1 : first; + int end = endGroup == null ? first : Integer.parseInt(endGroup); + if (start <= 0 || start > end || end > MAX_SCORE_LIST_COUNT) return null; + return new ScoreListRange(start, end); + } catch (NumberFormatException ignored) { + return null; + } + } + public void handleBp(Context ctx) { if (ctx.args().length == 0) { - ShortcutTarget target = resolver.parseTarget("bp1", ctx.senderUserId()); - if (target.isError()) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + target.errorMessage())); - return; - } + String player = resolver.player(null, ctx.senderUserId()); try (var _ = taskCoordinator.beginRequest(ctx, "Score")) { - var ids = history.resolve(ctx, SCORE, new TargetResolution(target, 0)); - history.remember(ctx, ids); - String scoreId = ids.scoreId(); - var response = APIHelper.getScoreResponse(scoreId); + long uid = OstellaApi.resolveUid(player); + String scoreId = OstellaApi.lookupPlayerScore(uid, "bp", 1, List.of(), null); + var response = OstellaApi.getScoreResponse(scoreId); + history.remember(ctx, null, null, scoreId); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.scoreMessage(ctx, response))); } return; @@ -81,31 +98,48 @@ public void handleBp(Context ctx) { return; } - ScoreListRequest request = parseScoreListRequest(ctx, CommandUsage.BP); + var request = parseScoreListRequest(ctx, CommandUsage.BP); if (request == null) return; + var range = request.range(); + var player = request.player(); + var filters = request.filters(); try (var _ = taskCoordinator.beginRequest(ctx, "Best Scores")) { - var response = APIHelper.getBoNResponse( - request.range().end(), - request.range().start(), - request.userRef(), - request.filters() + long uid = OstellaApi.resolveUid(player); + var response = OstellaApi.getBoNResponse( + range.end(), + range.start(), + uid, + filters.filters() ); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.bpMessage(ctx, response))); } } + public void handleRbp(Context ctx) { + if (ctx.argumentCount() > 1) { + ctx.sendReply(at(ctx) + "用法:/rbp [目标]"); + return; + } + + String player = resolver.player(ctx.argumentCount() == 1 ? ctx.argument(0) : null, ctx.senderUserId()); + + try (var _ = taskCoordinator.beginRequest(ctx, "Score")) { + long uid = OstellaApi.resolveUid(player); + String scoreId = OstellaApi.lookupPlayerScore(uid, "bp", ThreadLocalRandom.current().nextInt(200) + 1, List.of(), null); + var response = OstellaApi.getScoreResponse(scoreId); + history.remember(ctx, null, null, scoreId); + ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.scoreMessage(ctx, response))); + } + } + public void handleRs(Context ctx, boolean includeFail) { if (ctx.args().length == 0) { - ShortcutTarget target = resolver.parseTarget(ctx.command() + "1", ctx.senderUserId()); - if (target.isError()) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + target.errorMessage())); - return; - } + String player = resolver.player(null, ctx.senderUserId()); try (var _ = taskCoordinator.beginRequest(ctx, "Score")) { - var ids = history.resolve(ctx, SCORE, new TargetResolution(target, 0)); - history.remember(ctx, ids); - String scoreId = ids.scoreId(); - var response = APIHelper.getScoreResponse(scoreId); + long uid = OstellaApi.resolveUid(player); + String scoreId = OstellaApi.lookupPlayerScore(uid, ctx.command(), 1, List.of(), null); + var response = OstellaApi.getScoreResponse(scoreId); + history.remember(ctx, null, null, scoreId); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.scoreMessage(ctx, response))); } return; @@ -121,15 +155,19 @@ public void handleRs(Context ctx, boolean includeFail) { return; } - ScoreListRequest request = parseScoreListRequest(ctx, CommandUsage.RS); + var request = parseScoreListRequest(ctx, CommandUsage.RS); if (request == null) return; + var range = request.range(); + var player = request.player(); + var filters = request.filters(); try (var _ = taskCoordinator.beginRequest(ctx, "Recent Score")) { - var response = APIHelper.getRecentResponse( - request.range().end(), - request.range().start(), - request.userRef(), + long uid = OstellaApi.resolveUid(player); + var response = OstellaApi.getRecentResponse( + range.end(), + range.start(), + uid, includeFail, - request.filters() + filters.filters() ); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.rsMessage(ctx, response))); } @@ -142,45 +180,20 @@ public void handleTb(Context ctx) { return; } - UserRef userRef; - if (request.target() != null) { - UserRefResolution resolution = resolver.resolveUserRefArgument(request.target()); - if (resolution.errorMessage() != null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + resolution.errorMessage())); - return; - } - if (resolution.userRef() == null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.TB)); - return; - } - userRef = resolution.userRef(); - } else { - Long uid = resolver.resolveBoundUid(ctx.senderUserId()); - if (uid == null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.NO_BIND)); - return; - } - userRef = new UserRef.ByUid(uid); - } - - UserRef target = userRef; + String player = resolver.player(request.target(), ctx.senderUserId()); try (var _ = taskCoordinator.beginRequest(ctx, "Recent Best Scores")) { - var response = APIHelper.getTodayBestResponse(target, request.days()); + long uid = OstellaApi.resolveUid(player); + var response = OstellaApi.getTodayBestResponse(uid, request.days()); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.tbMessage(ctx, response))); } } private void handleFilteredSingleScore(Context ctx, String macroType) { - UserRef targetUser = null; + String targetUser = null; int startIndex = 0; if (resolver.looksLikeMention(ctx.args()[0]) || resolver.looksLikeUid(ctx.args()[0])) { - final UserRefResolution userRefResolution = resolver.resolveUserRefArgument(ctx.args()[0]); - if (userRefResolution.errorMessage() != null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + userRefResolution.errorMessage())); - return; - } - targetUser = userRefResolution.userRef(); + targetUser = resolver.player(ctx.args()[0], ctx.senderUserId()); startIndex = 1; } @@ -190,84 +203,122 @@ private void handleFilteredSingleScore(Context ctx, String macroType) { return; } - if (targetUser == null) { - Long uid = resolver.resolveBoundUid(ctx.senderUserId()); - if (uid == null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.NO_BIND)); - return; - } - targetUser = new UserRef.ByUid(uid); - } + if (targetUser == null) targetUser = resolver.player(null, ctx.senderUserId()); - ShortcutTarget target = new ShortcutTarget(null, targetUser, macroType, 1L, null); try (var _ = taskCoordinator.beginRequest(ctx, "Score")) { - var ids = history.resolve(ctx, SCORE, new TargetResolution(target, 0), filters.filters(), null); - history.remember(ctx, ids); - String scoreId = ids.scoreId(); - var response = APIHelper.getScoreResponse(scoreId); + long uid = OstellaApi.resolveUid(targetUser); + String scoreId = OstellaApi.lookupPlayerScore(uid, macroType, 1, filters.filters(), null); + var response = OstellaApi.getScoreResponse(scoreId); + history.remember(ctx, null, null, scoreId); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.scoreMessage(ctx, response))); } } - private ScoreListRequest parseScoreListRequest(Context ctx, String usage) { - String[] args = ctx.args(); - ScoreListRange range = parseScoreListRange(args[0]); - if (range == null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + usage - + "\n数量或范围必须在 1 到 " + MAX_SCORE_LIST_COUNT + " 之间,且范围起点不能大于终点。")); - return null; + public void handleS(Context ctx) { + String usage = CommandUsage.S; + TargetInput target; + String userOverride = null; + int optionIndex; + boolean playerOnly = ctx.argumentCount() > 0 && resolver.looksLikeMention(ctx.argument(0)) + && (ctx.argumentCount() == 1 || ctx.argument(1).startsWith("+")); + if (playerOnly || ctx.argumentCount() == 0 || ctx.argument(0).startsWith("+")) { + target = TargetInput.memory(); + } else { + target = TargetInput.read(ctx.args()); } - - int nextArg = 1; - UserRef userRef; - if (nextArg < args.length && (resolver.looksLikeMention(args[nextArg]) || resolver.looksLikeUid(args[nextArg]))) { - UserRefResolution resolution = resolver.resolveUserRefArgument(args[nextArg]); - if (resolution.errorMessage() != null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + resolution.errorMessage())); - return null; + optionIndex = target.consumedArgs(); + if (optionIndex < ctx.argumentCount() && !ctx.argument(optionIndex).startsWith("+")) { + userOverride = resolver.player(ctx.argument(optionIndex), ctx.senderUserId()); + optionIndex++; + } + var remembered = history.get(ctx); + if ((target.kind() == TargetInput.Kind.MEMORY && remembered == null) + || ctx.argumentCount() - optionIndex > 1 + || (optionIndex < ctx.argumentCount() && !ctx.argument(optionIndex).startsWith("+"))) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + usage)); + return; + } + String option = optionIndex < ctx.argumentCount() ? ctx.argument(optionIndex) : null; + String mod = option == null ? null : option.substring(1).toUpperCase(java.util.Locale.ROOT); + List filters = List.of(); + if (mod != null) { + var parsed = ScoreFilterArguments.parse(new String[]{"mod=" + mod}, 0); + if (parsed.isError()) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + parsed.errorMessage())); + return; } - if (resolution.userRef() == null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + usage)); - return null; + filters = parsed.filters(); + } + try (var _ = taskCoordinator.beginRequest(ctx, "Score")) { + var previous = target.kind() == TargetInput.Kind.MEMORY ? remembered : null; + Long beatmapId = previous == null ? null : previous.beatmapId(); + Long beatmapsetId = previous == null ? null : previous.beatmapsetId(); + String scoreId = previous == null ? null : previous.scoreId(); + boolean selectedPlayerScore = false; + switch (target.kind()) { + case ID, SCORE -> scoreId = target.id(); + case MAP -> beatmapId = Long.parseLong(target.id()); + case SET -> { + beatmapsetId = Long.parseLong(target.id()); + String player = userOverride == null ? resolver.player(null, ctx.senderUserId()) : userOverride; + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupBeatmapsetScore(beatmapsetId, target.index(), uid, filters, mod); + selectedPlayerScore = true; + } + case RS, RP, BP -> { + String player = userOverride == null ? resolver.player(target.player(), ctx.senderUserId()) : userOverride; + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupPlayerScore(uid, target.scoreList(), target.index(), filters, mod); + selectedPlayerScore = true; + } + case MP -> beatmapId = OstellaApi.lookupMultiplayerBeatmap(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> {} } - userRef = resolution.userRef(); - nextArg++; - } else { - Long uid = resolver.resolveBoundUid(ctx.senderUserId()); - if (uid == null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.NO_BIND)); - return null; + if ((userOverride != null || mod != null) && scoreId != null && !selectedPlayerScore) { + if (beatmapId == null) beatmapId = OstellaApi.getScoreBeatmapId(scoreId); + scoreId = null; } - userRef = new UserRef.ByUid(uid); - } - - ScoreFilterArguments.ParseResult filters = ScoreFilterArguments.parse(args, nextArg); - if (filters.isError()) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + filters.errorMessage() + "\n" + CommandUsage.SCORE_FILTERS)); - return null; + if (scoreId == null) { + if (beatmapId == null) throw new ResolutionException("请指定指令目标谱面喵"); + String player = userOverride == null ? resolver.player(null, ctx.senderUserId()) : userOverride; + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupBeatmapScore(beatmapId, uid, filters, mod); + } + var response = OstellaApi.getScoreResponse(scoreId); + history.remember(ctx, beatmapsetId, beatmapId, scoreId); + ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.scoreMessage(ctx, response))); + } catch (Exception e) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + TaskCoordinator.resolveErrorMessage(e))); + org.apache.logging.log4j.LogManager.getLogger(ScoreCommandHandler.class) + .error("Failed to execute /{}", ctx.command(), e); } - return new ScoreListRequest(range, userRef, filters.filters()); } - static ScoreListRange parseScoreListRange(String value) { - Matcher matcher = SCORE_LIST_RANGE_PATTERN.matcher(value); - if (!matcher.matches()) return null; - try { - int first = Integer.parseInt(matcher.group(1)); - String endGroup = matcher.group(2); - int start = endGroup == null ? 1 : first; - int end = endGroup == null ? first : Integer.parseInt(endGroup); - if (start <= 0 || start > end || end > MAX_SCORE_LIST_COUNT) return null; - return new ScoreListRange(start, end); - } catch (NumberFormatException ignored) { - return null; + public void handleSm(Context ctx) { + String usage = CommandUsage.SM; + TargetInput target; + String userOverride = null; + int optionIndex; + boolean playerOnly = ctx.argumentCount() > 0 && resolver.looksLikeMention(ctx.argument(0)) + && (ctx.argumentCount() == 1 || ctx.argument(1).startsWith("+")); + if (playerOnly || ctx.argumentCount() == 0 || ctx.argument(0).startsWith("+")) { + target = TargetInput.memory(); + } else { + target = TargetInput.read(ctx.args()); } - } - - public void handleS(Context ctx) { - var target = history.parseScoreArguments(ctx, CommandUsage.S, 1, arg -> arg.startsWith("+")); - if (target == null) return; - String option = target.nextArgument(ctx); + optionIndex = target.consumedArgs(); + if (optionIndex < ctx.argumentCount() && !ctx.argument(optionIndex).startsWith("+")) { + userOverride = resolver.player(ctx.argument(optionIndex), ctx.senderUserId()); + optionIndex++; + } + var remembered = history.get(ctx); + if ((target.kind() == TargetInput.Kind.MEMORY && remembered == null) + || ctx.argumentCount() - optionIndex > 1 + || (optionIndex < ctx.argumentCount() && !ctx.argument(optionIndex).startsWith("+"))) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + usage)); + return; + } + String option = optionIndex < ctx.argumentCount() ? ctx.argument(optionIndex) : null; String mod = option == null ? null : option.substring(1).toUpperCase(java.util.Locale.ROOT); List filters = List.of(); if (mod != null) { @@ -279,65 +330,145 @@ public void handleS(Context ctx) { filters = parsed.filters(); } try (var _ = taskCoordinator.beginRequest(ctx, "Score")) { - var ids = history.resolve(ctx, SCORE, target, filters, mod); - history.remember(ctx, ids); - var response = APIHelper.getScoreResponse(ids.scoreId()); + var previous = target.kind() == TargetInput.Kind.MEMORY ? remembered : null; + Long beatmapId = previous == null ? null : previous.beatmapId(); + Long beatmapsetId = previous == null ? null : previous.beatmapsetId(); + String scoreId = previous == null ? null : previous.scoreId(); + switch (target.kind()) { + case ID, MAP -> beatmapId = Long.parseLong(target.id()); + case SCORE -> scoreId = target.id(); + case SET -> { + beatmapsetId = Long.parseLong(target.id()); + beatmapId = OstellaApi.lookupBeatmapInSet(beatmapsetId, target.index(), accessTokenProvider.apply(ctx.senderUserId())); + } + case RS, RP, BP -> { + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupPlayerScore(uid, target.scoreList(), target.index(), List.of(), null); + } + case MP -> beatmapId = OstellaApi.lookupMultiplayerBeatmap(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> {} + } + if (beatmapId == null && scoreId != null) beatmapId = OstellaApi.getScoreBeatmapId(scoreId); + if (beatmapId == null) throw new ResolutionException("请指定指令目标谱面喵"); + String player = userOverride == null ? resolver.player(null, ctx.senderUserId()) : userOverride; + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupBeatmapScore(beatmapId, uid, filters, mod); + var response = OstellaApi.getScoreResponse(scoreId); + history.remember(ctx, beatmapsetId, beatmapId, scoreId); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.scoreMessage(ctx, response))); } catch (Exception e) { ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + TaskCoordinator.resolveErrorMessage(e))); -// ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + TaskCoordinator.resolveErrorMessage(e) -// + "\n> Tips: 若要查找指定谱面上的成绩,请使用 /s __m__`bid`")); org.apache.logging.log4j.LogManager.getLogger(ScoreCommandHandler.class) - .error("Failed to execute /s", e); + .error("Failed to execute /{}", ctx.command(), e); } } public void handleSa(Context ctx) { - var target = history.parseArguments(ctx, CommandUsage.SA, 0); - if (target == null) return; + var target = ctx.argumentCount() == 0 + ? TargetInput.memory() : TargetInput.read(ctx.args()); + var remembered = history.get(ctx); + if ((target.kind() == TargetInput.Kind.MEMORY && remembered == null) + || ctx.argumentCount() - target.consumedArgs() > 0) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.SA)); + return; + } try (var _ = taskCoordinator.beginRequest(ctx, "Score Analysis")) { - var ids = history.resolve(ctx, SCORE, target); - history.remember(ctx, ids); - String scoreId = ids.scoreId(); - var response = APIHelper.getScoreAnalyzeResponse(scoreId); + var previous = target.kind() == TargetInput.Kind.MEMORY ? remembered : null; + Long beatmapId = previous == null ? null : previous.beatmapId(); + Long beatmapsetId = previous == null ? null : previous.beatmapsetId(); + String scoreId = previous == null ? null : previous.scoreId(); + switch (target.kind()) { + case ID, SCORE -> scoreId = target.id(); + case MAP -> beatmapId = Long.parseLong(target.id()); + case SET -> { + beatmapsetId = Long.parseLong(target.id()); + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupBeatmapsetScore(beatmapsetId, target.index(), uid, List.of(), null); + } + case RS, RP, BP -> { + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupPlayerScore(uid, target.scoreList(), target.index(), List.of(), null); + } + case MP -> beatmapId = OstellaApi.lookupMultiplayerBeatmap(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> {} + } + if (scoreId == null) { + if (beatmapId == null) throw new ResolutionException("请指定指令目标谱面喵"); + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupBeatmapScore(beatmapId, uid, List.of(), null); + } + var response = OstellaApi.getScoreAnalyzeResponse(scoreId); + history.remember(ctx, beatmapsetId, beatmapId, scoreId); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.scoreAnalyzeMessage(ctx, response))); } } public void handleMa(Context ctx) { - var target = history.parseArguments(ctx, CommandUsage.MA, 1, arg -> arg.startsWith("#")); - if (target == null) return; - String indexArgument = target.nextArgument(ctx); - if (indexArgument != null) { - Integer index = parseMissIndex(indexArgument, target.consumedArgs() == 0); - if (index == null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.MA)); - return; - } - try (var _ = taskCoordinator.beginRequest(ctx, "Miss Visualize")) { - var ids = history.resolve(ctx, SCORE, target); - history.remember(ctx, ids); - String scoreId = ids.scoreId(); - var misses = APIHelper.getScoreMissesResponse(scoreId).getContent(); - if (misses.isEmpty()) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "本成绩没有Miss喵~")); - return; + var target = ctx.argumentCount() == 0 || ctx.argument(0).startsWith("#") + ? TargetInput.memory() : TargetInput.read(ctx.args()); + var remembered = history.get(ctx); + if ((target.kind() == TargetInput.Kind.MEMORY && remembered == null) + || ctx.argumentCount() - target.consumedArgs() > 1) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.MA)); + return; + } + String indexArgument = (ctx.argumentCount() > target.consumedArgs() ? ctx.argument(target.consumedArgs()) : null); + Integer index = indexArgument == null ? null : parseMissIndex(indexArgument, target.kind() == TargetInput.Kind.MEMORY); + if (indexArgument != null && index == null) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.MA)); + return; + } + try (var _ = taskCoordinator.beginRequest(ctx, "Score Misses")) { + var previous = target.kind() == TargetInput.Kind.MEMORY ? remembered : null; + Long beatmapId = previous == null ? null : previous.beatmapId(); + Long beatmapsetId = previous == null ? null : previous.beatmapsetId(); + String scoreId = previous == null ? null : previous.scoreId(); + switch (target.kind()) { + case ID, SCORE -> scoreId = target.id(); + case MAP -> beatmapId = Long.parseLong(target.id()); + case SET -> { + beatmapsetId = Long.parseLong(target.id()); + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupBeatmapsetScore(beatmapsetId, target.index(), uid, List.of(), null); } - if (index <= 0 || index > misses.size()) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "Miss序号不在范围内喵(1~" + misses.size() + ")")); - return; + case RS, RP, BP -> { + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupPlayerScore(uid, target.scoreList(), target.index(), List.of(), null); } - var response = APIHelper.getMissVisualizeResponse(scoreId, index); - ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.missImageMessage(ctx, scoreId, index, misses.size()))); + case MP -> beatmapId = OstellaApi.lookupMultiplayerBeatmap(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> {} } - return; - } - try (var _ = taskCoordinator.beginRequest(ctx, "Get Score Misses")) { - var ids = history.resolve(ctx, SCORE, target); - history.remember(ctx, ids); - String scoreId = ids.scoreId(); - var response = APIHelper.getScoreMissesResponse(scoreId); - ctx.sendReply(replyFactory.scoreMissesMessage(ctx, response)); + if (scoreId == null) { + if (beatmapId == null) throw new ResolutionException("请指定指令目标谱面喵"); + String player = resolver.player(target.player(), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + scoreId = OstellaApi.lookupBeatmapScore(beatmapId, uid, List.of(), null); + } + var response = OstellaApi.getScoreMissesResponse(scoreId); + history.remember(ctx, beatmapsetId, beatmapId, scoreId); + List misses = response.getContent(); + if (index == null && misses.size() != 1) { + ctx.sendReply(replyFactory.scoreMissesMessage(ctx, response)); + return; + } + if (misses.isEmpty()) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "本成绩没有Miss喵~")); + return; + } + int selectedIndex = index == null ? 1 : index; + if (selectedIndex > misses.size()) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "Miss序号不在范围内喵(1~" + misses.size() + ")")); + return; + } + var image = OstellaApi.getMissVisualizeResponse(scoreId, selectedIndex); + ctx.sendReply(taskCoordinator.imageMessage(image, + replyFactory.missImageMessage(ctx, scoreId, selectedIndex, misses.size()))); } } @@ -351,13 +482,37 @@ private Integer parseMissIndex(String arg, boolean requirePrefix) { return resolver.parsePositiveInt(value); } - record TbArguments(int days, String target) { - } + private ScoreListRequest parseScoreListRequest(Context ctx, String usage) { + String[] args = ctx.args(); + ScoreListRange range = parseScoreListRange(args[0]); + if (range == null) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + usage + + "\n数量或范围必须在 1 到 " + MAX_SCORE_LIST_COUNT + " 之间,且范围起点不能大于终点。")); + return null; + } - record ScoreListRange(int start, int end) { + int nextArg = 1; + String player; + if (nextArg < args.length && (resolver.looksLikeMention(args[nextArg]) || resolver.looksLikeUid(args[nextArg]))) { + player = resolver.player(args[nextArg], ctx.senderUserId()); + nextArg++; + } else { + player = resolver.player(null, ctx.senderUserId()); + } + + ScoreFilterArguments.ParseResult filters = ScoreFilterArguments.parse(args, nextArg); + if (filters.isError()) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + filters.errorMessage() + "\n" + CommandUsage.SCORE_FILTERS)); + return null; + } + return new ScoreListRequest(range, player, filters); } - private record ScoreListRequest(ScoreListRange range, UserRef userRef, java.util.List filters) { + private record ScoreListRequest(ScoreListRange range, String player, ScoreFilterArguments.ParseResult filters) {} + + record TbArguments(int days, String target) { } + record ScoreListRange(int start, int end) { + } } diff --git a/src/main/java/xyz/zcraft/seira/command/handler/SocialCommandHandler.java b/src/main/java/xyz/zcraft/seira/command/handler/SocialCommandHandler.java index 1288a63b..a7c10ade 100644 --- a/src/main/java/xyz/zcraft/seira/command/handler/SocialCommandHandler.java +++ b/src/main/java/xyz/zcraft/seira/command/handler/SocialCommandHandler.java @@ -2,20 +2,18 @@ import xyz.zcraft.osu.model.User; import xyz.zcraft.osu.model.UserExtended; -import xyz.zcraft.seira.api.APIHelper; +import xyz.zcraft.seira.api.OstellaApi; import xyz.zcraft.seira.api.data.FriendEntry; import xyz.zcraft.seira.api.data.OsuToken; import xyz.zcraft.seira.api.data.Response; import xyz.zcraft.seira.bot.data.PendingMessage; import xyz.zcraft.seira.command.Context; +import xyz.zcraft.seira.command.ResolutionException; import xyz.zcraft.seira.command.TaskCoordinator; import xyz.zcraft.seira.command.parse.Resolver; -import xyz.zcraft.seira.command.parse.ShortcutTarget; -import xyz.zcraft.seira.command.parse.TargetResolution; -import xyz.zcraft.seira.command.parse.UserRefResolution; +import xyz.zcraft.seira.command.parse.TargetInput; import xyz.zcraft.seira.command.reply.CommandUsage; import xyz.zcraft.seira.command.reply.ReplyFactory; -import xyz.zcraft.seira.data.UserRef; import xyz.zcraft.seira.db.UserDataStore; import xyz.zcraft.seira.util.OsuAuthHelper; @@ -32,19 +30,22 @@ public final class SocialCommandHandler { private final TaskCoordinator taskCoordinator; private final ReplyFactory replyFactory; private final Function accessTokenProvider; + private final Function avatarProvider; public SocialCommandHandler( Resolver resolver, OsuAuthHelper authHelper, TaskCoordinator taskCoordinator, ReplyFactory replyFactory, - Function accessTokenProvider + Function accessTokenProvider, + Function avatarProvider ) { this.resolver = resolver; this.authHelper = authHelper; this.taskCoordinator = taskCoordinator; this.replyFactory = replyFactory; this.accessTokenProvider = accessTokenProvider; + this.avatarProvider = avatarProvider; } public void handleMp(Context ctx) { @@ -61,7 +62,7 @@ public void handleMp(Context ctx) { } try (var _ = taskCoordinator.beginRequest(ctx, "Multiplayer Room")) { - var response = APIHelper.getMultiplayerRoom(token.accessToken()); + var response = OstellaApi.getMultiplayerRoom(token.accessToken()); ctx.sendReply(replyFactory.mpMessage(ctx, response)); } } @@ -90,36 +91,52 @@ public void handleFriendStatus(Context ctx) { if (ctx.argumentCount() == 0) { ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "用法:/mu @someone\n> 注: 读取@需要开启权限。")); - } - - final String s = resolver.extractMentionedUserId(ctx.argument(0)); - final Long targetId = resolver.resolveBoundUid(s); - - if (targetId == null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "对方还未绑定喵")); return; } + String player = resolver.player(ctx.argument(0), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + final String mention = UserDataStore.findGroupOpenIdByUid(ctx.groupId(), uid) + .map(ReplyFactory::at) + .orElse(""); + ctx.sendReply(PendingMessage.ofMarkdownRaw(mention + ": [%d(点击打开)](%s)".formatted(uid, "https://osu.ppy.sh/users/" + uid))); + final UserExtended targetUser = OstellaApi.getUserRaw(uid); + final String targetOsuAvatar = targetUser.getAvatarUrl(); + boolean selfFollowed; final AtomicReference targetFollowed = new AtomicReference<>(); - final OsuToken self = authHelper.updateTokenAndGet(ctx.senderUserId()); - final List selfFollowedList = APIHelper.getFollowed(self.accessToken()).getContent(); + final OsuToken selfToken = authHelper.updateTokenAndGet(ctx.senderUserId()); + final var selfUser = OstellaApi.getSelf(selfToken.accessToken()).getContent(); + final String selfOsuAvatar = selfUser.getAvatarUrl(); + + if (targetUser.getId() == selfUser.getId()) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + " 和 `" + targetUser.getUsername() + "` 是一个人喵。")); + return; + } + + final List selfFollowedList = OstellaApi.getFollowed(selfToken.accessToken()).getContent(); updateFriends(selfId, selfFollowedList); final Set users = new HashSet<>(selfFollowedList.stream().map(FriendEntry::user).toList()); - selfFollowed = selfFollowedList.stream().anyMatch(e -> e.user().getId() == targetId); + users.add(selfUser); + users.add(targetUser); - selfFollowedList.stream().filter(e -> e.user().getId() == targetId).findFirst().ifPresentOrElse( - e -> targetFollowed.set(e.mutual()), () -> { - } - ); + selfFollowed = selfFollowedList.stream() + .anyMatch(e -> e.user().getId() == targetUser.getId()); + + selfFollowedList.stream() + .filter(e -> e.user().getId() == targetUser.getId()) + .findFirst() + .ifPresent(e -> targetFollowed.set(e.mutual())); - if (targetFollowed.get() == null) { + final var targetOpenId = UserDataStore.findGroupOpenIdByUid(ctx.groupId(), targetUser.getId()).orElse(null); + + if (targetFollowed.get() == null && targetOpenId != null) { final List targetFollowedList; - final OsuToken target = authHelper.updateTokenAndGet(s); + final OsuToken target = authHelper.updateTokenAndGet(targetOpenId); if (target != null) { - targetFollowedList = APIHelper.getFollowed(target.accessToken()).getContent(); + targetFollowedList = OstellaApi.getFollowed(target.accessToken()).getContent(); targetFollowed.set(targetFollowedList.stream().anyMatch(e -> e.user().getId() == selfId)); users.addAll(targetFollowedList.stream().map(FriendEntry::user).toList()); } @@ -128,8 +145,14 @@ public void handleFriendStatus(Context ctx) { UserDataStore.storeUserInfo(users); ctx.sendReply(replyFactory.friendStatusMessage( - ctx.senderUserId(), selfId, UserDataStore.findUsername(selfId).orElse("未知"), - s, targetId, UserDataStore.findUsername(targetId).orElse("未知"), + ctx.senderUserId(), selfId, selfOsuAvatar, + UserDataStore.findUsername(selfId).orElse("未知"), + avatarProvider.apply(ctx.senderUserId()), + + targetOpenId, targetUser.getId(), targetOsuAvatar, + targetUser.getUsername(), + avatarProvider.apply(targetOpenId), + selfFollowed, targetFollowed.get() ) ); @@ -145,8 +168,8 @@ public void handleFriendList(Context ctx, boolean all) { } try (var _ = taskCoordinator.beginRequest(ctx, "Friend List")) { - final Response self = APIHelper.getSelf(token.accessToken()); - final Response> response = APIHelper.getFollowed(token.accessToken()); + final Response self = OstellaApi.getSelf(token.accessToken()); + final Response> response = OstellaApi.getFollowed(token.accessToken()); final List friendEntries = response.getContent(); final Predicate filter; @@ -246,7 +269,7 @@ public void handleLb(Context ctx) { } try (var _ = taskCoordinator.beginRequest(ctx, "Leaderboard")) { - var response = APIHelper.getLeaderboardResponse(groupBoundUids); + var response = OstellaApi.getLeaderboardResponse(groupBoundUids); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.lbMessage(ctx, response))); } return; @@ -258,70 +281,59 @@ public void handleLb(Context ctx) { } try (var _ = taskCoordinator.beginRequest(ctx, "Leaderboard")) { - var response = APIHelper.getLeaderboardResponse(List.of(uid)); + var response = OstellaApi.getLeaderboardResponse(List.of(uid)); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.lbMessage(ctx, response))); } } else if (ctx.args().length == 1 || ctx.args().length == 2) { - TargetResolution targetResolution = resolver.resolveTargetWithOptionalMention(ctx.args(), ctx.senderUserId()); - ShortcutTarget target = targetResolution.target(); - if (target.isError()) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + target.errorMessage())); - return; - } - - int remainingArgs = ctx.args().length - targetResolution.consumedArgs(); - if (remainingArgs == 0) { - if (ctx.groupId() != null && !ctx.groupId().isBlank()) { - List groupBoundUids = UserDataStore.findBoundUidsByGroup(ctx.groupId()); - if (groupBoundUids.isEmpty()) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "本群还没有已绑定的玩家,请先使用 /bind")); + var target = TargetInput.read(ctx.args()); + int remainingArgs = ctx.argumentCount() - target.consumedArgs(); + List uids = new LinkedList<>(); + if (remainingArgs == 1) { + String[] uidTokens = ctx.argument(target.consumedArgs()).split(","); + if (uidTokens.length == 0) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "玩家ID列表不能为空。")); + return; + } + for (String token : uidTokens) { + Long uid = resolver.parsePositiveLong(token.trim()); + if (uid == null) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "玩家ID列表包含非法值。")); return; } - try (var _ = taskCoordinator.beginRequest(ctx, "Map Leaderboard")) { - long beatmapId = APIHelper.lookupBeatmap(target, accessTokenProvider.apply(ctx.senderUserId())); - var response = APIHelper.getGroupLeaderboardResponse(beatmapId, groupBoundUids); - ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.lbMessage(ctx, response))); - } + uids.add(uid); + } + } else if (remainingArgs != 0) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "用法:/lb <谱面ID或快捷查询> [玩家ID列表(逗号分隔)]")); + return; + } else if (ctx.inGroup()) { + uids.addAll(UserDataStore.findBoundUidsByGroup(ctx.groupId())); + if (uids.isEmpty()) { + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "本群还没有已绑定的玩家,请先使用 /bind")); return; } + } else { Long uid = resolver.resolveBoundUid(ctx.senderUserId()); if (uid == null) { ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.NO_BIND)); return; } - - try (var _ = taskCoordinator.beginRequest(ctx, "Map Leaderboard")) { - long beatmapId = APIHelper.lookupBeatmap(target, accessTokenProvider.apply(ctx.senderUserId())); - var response = APIHelper.getGroupLeaderboardResponse(beatmapId, List.of(uid)); - ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.lbMessage(ctx, response))); - } - return; - } - - if (remainingArgs != 1) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "用法:/lb <谱面ID或快捷查询> [玩家ID列表(逗号分隔)]")); - return; - } - - String[] uidTokens = ctx.args()[targetResolution.consumedArgs()].split(","); - if (uidTokens.length == 0) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "玩家ID列表不能为空。用法:/lb <谱面ID或快捷查询> [玩家ID列表(逗号分隔)]")); - return; - } - - List uids = new LinkedList<>(); - for (String uidToken : uidTokens) { - Long uid = resolver.parsePositiveLong(uidToken.trim()); - if (uid == null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "玩家ID列表包含非法值。用法:/lb <谱面ID或快捷查询> [玩家ID列表(逗号分隔)]")); - return; - } uids.add(uid); } try (var _ = taskCoordinator.beginRequest(ctx, "Map Leaderboard")) { - long beatmapId = APIHelper.lookupBeatmap(target, accessTokenProvider.apply(ctx.senderUserId())); - var response = APIHelper.getGroupLeaderboardResponse(beatmapId, uids); + long beatmapId = switch (target.kind()) { + case ID, MAP -> Long.parseLong(target.id()); + case SCORE -> OstellaApi.getScoreBeatmapId(target.id()); + case SET -> OstellaApi.lookupBeatmapInSet(Long.parseLong(target.id()), target.index(), + accessTokenProvider.apply(ctx.senderUserId())); + case RS, RP, BP -> { + long uid = OstellaApi.resolveUid(resolver.player(target.player(), ctx.senderUserId())); + yield OstellaApi.lookupPlayerScoreBeatmap(uid, target.scoreList(), target.index(), accessTokenProvider.apply(ctx.senderUserId())); + } + case MP -> OstellaApi.lookupMultiplayerBeatmap(accessTokenProvider.apply(ctx.senderUserId())); + case MEMORY -> throw new ResolutionException("请指定指令目标谱面喵"); + }; + var response = OstellaApi.getGroupLeaderboardResponse(beatmapId, uids); ctx.sendReply(taskCoordinator.imageMessage(response, replyFactory.lbMessage(ctx, response))); } } else { @@ -330,28 +342,13 @@ public void handleLb(Context ctx) { } public void handleSup(Context ctx) { - final UserRef target; - - if (ctx.argumentCount() == 0) { - Long targetId = resolver.resolveBoundUid(ctx.senderUserId()); - if (targetId == null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.NO_BIND)); - return; - } - target = new UserRef.ByUid(targetId); - } else if (ctx.argumentCount() == 1) { - final UserRefResolution res = resolver.resolveUserRefArgument(ctx.argument(0)); - if (res.errorMessage() != null) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + res.errorMessage())); - return; - } - target = res.userRef(); - } else { + if (ctx.argumentCount() > 1) { ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + CommandUsage.SUP)); return; } - - final UserExtended user = APIHelper.getUserRaw(target); + String player = resolver.player(ctx.argumentCount() == 0 ? null : ctx.argument(0), ctx.senderUserId()); + long uid = OstellaApi.resolveUid(player); + final UserExtended user = OstellaApi.getUserRaw(uid); final String openId = UserDataStore.findGroupOpenIdByUid(ctx.groupId(), user.getId()).orElse(null); ctx.sendReply(replyFactory.supMessage(ctx, user.getUsername(), openId, user.isSupporter(), user.getHasSupported(), user.getSupportLevel())); diff --git a/src/main/java/xyz/zcraft/seira/command/handler/SpecificScoreWatchCommandHandler.java b/src/main/java/xyz/zcraft/seira/command/handler/SpecificScoreWatchCommandHandler.java index cce656b4..2f8e3f3f 100644 --- a/src/main/java/xyz/zcraft/seira/command/handler/SpecificScoreWatchCommandHandler.java +++ b/src/main/java/xyz/zcraft/seira/command/handler/SpecificScoreWatchCommandHandler.java @@ -11,6 +11,7 @@ import java.util.Objects; import java.util.Set; +import static xyz.zcraft.seira.command.reply.ReplyFactory.ExternalUrls.PERMISSION; import static xyz.zcraft.seira.command.reply.ReplyFactory.at; public final class SpecificScoreWatchCommandHandler { @@ -74,10 +75,10 @@ private void handleStart(Context ctx) { return; } - try (var timing = taskCoordinator.beginRequest(ctx, "Start Specific Score Watch")) { + try (var _ = taskCoordinator.beginRequest(ctx, "Start Specific Score Watch")) { if (!ctx.sendMessage(PendingMessage.ofString("正在尝试启动指定谱面成绩监视……")).success()) { ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + - "由于缺少主动消息权限,无法启动监视!权限配置请见:https://docs.seira.top/overview/use.html#extra-permission" + "由于缺少主动消息权限,无法启动监视!权限配置请见:" + PERMISSION )); return; } diff --git a/src/main/java/xyz/zcraft/seira/command/handler/WatchCommandHandler.java b/src/main/java/xyz/zcraft/seira/command/handler/WatchCommandHandler.java index 4360163a..4a4ef459 100644 --- a/src/main/java/xyz/zcraft/seira/command/handler/WatchCommandHandler.java +++ b/src/main/java/xyz/zcraft/seira/command/handler/WatchCommandHandler.java @@ -1,7 +1,7 @@ package xyz.zcraft.seira.command.handler; import xyz.zcraft.osu.model.User; -import xyz.zcraft.seira.api.APIHelper; +import xyz.zcraft.seira.api.OstellaApi; import xyz.zcraft.seira.bot.data.PendingMessage; import xyz.zcraft.seira.command.Context; import xyz.zcraft.seira.command.ResolutionException; @@ -16,35 +16,26 @@ import java.util.List; import java.util.Locale; import java.util.Objects; -import java.util.function.BiFunction; import java.util.function.Predicate; +import static xyz.zcraft.seira.command.reply.ReplyFactory.ExternalUrls.PERMISSION; import static xyz.zcraft.seira.command.reply.ReplyFactory.at; public final class WatchCommandHandler { private static final int DEFAULT_DURATION_MINUTES = 10; private static final int MAX_DURATION_MINUTES = 120; - private static final String USAGE = "用法:/watch add <玩家ID/用户名/@用户> [分钟];/watch del [玩家ID/用户名/@用户];/watch list"; + private static final String USAGE = "用法:/watch add <玩家ID/用户名/@用户> [分钟,1-120];/watch del [玩家ID/用户名/@用户];/watch list"; private final Resolver resolver; private final TaskCoordinator taskCoordinator; private final ScoreWatchService watchService; private final Predicate adminAuthorizer; - private final BiFunction targetResolver; public WatchCommandHandler(Resolver resolver, TaskCoordinator taskCoordinator, ScoreWatchService watchService, Predicate adminAuthorizer) { this.resolver = Objects.requireNonNull(resolver); this.taskCoordinator = Objects.requireNonNull(taskCoordinator); this.watchService = watchService; this.adminAuthorizer = Objects.requireNonNull(adminAuthorizer); - this.targetResolver = this::resolveTarget; - } - - private static User findUserById(long userId) { - return APIHelper.getUsers(List.of(userId)).stream() - .filter(user -> user.getId() == userId) - .findFirst() - .orElseThrow(() -> new ResolutionException("未找到指定的玩家。")); } private static PendingMessage removedMessage(WatchView removed) { @@ -134,13 +125,28 @@ private void handleAdd(Context ctx) { } String targetArgument = ctx.argument(1); - try (var timing = taskCoordinator.beginRequest(ctx, "Add Score Watch")) { - WatchTarget target = targetResolver.apply(ctx.groupId(), targetArgument); + try (var _ = taskCoordinator.beginRequest(ctx, "Add Score Watch")) { + String mentionedUser = resolver.extractMentionedUserId(targetArgument); + WatchTarget target; + if (mentionedUser != null) { + if (!UserDataStore.isGroupMember(ctx.groupId(), mentionedUser)) { + throw new ResolutionException("指定的用户不在当前群聊中。"); + } + Long userId = UserDataStore.findBoundUid(mentionedUser); + if (userId == null) throw new ResolutionException("被@的用户还没有绑定玩家ID,请先让对方使用 /bind。"); + User user = OstellaApi.getUsers(List.of(userId)).stream() + .filter(candidate -> candidate.getId() == userId) + .findFirst().orElseThrow(() -> new ResolutionException("未找到指定的玩家。")); + UserDataStore.storeUserInfo(user.getId(), user.getUsername()); + target = new WatchTarget(user.getId(), user.getUsername(), mentionedUser); + } else { + target = lookupGroupPlayer(ctx.groupId(), targetArgument); + } final boolean b = ctx.sendMessage(PendingMessage.ofMarkdownRaw( at(ctx) + "正在尝试添加监视..." )).success(); if (!b) { - ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "由于缺少主动消息权限,无法添加监视!权限配置请见[这里](https://docs.seira.top/overview/use.html#extra-permission)~")); + ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "由于缺少主动消息权限,无法添加监视!权限配置请见[这里](" + PERMISSION + ")~")); return; } @@ -178,12 +184,24 @@ private void handleDelete(Context ctx) { return; } - try (var timing = taskCoordinator.beginRequest(ctx, "Delete Score Watch")) { - WatchTarget target = resolveTarget(ctx.groupId(), targetArgument); + try (var _ = taskCoordinator.beginRequest(ctx, "Delete Score Watch")) { + WatchTarget target = lookupGroupPlayer(ctx.groupId(), targetArgument); ctx.sendReply(removedMessage(watchService.remove(ctx.groupId(), target.userId()))); } } + private WatchTarget lookupGroupPlayer(String groupId, String argument) { + Long uid = resolver.parsePositiveLong(argument); + User user = uid == null ? OstellaApi.lookupUser(argument).getContent() + : OstellaApi.getUsers(List.of(uid)).stream() + .filter(candidate -> candidate.getId() == uid) + .findFirst().orElseThrow(() -> new ResolutionException("未找到指定的玩家。")); + String openId = UserDataStore.findGroupOpenIdByUid(groupId, user.getId()) + .orElseThrow(() -> new ResolutionException("指定的玩家不在当前群聊中,或尚未在本群完成绑定。")); + UserDataStore.storeUserInfo(user.getId(), user.getUsername()); + return new WatchTarget(user.getId(), user.getUsername(), openId); + } + private void handleList(Context ctx) { if (ctx.argumentCount() != 1) { usage(ctx); @@ -206,28 +224,4 @@ private void handleList(Context ctx) { ctx.sendReply(PendingMessage.ofMarkdownRaw(content.toString().trim())); } - private WatchTarget resolveTarget(String groupId, String argument) { - String mentionedOpenId = resolver.extractMentionedUserId(argument); - if (mentionedOpenId != null) { - if (!UserDataStore.isGroupMember(groupId, mentionedOpenId)) { - throw new ResolutionException("指定的用户不在当前群聊中。"); - } - Long userId = UserDataStore.findBoundUid(mentionedOpenId); - if (userId == null) { - throw new ResolutionException("被@的用户还没有绑定玩家ID,请先让对方使用 /bind。"); - } - User user = findUserById(userId); - UserDataStore.storeUserInfo(user.getId(), user.getUsername()); - return new WatchTarget(user.getId(), user.getUsername(), mentionedOpenId); - } - - Long explicitUserId = resolver.parsePositiveLong(argument); - User user = explicitUserId == null - ? APIHelper.lookupUser(argument).getContent() - : findUserById(explicitUserId); - String qqOpenId = UserDataStore.findGroupOpenIdByUid(groupId, user.getId()) - .orElseThrow(() -> new ResolutionException("指定的玩家不在当前群聊中,或尚未在本群完成绑定。")); - UserDataStore.storeUserInfo(user.getId(), user.getUsername()); - return new WatchTarget(user.getId(), user.getUsername(), qqOpenId); - } } diff --git a/src/main/java/xyz/zcraft/seira/command/parse/CommandParser.java b/src/main/java/xyz/zcraft/seira/command/parse/CommandParser.java index 9c341d1d..108fafb3 100644 --- a/src/main/java/xyz/zcraft/seira/command/parse/CommandParser.java +++ b/src/main/java/xyz/zcraft/seira/command/parse/CommandParser.java @@ -52,7 +52,7 @@ public ParseResult parse(String rawContent, String senderUserId, String groupId, String normalized = rawContent.trim(); if (!normalized.startsWith(PREFIX)) { - return ParseResult.ignored(); + return ParseResult.text(new Context(senderUserId, groupId, messageId, null, null, rawContent, null)); } String body = normalized.substring(PREFIX.length()).trim(); @@ -75,8 +75,8 @@ public ParseResult parse(String rawContent, String senderUserId, String groupId, public record ParseResult(Status status, Context context) { public ParseResult { - if ((status == Status.PARSED) == (context == null)) { - throw new IllegalArgumentException("Only parsed results may contain a context"); + if ((status == Status.PARSED || status == Status.TEXT) == (context == null)) { + throw new IllegalArgumentException("Only parsed results or text may contain a context"); } } @@ -92,10 +92,15 @@ static ParseResult parsed(Context context) { return new ParseResult(Status.PARSED, Objects.requireNonNull(context)); } + static ParseResult text(Context context) { + return new ParseResult(Status.TEXT, Objects.requireNonNull(context)); + } + public enum Status { IGNORED, EMPTY_COMMAND, - PARSED + PARSED, + TEXT } } } diff --git a/src/main/java/xyz/zcraft/seira/command/parse/Resolver.java b/src/main/java/xyz/zcraft/seira/command/parse/Resolver.java index 65ce5b6b..d6583391 100644 --- a/src/main/java/xyz/zcraft/seira/command/parse/Resolver.java +++ b/src/main/java/xyz/zcraft/seira/command/parse/Resolver.java @@ -1,7 +1,7 @@ package xyz.zcraft.seira.command.parse; import xyz.zcraft.seira.api.data.SearchQuery; -import xyz.zcraft.seira.data.UserRef; +import xyz.zcraft.seira.command.ResolutionException; import xyz.zcraft.seira.db.UserDataStore; import java.util.*; @@ -10,7 +10,6 @@ import java.util.regex.Pattern; public final class Resolver { - private static final ArrayList USER_MACRO_TYPES = new ArrayList<>(List.of("rs", "bp", "rp")); private final java.util.function.Function boundUid; public Resolver() { @@ -21,10 +20,31 @@ public Resolver(java.util.function.Function boundUid) { this.boundUid = Objects.requireNonNull(boundUid); } + private static final Map, String> ALIASES = Map.of( + List.of("+", "+"), "+", + List.of("~", "~"), "~", + List.of("=", "="), "=" + ); + public String sanitize(String rawContent) { + for (Map.Entry, String> entry : ALIASES.entrySet()) { + for (String s : entry.getKey()) { + rawContent = rawContent.replace(s, entry.getValue()); + } + } + + if (rawContent.trim().equalsIgnoreCase("@") + || rawContent.trim().equalsIgnoreCase("//")) { + return "ux"; + } else if (looksLikeMention(rawContent.trim())) { + return "ux " + rawContent; + } + // Add surrounding space to <@> before expanding compact commands so /bp5<@...> is recognized. rawContent = Patterns.QQ_INLINE_AT_PATTERN.matcher(rawContent).replaceAll(r -> " " + r.group() + " "); + rawContent = rawContent.replace("+", " +"); + Matcher matcher = Patterns.COMPACT_SCORE_COMMAND_PATTERN.matcher(rawContent); if (matcher.find()) { String type = matcher.group(1).toLowerCase(Locale.ROOT); @@ -41,6 +61,14 @@ public String sanitize(String rawContent) { } } + matcher = Patterns.SPACE_MISSING_COMMAND_PATTERN.matcher(rawContent); + if (matcher.find()) { + String command = matcher.group(1).toLowerCase(Locale.ROOT); + String target = matcher.group(2); + String remaining = rawContent.substring(matcher.end()); + rawContent = command + " " + target + " " + remaining; + } + return rawContent; } @@ -62,110 +90,22 @@ public SearchQuery resolveSearchQuery(String arg) { return null; } - public TargetResolution resolveTargetWithOptionalMention(String[] args, String senderUserId) { - if (args.length >= 2 && isUserMacro(args[1])) { - String mentionedUserId = extractMentionedUserId(args[0]); - if (mentionedUserId != null) { - Long boundUid = resolveBoundUid(mentionedUserId); - if (boundUid == null) { - return new TargetResolution(new ShortcutTarget(null, null, null, null, - "被@的用户还没有绑定玩家ID,无法使用快捷查询。"), 2); - } - return new TargetResolution(parseTarget(args[1], new UserRef.ByUid(boundUid), true, false), 2); - } - - Long uid = parsePositiveLong(args[0]); - if (uid != null) { - return new TargetResolution(parseTarget(args[1], new UserRef.ByUid(uid), false, false), 2); - } - - if (!looksLikeMention(args[0]) && !args[0].isBlank()) { - return new TargetResolution(parseTarget(args[1], new UserRef.ByUsername(args[0]), false, false), 2); - } - -// if (looksLikeMention(args[0])) { -// return new TargetResolution(new ShortcutTarget(null, null, null, null, "@用户格式无效,请使用@用户后再输入快捷查询(如 rs2)。"), 2); -// } + public String player(String argument, String senderUserId) { + if (argument == null) { + Long uid = resolveBoundUid(senderUserId); + if (uid == null) throw new ResolutionException("你还没有绑定玩家ID,请先使用 /bind 喵"); + return uid.toString(); } - return new TargetResolution(parseTarget(args[0], senderUserId), 1); - } - - public ShortcutTarget parseTarget(String arg, String senderUserId) { - Long boundUid = resolveBoundUid(senderUserId); - UserRef userRef = boundUid == null ? null : new UserRef.ByUid(boundUid); - return parseTarget(arg, userRef, false, true); - } - - public UserRefResolution resolveUserRefArgument(String arg) { - Long explicitUid = parsePositiveLong(arg); - if (explicitUid != null) { - return new UserRefResolution(new UserRef.ByUid(explicitUid), null); + String mentioned = extractMentionedUserId(argument); + if (mentioned != null) { + Long uid = resolveBoundUid(mentioned); + if (uid == null) throw new ResolutionException("被@的用户还没有绑定玩家ID,请先让对方使用 /bind 喵"); + return uid.toString(); } - - String mentionedUserId = extractMentionedUserId(arg); - if (mentionedUserId != null) { - Long boundUid = resolveBoundUid(mentionedUserId); - if (boundUid == null) { - return new UserRefResolution(null, "被@的用户还没有绑定玩家ID,请先让对方使用 /bind"); - } - return new UserRefResolution(new UserRef.ByUid(boundUid), null); - } - -// if (looksLikeMention(arg)) { -// return new UserRefResolution(null, "@用户格式无效,请使用 @用户 后再输入指令。示例:/bo 5 @123456"); -// } - - String username = arg == null ? "" : arg.trim(); - if (username.startsWith("@")) { - username = username.substring(1); - } - - if (username.isEmpty()) { - return new UserRefResolution(null, null); - } - - return new UserRefResolution(new UserRef.ByUsername(username), null); - } - - public RscTarget resolveRscTarget(String groupId, String extraUidArg) { - Set merged = new LinkedHashSet<>(); - - if (extraUidArg == null || extraUidArg.trim().startsWith("+")) { - List groupBoundUids = UserDataStore.findBoundUidsByGroup(groupId); - if (groupBoundUids.isEmpty()) { - return new RscTarget(null, "本群还没有已绑定的玩家,请先使用 /bind"); - } - groupBoundUids.stream().map(String::valueOf).forEach(merged::add); - } - - if (extraUidArg == null) return new RscTarget(merged.toArray(String[]::new), null); - - String trimmed = extraUidArg.trim(); - String body = trimmed.substring(1).trim(); - if (body.isEmpty()) { - return new RscTarget(null, "追加ID列表不能为空。"); - } - - String[] extraTokens = body.split(","); - for (String token : extraTokens) { - if (Patterns.RSC_TARGET_PATTERN.matcher(token.trim()).matches()) { - merged.add(token); - } else if (looksLikeMention(token)) { - final UserRefResolution userRefResolution = resolveUserRefArgument(token); - if (userRefResolution.errorMessage() != null) { - return new RscTarget(null, "解析 " + token + " 时出错:" + userRefResolution.errorMessage()); - } - if (userRefResolution.userRef() instanceof UserRef.ByUid ref) { - merged.add("u" + ref.getUid()); - } else if (userRefResolution.userRef() instanceof UserRef.ByUsername ref) { - merged.add("@" + ref.getUsername()); - } - } else { - return new RscTarget(null, "追加ID列表包含非法值。"); - } - } - - return new RscTarget(merged.toArray(String[]::new), null); + String player = argument.trim(); + if (player.startsWith("@")) player = player.substring(1); + if (player.isBlank()) throw new ResolutionException("无法识别指定的玩家喵"); + return player; } public Long resolveBoundUid(String senderUserId) { @@ -184,93 +124,6 @@ public Integer parsePositiveInt(String value) { } } - private ShortcutTarget parseTarget(String arg, UserRef userRef, boolean mentionedUser, boolean needResolveBound) { - Matcher setMatcher = Patterns.SET_MACRO_PATTERN.matcher(arg.trim()); - if (setMatcher.matches()) { - Long setId = parsePositiveLong(setMatcher.group(1)); - Long index = parsePositiveLong(setMatcher.group(2)); - - if (setId == null || index == null || index < 1) { - return new ShortcutTarget(null, null, null, null, "谱面集索引无效。例如: 12345#2"); - } - - return new ShortcutTarget(setId, userRef, "ms", index, null); - } - - Matcher userMatcher = Patterns.USER_MACRO_PATTERN.matcher(arg.trim()); - if (userMatcher.matches()) { - String type = userMatcher.group(1).toLowerCase(); - - if (!USER_MACRO_TYPES.contains(type)) { - return new ShortcutTarget(null, null, null, null, "未知的快捷查询"); - } - - Long index = parsePositiveLong(userMatcher.group(2)); - - if (index == null) { - index = 1L; - } - - if (index < 1 || index > 200) { - return new ShortcutTarget(null, null, null, null, "快捷指令索引无效,请输入 1-200 之间的数字。例如: rs5"); - } - - if (userRef == null) { - String errorMessage; - if (needResolveBound) { - errorMessage = mentionedUser - ? "被@的用户还没有绑定玩家ID,无法使用快捷查询。" - : "你还没有绑定玩家ID,无法使用快捷查询。请先使用 /bind"; - } else { - errorMessage = "无法识别指定的玩家ID"; - } - - return new ShortcutTarget(null, null, null, null, errorMessage); - } - - return new ShortcutTarget(null, userRef, type, index, null); - } - - if ("mp".equalsIgnoreCase(arg.trim())) { - if (userRef == null) { - String errorMessage; - if (needResolveBound) { - errorMessage = mentionedUser - ? "被@的用户还没有绑定玩家ID,无法使用快捷查询。" - : "你还没有绑定玩家ID,无法使用快捷查询。请先使用 /bind"; - } else { - errorMessage = "无法识别指定的玩家ID"; - } - - return new ShortcutTarget(null, null, null, null, errorMessage); - } - - return new ShortcutTarget(null, userRef, "mp", null, null); - } - - Matcher beatmapMatcher = Patterns.BEATMAP_MACRO_PATTERN.matcher(arg.trim()); - if (beatmapMatcher.matches()) { - Long mapId = parsePositiveLong(beatmapMatcher.group(1)); - return new ShortcutTarget(mapId, userRef, "m", null, null); - } - - if (Patterns.LOCAL_SCORE_PATTERN.matcher(arg.trim()).matches()) { - return ShortcutTarget.localScore(arg.trim().toLowerCase(Locale.ROOT)); - } - - Long id = parsePositiveLong(arg); - if (id == null) { - return new ShortcutTarget(null, null, null, null, - "参数无效。请输入数字ID、本地成绩ID或快捷指令 (例如 loc123456789, rs1, 12345#2)。"); - } - - return new ShortcutTarget(id, null, null, null, null); - } - - private boolean isUserMacro(String arg) { - return Patterns.USER_MACRO_PATTERN.matcher(arg.trim()).matches(); - } - public boolean looksLikeMention(String token) { String trimmed = token == null ? "" : token.trim(); return trimmed.startsWith("@") @@ -297,6 +150,18 @@ public String extractMentionedUserId(String token) { return null; } + public Set extractAllMentionedIds(String token) { + if (token == null) { + return Set.of(); + } + + Set result = new HashSet<>(); + + Patterns.QQ_AT_IDS_PATTERN.matcher(token).results().forEach(m -> result.add(m.group(1))); + + return result; + } + public Long parsePositiveLong(String value) { try { long parsed = Long.parseLong(value); @@ -311,18 +176,17 @@ public boolean looksLikeUid(String arg) { } private static final class Patterns { - private static final Pattern USER_MACRO_PATTERN = Pattern.compile("(?i)^(rs|bo|rp|bp)(\\d+)?$"); private static final Pattern COMPACT_SCORE_COMMAND_PATTERN = Pattern.compile( "(?i)^(rs|rp|bp)(\\d+)(?:-(\\d+))?(?=\\s|$)" ); - private static final Pattern SET_MACRO_PATTERN = Pattern.compile("^(\\d+)#(\\d+)$"); - private static final Pattern BEATMAP_MACRO_PATTERN = Pattern.compile("^m(\\d+)$"); - private static final Pattern LOCAL_SCORE_PATTERN = Pattern.compile("(?i)^loc[1-9]\\d*$"); - private static final Pattern QQ_AT_PATTERN = Pattern.compile("^<@([A-Z|0-9]{32})>$"); - private static final Pattern QQ_INLINE_AT_PATTERN = Pattern.compile("(<@[A-Z|0-9]{32}>)"); + private static final Pattern SPACE_MISSING_COMMAND_PATTERN = Pattern.compile( + "^([a-zA-Z]+)(\\d+(?:#\\d+)?)" + ); + private static final Pattern QQ_AT_PATTERN = Pattern.compile("^<@([A-Z0-9]{32})>$"); + private static final Pattern QQ_AT_IDS_PATTERN = Pattern.compile("<@([A-Z0-9]{32})>"); + private static final Pattern QQ_INLINE_AT_PATTERN = Pattern.compile("(<@[A-Z0-9]{32}>)"); private static final Pattern PLAIN_AT_PATTERN = Pattern.compile("^@(\\d+)$"); private static final Pattern SEARCH_PATTERN = Pattern.compile("^(?:#(\\d+) )?(.+)$"); - private static final Pattern RSC_TARGET_PATTERN = Pattern.compile("^[us]?\\d+$"); } } diff --git a/src/main/java/xyz/zcraft/seira/command/parse/RscTarget.java b/src/main/java/xyz/zcraft/seira/command/parse/RscTarget.java deleted file mode 100644 index 9193b455..00000000 --- a/src/main/java/xyz/zcraft/seira/command/parse/RscTarget.java +++ /dev/null @@ -1,4 +0,0 @@ -package xyz.zcraft.seira.command.parse; - -public record RscTarget(String[] targets, String errorMessage) { -} diff --git a/src/main/java/xyz/zcraft/seira/command/parse/ScoreFilterArguments.java b/src/main/java/xyz/zcraft/seira/command/parse/ScoreFilterArguments.java index cce66a5a..195b2a43 100644 --- a/src/main/java/xyz/zcraft/seira/command/parse/ScoreFilterArguments.java +++ b/src/main/java/xyz/zcraft/seira/command/parse/ScoreFilterArguments.java @@ -13,7 +13,8 @@ public final class ScoreFilterArguments { private static final Pattern FILTER_PATTERN = Pattern.compile( "(?i)^(acc(?:uracy)?|combo|pp|time|length|len|star|stars|sr|bpm|miss|misses|score|mod|mods|rank|replay" - + "|any|title|artist|mapper|genre|language|video|storyboard|fullcombo|ar|od|cs|hp)" + + "|any|title|artist|mapper|genre|language|tag|source|nsfw|video|storyboard|fullcombo" + + "|ar|od|cs|hp|t|a|cb|cmb|m|vid|sb|fc|rep|rp)" + "(>=|<=|!=|!~|>|<|=|~)(.+)$" ); private static final Pattern MISS_SHORTHAND_PATTERN = Pattern.compile("(?i)^(!?)(\\d+)miss(?:es)?$"); @@ -71,7 +72,7 @@ private static String parseOne(String token) { if (!RANKS.contains(value)) { throw new IllegalArgumentException("rank 必须是 SSH/SS/XH/X/SH/S/A/B/C/D/F"); } - } else if (Set.of("any", "title", "artist", "mapper", "genre", "language").contains(field)) { + } else if (Set.of("any", "title", "artist", "mapper", "genre", "language", "tag", "source").contains(field)) { if (!Set.of("~", "!~", "=", "!=").contains(operator)) { throw new IllegalArgumentException(field + " 仅支持 ~、!~、=、!="); } @@ -95,7 +96,7 @@ private static String parseOne(String token) { throw new IllegalArgumentException("正则表达式无效:" + e.getDescription()); } } - } else if (Set.of("video", "storyboard", "fullcombo", "replay").contains(field)) { + } else if (Set.of("nsfw", "video", "storyboard", "fullcombo", "replay").contains(field)) { if (!Set.of("=", "!=").contains(operator)) { throw new IllegalArgumentException(field + " 仅支持 =、!="); } @@ -122,7 +123,7 @@ private static String parseOne(String token) { private static String normalizeField(String value) { return switch (value.toLowerCase(Locale.ROOT)) { case "acc", "accuracy" -> "acc"; - case "combo" -> "combo"; + case "combo", "cmb", "cb" -> "combo"; case "pp" -> "pp"; case "time", "length", "len" -> "time"; case "star", "stars", "sr" -> "star"; @@ -136,15 +137,18 @@ private static String normalizeField(String value) { case "mod", "mods" -> "mod"; case "rank" -> "rank"; case "any" -> "any"; - case "title" -> "title"; - case "artist" -> "artist"; - case "mapper" -> "mapper"; + case "title", "t" -> "title"; + case "artist", "a" -> "artist"; + case "mapper", "m" -> "mapper"; case "genre" -> "genre"; case "language" -> "language"; - case "video" -> "video"; - case "storyboard" -> "storyboard"; - case "fullcombo" -> "fullcombo"; - case "replay" -> "replay"; + case "tag", "tags" -> "tag"; + case "source" -> "source"; + case "nsfw" -> "nsfw"; + case "video", "vid" -> "video"; + case "storyboard", "sb" -> "storyboard"; + case "fullcombo", "fc" -> "fullcombo"; + case "replay", "rep", "rp" -> "replay"; default -> throw new IllegalArgumentException("未知字段 " + value); }; } @@ -160,6 +164,8 @@ private static String expandShorthand(String value) { case "!video" -> "video=false"; case "sb", "storyboard" -> "storyboard=true"; case "!sb", "!storyboard" -> "storyboard=false"; + case "nsfw" -> "nsfw=true"; + case "!nsfw" -> "nsfw=false"; case "fc", "fullcombo" -> "fullcombo=true"; case "replay" -> "replay=true"; case "!replay" -> "replay=false"; diff --git a/src/main/java/xyz/zcraft/seira/command/parse/ShortcutTarget.java b/src/main/java/xyz/zcraft/seira/command/parse/ShortcutTarget.java deleted file mode 100644 index 217e00d1..00000000 --- a/src/main/java/xyz/zcraft/seira/command/parse/ShortcutTarget.java +++ /dev/null @@ -1,36 +0,0 @@ -package xyz.zcraft.seira.command.parse; - -import xyz.zcraft.seira.data.UserRef; - -/** - * 用户输入或 API 目标。解析后的普通目标只包含 explicitId; - * 本地成绩使用 localScoreId。跨类型查找时 m/ms/s 标明原始 ID 的类型。 - */ -public record ShortcutTarget( - Long explicitId, - String localScoreId, - UserRef userRef, - String macroType, - Long macroIndex, - String errorMessage -) { - public ShortcutTarget(Long explicitId, UserRef userRef, String macroType, Long macroIndex, String errorMessage) { - this(explicitId, null, userRef, macroType, macroIndex, errorMessage); - } - - public static ShortcutTarget localScore(String localScoreId) { - return new ShortcutTarget(null, localScoreId, null, null, null, null); - } - - public boolean isMacro() { - return macroType != null; - } - - public boolean isError() { - return errorMessage != null; - } - - public boolean isLocalScore() { - return localScoreId != null; - } -} diff --git a/src/main/java/xyz/zcraft/seira/command/parse/TargetInput.java b/src/main/java/xyz/zcraft/seira/command/parse/TargetInput.java new file mode 100644 index 00000000..92cf495a --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/command/parse/TargetInput.java @@ -0,0 +1,61 @@ +package xyz.zcraft.seira.command.parse; + +import xyz.zcraft.seira.command.ResolutionException; + +import java.util.Locale; +import java.util.concurrent.ThreadLocalRandom; +import java.util.regex.Pattern; + +public record TargetInput(Kind kind, String id, long index, String player, int consumedArgs) { + private static final Pattern PLAYER_SCORE = Pattern.compile("(?i)^(rs|rp|bp)(\\d+)?$"); + private static final Pattern SET = Pattern.compile("^(\\d+)#(\\d+)$"); + + public enum Kind { MEMORY, ID, MAP, SET, SCORE, RS, RP, BP, MP } + + public static TargetInput memory() { + return new TargetInput(Kind.MEMORY, null, 1, null, 0); + } + + public static TargetInput read(String[] args) { + if (args.length == 0) return memory(); + int consumed = args.length >= 2 && PLAYER_SCORE.matcher(args[1]).matches() ? 2 : 1; + String player = consumed == 2 ? args[0] : null; + String value = args[consumed - 1].trim().toLowerCase(Locale.ROOT); + if (value.equals("rbp")) { + value = "bp" + (ThreadLocalRandom.current().nextInt(200) + 1); + } + var score = PLAYER_SCORE.matcher(value); + if (score.matches()) { + long index = score.group(2) == null ? 1 : positive(score.group(2), "快捷指令索引无效,请输入 1-200 之间的数字。例如: rp5"); + if (index > 200) throw new ResolutionException("快捷指令索引无效,请输入 1-200 之间的数字。例如: rp5"); + return new TargetInput(Kind.valueOf(score.group(1).toUpperCase(Locale.ROOT)), null, index, player, consumed); + } + var set = SET.matcher(value); + if (set.matches()) { + long id = positive(set.group(1), "谱面集索引无效。例如: 12345#2"); + long index = positive(set.group(2), "谱面集索引无效。例如: 12345#2"); + return new TargetInput(Kind.SET, Long.toString(id), index, null, consumed); + } + if (value.equals("mp")) return new TargetInput(Kind.MP, null, 1, null, consumed); + if (value.matches("loc[1-9]\\d*")) return new TargetInput(Kind.SCORE, value, 1, null, consumed); + if (value.matches("m\\d+")) { + long id = positive(value.substring(1), "谱面ID无效"); + return new TargetInput(Kind.MAP, Long.toString(id), 1, null, consumed); + } + long id = positive(value, "参数无效。请输入数字ID、本地成绩ID或快捷指令 (例如 loc123456789, rp1, 12345#2)。"); + return new TargetInput(Kind.ID, Long.toString(id), 1, null, consumed); + } + + public String scoreList() { + return kind.name().toLowerCase(Locale.ROOT); + } + + private static long positive(String value, String message) { + try { + long id = Long.parseLong(value); + if (id > 0) return id; + } catch (NumberFormatException ignored) { + } + throw new ResolutionException(message); + } +} diff --git a/src/main/java/xyz/zcraft/seira/command/parse/TargetResolution.java b/src/main/java/xyz/zcraft/seira/command/parse/TargetResolution.java deleted file mode 100644 index 39c58d5d..00000000 --- a/src/main/java/xyz/zcraft/seira/command/parse/TargetResolution.java +++ /dev/null @@ -1,17 +0,0 @@ -package xyz.zcraft.seira.command.parse; - -import xyz.zcraft.seira.command.Context; -import xyz.zcraft.seira.data.UserRef; - -/** - * target 为 null 表示省略目标;consumedArgs 之后是指令自己的可选参数。 - */ -public record TargetResolution(ShortcutTarget target, int consumedArgs, UserRef userOverride) { - public TargetResolution(ShortcutTarget target, int consumedArgs) { - this(target, consumedArgs, null); - } - - public String nextArgument(Context ctx) { - return ctx.argumentCount() > consumedArgs ? ctx.argument(consumedArgs) : null; - } -} diff --git a/src/main/java/xyz/zcraft/seira/command/parse/UserRefResolution.java b/src/main/java/xyz/zcraft/seira/command/parse/UserRefResolution.java deleted file mode 100644 index 56b5a953..00000000 --- a/src/main/java/xyz/zcraft/seira/command/parse/UserRefResolution.java +++ /dev/null @@ -1,6 +0,0 @@ -package xyz.zcraft.seira.command.parse; - -import xyz.zcraft.seira.data.UserRef; - -public record UserRefResolution(UserRef userRef, String errorMessage) { -} diff --git a/src/main/java/xyz/zcraft/seira/command/reply/CommandUsage.java b/src/main/java/xyz/zcraft/seira/command/reply/CommandUsage.java index a69468d6..809aa7b4 100644 --- a/src/main/java/xyz/zcraft/seira/command/reply/CommandUsage.java +++ b/src/main/java/xyz/zcraft/seira/command/reply/CommandUsage.java @@ -14,6 +14,7 @@ public final class CommandUsage { public static final String BGP = "用法:/bgp <谱面ID 或 快捷查询>"; public static final String DL = "用法:/dl <谱面集ID 或 快捷查询>"; public static final String S = "用法:/s [成绩ID 或 快捷查询] [用户] [+Mod](省略目标时使用记忆)"; + public static final String SM = "用法:/sm [谱面ID 或 快捷查询] [@玩家] [+Mod](省略目标时使用记忆,省略玩家时使用自己的绑定账号)"; public static final String SA = "用法:/sa <成绩ID 或 快捷查询>"; public static final String MA = "用法:/ma [成绩ID 或 快捷查询] [序号];省略目标并指定序号时请使用 #序号"; public static final String R = "用法:/r [成绩ID 或 快捷查询] [[mm:ss]-[mm:ss]]"; diff --git a/src/main/java/xyz/zcraft/seira/command/reply/ReplyFactory.java b/src/main/java/xyz/zcraft/seira/command/reply/ReplyFactory.java index 92061d15..84b14bf1 100644 --- a/src/main/java/xyz/zcraft/seira/command/reply/ReplyFactory.java +++ b/src/main/java/xyz/zcraft/seira/command/reply/ReplyFactory.java @@ -5,7 +5,8 @@ import com.google.gson.JsonObject; import org.jetbrains.annotations.NotNull; import xyz.zcraft.osu.model.*; -import xyz.zcraft.seira.api.APIHelper; +import xyz.zcraft.seira.api.AsteroidApi; +import xyz.zcraft.seira.api.OstellaApi; import xyz.zcraft.seira.api.data.*; import xyz.zcraft.seira.bot.data.Button; import xyz.zcraft.seira.bot.data.PendingMessage; @@ -27,6 +28,8 @@ import java.util.function.Supplier; import java.util.stream.Stream; +import static xyz.zcraft.seira.command.reply.ReplyFactory.ExternalUrls.*; + public final class ReplyFactory { private final Supplier configSupplier; @@ -100,13 +103,13 @@ public static String at(String openId) { } @SuppressWarnings("unused") - static String url(String text, String url) { + public static String url(String text, String url) { return "[" + text + "](" + url + ")"; } public static PendingMessage replayUploadMessage(ReplayUploadInfo info) { return PendingMessage.ofMarkdownRaw( - ("\n" + "## Replay上传成功~" + "\n" + + ("\n" + "__Replay上传成功~__" + "\n" + "> 成绩: " + s(info.scoreId()) + "\n" + "> 谱面: " + m(info.beatmapId()) + "\n" + "> 用户: " + u(info.userId(), info.username()) + "\n").trim(), @@ -119,13 +122,13 @@ private static boolean isCancelableReplayStatus(String status) { || "upload_queued".equals(status) || "uploading".equals(status); } - public PendingMessage friendStatusMessage(String selfOpenId, Long selfUid, String selfUsername, - String targetOpenId, Long targetUid, String targetUsername, + public PendingMessage friendStatusMessage(String selfOpenId, Long selfUid, String selfOsuAvatar, String selfUsername, String selfAvatar, + String targetOpenId, Long targetUid, String targetOsuAvatar, String targetUsername, String targetAvatar, boolean selfFollowed, Boolean targetFollowed) { return PendingMessage.ofMarkdownRaw( Contents.friendStatusContent( - selfOpenId, selfUid, selfUsername, - targetOpenId, targetUid, targetUsername, + selfOpenId, selfUid, selfOsuAvatar, selfUsername, selfAvatar, + targetOpenId, targetUid, targetOsuAvatar, targetUsername, targetAvatar, selfFollowed, targetFollowed ) ); @@ -317,7 +320,7 @@ public PendingMessage lbMessage(Context ctx, Response response) { ); } - public PendingMessage replayMessage(Context ctx, APIHelper.ReplayTaskInfo taskInfo) { + public PendingMessage replayMessage(Context ctx, OstellaApi.ReplayTaskInfo taskInfo) { return PendingMessage.ofMarkdownRaw( Contents.replayTaskContent(ctx, taskInfo), buttons().replayProgressButtons(taskInfo.taskId(), ctx.senderUserId()) @@ -424,9 +427,9 @@ public PendingMessage scoreMissesMessage(Context ctx, Response> s ); } - public PendingMessage statusMessage(Context ctx, APIHelper.ServerStatus status) { + public PendingMessage statusMessage(Context ctx, OstellaApi.ServerStatus status, AsteroidApi.ServerStatus asteroid) { return PendingMessage.ofMarkdownRaw( - Contents.statContent(ctx, status), null + Contents.statContent(ctx, status, asteroid), null ); } @@ -434,6 +437,10 @@ public PendingMessage helpMessage(Context ctx) { return PendingMessage.ofMarkdownRaw(Contents.helpContent(ctx)); } + public PendingMessage usagesMessage(Context ctx) { + return PendingMessage.ofMarkdownRaw(Contents.usagesContent(ctx)); + } + public PendingMessage faqMessage(Context ctx) { return PendingMessage.ofMarkdownRaw(Contents.faqContent(ctx)); } @@ -463,8 +470,21 @@ public PendingMessage supMessage(Context ctx, String username, String openId, Bo return PendingMessage.ofMarkdownRaw(Contents.supContent(ctx, username, openId, isSupporter, hasSupported, supportLevel)); } + public PendingMessage userInfoShortMessage(Context ctx, UserExtended user) { + return PendingMessage.ofMarkdownRaw(Contents.userInfoShortContent(ctx, user)); + } + + public static class ExternalUrls { + public static final String COMMANDS = "https://docs.seira.top/overview/commands.html"; + public static final String PERMISSION = "https://docs.seira.top/overview/use.html#extra-permission"; + public static final String CHANNEL = "https://pd.qq.com/s/f9icas5gj?b=5"; + public static final String CHANGELOG = "https://docs.seira.top/overview/changelog.html"; + public static final String FAQ = "https://docs.seira.top/overview/faq.html"; + public static final String GITHUB = "https://github.com/BotSeira"; + } + private static final class Contents { - static String replayTaskContent(Context ctx, APIHelper.ReplayTaskInfo taskInfo) { + static String replayTaskContent(Context ctx, OstellaApi.ReplayTaskInfo taskInfo) { StringBuilder sb = new StringBuilder(); sb.append(at(ctx)).append("回放生成请求已提交").append("\n"); @@ -643,16 +663,16 @@ public static String friendContent(Context ctx, boolean collapsed = false; - sb.append("> 好友←→ (").append(mutual.size()).append(")\n>"); + sb.append("> __好友←→ (").append(mutual.size()).append(")__ \n>"); collapsed |= appendFriends(ctx, mutual, sb); - sb.append("\n> 仅关注→ (").append(onlyFollowed.size()).append(")\n>"); + sb.append("\n> __仅关注→ (").append(onlyFollowed.size()).append(")__ \n>"); collapsed |= appendFriends(ctx, onlyFollowed, sb); - sb.append("\n> 仅粉丝← ("); + sb.append("\n> __仅粉丝← ("); sb.append(onlyFollower.size()).append(" 已知"); - if (all) sb.append(" 共 ").append(self.getFollowerCount() - allMutualCount); - sb.append(")\n>"); + if (all) sb.append(" 共 ").append(Math.max(self.getFollowerCount() - allMutualCount, 0)); + sb.append(")__ \n>"); collapsed |= appendFriends(ctx, onlyFollower, sb); if (ctx.inGroup() && collapsed) { @@ -720,16 +740,20 @@ public static String scoreMissesContent(Context ctx, Response> sc return sb.toString().trim(); } - public static String statContent(Context ctx, APIHelper.ServerStatus status) { + public static String statContent(Context ctx, OstellaApi.ServerStatus status, AsteroidApi.ServerStatus asteroid) { String stat = at(ctx) + "\n" + "## 服务器状态\n" + "> 消息网关: ✅ 正常\n" + "> oStella API: " + (status.oStella() ? "✅ 正常" : "❌ 无法访问") + "\n"; if (status.oStella()) { - stat += "> osu! API: " + (status.osu() ? "✅ 正常" : "❌ 无法访问") + "\n"; + stat += "> ↳ osuRenderer: " + (status.onlineWorkers() + " / " + status.allWorkers()) + + (status.onlineWorkers() > 0 ? " (✅在线)" : " (❌全部离线)") + "\n"; + stat += "> ↳ osu! API: " + (status.osu() ? "✅ 正常" : "❌ 无法访问") + "\n"; } + stat += "> Asteroid API: " + (asteroid.online() ? "✅ 正常" : "❌ 无法访问") + "\n"; + String version = "## 版本信息" + "\n" + "> SeiraCore: " + VersionInfo.getVersion() + "\n"; @@ -737,6 +761,10 @@ public static String statContent(Context ctx, APIHelper.ServerStatus status) { version += "> oStella: " + status.oStellaVersion() + "\n"; } + if (asteroid.online() && asteroid.version() != null) { + version += "> Asteroid: " + asteroid.version() + "\n"; + } + String res = "## 统计信息\n" + "> Seira已经" + "\n" + @@ -750,32 +778,51 @@ public static String statContent(Context ctx, APIHelper.ServerStatus status) { } public static String helpContent(Context ctx) { - return at(ctx) + "\n" + - """ - 常用指令: - > /bind - 绑定你的玩家ID - > /rp - 获取最近通过的一个成绩 - > /bo [个数] [玩家ID] - 获取一个或多个最佳成绩 - > /rp [个数] [玩家ID] - 获取最近通过一个或多个成绩 - > /tb [#天数] [玩家ID] - 获取近N天达成的BP - > /s <成绩ID或快捷查询> - 获取指定成绩 - > /m <谱面ID或快捷查询> - 获取谱面 - > /bma <谱面ID或快捷查询> [Mod] - 分析谱面PP构成和类型 - > /ms <谱面集ID或快捷查询> - 获取谱面集 - > /r [成绩ID或快捷查询] [[mm:ss]-[mm:ss]] - 生成成绩高光视频或指定片段 - > /rcancel <任务ID> - 取消回放渲染任务 - > /rg - 猜 Rank 游戏与个人战绩 - > /lb <谱面ID> [玩家ID列表] - 获取指定谱面排行榜 - > /watch add <玩家ID/用户名/@用户> [分钟] - 监视群友的新成绩 - > /wx start <谱面ID列表> - 监视指定玩家在指定谱面的成绩 - > /mpwatch [start] <房间ID> [stable|lazer] - 监视多人房间的逐图结果(链接可自动识别版本,stop all 停止本群全部监视) - > /f - 获取好友列表 - - 详细指令列表请在 [这里](https://docs.seira.top/overview/commands.html) 查看 - 配置额外权限请在 [这里](https://docs.seira.top/overview/use.html#extra-permission) 查看 - """ + "\n" - + "当前版本: " + VersionInfo.getVersion() + " [更新日志](https://docs.seira.top/overview/changelog.html)" + "\n" - + "[常见问题](https://docs.seira.top/overview/faq.html)" + " " + cmd("/stat", "状态信息").trim(); + return at(ctx) + "常用指令: \n" + """ + > /rp - 获取最近通过的一个或多个成绩 + > /bp - 获取一个或多个最佳成绩 + > /tb - 获取近日BP + > /s - 获取指定成绩 + > /m - 获取谱面 + > /r - 生成成绩高光视频或指定片段 + > /rg - 猜 Rank 游戏 + > /watch - 监视群友的新成绩 + > /mpw - 监视多人房间的逐图结果 + > /f - 获取好友列表 + + [详细指令列表](%s) | [配置额外权限](%s) + [加入官方频道](%s) | [查看常见问题](%s) + [查看更新日志](%s) | %s + %s | [Github主页](%s) + + 当前版本: %s + """.formatted( + COMMANDS, PERMISSION, + CHANNEL, FAQ, + CHANGELOG, cmd("/stat", "查看状态信息"), + cmd("/usages", "查看用法示例"), GITHUB, + VersionInfo.getVersion() + ) + "\n"; + } + + public static String usagesContent(Context ctx) { + return at(ctx) + "部分指令示例\n" + + "> 注意:所有指令中的@均需要开启权限才能正常读取。权限配置见 [这里]( " + PERMISSION + " )~\n" + """ + > /rp -> 查看最近通过的一个成绩 + > /rp1-20 -> 查看最近通过的1到20个成绩 + > /bp1-20 -> 查看20个最佳成绩 + > /bp1-20 @peppy acc>95 -> 查看指定玩家BP1-20中准确率大于95%的成绩 + > /sa bp2 -> 查看BP2的成绩分析 + > /tb #7 @peppy -> 查看指定玩家近7天的新BP + > /@peppy -> 查看指定玩家的基本信息 + > /rg group -> 开始群组猜 Rank 游戏 + > /m @peppy rp2 -> 查看指定玩家最近第2条成绩的谱面 + > /dl mp -> 获取所在lazer多人房间当前谱面的镜像下载链接 + > /r rp -> 渲染最近通过的成绩的高光片段回放视频 + > /r @peppy bp2 90- -> 渲染指定玩家BP2从1:30开始的回放视频 + > /mpw -> 开始多人房间监视 + > /romai @peppy -> 开始监视指定玩家所在的RomAI对局 + """; } public static String faqContent(Context ctx) { @@ -790,8 +837,7 @@ public static String bgpContent(Context context, Response response) { public static String luckContent(Context ctx, DailyLuck.Luck luck, Beatmapset mapset, UploadedImage cover) { final List list = mapset.getBeatmaps().stream().map(Beatmap::getDifficultyRating).sorted().toList(); - String sb = at(ctx) + "\n" + - "## 今日运势" + "\n" + + String sb = at(ctx) + "你的今日运势" + "\n" + "> 人品值: **" + luck.luck() + "**/100\n" + "> 宜: " + luck.ups() + "\n" + "> 忌: " + luck.downs() + "\n\n" + @@ -801,8 +847,8 @@ public static String luckContent(Context ctx, DailyLuck.Luck luck, Beatmapset ma return sb.trim(); } - public static String friendStatusContent(String selfOpenId, Long selfUid, String selfUsername, - String targetOpenId, Long targetUid, String targetUsername, + public static String friendStatusContent(String selfOpenId, Long selfUid, String selfOsuAvatar, String selfUsername, String selfAvatar, + String targetOpenId, Long targetUid, String targetOsuAvatar, String targetUsername, String targetAvatar, boolean selfFollowed, Boolean targetFollowed) { final String status; if (targetFollowed == null) { @@ -822,10 +868,20 @@ public static String friendStatusContent(String selfOpenId, Long selfUid, String status = "✕ 路人 ✕"; } } - return at(selfOpenId) + "你们的好友状态(点击打开个人主页):" + "\n" + - url(selfUsername, "https://osu.ppy.sh/users/" + selfUid) + " (" + at(selfOpenId) + ")\n" + - " " + status + "\n" + - url(targetUsername, "https://osu.ppy.sh/users/" + targetUid) + " (" + at(targetOpenId) + ")"; + + return """ + %s 和 %s 的好友状态: + > # ![image #30px #30px](%s) __|__ %s + + # %s + + > # ![image #30px #30px](%s) __|__ %s + """.formatted( + at(selfOpenId), targetAvatar == null ? targetUsername : at(targetOpenId), + selfOsuAvatar, url(selfUsername, "https://osu.ppy.sh/users/" + selfUid), + status, + targetOsuAvatar, url(targetUsername, "https://osu.ppy.sh/users/" + targetUid) + ).trim(); } public static String missImageContent(Context ctx, String scoreId, Integer index, int size) { @@ -864,6 +920,27 @@ public static String supContent( return sb.toString().trim(); } + + public static String userInfoShortContent(Context ctx, UserExtended user) { + final Duration playTime = Duration.ofSeconds(user.getStatistics().getPlayTime()); + return """ + %s `%s` 的用户信息 + > - PP: %.2f + > - Rank: #%,d (%s #%,d) + > - 准确率: %.2f%% + > - 游玩次数: %,d + > - 获得总分: %,d + > - 游玩时间: %dd %dh %dm + """.formatted( + at(ctx), user.getUsername(), + user.getStatistics().getPp(), + user.getStatistics().getGlobalRank(), user.getCountry().getCode(), user.getStatistics().getRank().getCountry(), + user.getStatistics().getAccuracy() * 100, + user.getStatistics().getPlayCount(), + user.getStatistics().getRankedScore(), + playTime.toDaysPart(), playTime.toHoursPart(), playTime.toMinutesPart() + ); + } } private record Buttons(String directUrl) { @@ -1047,7 +1124,9 @@ List> replayProgressButtons(String jobId, boolean cancelable, Strin return Button.keyboard(cancelable ? Button.row( Button.command(1, "查询渲染进度", "/rstat " + jobId), - Button.command(2, "取消渲染", "/rcancel " + jobId).permit(userId) + Button.command(2, "取消渲染", "/rcancel " + jobId) + .permit(userId) + .modal("确定要取消渲染吗") ) : Button.row(Button.command(1, "查询渲染进度", "/rstat " + jobId))); } diff --git a/src/main/java/xyz/zcraft/seira/command/route/DebugRoutes.java b/src/main/java/xyz/zcraft/seira/command/route/DebugRoutes.java index d80febc3..6ce50ce2 100644 --- a/src/main/java/xyz/zcraft/seira/command/route/DebugRoutes.java +++ b/src/main/java/xyz/zcraft/seira/command/route/DebugRoutes.java @@ -4,7 +4,7 @@ import org.apache.logging.log4j.Logger; import org.bouncycastle.util.encoders.Base64Encoder; import xyz.zcraft.osu.model.UserExtended; -import xyz.zcraft.seira.api.APIHelper; +import xyz.zcraft.seira.api.OstellaApi; import xyz.zcraft.seira.api.data.FriendEntry; import xyz.zcraft.seira.api.data.OsuToken; import xyz.zcraft.seira.api.data.Response; @@ -148,8 +148,8 @@ public void handleUpdateUserInfo(Context ctx) { ctx.sendReply(PendingMessage.ofMarkdownRaw(at(ctx) + "用户信息更新失败")); return; } - try (var timing = taskCoordinator.beginRequest(ctx, "Update All User Info")) { - var users = APIHelper.getUsers(allUsers); + try (var _ = taskCoordinator.beginRequest(ctx, "Update All User Info")) { + var users = OstellaApi.getUsers(allUsers); for (var user : users) { UserDataStore.storeUserInfo(user.getId(), user.getUsername()); } @@ -158,7 +158,7 @@ public void handleUpdateUserInfo(Context ctx) { } public void handleGetAllFriends(Context ctx) { - try (var timing = taskCoordinator.beginRequest(ctx, "Get All Friends")) { + try (var _ = taskCoordinator.beginRequest(ctx, "Get All Friends")) { try { final List allOsuTokens = UserDataStore.getAllOsuTokens(); allOsuTokens @@ -167,8 +167,8 @@ public void handleGetAllFriends(Context ctx) { .map(authHelper::updateTokenAndGet) .map(OsuToken::accessToken) .forEach(accessToken -> { - final Response self = APIHelper.getSelf(accessToken); - final Response> response = APIHelper.getFollowed(accessToken); + final Response self = OstellaApi.getSelf(accessToken); + final Response> response = OstellaApi.getFollowed(accessToken); final List content = response.getContent(); final List ids = content.stream().map(e -> e.user().getId()).toList(); @@ -207,7 +207,7 @@ public void handleGetAllFriends(Context ctx) { } public void handleValidateToken(Context ctx) { - try (var timing = taskCoordinator.beginRequest(ctx, "Validate Token")) { + try (var _ = taskCoordinator.beginRequest(ctx, "Validate Token")) { int updated = 0, removed = 0; try { final List allOsuTokens = UserDataStore.getAllOsuTokens(); diff --git a/src/main/java/xyz/zcraft/seira/command/route/Router.java b/src/main/java/xyz/zcraft/seira/command/route/Router.java index 2e139b25..f4082c71 100644 --- a/src/main/java/xyz/zcraft/seira/command/route/Router.java +++ b/src/main/java/xyz/zcraft/seira/command/route/Router.java @@ -3,30 +3,37 @@ import lombok.Getter; import org.apache.logging.log4j.LogManager; import org.apache.logging.log4j.Logger; +import xyz.zcraft.seira.ai.provider.ChatProvider; +import xyz.zcraft.seira.ai.AiChatHandler; import xyz.zcraft.seira.api.data.OsuToken; import xyz.zcraft.seira.api.data.VideoRenderRecord; import xyz.zcraft.seira.bot.MessageSender; -import xyz.zcraft.seira.bot.data.PendingMessage; +import xyz.zcraft.seira.bot.QQApi; +import xyz.zcraft.seira.bot.data.*; import xyz.zcraft.seira.command.*; import xyz.zcraft.seira.command.handler.*; import xyz.zcraft.seira.command.parse.CommandParser; import xyz.zcraft.seira.command.parse.Resolver; import xyz.zcraft.seira.command.reply.ReplyFactory; import xyz.zcraft.seira.config.AppConfig; +import xyz.zcraft.seira.data.UploadedImage; import xyz.zcraft.seira.db.UserDataStore; import xyz.zcraft.seira.discord.DiscordBridgeService; import xyz.zcraft.seira.rankguess.RankGuessGameService; +import xyz.zcraft.seira.services.AiPermission; import xyz.zcraft.seira.services.BindingService; import xyz.zcraft.seira.util.AdminRegistry; import xyz.zcraft.seira.util.NoticesHelper; import xyz.zcraft.seira.util.OsuAuthHelper; -import xyz.zcraft.seira.watch.MultiplayerRoomWatchService; +import xyz.zcraft.seira.watch.MPWatchService; import xyz.zcraft.seira.watch.ScoreWatchService; +import java.util.List; import java.util.Optional; import java.util.Set; import java.util.concurrent.Executor; import java.util.concurrent.atomic.AtomicInteger; +import java.util.function.Function; import java.util.function.Supplier; import static xyz.zcraft.seira.command.reply.ReplyFactory.at; @@ -44,103 +51,82 @@ public class Router { private final Runnable commandMetric; private final CommandHandler unknownCommand; private final Executor commandExecutor; + private final Supplier selfSupplier; + private final AiChatHandler aiChatHandler; public Router( - MessageSender messageSender, - Supplier configSupplier, - AdminRegistry admins, - BindingService bindingService, - ScoreWatchService watchService, - MultiplayerRoomWatchService multiplayerRoomWatchService, - DiscordBridgeService discordBridgeService, - RankGuessGameService rankGuessGameService, - Executor commandExecutor, - Runnable commandMetric + MessageSender messageSender, Supplier configSupplier, AdminRegistry admins, + BindingService bindingService, ScoreWatchService watchService, MPWatchService mpWatchService, + DiscordBridgeService discordBridgeService, RankGuessGameService rankGuessGameService, Executor commandExecutor, + Runnable commandMetric, Function imageUploader, Supplier selfSupplier, + ChatProvider chatProvider, Function botStateGetter ) { this.configSupplier = java.util.Objects.requireNonNull(configSupplier); this.commandExecutor = commandExecutor; this.commandMetric = java.util.Objects.requireNonNull(commandMetric); + this.selfSupplier = selfSupplier; AppConfig startupConfig = configSupplier.get(); ReplyFactory replyFactory = new ReplyFactory(configSupplier); Resolver resolver = new Resolver(); - TargetHistory history = new TargetHistory(resolver, this::getAccessTokenFor); + TargetHistory history = new TargetHistory(); ReplayResultStore replayResults = new ReplayResultStore(); VideoRenderRecord videoRenderRecord = new VideoRenderRecord(); this.taskCoordinator = new TaskCoordinator(messageSender, replayResults, discordBridgeService); this.authHelper = new OsuAuthHelper(startupConfig.binding()); BindingCommandHandler bindingCommands = new BindingCommandHandler(startupConfig, replyFactory, bindingService); ScoreCommandHandler scoreCommands = new ScoreCommandHandler( - resolver, history, taskCoordinator, replyFactory + resolver, history, taskCoordinator, replyFactory, this::getAccessTokenFor ); BeatmapCommandHandler beatmapCommands = new BeatmapCommandHandler( resolver, history, taskCoordinator, replyFactory, videoRenderRecord, this::getAccessTokenFor ); SocialCommandHandler socialCommands = new SocialCommandHandler( - resolver, authHelper, taskCoordinator, replyFactory, this::getAccessTokenFor + resolver, authHelper, taskCoordinator, replyFactory, this::getAccessTokenFor, this::getAvatar ); ReplayCommandHandler replayCommands = new ReplayCommandHandler( - resolver, - history, - taskCoordinator, - replyFactory, - videoRenderRecord, - replayResults + resolver, history, taskCoordinator, replyFactory, videoRenderRecord, replayResults, admins::isAdmin, this::getAccessTokenFor ); GeneralCommandHandler generalCommands = new GeneralCommandHandler( messageSender, taskCoordinator, replyFactory, resolver, admins::isAdmin ); + this.aiChatHandler = new AiChatHandler( + resolver, chatProvider, admins::isAdmin, botStateGetter + ); WatchCommandHandler watchCommands = new WatchCommandHandler(resolver, taskCoordinator, watchService, admins::isAdmin); SpecificScoreWatchCommandHandler specificScoreWatchCommands = new SpecificScoreWatchCommandHandler(taskCoordinator, watchService); - MultiplayerRoomWatchCommandHandler multiplayerRoomWatchCommands = - new MultiplayerRoomWatchCommandHandler(taskCoordinator, multiplayerRoomWatchService); + MPWatchCommandHandler multiplayerRoomWatchCommands = + new MPWatchCommandHandler(taskCoordinator, resolver, mpWatchService); DcsCommandHandler dcsCommands = new DcsCommandHandler(discordBridgeService); RankGuessCommandHandler rankGuessCommands = new RankGuessCommandHandler( - taskCoordinator, replyFactory, rankGuessGameService, resolver, admins::isAdmin + taskCoordinator, replyFactory, rankGuessGameService, resolver, admins::isAdmin, this::getAvatar, imageUploader ); this.unknownCommand = generalCommands::handleUnknown; this.commandParser = new CommandParser(resolver::sanitize); this.commandRegistry = createCommandRegistry( - bindingCommands, - scoreCommands, - beatmapCommands, - socialCommands, - replayCommands, - generalCommands, - watchCommands, - specificScoreWatchCommands, - multiplayerRoomWatchCommands, - dcsCommands, - rankGuessCommands + bindingCommands, scoreCommands, beatmapCommands, socialCommands, + replayCommands, generalCommands, watchCommands, specificScoreWatchCommands, + multiplayerRoomWatchCommands, dcsCommands, rankGuessCommands, aiChatHandler ); this.debugRoutes = new DebugRoutes( - configSupplier, - messageSender, - replyFactory, - taskCoordinator, - authHelper, - admins::isAdmin, - unknownCommand + configSupplier, messageSender, replyFactory, taskCoordinator, + authHelper, admins::isAdmin, unknownCommand ); } private static CommandRegistry createCommandRegistry( - BindingCommandHandler bindingCommands, - ScoreCommandHandler scoreCommands, - BeatmapCommandHandler beatmapCommands, - SocialCommandHandler socialCommands, - ReplayCommandHandler replayCommands, - GeneralCommandHandler generalCommands, - WatchCommandHandler watchCommands, - SpecificScoreWatchCommandHandler specificScoreWatchCommands, - MultiplayerRoomWatchCommandHandler multiplayerRoomWatchCommands, - DcsCommandHandler dcsCommands, - RankGuessCommandHandler rankGuessCommands + BindingCommandHandler bindingCommands, ScoreCommandHandler scoreCommands, + BeatmapCommandHandler beatmapCommands, SocialCommandHandler socialCommands, + ReplayCommandHandler replayCommands, GeneralCommandHandler generalCommands, + WatchCommandHandler watchCommands, SpecificScoreWatchCommandHandler specificScoreWatchCommands, + MPWatchCommandHandler multiplayerRoomWatchCommands, DcsCommandHandler dcsCommands, + RankGuessCommandHandler rankGuessCommands, AiChatHandler aiChatHandler ) { return CommandRegistry.builder() .register(bindingCommands::handleBind, "bind") .register(bindingCommands::handleUnbind, "unbind") .register(bindingCommands::handleClearHistory, "clearhistory") + .register(scoreCommands::handleRbp, "rbp") .register(scoreCommands::handleBp, "bp") .register(beatmapCommands::handleDaily, "daily") .register(socialCommands::handleMp, "mp") @@ -159,6 +145,7 @@ private static CommandRegistry createCommandRegistry( .register(socialCommands::handleFclear, "fclear") .register(beatmapCommands::handleDl, "dl") .register(scoreCommands::handleS, "s") + .register(scoreCommands::handleSm, "sm") .register(scoreCommands::handleSa, "sa") .register(scoreCommands::handleMa, "ma") .register(replayCommands::handleR, "r") @@ -168,30 +155,51 @@ private static CommandRegistry createCommandRegistry( .register(socialCommands::handleLb, "lb") .register(generalCommands::handleStat, "stat") .register(generalCommands::handleU, "u") + .register(generalCommands::handleUx, "ux") .register(generalCommands::handleLuck, "luck") + .register(generalCommands::handleRoll, "roll") .register(replayCommands::handleRstat, "rstat") .register(replayCommands::handleRcancel, "rcancel") .register(generalCommands::handleInspect, "inspect") .register(generalCommands::handleHelp, "help") + .register(generalCommands::handleUsages, "usages") .register(generalCommands::handleFaq, "faq") .register(watchCommands::handleWatch, "watch") .register(specificScoreWatchCommands::handleWx, "wx") .register(multiplayerRoomWatchCommands::handleMpWatch, "mpwatch", "mpw") + .register(multiplayerRoomWatchCommands::handleRomAI, "romai") .register(dcsCommands::handleDcs, "dcs") .register(rankGuessCommands::handleRankGuess, "rg") .register(generalCommands::handleNotice, "notice") + .register(aiChatHandler::handleAi, "ai") + .register(generalCommands::handleMc, "mc") .build(); } - public void onPrivateMessageReceived(String userId, String messageId, String rawContent) { - handleMessageReceived(userId, null, userId, messageId, rawContent, false); + private String getAvatar(String openId) { + return QQApi.getAvatarUrl(configSupplier.get().qq().appId(), openId); } - public void onGroupMessageReceived(String groupId, String senderUserId, String messageId, String rawContent) { - handleMessageReceived(groupId, groupId, senderUserId, messageId, rawContent, true); + public void onPrivateMessageReceived( + String userId, String messageId, String rawContent, + String msgIdx ,List attachments, List msgElems + ) { + handleMessageReceived(userId, null, userId, messageId, rawContent, false, msgIdx, attachments, msgElems); } - private void handleMessageReceived(String targetId, String groupId, String userId, String messageId, String rawContent, boolean groupMessage) { + public void onGroupMessageReceived( + String groupId, String senderUserId, + String messageId, String rawContent, + String msgIdx, List attachments, + List msgElems + ) { + handleMessageReceived(groupId, groupId, senderUserId, messageId, rawContent, true, msgIdx, attachments, msgElems); + } + + private void handleMessageReceived( + String targetId, String groupId, String userId, String messageId, String rawContent, + boolean groupMessage, String msgIdx, List attachments, List msgElems + ) { AtomicInteger messageSeqCounter = new AtomicInteger(1); try { final boolean group = groupMessage && groupId != null && !groupId.isBlank(); @@ -201,34 +209,82 @@ private void handleMessageReceived(String targetId, String groupId, String userI rawContent = rawContent == null ? "" : rawContent.trim(); + String msgToRecord = rawContent; + + boolean beingAt = false; + AppConfig config = configSupplier.get(); final String selfAt = "<@" + config.qq().selfId() + ">"; + + if (rawContent.contains(selfAt)) { + beingAt = true; + msgToRecord = msgToRecord.replace(selfAt, "@Seira"); + } + if (rawContent.startsWith(selfAt)) { rawContent = rawContent.substring(selfAt.length()).trim(); } + final QQUser qqUser = selfSupplier.get(); + if (qqUser != null) { + final String selfLiteralAt = "@" + qqUser.username(); + + if (rawContent.contains(selfLiteralAt)) { + beingAt = true; + } + + if (rawContent.startsWith(selfLiteralAt)) { + rawContent = rawContent.substring(selfLiteralAt.length()).trim(); + } + } + CommandParser.ParseResult parseResult = commandParser.parse( rawContent, userId, groupId, messageId ); + if (parseResult.status() == CommandParser.ParseResult.Status.IGNORED) { return; } - CommandReplyChannel replies = taskCoordinator.openReplyChannel( - targetId, - messageId, - groupMessage, - config.seira().queueMessageInGroup() + final boolean permitAi = AiPermission.permits(groupId); + final boolean activatedAi = AiPermission.isActivated(groupId); + + if (parseResult.status() == CommandParser.ParseResult.Status.TEXT) { + ReplyChannel replies = taskCoordinator.openReplyChannel( + targetId, messageId, groupMessage, false, msgIdx + ); + + if (beingAt && permitAi && activatedAi) { + aiChatHandler.handleChat(parseResult.context().withReplies(replies), msgToRecord, msgElems); + } else if (permitAi) { + aiChatHandler.recordHistory(groupId, userId, msgToRecord, attachments); + } + return; + } + + ReplyChannel replies = taskCoordinator.openReplyChannel( + targetId, messageId, groupMessage, config.seira().queueMessageInGroup(), msgIdx ); + if (parseResult.status() == CommandParser.ParseResult.Status.EMPTY_COMMAND && !group) { replies.sendReply(PendingMessage.ofString("请输入指令。使用/help获取帮助。")); return; } - Context context = parseResult.context().withReplies(replies); + Context context = parseResult.context() + .withReplies(replies) + .withRecorder(s -> { + if (permitAi) { + aiChatHandler.recordHistory(groupId, "Seira(你,回复" + userId + "的消息)", s); + } + }); + commandExecutor.execute(() -> { try { + if (permitAi) { + aiChatHandler.recordHistory(groupId, userId, context.rawContent()); + } LOG.info("Routing {} message : {}", groupMessage ? "group" : "private", context.rawContent()); dispatch(context); NoticesHelper.checkNotices(context); diff --git a/src/main/java/xyz/zcraft/seira/config/AppConfig.java b/src/main/java/xyz/zcraft/seira/config/AppConfig.java index 835f45ee..a1a7264b 100644 --- a/src/main/java/xyz/zcraft/seira/config/AppConfig.java +++ b/src/main/java/xyz/zcraft/seira/config/AppConfig.java @@ -7,21 +7,14 @@ public record AppConfig( QqConfig qq, CosConfig cos, DiscordConfig discord, - BridgeConfig bridge + BridgeConfig bridge, + LLMConfig llm, + AsteroidConfig asteroid ) { public AppConfig { discord = discord == null ? DiscordConfig.disabled() : discord; bridge = bridge == null ? BridgeConfig.defaults() : bridge; - } - - public AppConfig( - SeiraConfig seira, - OstellaConfig ostella, - BindingConfig binding, - QqConfig qq, - CosConfig cos - ) { - this(seira, ostella, binding, qq, cos, null, null); + llm = llm == null ? new LLMConfig(null, null) : llm; } } diff --git a/src/main/java/xyz/zcraft/seira/config/AsteroidConfig.java b/src/main/java/xyz/zcraft/seira/config/AsteroidConfig.java new file mode 100644 index 00000000..57b0ab9e --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/config/AsteroidConfig.java @@ -0,0 +1,7 @@ +package xyz.zcraft.seira.config; + +public record AsteroidConfig( + String endpoint, + String token +) { +} diff --git a/src/main/java/xyz/zcraft/seira/config/LLMConfig.java b/src/main/java/xyz/zcraft/seira/config/LLMConfig.java new file mode 100644 index 00000000..aaceaa26 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/config/LLMConfig.java @@ -0,0 +1,7 @@ +package xyz.zcraft.seira.config; + +public record LLMConfig( + String baseUrl, + String apiKey +) { +} diff --git a/src/main/java/xyz/zcraft/seira/config/OstellaConfig.java b/src/main/java/xyz/zcraft/seira/config/OstellaConfig.java index f1b1b537..a27df811 100644 --- a/src/main/java/xyz/zcraft/seira/config/OstellaConfig.java +++ b/src/main/java/xyz/zcraft/seira/config/OstellaConfig.java @@ -1,6 +1,7 @@ package xyz.zcraft.seira.config; public record OstellaConfig( - String endpoint + String endpoint, + String token ) { } diff --git a/src/main/java/xyz/zcraft/seira/config/RuntimeConfig.java b/src/main/java/xyz/zcraft/seira/config/RuntimeConfig.java index 3252f236..a040769d 100644 --- a/src/main/java/xyz/zcraft/seira/config/RuntimeConfig.java +++ b/src/main/java/xyz/zcraft/seira/config/RuntimeConfig.java @@ -54,7 +54,9 @@ private static AppConfig mergeReloadable(AppConfig previous, AppConfig loaded) { previous.qq(), previous.cos(), previous.discord(), - previous.bridge() + previous.bridge(), + previous.llm(), + previous.asteroid() ); } diff --git a/src/main/java/xyz/zcraft/seira/console/ConsoleCommandProcessor.java b/src/main/java/xyz/zcraft/seira/console/ConsoleCommandProcessor.java index d4d2ae1f..5f8b51ef 100644 --- a/src/main/java/xyz/zcraft/seira/console/ConsoleCommandProcessor.java +++ b/src/main/java/xyz/zcraft/seira/console/ConsoleCommandProcessor.java @@ -9,6 +9,7 @@ import xyz.zcraft.seira.command.route.Router; import xyz.zcraft.seira.config.RuntimeConfig; import xyz.zcraft.seira.db.SqliteDatabase; +import xyz.zcraft.seira.services.AiPermission; import xyz.zcraft.seira.services.BotStat; import xyz.zcraft.seira.services.handler.ConfigHandler; import xyz.zcraft.seira.services.handler.NoticeHandler; @@ -21,6 +22,7 @@ import java.util.*; import java.util.concurrent.CompletableFuture; import java.util.concurrent.TimeUnit; +import java.util.function.Supplier; import java.util.regex.Pattern; public final class ConsoleCommandProcessor { @@ -30,27 +32,34 @@ public final class ConsoleCommandProcessor { private static final Pattern SQL_IDENTIFIER = Pattern.compile("[A-Za-z_][A-Za-z0-9_]*"); private static final List ROOT_COMMANDS = List.of( "help", "status", "metrics", "system", "config", "admin", "data", "send", - "watch", "cache", "gateway", "log", "inspect", "stop", "panel", "notice" - ); - private static final Map> SUBCOMMANDS = Map.of( - "config", List.of("show", "check", "reload"), - "admin", List.of("list", "check", "add", "remove"), - "data", List.of("stats", "tables", "describe", "query"), - "send", List.of("group", "private"), - "watch", List.of("status", "list", "poll", "remove", "clear"), - "cache", List.of("query", "delete", "get", "fetch"), - "gateway", List.of("status", "reconnect"), - "log", List.of("show", "level"), - "panel", List.of("list", "create", "delete", "edit", "get"), - "notice", List.of("new", "reload", "publish", "revoke", "list") + "watch", "cache", "gateway", "log", "inspect", "stop", "panel", "notice", "group", + "ai" ); + private static final Map> SUBCOMMANDS; + private static final Pattern ID_PATTERN = Pattern.compile("^[A-Z0-9]{32}$"); + + static { + SUBCOMMANDS = new HashMap<>(); + SUBCOMMANDS.put("config", List.of("show", "check", "reload")); + SUBCOMMANDS.put("admin", List.of("list", "check", "add", "remove")); + SUBCOMMANDS.put("data", List.of("stats", "tables", "describe", "query")); + SUBCOMMANDS.put("send", List.of("group", "private")); + SUBCOMMANDS.put("watch", List.of("status", "list", "poll", "remove", "clear")); + SUBCOMMANDS.put("cache", List.of("query", "delete", "get", "fetch")); + SUBCOMMANDS.put("gateway", List.of("status", "reconnect")); + SUBCOMMANDS.put("log", List.of("show", "level")); + SUBCOMMANDS.put("panel", List.of("list", "create", "delete", "edit", "get")); + SUBCOMMANDS.put("notice", List.of("new", "reload", "publish", "revoke", "list")); + SUBCOMMANDS.put("group", List.of("info", "state")); + SUBCOMMANDS.put("ai", List.of("reload")); + } + private final RuntimeConfig runtimeConfig; private final AdminRegistry admins; private final ConsoleDataAccess dataAccess; private final MessageSender messenger; private final ConsoleRuntimeControl runtimeControl; - private final ConfigHandler configHandler; private final NoticeHandler noticeHandler; @@ -93,10 +102,7 @@ private static String formatCacheControl(ConsoleRuntimeControl.CacheControlResul } private static ConsoleResult exact( - ConsoleInputParser.ParsedInput input, - int size, - java.util.function.Supplier action, - String usage + ConsoleInputParser.ParsedInput input, int size, Supplier action, String usage ) { return input.size() == size ? action.get() : ConsoleResult.failure(usage); } @@ -149,6 +155,17 @@ private static long positiveLong(String value, String name) { } } catch (NumberFormatException ignored) { } + throw new IllegalArgumentException(name + " must be a positive long."); + } + + private static int positiveInt(String value, String name) { + try { + int parsed = Integer.parseInt(value); + if (parsed > 0) { + return parsed; + } + } catch (NumberFormatException ignored) { + } throw new IllegalArgumentException(name + " must be a positive integer."); } @@ -236,7 +253,9 @@ public ConsoleResult execute(String line) { case "watch" -> watch(input); case "cache" -> cache(input); case "panel" -> panel(input); + case "group" -> group(input); case "gateway" -> gateway(input); + case "ai" -> ai(input); case "log" -> log(input); case "inspect" -> exact(input, 1, this::inspect, "Usage: inspect"); case "stop", "shutdown", "exit", "quit" -> stop(input); @@ -252,6 +271,61 @@ public ConsoleResult execute(String line) { } } + private ConsoleResult ai(ConsoleInputParser.ParsedInput input) { + if (input.size() == 1) { + final String value = input.value(0).toLowerCase(Locale.ROOT); + if (value.equalsIgnoreCase("reload")) { + AiPermission.loadFromFile(); + return ConsoleResult.success("AI permission reloaded."); + } else if (ID_PATTERN.matcher(value).matches()) { + return ConsoleResult.success("AI chat for group " + value + " is: " + AiPermission.permits(value)); + } else { + return ConsoleResult.failure("Group id " + value + " not a valid id."); + } + } else if (input.size() == 2) { + final String target = input.value(0).toLowerCase(Locale.ROOT); + final String option = input.value(1).toLowerCase(Locale.ROOT); + + if (!ID_PATTERN.matcher(target).matches()) { + return ConsoleResult.failure("Group id " + target + " not a valid id."); + } + + if ("on".equalsIgnoreCase(option)) { + AiPermission.permits(target); + AiPermission.activate(target); + return ConsoleResult.success("AI chat for group " + target + " is activated."); + } else if ("off".equalsIgnoreCase(option)) { + AiPermission.deactivate(target); + return ConsoleResult.success("AI chat for group " + target + " is deactivated."); + } else if ("grant".equalsIgnoreCase(option)) { + AiPermission.grant(target); + return ConsoleResult.success("AI chat for group " + target + " is granted."); + } else if ("revoke".equalsIgnoreCase(option)) { + AiPermission.revoke(target); + return ConsoleResult.success("AI chat for group " + target + " is revoked."); + } else if ("parallel".equalsIgnoreCase(option)) { + final int parallel = AiPermission.getParallel(target); + return ConsoleResult.success("Parallel count of AI chat for group " + target + " is " + parallel + "."); + } + } else if (input.size() == 3) { + final String target = input.value(0).toLowerCase(Locale.ROOT); + final String option = input.value(1).toLowerCase(Locale.ROOT); + final String value = input.value(2).toLowerCase(Locale.ROOT); + + if (!ID_PATTERN.matcher(target).matches()) { + return ConsoleResult.failure("Group id " + target + " not a valid id."); + } + + if (option.equalsIgnoreCase("parallel")) { + final int l = positiveInt(value, "parallel count"); + AiPermission.setParallel(target, l); + return ConsoleResult.success("Parallel count of AI chat for group " + target + " is now set to " + l + "."); + } + } + + return ConsoleResult.failure("Usage: ai "); + } + private ConsoleResult help(ConsoleInputParser.ParsedInput input) { if (input.size() > 2) { return ConsoleResult.failure("Usage: help [command]"); @@ -273,6 +347,7 @@ private ConsoleResult help(ConsoleInputParser.ParsedInput input) { Manage command panels cache Inspect/delete cache through oStella and workers + ai Manage AI chat status gateway Inspect or reconnect the QQ gateway log Inspect or change the runtime log level inspect Show the last dispatched message context @@ -318,9 +393,14 @@ private ConsoleResult help(ConsoleInputParser.ParsedInput input) { watch clear confirm Polling runs in the watch worker. Clearing a group requires explicit confirmation."""; case "cache" -> """ - cache + cache query reports presence at oStella and every osuRenderer worker. get includes metadata; fetch populates oStella and downstream workers; delete removes every reachable copy."""; + case "ai" -> """ + ai [on|off] + ai reload + Turn on or off AI chat function for specific group. + """; case "gateway" -> """ gateway status gateway reconnect @@ -442,6 +522,20 @@ private ConsoleResult admin(ConsoleInputParser.ParsedInput input) { }; } + private ConsoleResult group(ConsoleInputParser.ParsedInput input) { + if (input.size() < 2) { + return ConsoleResult.failure("Usage: group [args]"); + } + + return switch (input.value(1).toLowerCase(Locale.ROOT)) { + case "info" -> + input.size() == 3 ? runtimeControl.getGroupInfo(input.value(2)) : ConsoleResult.failure("Usage: group info "); + case "state" -> + input.size() == 3 ? runtimeControl.getGroupBotState(input.value(2)) : ConsoleResult.failure("Usage: group state "); + default -> ConsoleResult.failure("Usage: group [args]"); + }; + } + private ConsoleResult panel(ConsoleInputParser.ParsedInput input) { if (input.size() < 2) { return ConsoleResult.failure("Usage: panel [args]"); @@ -555,8 +649,8 @@ private ConsoleResult send(ConsoleInputParser.ParsedInput input) { } boolean sent = switch (targetType) { - case "group" -> messenger.sendGroupText(targetId, content) != null; - case "private" -> messenger.sendPrivateText(targetId, content) != null; + case "group" -> messenger.sendGroupMarkdown(targetId, content) != null; + case "private" -> messenger.sendPrivateMarkdown(targetId, content) != null; default -> throw new IllegalArgumentException("Message type must be 'group' or 'private'."); }; return sent @@ -655,7 +749,7 @@ private ConsoleResult gateway(ConsoleInputParser.ParsedInput input) { private ConsoleResult cache(ConsoleInputParser.ParsedInput input) { if (input.size() != 4) { return ConsoleResult.failure( - "Usage: cache " + "Usage: cache " ); } String operation = input.value(1).toLowerCase(Locale.ROOT); @@ -663,8 +757,8 @@ private ConsoleResult cache(ConsoleInputParser.ParsedInput input) { return ConsoleResult.failure("Cache operation must be query, delete, get, or fetch."); } String type = input.value(2).toUpperCase(Locale.ROOT); - if (!List.of("SCORE", "BEATMAP", "BEATMAPSET", "REPLAY").contains(type)) { - return ConsoleResult.failure("Cache type must be score, beatmap, beatmapset, or replay."); + if (!List.of("SCORE", "BEATMAP", "BEATMAPSET", "REPLAY", "BEATMAP-JSON", "BEATMAPSET-JSON").contains(type)) { + return ConsoleResult.failure("Cache type must be score, beatmap, beatmapset, replay, beatmap-json, or beatmapset-json."); } long id = positiveLong(input.value(3), "id"); return ConsoleResult.success(formatCacheControl(runtimeControl.controlCache(operation, type, id))); @@ -723,6 +817,10 @@ private ConsoleResult stop(ConsoleInputParser.ParsedInput input) { return ConsoleResult.success("Graceful shutdown requested."); } + public RuntimeConfig getRuntimeConfig() { + return runtimeConfig; + } + public record ConsoleResult(boolean success, String message) { public static ConsoleResult success(String message) { return new ConsoleResult(true, message); diff --git a/src/main/java/xyz/zcraft/seira/console/ConsoleRuntimeControl.java b/src/main/java/xyz/zcraft/seira/console/ConsoleRuntimeControl.java index 379a8e2c..b53fae1a 100644 --- a/src/main/java/xyz/zcraft/seira/console/ConsoleRuntimeControl.java +++ b/src/main/java/xyz/zcraft/seira/console/ConsoleRuntimeControl.java @@ -37,6 +37,10 @@ public interface ConsoleRuntimeControl { ConsoleCommandProcessor.ConsoleResult editPanel(String panelId, String jsonPath); + ConsoleCommandProcessor.ConsoleResult getGroupInfo(String groupId); + + ConsoleCommandProcessor.ConsoleResult getGroupBotState(String groupId); + record RuntimeStatus( boolean running, boolean gatewayConnected, diff --git a/src/main/java/xyz/zcraft/seira/console/OstellaCacheControlClient.java b/src/main/java/xyz/zcraft/seira/console/OstellaCacheControlClient.java index 26cb59cf..485cbfce 100644 --- a/src/main/java/xyz/zcraft/seira/console/OstellaCacheControlClient.java +++ b/src/main/java/xyz/zcraft/seira/console/OstellaCacheControlClient.java @@ -13,19 +13,27 @@ public final class OstellaCacheControlClient { private static final Gson GSON = new Gson(); private final HttpClient client = HttpClient.newBuilder().connectTimeout(Duration.ofSeconds(10)).build(); private final URI endpoint; + private final String serviceToken; public OstellaCacheControlClient(String endpoint) { + this(endpoint, null); + } + + public OstellaCacheControlClient(String endpoint, String serviceToken) { String normalized = endpoint.replaceAll("/+$", ""); this.endpoint = URI.create(normalized + "/cache/control"); + this.serviceToken = serviceToken; } public ConsoleRuntimeControl.CacheControlResult control(String operation, String type, long id) { String body = GSON.toJson(new Request(operation, type, id)); - HttpRequest request = HttpRequest.newBuilder(endpoint) + HttpRequest.Builder builder = HttpRequest.newBuilder(endpoint) .timeout(Duration.ofSeconds(40)) - .header("Content-Type", "application/json") - .POST(HttpRequest.BodyPublishers.ofString(body, StandardCharsets.UTF_8)) - .build(); + .header("Content-Type", "application/json"); + if (serviceToken != null && !serviceToken.isBlank()) { + builder.header("Authorization", "Bearer " + serviceToken); + } + HttpRequest request = builder.POST(HttpRequest.BodyPublishers.ofString(body, StandardCharsets.UTF_8)).build(); try { HttpResponse response = client.send(request, HttpResponse.BodyHandlers.ofString(StandardCharsets.UTF_8)); if (response.statusCode() < 200 || response.statusCode() >= 300) { diff --git a/src/main/java/xyz/zcraft/seira/data/UploadedImage.java b/src/main/java/xyz/zcraft/seira/data/UploadedImage.java index 62c41099..a4f2d0d3 100644 --- a/src/main/java/xyz/zcraft/seira/data/UploadedImage.java +++ b/src/main/java/xyz/zcraft/seira/data/UploadedImage.java @@ -4,4 +4,8 @@ public record UploadedImage(String url, int width, int height) { public String toMarkdown() { return "![image #" + width + "px #" + height + "px](" + url + ")"; } + + public String toMarkdown(int width, int height) { + return "![image #" + width + "px #" + height + "px](" + url + ")"; + } } diff --git a/src/main/java/xyz/zcraft/seira/data/UserRef.java b/src/main/java/xyz/zcraft/seira/data/UserRef.java deleted file mode 100644 index 87dc0cc1..00000000 --- a/src/main/java/xyz/zcraft/seira/data/UserRef.java +++ /dev/null @@ -1,26 +0,0 @@ -package xyz.zcraft.seira.data; - -import lombok.Getter; - -public class UserRef { - private UserRef() { - } - - public static class ByUid extends UserRef { - @Getter - private final long uid; - - public ByUid(long uid) { - this.uid = uid; - } - } - - public static class ByUsername extends UserRef { - @Getter - private final String username; - - public ByUsername(String username) { - this.username = username; - } - } -} diff --git a/src/main/java/xyz/zcraft/seira/db/UserDataStore.java b/src/main/java/xyz/zcraft/seira/db/UserDataStore.java index 19b204c3..56e29129 100644 --- a/src/main/java/xyz/zcraft/seira/db/UserDataStore.java +++ b/src/main/java/xyz/zcraft/seira/db/UserDataStore.java @@ -10,7 +10,6 @@ import java.sql.*; import java.util.*; -import java.util.stream.Collectors; public final class UserDataStore { private static final Logger LOG = LogManager.getLogger(UserDataStore.class); diff --git a/src/main/java/xyz/zcraft/seira/rankguess/HintUtil.java b/src/main/java/xyz/zcraft/seira/rankguess/HintUtil.java index 0b499022..4aa0cba3 100644 --- a/src/main/java/xyz/zcraft/seira/rankguess/HintUtil.java +++ b/src/main/java/xyz/zcraft/seira/rankguess/HintUtil.java @@ -106,9 +106,9 @@ private static RankGuessGame.Hint selectWeightedByStrength( } private static EnumMap strengthWeights(double progress) { - double[] first = {0, 48, 34, 15, 0, 0, 3}; - double[] middle = {0, 9, 28, 44, 15, 0, 4}; - double[] late = {0, 0, 4, 38, 48, 5, 5}; + double[] first = {0, 48, 34, 14, 0, 0, 4}; + double[] middle = {0, 7, 28, 44, 15, 0, 6}; + double[] late = {0, 0, 1, 38, 48, 5, 8}; double phase = progress <= 0.5 ? progress * 2 : (progress - 0.5) * 2; diff --git a/src/main/java/xyz/zcraft/seira/rankguess/RankGuessGameService.java b/src/main/java/xyz/zcraft/seira/rankguess/RankGuessGameService.java index cfc1a7b6..751db64c 100644 --- a/src/main/java/xyz/zcraft/seira/rankguess/RankGuessGameService.java +++ b/src/main/java/xyz/zcraft/seira/rankguess/RankGuessGameService.java @@ -85,12 +85,12 @@ public static boolean isOutstandingGuess( return Math.abs(guess - actualRank) <= allowedDifference; } - public static List getUsernameFeature(String username) { + public static List> getUsernameFeature(String username) { if (username == null || username.isBlank()) { return List.of(); } - final List features = new ArrayList<>(); + final List> features = new ArrayList<>(); final int leftBracket = username.indexOf("["); final int rightBracket = username.indexOf("]"); @@ -103,42 +103,42 @@ public static List getUsernameFeature(String username) { if (leftBracket == 0 && rightBracket < username.length() - 1) { // [Prefix]Example final String prefix = username.substring(0, rightBracket + 1); - features.add("有前缀 `" + prefix + "`"); + features.add(Map.entry("有前缀 `" + prefix + "`", 0.5)); } } if (username.charAt(0) == username.charAt(username.length() - 1)) { - features.add("为 `首尾一样`"); + features.add(Map.entry("为 `首尾一样`", 0.4)); } boolean hasLetter = username.chars().anyMatch(Character::isLetter); if (hasLetter) { if (Objects.equals(username, username.toUpperCase())) { - features.add("为 `全大写`"); + features.add(Map.entry("为 `全大写`", 0.2)); } else if (Objects.equals(username, username.toLowerCase())) { - features.add("为 `全小写`"); + features.add(Map.entry("为 `全小写`", 0.2)); } } - features.add("长度为 `" + username.length() + "`"); + features.add(Map.entry("长度为 `" + username.length() + "`", 0.1)); if (username.contains(" ")) { - features.add("有 `空格`"); + features.add(Map.entry("有 `空格`", 0.15)); } if (username.contains("_")) { - features.add("有 `下划线(_)`"); + features.add(Map.entry("有 `下划线(_)`", 0.15)); } if (username.contains("-")) { - features.add("有 `横杠(-)`"); + features.add(Map.entry("有 `横杠(-)`", 0.15)); } if (PREFIX_NUMBER_PATTERN.matcher(username).matches()) { - features.add("是 `一串数字一串字母`"); + features.add(Map.entry("是 `一串数字一串字母`", 0.25)); } else if (SUFFIX_NUMBER_PATTERN.matcher(username).matches()) { - features.add("是 `一串字母一串数字`"); + features.add(Map.entry("是 `一串字母一串数字`", 0.25)); } return features; diff --git a/src/main/java/xyz/zcraft/seira/rankguess/data/Round.java b/src/main/java/xyz/zcraft/seira/rankguess/data/Round.java index 0942c3b5..4fc4eeb2 100644 --- a/src/main/java/xyz/zcraft/seira/rankguess/data/Round.java +++ b/src/main/java/xyz/zcraft/seira/rankguess/data/Round.java @@ -3,13 +3,18 @@ import xyz.zcraft.osu.model.Score; import xyz.zcraft.osu.model.UserExtended; import xyz.zcraft.seira.api.data.RandomScore; +import xyz.zcraft.seira.data.UploadedImage; import xyz.zcraft.seira.rankguess.RankGuessGame; import xyz.zcraft.seira.rankguess.RankGuessGameService; +import xyz.zcraft.seira.util.ImageUtil; +import xyz.zcraft.seira.util.WeightedRandom; +import java.awt.image.BufferedImage; import java.util.ArrayList; import java.util.Collections; import java.util.LinkedList; import java.util.List; +import java.util.function.Function; public record Round(long userId, long scoreId, int bestIndex, long actualRank, Double pp, RandomScore randomScore, boolean standard) { @@ -191,80 +196,134 @@ public LinkedList getNormalHints() { return hints; } - public LinkedList getGroupHints() { - LinkedList hints = new LinkedList<>(); + public List getGroupHints(String qqAvatarUrl, Function imageUploader) { + WeightedRandom hintRandom = new WeightedRandom<>(); final UserExtended user = this.randomScore.user(); final UserExtended.Team team = user.getTeam(); if (team != null && team.getName() != null && team.getShortName() != null) { - hints.add(new RankGuessGame.Hint( + hintRandom.add(new RankGuessGame.Hint( "该玩家所处的队伍缩写为 `%s`".formatted(team.getShortName()), "玩家队伍缩写", RankGuessGame.Hint.HintCategory.USER, RankGuessGame.Hint.HintStrength.SPECIAL - )); + ), 2.0); } if (user.getHasSupported() && user.isSupporter()) { - hints.add(new RankGuessGame.Hint( + hintRandom.add(new RankGuessGame.Hint( "该玩家是尊贵的撒泼特!", "支持者状态", RankGuessGame.Hint.HintCategory.USER, RankGuessGame.Hint.HintStrength.SPECIAL - )); + ), 1.5); } else if (user.getHasSupported() && !user.isSupporter()) { - hints.add(new RankGuessGame.Hint( + hintRandom.add(new RankGuessGame.Hint( "该玩家的撒泼特已经过期了。", "支持者状态", RankGuessGame.Hint.HintCategory.USER, RankGuessGame.Hint.HintStrength.SPECIAL - )); + ), 1.5); } if (user.getInterests() != null && !user.getInterests().isBlank()) { - hints.add(new RankGuessGame.Hint( + hintRandom.add(new RankGuessGame.Hint( "该玩家填写的兴趣爱好为 `%s`".formatted(user.getInterests()), "玩家自述兴趣", RankGuessGame.Hint.HintCategory.USER, RankGuessGame.Hint.HintStrength.SPECIAL - )); + ), 0.4); } if (user.getLocation() != null && !user.getLocation().isBlank()) { - hints.add(new RankGuessGame.Hint( + hintRandom.add(new RankGuessGame.Hint( "该玩家填写的位置为 `%s`".formatted(user.getLocation()), "玩家自述位置", RankGuessGame.Hint.HintCategory.USER, RankGuessGame.Hint.HintStrength.SPECIAL - )); + ), 0.4); } if (user.getOccupation() != null && !user.getOccupation().isBlank()) { - hints.add(new RankGuessGame.Hint( + hintRandom.add(new RankGuessGame.Hint( "该玩家填写的职业为 `%s`".formatted(user.getOccupation()), "玩家自述职业", RankGuessGame.Hint.HintCategory.USER, RankGuessGame.Hint.HintStrength.SPECIAL - )); + ), 0.4); } - final List features = new ArrayList<>(RankGuessGameService.getUsernameFeature(user.getUsername())); + final var features = new ArrayList<>(RankGuessGameService.getUsernameFeature(user.getUsername())); if (!features.isEmpty()) { Collections.shuffle(features); for (int i = 0; i < Math.min(features.size(), 4); i++) { - final String s = features.get(i); - hints.add(new RankGuessGame.Hint( - "该玩家用户名" + s, + final var s = features.get(i); + hintRandom.add(new RankGuessGame.Hint( + "该玩家用户名" + s.getKey(), "用户名特征", RankGuessGame.Hint.HintCategory.USER, RankGuessGame.Hint.HintStrength.SPECIAL - )); + ), s.getValue()); } } - return hints; + if (qqAvatarUrl != null) { + try { + final BufferedImage original = ImageUtil.readImage(qqAvatarUrl); + + final byte[] mosaicBytes = ImageUtil.toPngBytes(ImageUtil.mosaic(original, 18)); + final byte[] blurBytes = ImageUtil.toPngBytes(ImageUtil.gaussianBlur(original, 60)); + + final var mosaicImage = imageUploader.apply(mosaicBytes); + final var blurImage = imageUploader.apply(blurBytes); + + hintRandom.add(new RankGuessGame.Hint( + "该玩家 QQ 头像: " + mosaicImage.toMarkdown(25, 25), + "玩家 QQ 头像", + RankGuessGame.Hint.HintCategory.USER, + RankGuessGame.Hint.HintStrength.SPECIAL + ), 1.25); + + hintRandom.add(new RankGuessGame.Hint( + "该玩家 QQ 头像: " + blurImage.toMarkdown(25, 25), + "玩家 QQ 头像", + RankGuessGame.Hint.HintCategory.USER, + RankGuessGame.Hint.HintStrength.SPECIAL + ), 1.25); + } catch (Exception ignored) { + // Ignored + } + } + + try { + final BufferedImage original = ImageUtil.readImage(randomScore.user().getAvatarUrl()); + + final byte[] mosaicBytes = ImageUtil.toPngBytes(ImageUtil.mosaic(original, 18)); + final byte[] blurBytes = ImageUtil.toPngBytes(ImageUtil.gaussianBlur(original, 60)); + + final var mosaicImage = imageUploader.apply(mosaicBytes); + final var blurImage = imageUploader.apply(blurBytes); + + hintRandom.add(new RankGuessGame.Hint( + "该玩家 osu! 头像: " + mosaicImage.toMarkdown(25, 25), + "玩家 osu! 头像", + RankGuessGame.Hint.HintCategory.USER, + RankGuessGame.Hint.HintStrength.SPECIAL + ), 1.25); + + hintRandom.add(new RankGuessGame.Hint( + "该玩家 osu! 头像: " + blurImage.toMarkdown(25, 25), + "玩家 osu! 头像", + RankGuessGame.Hint.HintCategory.USER, + RankGuessGame.Hint.HintStrength.SPECIAL + ), 1.25); + } catch (Exception ignored) { + // Ignored + } + + return List.of(hintRandom.next()); } } diff --git a/src/main/java/xyz/zcraft/seira/services/AiPermission.java b/src/main/java/xyz/zcraft/seira/services/AiPermission.java new file mode 100644 index 00000000..93e32608 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/services/AiPermission.java @@ -0,0 +1,182 @@ +package xyz.zcraft.seira.services; + +import com.google.gson.Gson; +import com.google.gson.GsonBuilder; +import com.google.gson.JsonObject; +import com.google.gson.JsonParser; +import lombok.Getter; + +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.HashMap; +import java.util.HashSet; +import java.util.Map; +import java.util.Set; +import java.util.concurrent.ConcurrentHashMap; + +public class AiPermission { + private static final Object LOCK = new Object(); + private static final Gson GSON = new GsonBuilder().setPrettyPrinting().create(); + private static final Path STORE = Path.of("data", "ai-permission.json"); + private static Set groups = null; + @Getter + private static Mode mode = Mode.WHITELIST; + private static Set activated = null; + private static Map parallel = null; + + public static void initialize() { + loadFromFile(); + } + + public static void loadFromFile() { + synchronized (LOCK) { + try { + if (Files.exists(STORE)) { + String json = Files.readString(STORE); + + if (!json.isBlank()) { + JsonObject obj = JsonParser.parseString(json).getAsJsonObject(); + + final PermissionSnapshot snapshot = GSON.fromJson(obj, PermissionSnapshot.class); + + AiPermission.mode = snapshot.mode; + AiPermission.groups = new HashSet<>(); + AiPermission.activated = new HashSet<>(); + AiPermission.parallel = new HashMap<>(); + + final Set groups = snapshot.groups; + if (groups != null) { + AiPermission.groups.addAll(groups); + } + + final Set activated = snapshot.activated; + if (activated != null) { + AiPermission.activated.addAll(activated); + } + + final Map parallel = snapshot.parallel; + if (parallel != null) { + AiPermission.parallel.putAll(parallel); + } + } + } + } catch (Exception e) { + throw new RuntimeException("Failed to load ai permission", e); + } + + if (mode == null) { + mode = Mode.WHITELIST; + } + + if (groups == null) { + groups = new HashSet<>(); + } + + if (activated == null) { + activated = new HashSet<>(); + } + + if (parallel == null) { + parallel = new ConcurrentHashMap<>(); + } + } + } + + public static void saveToFile() { + synchronized (LOCK) { + try { + Files.createDirectories(STORE.getParent()); + + PermissionSnapshot snapshot = new PermissionSnapshot(mode, groups, activated, parallel); + + Files.writeString(STORE, GSON.toJson(snapshot)); + } catch (Exception e) { + throw new RuntimeException("Failed to load ai permission", e); + } + } + } + + public static boolean permits(String groupId) { + if (mode == null || groups == null) return false; + + if (mode == Mode.WHITELIST) { + return groups.contains(groupId); + } else if (mode == Mode.BLACKLIST) { + return !groups.contains(groupId); + } else { + return false; + } + } + + public static void grant(String groupId) { + if (mode == null || groups == null) return; + + if (mode == Mode.WHITELIST) { + groups.add(groupId); + } else if (mode == Mode.BLACKLIST) { + groups.remove(groupId); + } + + saveToFile(); + } + + public static void revoke(String groupId) { + if (mode == null || groups == null) return; + + if (mode == Mode.WHITELIST) { + groups.remove(groupId); + } else if (mode == Mode.BLACKLIST) { + groups.add(groupId); + } + + if (activated != null) { + activated.remove(groupId); + } + + saveToFile(); + } + + public static int getParallel(String groupId) { + if (parallel == null) return 1; + return parallel.getOrDefault(groupId, 1); + } + + public static void activate(String groupId) { + if (activated == null) return; + activated.add(groupId); + + saveToFile(); + } + + public static void deactivate(String groupId) { + if (activated == null) return; + activated.remove(groupId); + + saveToFile(); + } + + public static boolean isActivated(String groupId) { + if (activated == null) return false; + return activated.contains(groupId); + } + + public static void setParallel(String groupId, int parallel) { + if (AiPermission.parallel == null) return; + AiPermission.parallel.put(groupId, parallel); + + saveToFile(); + } + + public enum Mode { + WHITELIST, + BLACKLIST + } + + public record PermissionSnapshot( + Mode mode, + Set groups, + Set activated, + Map parallel + ) { + } +} diff --git a/src/main/java/xyz/zcraft/seira/util/ImageUtil.java b/src/main/java/xyz/zcraft/seira/util/ImageUtil.java new file mode 100644 index 00000000..84211ccc --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/util/ImageUtil.java @@ -0,0 +1,222 @@ +package xyz.zcraft.seira.util; + +import javax.imageio.ImageIO; +import java.awt.*; +import java.awt.image.BufferedImage; +import java.io.ByteArrayOutputStream; +import java.net.URI; + +import static java.lang.Math.clamp; + +public class ImageUtil { + public static BufferedImage readImage(String url) throws Exception { + return ImageIO.read(URI.create(url).toURL()); + } + + public static BufferedImage crop(BufferedImage image, int x, int y, int width, int height) { + BufferedImage sub = image.getSubimage(x, y, width, height); + + BufferedImage result = new BufferedImage( + width, height, BufferedImage.TYPE_INT_ARGB + ); + + result.getGraphics().drawImage(sub, 0, 0, null); + + return result; + } + + public static BufferedImage gaussianBlur(BufferedImage image, int radius) { + if (radius <= 0) { + return image; + } + + int width = image.getWidth(); + int height = image.getHeight(); + + float[] kernel = createGaussianKernel(radius); + + int[] source = image.getRGB( + 0, 0, + width, height, + null, + 0, width + ); + + int[] temp = new int[source.length]; + int[] result = new int[source.length]; + + blurHorizontal( + source, + temp, + width, + height, + kernel, + radius + ); + + blurVertical( + temp, + result, + width, + height, + kernel, + radius + ); + + BufferedImage output = new BufferedImage( + width, + height, + BufferedImage.TYPE_INT_ARGB + ); + + output.setRGB( + 0, 0, + width, height, + result, + 0, width + ); + + return output; + } + + private static float[] createGaussianKernel(int radius) { + int size = radius * 2 + 1; + float[] kernel = new float[size]; + + // 常见经验值 + double sigma = Math.max(radius / 3.0, 0.1); + + double sum = 0; + + for (int i = -radius; i <= radius; i++) { + double value = Math.exp( + -(i * i) / (2.0 * sigma * sigma) + ); + + kernel[i + radius] = (float) value; + sum += value; + } + + // 归一化 + for (int i = 0; i < size; i++) { + kernel[i] /= (float) sum; + } + + return kernel; + } + + private static void blurHorizontal( + int[] source, + int[] target, + int width, + int height, + float[] kernel, + int radius + ) { + for (int y = 0; y < height; y++) { + int row = y * width; + + for (int x = 0; x < width; x++) { + float a = 0; + float r = 0; + float g = 0; + float b = 0; + + for (int k = -radius; k <= radius; k++) { + int sampleX = clamp(x + k, 0, width - 1); + int argb = source[row + sampleX]; + + float weight = kernel[k + radius]; + + a += ((argb >>> 24) & 0xff) * weight; + r += ((argb >>> 16) & 0xff) * weight; + g += ((argb >>> 8) & 0xff) * weight; + b += (argb & 0xff) * weight; + } + + target[row + x] = + ((clamp(Math.round(a), 0, 255)) << 24) + | ((clamp(Math.round(r), 0, 255)) << 16) + | ((clamp(Math.round(g), 0, 255)) << 8) + | clamp(Math.round(b), 0, 255); + } + } + } + + private static void blurVertical( + int[] source, + int[] target, + int width, + int height, + float[] kernel, + int radius + ) { + for (int y = 0; y < height; y++) { + for (int x = 0; x < width; x++) { + float a = 0; + float r = 0; + float g = 0; + float b = 0; + + for (int k = -radius; k <= radius; k++) { + int sampleY = clamp(y + k, 0, height - 1); + int argb = source[sampleY * width + x]; + + float weight = kernel[k + radius]; + + a += ((argb >>> 24) & 0xff) * weight; + r += ((argb >>> 16) & 0xff) * weight; + g += ((argb >>> 8) & 0xff) * weight; + b += (argb & 0xff) * weight; + } + + target[y * width + x] = + ((clamp(Math.round(a), 0, 255)) << 24) + | ((clamp(Math.round(r), 0, 255)) << 16) + | ((clamp(Math.round(g), 0, 255)) << 8) + | clamp(Math.round(b), 0, 255); + } + } + } + + public static BufferedImage mosaic(BufferedImage image, int blockSize) { + int width = image.getWidth(); + int height = image.getHeight(); + + int smallWidth = Math.max(1, width / blockSize); + int smallHeight = Math.max(1, height / blockSize); + + BufferedImage small = new BufferedImage(smallWidth, smallHeight, BufferedImage.TYPE_INT_ARGB); + + Graphics2D g1 = small.createGraphics(); + g1.setRenderingHint(RenderingHints.KEY_INTERPOLATION, RenderingHints.VALUE_INTERPOLATION_NEAREST_NEIGHBOR); + g1.drawImage(image, 0, 0, smallWidth, smallHeight, null); + g1.dispose(); + + BufferedImage result = new BufferedImage(width, height, BufferedImage.TYPE_INT_ARGB); + + Graphics2D g2 = result.createGraphics(); + g2.setRenderingHint(RenderingHints.KEY_INTERPOLATION, RenderingHints.VALUE_INTERPOLATION_NEAREST_NEIGHBOR); + g2.drawImage(small, 0, 0, width, height, null); + g2.dispose(); + + return result; + } + + public static void mosaicRegion(BufferedImage image, int x, int y, int width, int height, int blockSize) { + BufferedImage region = image.getSubimage(x, y, width, height); + + BufferedImage processed = mosaic(region, blockSize); + + Graphics2D g = image.createGraphics(); + g.drawImage(processed, x, y, null); + g.dispose(); + } + + public static byte[] toPngBytes(BufferedImage image) throws Exception { + try (ByteArrayOutputStream out = new ByteArrayOutputStream()) { + ImageIO.write(image, "png", out); + return out.toByteArray(); + } + } +} diff --git a/src/main/java/xyz/zcraft/seira/util/TokenManager.java b/src/main/java/xyz/zcraft/seira/util/TokenManager.java index 87fb90bb..aecc987e 100644 --- a/src/main/java/xyz/zcraft/seira/util/TokenManager.java +++ b/src/main/java/xyz/zcraft/seira/util/TokenManager.java @@ -22,6 +22,7 @@ public class TokenManager implements AutoCloseable { private final AtomicBoolean started = new AtomicBoolean(); private final AtomicBoolean closed = new AtomicBoolean(); + @Getter private final String clientId; private final String clientSecret; diff --git a/src/main/java/xyz/zcraft/seira/util/WeightedRandom.java b/src/main/java/xyz/zcraft/seira/util/WeightedRandom.java new file mode 100644 index 00000000..33889aaf --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/util/WeightedRandom.java @@ -0,0 +1,70 @@ +package xyz.zcraft.seira.util; + +import java.util.ArrayList; +import java.util.List; +import java.util.NoSuchElementException; +import java.util.concurrent.ThreadLocalRandom; + +public class WeightedRandom { + private final List items = new ArrayList<>(); + private final List cumulative = new ArrayList<>(); + + public void add(T item, double weight) { + if (weight <= 0 || Double.isNaN(weight) || Double.isInfinite(weight)) { + throw new IllegalArgumentException("Weight must be positive and finite"); + } + + items.add(item); + + double previous = cumulative.isEmpty() ? 0.0 : cumulative.getLast(); + + cumulative.add(previous + weight); + } + + public T next() { + return items.get(randomIndex()); + } + + public T getAndRemove() { + int index = randomIndex(); + + T result = items.remove(index); + + double previous = index == 0 ? 0.0 : cumulative.get(index - 1); + + double removedWeight = cumulative.get(index) - previous; + + cumulative.remove(index); + + for (int i = index; i < cumulative.size(); i++) { + cumulative.set(i, cumulative.get(i) - removedWeight); + } + + return result; + } + + private int randomIndex() { + if (items.isEmpty()) { + throw new NoSuchElementException("WeightedRandom is empty"); + } + + double total = cumulative.getLast(); + double random = ThreadLocalRandom.current().nextDouble(total); + + for (int i = 0; i < cumulative.size(); i++) { + if (random < cumulative.get(i)) { + return i; + } + } + + throw new IllegalStateException(); + } + + public boolean isEmpty() { + return items.isEmpty(); + } + + public int size() { + return items.size(); + } +} \ No newline at end of file diff --git a/src/main/java/xyz/zcraft/seira/util/dice/Dice.java b/src/main/java/xyz/zcraft/seira/util/dice/Dice.java new file mode 100644 index 00000000..59b6dc14 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/util/dice/Dice.java @@ -0,0 +1,15 @@ +package xyz.zcraft.seira.util.dice; + +import xyz.zcraft.seira.util.dice.expr.DiceExpr; +import xyz.zcraft.seira.util.dice.result.DiceResult; + +import java.util.LinkedList; +import java.util.List; + +public class Dice { + public DiceResult roll(DiceExpr expression) { + final List> results = new LinkedList<>(); + expression.parts().forEach(e -> results.add(e.calculate())); + return new DiceResult(results); + } +} diff --git a/src/main/java/xyz/zcraft/seira/util/dice/expr/DiceExpr.java b/src/main/java/xyz/zcraft/seira/util/dice/expr/DiceExpr.java new file mode 100644 index 00000000..a462417b --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/util/dice/expr/DiceExpr.java @@ -0,0 +1,83 @@ +package xyz.zcraft.seira.util.dice.expr; + +import org.jetbrains.annotations.NotNull; +import xyz.zcraft.seira.util.dice.expr.part.DiceExprPart; +import xyz.zcraft.seira.util.dice.expr.part.DiceExprPartComplex; +import xyz.zcraft.seira.util.dice.expr.part.DiceExprPartNumber; +import xyz.zcraft.seira.util.dice.expr.part.DiceExprPartSingle; + +import java.util.LinkedList; +import java.util.List; + +public record DiceExpr(List parts) { + public static final DiceExpr CHECK; + public static final DiceExpr HUNDRED; + public static final DiceExpr TEN; + public static final DiceExpr SIX; + + static { + CHECK = new DiceExpr(List.of(new DiceExprPartSingle("1d100"))); + HUNDRED = new DiceExpr(List.of(new DiceExprPartSingle("1d100"))); + TEN = new DiceExpr(List.of(new DiceExprPartSingle("1d10"))); + SIX = new DiceExpr(List.of(new DiceExprPartSingle("1d6"))); + } + + public static DiceExpr parse(String str) { + str = str.replace(" ", ""); + + if (isInvalid(str)) { + throw new IllegalArgumentException("骰子表达式有误"); + } + + if (!str.startsWith("-")) str = "+" + str; + final LinkedList compStr = new LinkedList<>(); + int last = 0; + for (int i = 1; i < str.length(); i++) { + if (str.charAt(i) == '+' || str.charAt(i) == '-') { + compStr.add(str.substring(last, i)); + last = i; + } + } + + if (last != str.length() - 1) compStr.add(str.substring(last)); + + final LinkedList parts = new LinkedList<>(); + + compStr.forEach(s -> { + if (isDigit(s)) { + parts.add(new DiceExprPartNumber(s)); + } else if (s.contains("*")) { + parts.add(new DiceExprPartComplex(s)); + } else { + parts.add(new DiceExprPartSingle(s)); + } + }); + + return new DiceExpr(parts); + } + + public static boolean isDigit(String str) { + if (str == null) return false; + if (str.startsWith("+") || str.startsWith("-")) str = str.substring(1); + return str.chars().allMatch(value -> value >= '0' && value <= '9'); + } + + public static boolean isInvalid(String exp) { + if (exp == null || exp.isBlank()) return true; + return !exp.chars().allMatch(value -> (value >= '0' && value <= '9') || + value == 'd' || value == '+' || value == '*'); + } + + @NotNull + @Override + public String toString() { + final StringBuilder sb = new StringBuilder(); + if (parts.getFirst().isNegative()) sb.append("-"); + sb.append(parts.getFirst().toString(false)); + for (int i = 1; i < parts.size(); i++) { + sb.append(parts.get(i).toString(true)); + } + + return sb.toString(); + } +} diff --git a/src/main/java/xyz/zcraft/seira/util/dice/expr/part/DiceExprPart.java b/src/main/java/xyz/zcraft/seira/util/dice/expr/part/DiceExprPart.java new file mode 100644 index 00000000..edaf43fd --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/util/dice/expr/part/DiceExprPart.java @@ -0,0 +1,25 @@ +package xyz.zcraft.seira.util.dice.expr.part; + +import lombok.Getter; + +import java.util.List; + +@Getter +public abstract class DiceExprPart { + protected final boolean negative; + protected final String origString; + + protected DiceExprPart(boolean negative, String origString) { + this.negative = negative; + this.origString = origString; + } + + public abstract List calculate(); + + public String toString(boolean signed) { + if ((origString.startsWith("-") || origString.startsWith("+")) && !signed) { + return origString.substring(1); + } + return origString; + } +} diff --git a/src/main/java/xyz/zcraft/seira/util/dice/expr/part/DiceExprPartComplex.java b/src/main/java/xyz/zcraft/seira/util/dice/expr/part/DiceExprPartComplex.java new file mode 100644 index 00000000..4ec2c542 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/util/dice/expr/part/DiceExprPartComplex.java @@ -0,0 +1,39 @@ +package xyz.zcraft.seira.util.dice.expr.part; + +import xyz.zcraft.seira.util.dice.expr.DiceExpr; + +import java.util.LinkedList; +import java.util.List; + +public class DiceExprPartComplex extends DiceExprPart { + final LinkedList factors; + + public DiceExprPartComplex(String origString) { + super(origString.startsWith("-"), origString); + + if (origString.startsWith("-") || origString.startsWith("+")) { + origString = origString.substring(1); + } + + final String[] split = origString.split("\\*"); + + factors = new LinkedList<>(); + + for (String s : split) { + if (DiceExpr.isDigit(s)) { + factors.add(new DiceExprPartNumber(s)); + } else { + factors.add(new DiceExprPartSingle(s)); + } + } + } + + @Override + public List calculate() { + LinkedList result = new LinkedList<>(); + for (DiceExprPart factor : factors) { + result.add(factor.calculate().getFirst()); + } + return result; + } +} diff --git a/src/main/java/xyz/zcraft/seira/util/dice/expr/part/DiceExprPartNumber.java b/src/main/java/xyz/zcraft/seira/util/dice/expr/part/DiceExprPartNumber.java new file mode 100644 index 00000000..e2e4390f --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/util/dice/expr/part/DiceExprPartNumber.java @@ -0,0 +1,17 @@ +package xyz.zcraft.seira.util.dice.expr.part; + +import java.util.List; + +public class DiceExprPartNumber extends DiceExprPart { + final int num; + + public DiceExprPartNumber(String string) { + super(Integer.parseInt(string) < 0, string); + this.num = Integer.parseInt(string); + } + + @Override + public List calculate() { + return List.of(num); + } +} diff --git a/src/main/java/xyz/zcraft/seira/util/dice/expr/part/DiceExprPartSingle.java b/src/main/java/xyz/zcraft/seira/util/dice/expr/part/DiceExprPartSingle.java new file mode 100644 index 00000000..93a8dba6 --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/util/dice/expr/part/DiceExprPartSingle.java @@ -0,0 +1,31 @@ +package xyz.zcraft.seira.util.dice.expr.part; + +import java.util.List; +import java.util.Random; + +public class DiceExprPartSingle extends DiceExprPart { + final int number; + final int faces; + + public DiceExprPartSingle(String str) { + super(str.startsWith("-"), str); + if (str.startsWith("-") || str.startsWith("+")) + str = str.substring(1); + + final int split = str.indexOf("d"); + this.number = Integer.parseInt(str.substring(0, split)); + this.faces = Integer.parseInt(str.substring(split + 1)); + } + + @Override + public List calculate() { + int result = 0; + for (int i = 0; i < number; i++) { + result += new Random().nextInt(faces) + 1; + } + + if (negative) result = -result; + + return List.of(result); + } +} diff --git a/src/main/java/xyz/zcraft/seira/util/dice/result/DiceResult.java b/src/main/java/xyz/zcraft/seira/util/dice/result/DiceResult.java new file mode 100644 index 00000000..d469086a --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/util/dice/result/DiceResult.java @@ -0,0 +1,64 @@ +package xyz.zcraft.seira.util.dice.result; + +import org.jetbrains.annotations.NotNull; + +import java.util.List; + +public class DiceResult { + private final List> parts; + + public DiceResult(List> parts) { + this.parts = parts; + } + + public DiceResult(int num) { + this.parts = List.of(List.of(num)); + } + + public int total() { + int sum = 0; + for (var factors : parts) { + int tmp = 1; + for (var factor : factors) { + tmp *= factor; + } + sum += tmp; + } + + return sum; + } + + @NotNull + @Override + public String toString() { + final StringBuilder sb = new StringBuilder(); + for (int i = 0; i < parts.size(); i++) { + if(i == 0) { + if (parts.get(i).size() < 2) { + sb.append(parts.get(i).getFirst()); + continue; + } + } else { + if (parts.get(i).size() >= 2) sb.append("+"); + else { + if(parts.get(i).getFirst() >= 0) { + sb.append("+"); + } + sb.append(parts.get(i).getFirst()); + continue; + } + } + + for (int j = 0; j < parts.get(i).size(); j++) { + final int cur = parts.get(i).get(j); + if (j != 0) { + sb.append("*"); + } + if (cur >= 0) sb.append(cur); + else sb.append("(-").append(cur).append(")"); + } + } + + return sb.toString(); + } +} diff --git a/src/main/java/xyz/zcraft/seira/watch/MPNotifier.java b/src/main/java/xyz/zcraft/seira/watch/MPNotifier.java new file mode 100644 index 00000000..0bc4543b --- /dev/null +++ b/src/main/java/xyz/zcraft/seira/watch/MPNotifier.java @@ -0,0 +1,36 @@ +package xyz.zcraft.seira.watch; + +import xyz.zcraft.seira.bot.MessageSender; +import xyz.zcraft.seira.data.UploadedImage; + +import java.util.Objects; + +public final class MPNotifier { + private final MessageSender messageSender; + + public MPNotifier(MessageSender messageSender) { + this.messageSender = Objects.requireNonNull(messageSender); + } + + public boolean sendResult(MPWatchService.WatchEntry watch, String groupId, byte[] imageBytes) { + final UploadedImage uploadedImage = messageSender.uploadImageToCos(imageBytes); + + if (uploadedImage == null) return false; + + return messageSender.sendGroupMarkdown( + groupId, + """ + %s + > __%s__ + > %s mp - %s + """.formatted(uploadedImage.toMarkdown(), watch.roomName(), watch.version().value(), watch.roomId()).trim() + ) != null; + } + + public boolean sendRoomEnded(String groupId, RoomWatchSnapshot snapshot) { + return messageSender.sendGroupText( + groupId, + "多人房间“" + snapshot.roomName() + "” (#" + snapshot.roomId() + ") 已结束,监视已自动停止。" + ) != null; + } +} diff --git a/src/main/java/xyz/zcraft/seira/watch/MultiplayerRoomVersion.java b/src/main/java/xyz/zcraft/seira/watch/MPVersion.java similarity index 76% rename from src/main/java/xyz/zcraft/seira/watch/MultiplayerRoomVersion.java rename to src/main/java/xyz/zcraft/seira/watch/MPVersion.java index 0ab2ac8f..42db4c1c 100644 --- a/src/main/java/xyz/zcraft/seira/watch/MultiplayerRoomVersion.java +++ b/src/main/java/xyz/zcraft/seira/watch/MPVersion.java @@ -2,17 +2,17 @@ import java.util.Locale; -public enum MultiplayerRoomVersion { +public enum MPVersion { LAZER("lazer"), STABLE("stable"); private final String value; - MultiplayerRoomVersion(String value) { + MPVersion(String value) { this.value = value; } - public static MultiplayerRoomVersion parse(String value) { + public static MPVersion parse(String value) { if (value == null) { return null; } diff --git a/src/main/java/xyz/zcraft/seira/watch/OstellaMultiplayerRoomWatchApi.java b/src/main/java/xyz/zcraft/seira/watch/MPWatchApi.java similarity index 77% rename from src/main/java/xyz/zcraft/seira/watch/OstellaMultiplayerRoomWatchApi.java rename to src/main/java/xyz/zcraft/seira/watch/MPWatchApi.java index 57c69ac3..cf6b7fd0 100644 --- a/src/main/java/xyz/zcraft/seira/watch/OstellaMultiplayerRoomWatchApi.java +++ b/src/main/java/xyz/zcraft/seira/watch/MPWatchApi.java @@ -12,17 +12,27 @@ import java.nio.charset.StandardCharsets; import java.time.Duration; -public final class OstellaMultiplayerRoomWatchApi implements MultiplayerRoomWatchApi { +public final class MPWatchApi { private final String endpoint; + private final String serviceToken; private final HttpClient client; private final Gson gson; - public OstellaMultiplayerRoomWatchApi(String endpoint) { - this(endpoint, HttpClient.newBuilder().connectTimeout(Duration.ofSeconds(30)).build(), new Gson()); + public MPWatchApi(String endpoint) { + this(endpoint, null); } - OstellaMultiplayerRoomWatchApi(String endpoint, HttpClient client, Gson gson) { + public MPWatchApi(String endpoint, String serviceToken) { + this(endpoint, serviceToken, HttpClient.newBuilder().connectTimeout(Duration.ofSeconds(30)).build(), new Gson()); + } + + MPWatchApi(String endpoint, HttpClient client, Gson gson) { + this(endpoint, null, client, gson); + } + + MPWatchApi(String endpoint, String serviceToken, HttpClient client, Gson gson) { this.endpoint = endpoint.endsWith("/") ? endpoint.substring(0, endpoint.length() - 1) : endpoint; + this.serviceToken = serviceToken; this.client = client; this.gson = gson; } @@ -61,15 +71,14 @@ private static long requirePositive(long value, String name) { return value; } - private static MultiplayerRoomVersion requireVersion(MultiplayerRoomVersion version) { + private static MPVersion requireVersion(MPVersion version) { if (version == null) { throw new IllegalArgumentException("version is required"); } return version; } - @Override - public RoomWatchSnapshot getSnapshot(MultiplayerRoomVersion version, long roomId) { + public RoomWatchSnapshot getSnapshot(MPVersion version, long roomId) { HttpResponse response = get( "/multiplayer/rooms/" + requirePositive(roomId, "roomId") + "/watch?version=" + requireVersion(version).value(), @@ -88,12 +97,12 @@ public RoomWatchSnapshot getSnapshot(MultiplayerRoomVersion version, long roomId return snapshot; } - @Override - public byte[] renderResult(MultiplayerRoomVersion version, long roomId, long playlistItemId) { + public byte[] renderResult(MPVersion version, long roomId, long playlistItemId, Integer customBo) { HttpResponse response = get( "/multiplayer/rooms/" + requirePositive(roomId, "roomId") + "/playlist/" + requirePositive(playlistItemId, "playlistItemId") + "/result?version=" - + requireVersion(version).value(), + + requireVersion(version).value() + + (customBo == null ? "" : "&bo=" + customBo), HttpResponse.BodyHandlers.ofByteArray() ); ensureSuccessfulStatus(response.statusCode(), response.body(), "生成多人房间结果图片"); @@ -104,12 +113,14 @@ public byte[] renderResult(MultiplayerRoomVersion version, long roomId, long pla } private HttpResponse get(String path, HttpResponse.BodyHandler handler) { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest.Builder builder = HttpRequest.newBuilder() .uri(URI.create(endpoint + path)) .timeout(Duration.ofMinutes(2)) - .header("Accept", "application/json, image/*") - .GET() - .build(); + .header("Accept", "image/png"); + if (serviceToken != null && !serviceToken.isBlank()) { + builder.header("Authorization", "Bearer " + serviceToken); + } + HttpRequest request = builder.GET().build(); try { return client.send(request, handler); } catch (InterruptedException e) { diff --git a/src/main/java/xyz/zcraft/seira/watch/MultiplayerRoomWatchService.java b/src/main/java/xyz/zcraft/seira/watch/MPWatchService.java similarity index 85% rename from src/main/java/xyz/zcraft/seira/watch/MultiplayerRoomWatchService.java rename to src/main/java/xyz/zcraft/seira/watch/MPWatchService.java index 27c2c466..62b8a165 100644 --- a/src/main/java/xyz/zcraft/seira/watch/MultiplayerRoomWatchService.java +++ b/src/main/java/xyz/zcraft/seira/watch/MPWatchService.java @@ -10,22 +10,20 @@ import java.util.concurrent.TimeUnit; import java.util.concurrent.atomic.AtomicBoolean; -public final class MultiplayerRoomWatchService implements AutoCloseable { - private static final Logger LOG = LogManager.getLogger(MultiplayerRoomWatchService.class); +public final class MPWatchService implements AutoCloseable { + private static final Logger LOG = LogManager.getLogger(MPWatchService.class); private final Object lock = new Object(); private final Map> watchesByGroup = new LinkedHashMap<>(); - private final MultiplayerRoomWatchApi api; - private final MultiplayerRoomNotifier notifier; + private final MPWatchApi api; + private final MPNotifier notifier; private final Duration pollInterval; private final ScheduledExecutorService scheduler; private final AtomicBoolean started = new AtomicBoolean(); private final AtomicBoolean closed = new AtomicBoolean(); - public MultiplayerRoomWatchService( - MultiplayerRoomWatchApi api, - MultiplayerRoomNotifier notifier, - Duration pollInterval + public MPWatchService( + MPWatchApi api, MPNotifier notifier, Duration pollInterval ) { this.api = Objects.requireNonNull(api); this.notifier = Objects.requireNonNull(notifier); @@ -69,30 +67,40 @@ public void start() { ); } - public RoomWatchView watch( - String groupId, - String userId, - MultiplayerRoomVersion version, - long roomId - ) { + public RoomWatchView watch(String groupId, String userId, MPVersion version, long roomId, Integer customBo) { requireIdentifier(groupId, "groupId"); requireIdentifier(userId, "userId"); - Objects.requireNonNull(version); + if (roomId <= 0) { throw new IllegalArgumentException("房间 ID 必须为正整数。"); } - RoomKey room = new RoomKey(version, roomId); + RoomKey room; + RoomWatchSnapshot snapshot; + + if (version == null) { + try { + version = MPVersion.STABLE; + snapshot = api.getSnapshot(MPVersion.STABLE, roomId); + } catch (IllegalStateException e) { + version = MPVersion.LAZER; + snapshot = api.getSnapshot(MPVersion.LAZER, roomId); + } + } else { + snapshot = api.getSnapshot(version, roomId); + } + + room = new RoomKey(version, roomId); synchronized (lock) { ensureRoomAvailable(groupId, userId, room); } - RoomWatchSnapshot snapshot = api.getSnapshot(version, roomId); + if (!snapshot.active()) { throw new IllegalStateException("该多人房间已经结束,无法开始监视。"); } Set baseline = new LinkedHashSet<>(); snapshot.completedPlays().forEach(play -> baseline.add(play.playlistItemId())); - WatchEntry entry = new WatchEntry(version, snapshot.roomId(), snapshot.roomName(), baseline); + WatchEntry entry = new WatchEntry(version, snapshot.roomId(), snapshot.roomName(), baseline, customBo); synchronized (lock) { ensureRoomAvailable(groupId, userId, room); watchesByGroup.computeIfAbsent(groupId, ignored -> new LinkedHashMap<>()) @@ -124,7 +132,7 @@ public List stopAll(String groupId) { if (removed == null) { return List.of(); } - return removed.values().stream().map(MultiplayerRoomWatchService::view).toList(); + return removed.values().stream().map(MPWatchService::view).toList(); } } @@ -181,9 +189,7 @@ private void processRoom(RoomKey room, List watches) { } private boolean sendPendingResults( - WatchRef watch, - List completed, - Map rendered + WatchRef watch, List completed, Map rendered ) { for (CompletedRoomPlay play : completed) { if (wasSent(watch, play.playlistItemId())) { @@ -192,9 +198,9 @@ private boolean sendPendingResults( try { byte[] image = rendered.computeIfAbsent( play.playlistItemId(), - itemId -> api.renderResult(watch.entry().version, watch.entry().roomId, itemId) + itemId -> api.renderResult(watch.entry().version, watch.entry().roomId, itemId, watch.entry().customBo()) ); - if (!notifier.sendResult(watch.groupId(), image)) { + if (!notifier.sendResult(watch.entry(), watch.groupId(), image)) { LOG.warn( "Failed to send room {} playlist item {} to group {}", watch.entry().roomId, play.playlistItemId(), watch.groupId() @@ -290,24 +296,23 @@ public void close() { } } - private record WatchEntry(MultiplayerRoomVersion version, long roomId, String roomName, - Set sentPlaylistItemIds) { - private WatchEntry( - MultiplayerRoomVersion version, - long roomId, - String roomName, - Set sentPlaylistItemIds + public record WatchEntry( + MPVersion version, long roomId, String roomName, Set sentPlaylistItemIds, Integer customBo + ) { + public WatchEntry( + MPVersion version, long roomId, String roomName, Set sentPlaylistItemIds, Integer customBo ) { this.version = version; this.roomId = roomId; this.roomName = roomName; this.sentPlaylistItemIds = new LinkedHashSet<>(sentPlaylistItemIds); + this.customBo = customBo; } } - private record WatchRef(String groupId, String userId, WatchEntry entry) { + public record WatchRef(String groupId, String userId, WatchEntry entry) { } - private record RoomKey(MultiplayerRoomVersion version, long roomId) { + public record RoomKey(MPVersion version, long roomId) { } } diff --git a/src/main/java/xyz/zcraft/seira/watch/MultiplayerRoomNotifier.java b/src/main/java/xyz/zcraft/seira/watch/MultiplayerRoomNotifier.java deleted file mode 100644 index 81394d5c..00000000 --- a/src/main/java/xyz/zcraft/seira/watch/MultiplayerRoomNotifier.java +++ /dev/null @@ -1,7 +0,0 @@ -package xyz.zcraft.seira.watch; - -public interface MultiplayerRoomNotifier { - boolean sendResult(String groupId, byte[] imageBytes); - - boolean sendRoomEnded(String groupId, RoomWatchSnapshot snapshot); -} diff --git a/src/main/java/xyz/zcraft/seira/watch/MultiplayerRoomWatchApi.java b/src/main/java/xyz/zcraft/seira/watch/MultiplayerRoomWatchApi.java deleted file mode 100644 index c724e1b2..00000000 --- a/src/main/java/xyz/zcraft/seira/watch/MultiplayerRoomWatchApi.java +++ /dev/null @@ -1,7 +0,0 @@ -package xyz.zcraft.seira.watch; - -public interface MultiplayerRoomWatchApi { - RoomWatchSnapshot getSnapshot(MultiplayerRoomVersion version, long roomId); - - byte[] renderResult(MultiplayerRoomVersion version, long roomId, long playlistItemId); -} diff --git a/src/main/java/xyz/zcraft/seira/watch/QqMultiplayerRoomNotifier.java b/src/main/java/xyz/zcraft/seira/watch/QqMultiplayerRoomNotifier.java deleted file mode 100644 index bd2f6639..00000000 --- a/src/main/java/xyz/zcraft/seira/watch/QqMultiplayerRoomNotifier.java +++ /dev/null @@ -1,42 +0,0 @@ -package xyz.zcraft.seira.watch; - -import xyz.zcraft.seira.bot.MessageSender; -import xyz.zcraft.seira.bot.data.FileInfo; -import xyz.zcraft.seira.bot.data.Message; -import xyz.zcraft.seira.bot.data.PendingMessage; - -import java.util.Base64; -import java.util.Objects; - -public final class QqMultiplayerRoomNotifier implements MultiplayerRoomNotifier { - private final MessageSender messageSender; - - public QqMultiplayerRoomNotifier(MessageSender messageSender) { - this.messageSender = Objects.requireNonNull(messageSender); - } - - @Override - public boolean sendResult(String groupId, byte[] imageBytes) { - FileInfo media = messageSender.uploadGroupMediaBase64( - groupId, - PendingMessage.FILE_TYPE_IMAGE, - Base64.getEncoder().encodeToString(imageBytes) - ); - if (media == null) { - return false; - } - - Message message = new Message(); - message.setMsgType(PendingMessage.MSG_TYPE_MEDIA); - message.setMedia(media); - return messageSender.sendGroupMessage(groupId, message) != null; - } - - @Override - public boolean sendRoomEnded(String groupId, RoomWatchSnapshot snapshot) { - return messageSender.sendGroupText( - groupId, - "多人房间“" + snapshot.roomName() + "” (#" + snapshot.roomId() + ") 已结束,监视已自动停止。" - ) != null; - } -} diff --git a/src/main/java/xyz/zcraft/seira/watch/RoomWatchView.java b/src/main/java/xyz/zcraft/seira/watch/RoomWatchView.java index bf906517..8168c6f8 100644 --- a/src/main/java/xyz/zcraft/seira/watch/RoomWatchView.java +++ b/src/main/java/xyz/zcraft/seira/watch/RoomWatchView.java @@ -1,4 +1,4 @@ package xyz.zcraft.seira.watch; -public record RoomWatchView(MultiplayerRoomVersion version, long roomId, String roomName) { +public record RoomWatchView(MPVersion version, long roomId, String roomName) { } diff --git a/src/main/java/xyz/zcraft/seira/watch/OstellaWatchApi.java b/src/main/java/xyz/zcraft/seira/watch/ScoreWatchApi.java similarity index 82% rename from src/main/java/xyz/zcraft/seira/watch/OstellaWatchApi.java rename to src/main/java/xyz/zcraft/seira/watch/ScoreWatchApi.java index e06a8f4d..62bb1796 100644 --- a/src/main/java/xyz/zcraft/seira/watch/OstellaWatchApi.java +++ b/src/main/java/xyz/zcraft/seira/watch/ScoreWatchApi.java @@ -18,20 +18,30 @@ import java.util.List; import java.util.Map; -public class OstellaWatchApi implements WatchApi { +public class ScoreWatchApi { private static final Type SCORE_MAP_TYPE = new TypeToken>>() { }.getType(); private final String endpoint; + private final String serviceToken; private final HttpClient client; private final Gson gson; - public OstellaWatchApi(String endpoint) { - this(endpoint, HttpClient.newBuilder().connectTimeout(Duration.ofSeconds(30)).build(), new Gson()); + public ScoreWatchApi(String endpoint) { + this(endpoint, null); } - OstellaWatchApi(String endpoint, HttpClient client, Gson gson) { + public ScoreWatchApi(String endpoint, String serviceToken) { + this(endpoint, serviceToken, HttpClient.newBuilder().connectTimeout(Duration.ofSeconds(30)).build(), new Gson()); + } + + ScoreWatchApi(String endpoint, HttpClient client, Gson gson) { + this(endpoint, null, client, gson); + } + + ScoreWatchApi(String endpoint, String serviceToken, HttpClient client, Gson gson) { this.endpoint = endpoint.endsWith("/") ? endpoint.substring(0, endpoint.length() - 1) : endpoint; + this.serviceToken = serviceToken; this.client = client; this.gson = gson; } @@ -58,7 +68,6 @@ private static void ensureSuccessfulStatus(int statusCode, Object body, String a throw new IllegalStateException(action + "失败: HTTP " + statusCode + " " + detail); } - @Override public Map> getRecentScores(Collection userIds, int limit) { JsonObject body = new JsonObject(); body.add("user_ids", gson.toJsonTree(userIds)); @@ -91,7 +100,6 @@ public Map> getRecentScores(Collection userIds, in return Map.copyOf(scores); } - @Override public byte[] renderScore(long userId, long scoreId) { JsonObject body = new JsonObject(); body.addProperty("name", Long.toString(userId)); @@ -110,13 +118,14 @@ public byte[] renderScore(long userId, long scoreId) { } private HttpResponse send(String path, String body, HttpResponse.BodyHandler handler) { - HttpRequest request = HttpRequest.newBuilder() + HttpRequest.Builder builder = HttpRequest.newBuilder() .uri(URI.create(endpoint + path)) .timeout(Duration.ofMinutes(2)) - .header("Content-Type", "application/json") - .header("Accept", "application/json, image/*") - .POST(HttpRequest.BodyPublishers.ofString(body, StandardCharsets.UTF_8)) - .build(); + .header("Content-Type", "application/json"); + if (serviceToken != null && !serviceToken.isBlank()) { + builder.header("Authorization", "Bearer " + serviceToken); + } + HttpRequest request = builder.POST(HttpRequest.BodyPublishers.ofString(body, StandardCharsets.UTF_8)).build(); try { return client.send(request, handler); } catch (InterruptedException e) { diff --git a/src/main/java/xyz/zcraft/seira/watch/ScoreWatchService.java b/src/main/java/xyz/zcraft/seira/watch/ScoreWatchService.java index d0631a6d..86c04c4a 100644 --- a/src/main/java/xyz/zcraft/seira/watch/ScoreWatchService.java +++ b/src/main/java/xyz/zcraft/seira/watch/ScoreWatchService.java @@ -20,7 +20,7 @@ public final class ScoreWatchService implements AutoCloseable { private final Object lock = new Object(); private final Map> watchesByGroup = new LinkedHashMap<>(); private final Map specificWatchesByGroup = new LinkedHashMap<>(); - private final WatchApi api; + private final ScoreWatchApi api; private final WatchScoreNotifier notifier; private final SpecificScoreNotifier specificNotifier; private final SpecificScoreWatchStore specificWatchStore; @@ -31,7 +31,7 @@ public final class ScoreWatchService implements AutoCloseable { private final AtomicBoolean closed = new AtomicBoolean(); public ScoreWatchService( - WatchApi api, + ScoreWatchApi api, WatchScoreNotifier notifier, SpecificScoreNotifier specificNotifier, SpecificScoreWatchStore specificWatchStore, @@ -320,7 +320,7 @@ private void sendNewScores(WatchRef watch, List scores, Map api.renderScore(watch.entry.target.userId(), score.scoreId()) ); - if (!notifier.sendScore(watch.groupId, image)) { + if (!notifier.sendScore(watch.groupId, score, image)) { LOG.warn("Failed to send watched score {} to group {}", score.scoreId(), watch.groupId); return; } diff --git a/src/main/java/xyz/zcraft/seira/watch/WatchApi.java b/src/main/java/xyz/zcraft/seira/watch/WatchApi.java deleted file mode 100644 index 30bc9d11..00000000 --- a/src/main/java/xyz/zcraft/seira/watch/WatchApi.java +++ /dev/null @@ -1,14 +0,0 @@ -package xyz.zcraft.seira.watch; - -import java.util.Collection; -import java.util.List; -import java.util.Map; - -/** - * Backend boundary used by the score watch domain service. - */ -public interface WatchApi { - Map> getRecentScores(Collection userIds, int limit); - - byte[] renderScore(long userId, long scoreId); -} diff --git a/src/main/java/xyz/zcraft/seira/watch/WatchScoreNotifier.java b/src/main/java/xyz/zcraft/seira/watch/WatchScoreNotifier.java index 9003f73c..4181115c 100644 --- a/src/main/java/xyz/zcraft/seira/watch/WatchScoreNotifier.java +++ b/src/main/java/xyz/zcraft/seira/watch/WatchScoreNotifier.java @@ -1,13 +1,12 @@ package xyz.zcraft.seira.watch; import xyz.zcraft.seira.bot.MessageSender; -import xyz.zcraft.seira.bot.data.FileInfo; -import xyz.zcraft.seira.bot.data.Message; -import xyz.zcraft.seira.bot.data.PendingMessage; +import xyz.zcraft.seira.data.UploadedImage; -import java.util.Base64; import java.util.Objects; +import static xyz.zcraft.seira.command.reply.ReplyFactory.s; + public final class WatchScoreNotifier { private final MessageSender messageSender; @@ -15,20 +14,17 @@ public WatchScoreNotifier(MessageSender messageSender) { this.messageSender = Objects.requireNonNull(messageSender); } - public boolean sendScore(String groupId, byte[] imageBytes) { - String base64 = Base64.getEncoder().encodeToString(imageBytes); - FileInfo media = messageSender.uploadGroupMediaBase64( + public boolean sendScore(String groupId, RecentScore score, byte[] imageBytes) { + final UploadedImage uploadedImage = messageSender.uploadImageToCos(imageBytes); + + if (uploadedImage == null) return false; + + return messageSender.sendGroupMarkdown( groupId, - PendingMessage.FILE_TYPE_IMAGE, - base64 - ); - if (media == null) { - return false; - } - - Message message = new Message(); - message.setMsgType(PendingMessage.MSG_TYPE_MEDIA); - message.setMedia(media); - return messageSender.sendGroupMessage(groupId, message) != null; + """ + %s + > ID %s + """.formatted(uploadedImage.toMarkdown(), s(score.scoreId())).trim() + ) != null; } } diff --git a/src/main/resources/seira-example-config.yml b/src/main/resources/seira-example-config.yml index c83dd5a8..b3c9d8e1 100644 --- a/src/main/resources/seira-example-config.yml +++ b/src/main/resources/seira-example-config.yml @@ -29,6 +29,7 @@ binding: ostella: # oStella API 的地址,由 Seira 访问,默认为本地运行的 oStella 实例 endpoint: "http://localhost:8721" + token: # === QQ机器人 配置 === qq: @@ -70,3 +71,13 @@ bridge: # 可用占位符:{name}、{id}、{message} qqToDiscordFormat: "[{name}] {message}" discordToQqFormat: "[{name}] {message}" + +llm: + # LLM API 的地址 + baseUrl: "http://localhost:8723" + apiKey: "" + +asteroid: + # asteroid API 的地址 + endpoint: "http://localhost:8728" + token: