From 678006e4dc93e0e1805fb9c1a6c79ec1d2efc2bc Mon Sep 17 00:00:00 2001 From: This-is-XiaoDeng <1744793737@qq.com> Date: Fri, 14 Aug 2026 07:55:32 +0800 Subject: [PATCH 1/5] =?UTF-8?q?feat:=20=E7=BC=93=E5=AD=98=E7=9B=AE?= =?UTF-8?q?=E5=BD=95=E5=8F=AF=E9=85=8D=E7=BD=AE=E4=B8=8E=E8=BF=87=E6=9C=9F?= =?UTF-8?q?=E6=B8=85=E7=90=86=E3=80=81reply=20=E4=B8=89=E5=B1=82=E8=A7=A3?= =?UTF-8?q?=E6=9E=90=E3=80=81=E6=8E=A5=E6=94=B6=E6=96=B9=E5=90=91=E8=BD=AC?= =?UTF-8?q?=E5=8F=91=E8=87=AA=E5=8A=A8=E5=90=88=E5=B9=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - system.cache_dir 支持修改缓存目录(默认 .cache),database 未配置时自动使用缓存目录下的 onedisc.db - system.cache 支持配置 files/url_cache/db 三类缓存过期时间(默认不过期)与清理间隔,启动后定时自动清理 - reply 段解析升级为三层:cached_messages → 本地 DB → 同频道 REST 预检 - 接收方向(Discord → OneBot)转发消息自动合并:同一频道同一发送者窗口内(默认 500ms)的转发合并为一条 forward 上报,get_forward_msg 通过第一条消息调 Discord API 拉取上下文无状态重建 - 测试:test_cache_and_reply(10 例)、test_forward_merge(33 例)、test_native_forward(13 例)全部通过 - docs/config.md 与 docs/diff-v11.md 同步更新 --- actions/v11/basic.py | 15 +- actions/v11/get_image.py | 4 +- actions/v12/basic.py | 2 +- actions/v12/file.py | 56 ++++---- docs/config.md | 47 +++++- docs/diff-v11.md | 2 +- pyproject.toml | 2 +- test_cache_and_reply.py | 189 ++++++++++++++++++++++++ test_forward_merge.py | 266 ++++++++++++++++++++++++++++++++++ utils/cache.py | 112 +++++++++++++++ utils/db.py | 34 ++++- utils/event/discord_event.py | 9 ++ utils/forward_merge.py | 270 +++++++++++++++++++++++++++++++++++ utils/message/v12/parser.py | 41 +++++- 14 files changed, 1004 insertions(+), 45 deletions(-) create mode 100644 test_cache_and_reply.py create mode 100644 test_forward_merge.py create mode 100644 utils/cache.py create mode 100644 utils/forward_merge.py diff --git a/actions/v11/basic.py b/actions/v11/basic.py index 0542060..91ec439 100644 --- a/actions/v11/basic.py +++ b/actions/v11/basic.py @@ -13,6 +13,8 @@ import utils.translator as translator import utils.message.v11.parser as parser import utils.return_object as return_object +from utils.cache import get_cache_dir +import utils.forward_merge as forward_merge from utils.config import config from utils.client import client import utils.message.v11.parser as parser @@ -71,9 +73,9 @@ async def delete_msg(message_id: int) -> dict: def clean_node_cache() -> None: - for file in os.listdir(".cache"): + for file in os.listdir(get_cache_dir()): if file.startswith("node."): - os.remove(os.path.join(".cache", file)) + os.remove(os.path.join(get_cache_dir(), file)) @register_action("v11") @@ -352,6 +354,15 @@ async def send_private_forward_msg(user_id: int, messages: list) -> dict: ) +@register_action("v11") +async def get_forward_msg(message_id: str) -> dict: + """获取合并转发消息(接收方向:转发消息自动合并后由框架按 id 取回)""" + nodes = await forward_merge.get_forward(message_id) + if nodes is None: + return return_object.get(400, f"合并转发消息 {message_id} 不存在") + return return_object.get(0, message_id=message_id, message=nodes) + + async def _restart() -> None: script = sys.argv[0] args = sys.argv[1:] diff --git a/actions/v11/get_image.py b/actions/v11/get_image.py index e471ea2..df38e89 100644 --- a/actions/v11/get_image.py +++ b/actions/v11/get_image.py @@ -3,7 +3,7 @@ import httpx from ..v12.file import get_file_name_by_id from utils import return_object -from pathlib import Path +from utils.cache import get_file_path @register_action("v11") @@ -11,4 +11,4 @@ async def get_image(file: str) -> dict: file_name = await get_file_name_by_id(file.split("_")[0]) if not file_name: return return_object.get(31001, f"文件 {file} 不存在") - return return_object.get(0, file=Path(".cache/file").joinpath(file_name).as_posix()) + return return_object.get(0, file=get_file_path(file_name)) diff --git a/actions/v12/basic.py b/actions/v12/basic.py index b0ba6f4..771071d 100644 --- a/actions/v12/basic.py +++ b/actions/v12/basic.py @@ -54,7 +54,7 @@ async def send_message( if not (channel := client.get_channel(int(_channel_id))): logger.warning(f"频道 {group_id} 不存在") return return_object.get(35001, "频道(群号)不存在") - parsed_message = await parser.parse_message(message) + parsed_message = await parser.parse_message(message, channel.id) if _channel_id not in commands.deferred_sessions: try: msg = await channel.send(**parsed_message) # type: ignore diff --git a/actions/v12/file.py b/actions/v12/file.py index e2e429e..5f879f8 100644 --- a/actions/v12/file.py +++ b/actions/v12/file.py @@ -10,21 +10,23 @@ import hashlib import json import httpx +import time import utils.return_object as return_object +from utils.cache import files_dir, file_list_path, cached_url_path, get_file_path try: - os.makedirs(".cache/files") + files_dir() except OSError: pass try: - json.load(open(".cache/file_list.json", "r", encoding="utf-8")) + json.load(open(file_list_path(), "r", encoding="utf-8")) except Exception: - json.dump({}, open(".cache/file_list.json", "w", encoding="utf-8")) + json.dump({}, open(file_list_path(), "w", encoding="utf-8")) try: - json.load(open(".cache/cached_url.json", "r", encoding="utf-8")) + json.load(open(cached_url_path(), "r", encoding="utf-8")) except Exception: - json.dump({}, open(".cache/cached_url.json", "w", encoding="utf-8")) + json.dump({}, open(cached_url_path(), "w", encoding="utf-8")) logger = get_logger() @@ -36,10 +38,12 @@ def verify_sha256(content: bytes, sha256: str | None) -> bool: def create_url_cache(name: str, url: str) -> str: - with open(".cache/cached_url.json", "r", encoding="utf-8") as f: + with open(cached_url_path(), "r", encoding="utf-8") as f: cache = json.load(f) - cache[file_id := create_file_id()] = {"name": name, "url": url} - with open(".cache/cached_url.json", "w", encoding="utf-8") as f: + cache[file_id := create_file_id()] = { + "name": name, "url": url, "time": int(time.time()) + } + with open(cached_url_path(), "w", encoding="utf-8") as f: json.dump(cache, f) return file_id @@ -56,7 +60,7 @@ async def upload_file_from_url( async with httpx.AsyncClient(proxies=proxy) as client: response = await client.get(url, headers=headers) if response.status_code == 200 and verify_sha256(response.content, sha256): - with open(f".cache/files/{name}", "wb") as f: + with open(get_file_path(name), "wb") as f: f.write(response.content) return True logger.warning( @@ -75,7 +79,7 @@ async def upload_file_from_url( def upload_file_from_data(name: str, data: str) -> tuple[bool, str]: try: - with open(f".cache/files/{name}", "wb") as f: + with open(get_file_path(name), "wb") as f: f.write(base64.b64decode(data)) return True, "" except Exception as e: @@ -85,10 +89,10 @@ def upload_file_from_data(name: str, data: str) -> tuple[bool, str]: def create_file_id() -> str: file_id = str(uuid.uuid1()) - with open(".cache/file_list.json", "r", encoding="utf-8") as f: + with open(file_list_path(), "r", encoding="utf-8") as f: if file_id in json.load(f).keys(): return create_file_id() - with open(".cache/cached_url.json", "r", encoding="utf-8") as f: + with open(cached_url_path(), "r", encoding="utf-8") as f: if file_id in json.load(f).keys(): return create_file_id() return file_id @@ -96,10 +100,10 @@ def create_file_id() -> str: def register_saved_file(name: str, _file_id: str | None = None) -> str: file_id = _file_id or create_file_id() - with open(".cache/file_list.json", "r", encoding="utf-8") as f: + with open(file_list_path(), "r", encoding="utf-8") as f: file_list = json.load(f) file_list[file_id] = name - with open(".cache/file_list.json", "w", encoding="utf-8") as f: + with open(file_list_path(), "w", encoding="utf-8") as f: json.dump(file_list, f, ensure_ascii=False, indent=4) return file_id @@ -107,7 +111,7 @@ def register_saved_file(name: str, _file_id: str | None = None) -> str: def upload_file_from_path(name: str, path: str) -> tuple[bool, str]: try: with open(path, "rb") as from_f: - with open(f".cache/files/{name}", "wb") as to_f: + with open(get_file_path(name), "wb") as to_f: to_f.write(from_f.read()) return True, "" except Exception as e: @@ -177,7 +181,7 @@ async def upload_file_fragmented( case "finish": type_checker.check_arguments(file_id, offset, sha256) file_name = uploading_files[file_id]["name"] - with open(f".cache/files/{file_name}", "wb") as f: + with open(get_file_path(file_name), "wb") as f: f.write(bytes(uploading_files.pop(file_id)["content"])) # TODO sha256 校验 return return_object.get(file_id=register_saved_file(file_name)) @@ -234,11 +238,11 @@ async def get_file_name_by_id(file_id: str) -> str | None: """ 根据文件 ID 获取文件名 """ - with open(f".cache/file_list.json", "r", encoding="utf-8") as f: + with open(ffile_list_path(), "r", encoding="utf-8") as f: file_list = json.load(f) if _id := file_list.get(file_id): return _id - with open(".cache/cached_url.json", "r", encoding="utf-8") as f: + with open(cached_url_path(), "r", encoding="utf-8") as f: cached_url_list = json.load(f) if cache_data := cached_url_list.get(file_id): return await get_file_name_by_id( @@ -249,9 +253,9 @@ async def get_file_name_by_id(file_id: str) -> str | None: async def clean_files() -> None: - with open(".cache/file_list.json", "r", encoding="utf-8") as f: + with open(file_list_path(), "r", encoding="utf-8") as f: file_list = json.load(f) - with open(".cache/cached_url.json", "r", encoding="utf-8") as f: + with open(cached_url_path(), "r", encoding="utf-8") as f: cached_url_list = json.load(f) for file_id in list(file_list.keys()): if not os.path.exists(get_file_path(file_list[file_id])): @@ -272,16 +276,12 @@ async def clean_files() -> None: cached_url_list.pop(file_id) logger.debug(file_list) logger.debug(cached_url_list) - with open(".cache/file_list.json", "w", encoding="utf-8") as f: + with open(file_list_path(), "w", encoding="utf-8") as f: json.dump(file_list, f, ensure_ascii=False, indent=4) - with open(".cache/cached_url.json", "w", encoding="utf-8") as f: + with open(cached_url_path(), "w", encoding="utf-8") as f: json.dump(cached_url_list, f, ensure_ascii=False, indent=4) -def get_file_path(file_name: str) -> str: - return os.path.abspath(f".cache/files/{file_name}") - - @register_action() async def get_file(file_id: str, type: str) -> dict: """ @@ -297,12 +297,12 @@ async def get_file(file_id: str, type: str) -> dict: case "path": return return_object.get( - 0, name=file, path=os.path.abspath(f".cache/files/{file}") + 0, name=file, path=get_file_path(file) ) # TODO 返回 sha256 case "data": - with open(f".cache/files/{file}", "rb") as f: + with open(get_file_path(file), "rb") as f: return return_object.get( 0, name=file, data=base64.b64encode(f.read()).decode("utf-8") ) diff --git a/docs/config.md b/docs/config.md index 3beeed0..a2ac6b3 100644 --- a/docs/config.md +++ b/docs/config.md @@ -87,16 +87,57 @@ OneDisc 高级设置(无特殊需要不建议更改) | 类型 | 必须 | 默认值 | |:----------:|:----:|:----------------------:| -| 字符串 | 否 | `sqlite+aiosqlite:///:memory:` | +| 字符串 | 否 | `sqlite+aiosqlite:///缓存目录/onedisc.db` | OneDisc 缓存消息使用的数据库地址 参考 [Engine Configuration — SQLAlchemy 2.0 Documentation](https://docs.sqlalchemy.org/en/20/core/engines.html#database-urls) -不支持自动创建数据库 +未配置(或为 `null`)时,自动使用缓存目录(`cache_dir`)下的 `onedisc.db`,目录不存在会自动创建 > 目前可执行版只支持 SQLite3,源码版使用其他数据库需要手动安装依赖 +### 缓存目录(`cache_dir`) + +| 类型 | 必须 | 默认值 | +|:----------:|:----:|:----------------------:| +| 字符串 | 否 | `.cache` | + +OneDisc 缓存文件、缓存索引(`file_list.json` / `cached_url.json`)与默认数据库的存放目录 + +### 缓存过期清理(`cache`) + +| 类型 | 必须 | 默认值 | +|:----------:|:----:|:----------------------:| +| 对象 | 否 | `{}`(全部不过期) | + +配置各类缓存的过期时间(秒)与自动清理间隔: + +| 字段 | 说明 | 默认值 | +|:------------------:|:-----------------------------------------:|:-----------:| +| `files_ttl` | 文件缓存(`files/` 目录与 `node.*` 节点)过期秒数 | `0`(不过期) | +| `url_cache_ttl` | URL 缓存索引(`cached_url.json`)过期秒数 | `0`(不过期) | +| `db_ttl` | 本地消息数据库记录过期秒数 | `0`(不过期) | +| `cleanup_interval` | 自动清理检查间隔秒数 | `3600` | + +所有 TTL 为 `0`(默认)时表示永不过期,不会自动清理任何缓存;只有配置了 TTL,程序才会在每个清理周期删除过期条目 + +### 合并转发消息(`merge_forward`) + +| 类型 | 必须 | 默认值 | +|:----------:|:----:|:----------------------:| +| 布尔 | 否 | `true` | + +接收方向(Discord → OneBot):将同一频道同一发送者在 `merge_forward_interval` 内连续发送的「转发」消息自动合并为一条 OneBot V11 合并转发(forward)消息上报,框架收到后可通过 `get_forward_msg` 动作按 id 取回完整消息列表 + +### 合并转发窗口时长(`merge_forward_interval`) + +| 类型 | 必须 | 默认值 | +|:----------:|:----:|:----------------------:| +| 整数(毫秒)| 否 | `500` | + +合并转发的收集窗口时长,窗口内到达的满足条件的转发消息会被吸收进同一条合并转发,不再单独上报 + ### 使用静态表情(`use_static_face`) | 类型 | 必须 | 默认值 | @@ -271,7 +312,7 @@ OneBot V11 中,私聊消息事件(`message.private`)的 `sub_type` 字段 |:---------:|:----:|:----------------:| | 布尔 | 否 | `false` | -此项为 `true` 时,当同一文件同时存在于缓存 URL 索引和本地储存库(`.cache/files`)时,优先保留缓存(删除本地储存文件) +此项为 `true` 时,当同一文件同时存在于缓存 URL 索引和本地储存库(`cache_dir` 下的 `files/` 目录)时,优先保留缓存(删除本地储存文件) 此项为 `false` 时,将删除缓存库的索引 diff --git a/docs/diff-v11.md b/docs/diff-v11.md index 46317d1..b47b2a9 100644 --- a/docs/diff-v11.md +++ b/docs/diff-v11.md @@ -39,7 +39,7 @@ | 接口名称 | 终结点 | 说明 | |------------------|---------------------------|--------------------------| -| 获取合并转发消息 | `get_forward_msg` | Discord 不支持相关功能 | +| 获取合并转发消息 | `get_forward_msg` | 仅支持查询 OneDisc 自动合并转发产生的消息(id 格式为「频道id_消息id」,见 `merge_forward` 配置项),其他 id 返回错误;Discord 本身没有合并转发概念 | | 发送好友赞 | `send_like` | Discord 不支持相关功能 | | 群组匿名用户禁言 | `set_group_anonymous_ban` | Discord 不支持相关功能 | | 群组全员禁言 | `set_group_whole_ban` | Discord 不支持相关功能 | diff --git a/pyproject.toml b/pyproject.toml index 4d83019..5bb4457 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "onedisc" -version = "1.0.2" +version = "1.0.3" description = "OneBot implement for Discord" authors = [ {name = "XiaoDeng3386",email = "1744793737@qq.com"} diff --git a/test_cache_and_reply.py b/test_cache_and_reply.py new file mode 100644 index 0000000..aea2909 --- /dev/null +++ b/test_cache_and_reply.py @@ -0,0 +1,189 @@ +"""缓存目录 / DB 默认路径 / reply 三层解析 / 过期清理 单元测试(无需真实 Discord)""" +import asyncio +import json +import os +import shutil +import sys +import tempfile +import time +from types import SimpleNamespace +from unittest.mock import patch, AsyncMock + +sys.path.insert(0, "/vol2/@apphome/trim.openclaw/data/workspace/onedisc-dev") + +import discord + +from utils import cache +from utils import db +from utils.message.v12 import parser + + +def run(coro): + return asyncio.get_event_loop().run_until_complete(coro) + + +results = [] + + +def check(name, cond): + results.append((name, bool(cond))) + print(f"{'✅' if cond else '❌'} {name}") + + +class FakeSession: + def __init__(self, record): + self.record = record + + async def __aenter__(self): + return self + + async def __aexit__(self, *a): + pass + + async def get(self, model, id_): + return self.record + + +def make_msg(mid, cid): + return SimpleNamespace(id=mid, channel=SimpleNamespace(id=cid)) + + +# ---------- reply 三层解析 ---------- + +# 1. 第一层:cached_messages 命中 → 返回原消息对象 +with patch.object(parser, "client", SimpleNamespace(cached_messages=[make_msg(1001, 2001)])), patch.object( + parser, "get_session", lambda: FakeSession(None) +), patch.object(parser.discord_api, "call", AsyncMock()): + r = run(parser._resolve_reply(1001, 3001)) + check("reply: cached 命中返回原消息对象", getattr(r, "id", None) == 1001) + +# 2. 第二层:cached 空 + DB 命中 → MessageReference(channel=record.channel) +with patch.object(parser, "client", SimpleNamespace(cached_messages=[])), patch.object( + parser, "get_session", lambda: FakeSession(SimpleNamespace(channel=2001)) +), patch.object(parser.discord_api, "call", AsyncMock()): + r = run(parser._resolve_reply(1001, None)) + check( + "reply: DB 命中返回 MessageReference(channel=2001)", + isinstance(r, discord.MessageReference) and r.channel_id == 2001, + ) + +# 3. 第三层:cached/DB 空 + REST 预检成功 → MessageReference(channel=传入) +with patch.object(parser, "client", SimpleNamespace(cached_messages=[])), patch.object( + parser, "get_session", lambda: FakeSession(None) +), patch.object(parser.discord_api, "call", AsyncMock(return_value={})): + r = run(parser._resolve_reply(1001, 3001)) + check( + "reply: REST 预检成功返回 MessageReference(channel=3001)", + isinstance(r, discord.MessageReference) and r.channel_id == 3001, + ) + +# 4. 全 miss(REST 抛异常)→ None +async def boom(*a, **k): + raise RuntimeError("404") + + +with patch.object(parser, "client", SimpleNamespace(cached_messages=[])), patch.object( + parser, "get_session", lambda: FakeSession(None) +), patch.object(parser.discord_api, "call", boom): + r = run(parser._resolve_reply(1001, 3001)) + check("reply: 全 miss 返回 None", r is None) + +# 5. parse_message 集成:reply 段 + DB 命中 → message_data["reference"] +with patch.object(parser, "client", SimpleNamespace(cached_messages=[])), patch.object( + parser, "get_session", lambda: FakeSession(SimpleNamespace(channel=2001)) +), patch.object(parser.discord_api, "call", AsyncMock()): + data = run(parser.parse_message([{"type": "reply", "data": {"message_id": "1001"}}])) + check( + "parse_message: reply 段生成 reference(channel=2001)", + data.get("reference") is not None and data["reference"].channel_id == 2001, + ) + +# ---------- DB 默认路径 ---------- + +# 6. database 未配置 → 拼接缓存目录 + onedisc.db +with patch.dict(db.config, {"system": {"cache_dir": "/tmp/onedisc_test_cache", "database": None}}): + url = db._resolve_db_url() + check( + "db_url: 默认拼接缓存目录/onedisc.db", + url.endswith("onedisc.db") and "onedisc_test_cache" in url, + ) + +# 7. database 显式配置 → 原样使用 +with patch.dict(db.config, {"system": {"cache_dir": ".cache", "database": "sqlite+aiosqlite:///:memory:"}}): + check( + "db_url: 显式配置原样使用", + db._resolve_db_url() == "sqlite+aiosqlite:///:memory:", + ) + +# ---------- 过期清理 ---------- + +# 8. files_ttl=1:过期文件删除,新文件保留(files/ 目录 + node.* 节点文件) +tmp = tempfile.mkdtemp() +try: + files = os.path.join(tmp, "files") + os.makedirs(files) + old = os.path.join(files, "old.bin") + new = os.path.join(files, "new.bin") + open(old, "w").write("x") + open(new, "w").write("x") + os.utime(old, (time.time() - 100, time.time() - 100)) + old_node = os.path.join(tmp, "node.123") + open(old_node, "w").write("x") + os.utime(old_node, (time.time() - 100, time.time() - 100)) + with patch.object(cache, "get_cache_dir", lambda: tmp), patch.object( + cache, "cache_config", lambda: {"files_ttl": 1, "url_cache_ttl": 0, "db_ttl": 0} + ): + run(cache.clean_expired()) + check( + "clean: files_ttl=1 删旧留新", + not os.path.exists(old) and os.path.exists(new) and not os.path.exists(old_node), + ) +finally: + shutil.rmtree(tmp, ignore_errors=True) + +# 9. TTL=0(默认)→ 不过期不清理 +tmp = tempfile.mkdtemp() +try: + files = os.path.join(tmp, "files") + os.makedirs(files) + old = os.path.join(files, "old.bin") + open(old, "w").write("x") + os.utime(old, (time.time() - 10000, time.time() - 10000)) + with patch.object(cache, "get_cache_dir", lambda: tmp), patch.object( + cache, "cache_config", lambda: {"files_ttl": 0, "url_cache_ttl": 0, "db_ttl": 0} + ): + run(cache.clean_expired()) + check("clean: TTL=0 不过期不清理", os.path.exists(old)) +finally: + shutil.rmtree(tmp, ignore_errors=True) + +# 10. url_cache_ttl:过期条目删除、未过期保留、无 time 字段的旧条目视为不过期 +tmp = tempfile.mkdtemp() +try: + url_path = os.path.join(tmp, "cached_url.json") + json.dump( + { + "old": {"name": "a", "url": "u", "time": int(time.time()) - 100}, + "new": {"name": "b", "url": "u", "time": int(time.time())}, + "legacy": {"name": "c", "url": "u"}, + }, + open(url_path, "w"), + ) + with patch.object(cache, "cached_url_path", lambda: url_path), patch.object( + cache, "cache_config", lambda: {"files_ttl": 0, "url_cache_ttl": 10, "db_ttl": 0} + ): + run(cache.clean_expired()) + data = json.load(open(url_path)) + check( + "clean: url_cache_ttl 删过期/留新/留旧格式条目", + "old" not in data and "new" in data and "legacy" in data, + ) +finally: + shutil.rmtree(tmp, ignore_errors=True) + +# ---------- 清理 import 副作用(模块级创建的真实 .cache 目录) ---------- +shutil.rmtree(".cache", ignore_errors=True) + +failed = [n for n, c in results if not c] +print(f"\n共 {len(results)} 例,失败 {len(failed)} 例") +sys.exit(1 if failed else 0) diff --git a/test_forward_merge.py b/test_forward_merge.py new file mode 100644 index 0000000..23de33c --- /dev/null +++ b/test_forward_merge.py @@ -0,0 +1,266 @@ +"""转发消息自动合并(forward merge)单元测试:无状态 + REST 上下文重建""" +import asyncio +import sys +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +sys.path.insert(0, "/vol2/@apphome/trim.openclaw/data/workspace/onedisc-dev") + +import discord + +from utils import forward_merge as fm + + +def run(coro): + return asyncio.get_event_loop().run_until_complete(coro) + + +results = [] + + +def check(name, cond): + results.append((name, bool(cond))) + print(f"{'✅' if cond else '❌'} {name}") + + +NOW = datetime(2026, 8, 14, 7, 0, 0, tzinfo=timezone.utc) + + +def make_msg( + mid, + cid=2001, + author=("1001", "Alice"), + ref_type=discord.MessageReferenceType.forward, + content="", + snapshot=None, + guild=True, + created_at=NOW, +): + ref = ( + SimpleNamespace(type=ref_type, message_id=123, channel_id=cid) + if ref_type is not None + else None + ) + snapshots = [] + if snapshot is not None: + snapshots = [ + SimpleNamespace( + message=SimpleNamespace( + content=snapshot.get("content", "快照内容"), + attachments=snapshot.get("attachments", []), + author=SimpleNamespace( + id=snapshot.get("user_id", 9999), + name=snapshot.get("nickname", "原作者"), + ), + ) + ) + ] + return SimpleNamespace( + id=mid, + channel=SimpleNamespace(id=cid), + author=SimpleNamespace(id=author[0], name=author[1]), + content=content, + attachments=[], + reference=ref, + message_snapshots=snapshots, + guild=SimpleNamespace(id=111) if guild else None, + created_at=created_at, + ) + + +def msg_json(mid, ts, author="1001", ref_type=1, content="", snapshot=None): + """REST 返回的消息 JSON""" + ref = {"type": ref_type, "message_id": 123, "channel_id": 2001} if ref_type is not None else None + snapshots = [] + if snapshot is not None: + snapshots = [ + { + "message": { + "content": snapshot.get("content", "快照内容"), + "attachments": snapshot.get("attachments", []), + "author": { + "id": str(snapshot.get("user_id", 9999)), + "username": snapshot.get("nickname", "原作者"), + }, + } + } + ] + return { + "id": str(mid), + "channel_id": "2001", + "author": {"id": str(author), "username": f"用户{author}"}, + "content": content, + "attachments": [], + "message_reference": ref, + "message_snapshots": snapshots, + "timestamp": ts.isoformat(), + } + + +def patch_config(**system): + merged = { + **fm.config["system"], + "merge_forward": True, + "merge_forward_interval": 500, + **system, + } + return patch.dict(fm.config, {"system": merged}) + + +def reset(): + fm._first_records.clear() + + +# ---------- is_forward_message ---------- +check( + "is_forward: forward 类型 → True", + fm.is_forward_message(make_msg(1, ref_type=discord.MessageReferenceType.forward)), +) +check( + "is_forward: reply 类型 → False", + not fm.is_forward_message(make_msg(2, ref_type=discord.MessageReferenceType.default)), +) +check("is_forward: 无 reference → False", not fm.is_forward_message(make_msg(3, ref_type=None))) + +# ---------- build/parse forward id ---------- +fid = fm.build_forward_id(2001, 12345) +check("build_forward_id 编码频道+消息", fid == "2001_12345") +check("parse_forward_id 解码正常", fm.parse_forward_id(fid) == (2001, 12345)) +check("parse_forward_id 非法返回 None", fm.parse_forward_id("no_such_id") is None) +check("parse_forward_id 空串返回 None", fm.parse_forward_id("") is None) + +# ---------- handle_forward_message ---------- +reset() +with patch_config(), patch("utils.event.new_event") as new_event: + r = fm.handle_forward_message(make_msg(10)) + check("handle: 第一条返回 True", r is True) + check("handle: 第一条上报 forward 段", new_event.called) + if new_event.called: + kwargs = new_event.call_args.kwargs + seg = kwargs["message"][0] + check( + "handle: 上报内容为 forward 段且 id=频道_消息", + seg["type"] == "forward" and seg["data"]["id"] == "2001_10", + ) + + # 窗口内第二条(created_at +300ms):吸收、不重复上报 + r2 = fm.handle_forward_message(make_msg(11, created_at=NOW + timedelta(milliseconds=300))) + check("handle: 窗口内第二条被吸收返回 True", r2 is True) + check("handle: 窗口内第二条不重复上报", new_event.call_count == 1) + + # 超过窗口(+600ms):开新窗口并上报 + r3 = fm.handle_forward_message(make_msg(12, created_at=NOW + timedelta(milliseconds=600))) + check("handle: 超窗口消息开新窗口", r3 is True and new_event.call_count == 2) + check( + "handle: 新窗口 id 用第二条消息", + new_event.call_args.kwargs["message"][0]["data"]["id"] == "2001_12", + ) + + # 同频道不同发送者:各自独立窗口 + r4 = fm.handle_forward_message(make_msg(13, author=("2002", "Bob"))) + check("handle: 不同发送者开新窗口", r4 is True and new_event.call_count == 3) + +reset() +with patch_config(merge_forward=False), patch("utils.event.new_event") as new_event: + r5 = fm.handle_forward_message(make_msg(20)) + check("handle: merge_forward=false 不处理", r5 is False and not new_event.called) + +with patch_config(), patch("utils.event.new_event") as new_event: + r6 = fm.handle_forward_message(make_msg(21, ref_type=None)) + check("handle: 非转发消息不处理", r6 is False and not new_event.called) + +# ---------- get_forward(REST 上下文重建) ---------- +# 场景:第一条 100(t0),窗口内 101(+300ms)同作者转发;窗口外 102(+1s)同作者转发 +# 应被排除:103(+200ms,reply 类型)、104(+200ms,不同作者转发)、105(+200ms,普通消息) +messages = [ + msg_json(100, NOW, snapshot={"content": "原始内容", "user_id": 9999, "nickname": "原作者"}), + msg_json(101, NOW + timedelta(milliseconds=300), snapshot={"content": "第二条", "user_id": 8888, "nickname": "二号"}), + msg_json(102, NOW + timedelta(seconds=1), snapshot={"content": "超窗", "user_id": 7777}), + msg_json(103, NOW + timedelta(milliseconds=200), ref_type=0, content="回复"), + msg_json(104, NOW + timedelta(milliseconds=200), author="2002", snapshot={"content": "别人"}), + msg_json(105, NOW + timedelta(milliseconds=200), ref_type=None, content="普通"), +] +with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock(return_value=messages) +) as call_mock: + nodes = run(fm.get_forward("2001_100")) + check("get_forward: 重建出 2 个节点", nodes is not None and len(nodes) == 2) + check("get_forward: 按时间升序(第一条在前)", nodes[0]["user_id"] == "9999" and nodes[1]["user_id"] == "8888") + check( + "get_forward: 节点用快照原作者/内容", + nodes[0]["nickname"] == "原作者" and nodes[0]["content"][0]["data"]["text"] == "原始内容", + ) + check( + "get_forward: 调用 REST around 接口", + call_mock.await_args.args[1] == "/channels/2001/messages" + and call_mock.await_args.kwargs["params"]["around"] == "100", + ) + +# REST 失败 → None +with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock(side_effect=Exception("网络错误")) +): + check("get_forward: REST 异常返回 None", run(fm.get_forward("2001_100")) is None) + +# 第一条不在上下文中 → None +with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock(return_value=[msg_json(999, NOW)]) +): + check("get_forward: 第一条缺失返回 None", run(fm.get_forward("2001_100")) is None) + +# 非法 id → None(不调 REST) +with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock() +) as call_mock: + check("get_forward: 非法 id 返回 None", run(fm.get_forward("no_such_id")) is None) + check("get_forward: 非法 id 不调 REST", not call_mock.called) + +# ---------- 翻页:最后一条仍在窗口内且全部符合 → after 继续拉 ---------- +# 第一页(倒序):101(+300ms) 100(t0),最后一条 101 在窗口内且全部符合 → 翻页 +# 第二页(after=101,倒序):103(+600ms 超窗) 102(+400ms) → 停止 +with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock(side_effect=[[ + msg_json(101, NOW + timedelta(milliseconds=300), snapshot={"content": "第二", "user_id": 8888}), + msg_json(100, NOW, snapshot={"content": "原始内容", "user_id": 9999, "nickname": "原作者"}), + ], [ + msg_json(103, NOW + timedelta(milliseconds=600), snapshot={"content": "超窗", "user_id": 9999}), + msg_json(102, NOW + timedelta(milliseconds=400), snapshot={"content": "第三", "user_id": 7777}), + ]])) as call_mock: + + nodes = run(fm.get_forward("2001_100")) + check("翻页: 合并 3 条(100/101/102)", nodes is not None and len(nodes) == 3) + check("翻页: 按时间升序", [n["user_id"] for n in nodes] == ["9999", "8888", "7777"]) + check("翻页: 共调用 2 次 REST", call_mock.await_count == 2) + check( + "翻页: 第二次用 after=最后一条(101)", + call_mock.await_args.args[1] == "/channels/2001/messages" + and call_mock.await_args.kwargs["params"] == {"after": "101", "limit": "100"}, + ) + +# 第一页最后一条已超窗 → 不翻页 +with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock(return_value=[ + msg_json(102, NOW + timedelta(milliseconds=600), snapshot={"content": "超窗"}), + msg_json(101, NOW + timedelta(milliseconds=300), snapshot={"content": "第二"}), + msg_json(100, NOW, snapshot={"content": "原始内容", "user_id": 9999}), + ]) +) as call_mock: + nodes = run(fm.get_forward("2001_100")) + check("超窗即停: 合并 2 条", nodes is not None and len(nodes) == 2) + check("超窗即停: 只调用 1 次 REST", call_mock.await_count == 1) + +# 窗口内出现不符合条件的消息(不同作者转发)→ 不翻页,序列中断 +with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock(return_value=[ + msg_json(101, NOW + timedelta(milliseconds=300), author="2002", snapshot={"content": "别人"}), + msg_json(100, NOW, snapshot={"content": "原始内容", "user_id": 9999}), + ]) +) as call_mock: + nodes = run(fm.get_forward("2001_100")) + check("序列中断: 只合并第一条", nodes is not None and len(nodes) == 1) + check("序列中断: 不翻页", call_mock.await_count == 1) + +failed = [n for n, c in results if not c] +print(f"\n共 {len(results)} 例,失败 {len(failed)} 例") +sys.exit(1 if failed else 0) diff --git a/utils/cache.py b/utils/cache.py new file mode 100644 index 0000000..ce3122e --- /dev/null +++ b/utils/cache.py @@ -0,0 +1,112 @@ +"""缓存目录工具与过期清理 + +- 缓存目录可通过 system.cache_dir 配置(默认 .cache),所有模块统一从这里取路径 +- 缓存过期清理:system.cache 下的 files_ttl / url_cache_ttl / db_ttl(秒), + 0 或缺省表示不过期(默认不清理);start_cleanup_loop() 按 cleanup_interval 周期执行 +""" +import os +import time +import json +import asyncio +from .config import config +from .logger import get_logger + +logger = get_logger() + + +def get_cache_dir() -> str: + """返回缓存目录(system.cache_dir,默认 .cache),确保目录存在""" + cache_dir = config["system"].get("cache_dir", ".cache") + os.makedirs(cache_dir, exist_ok=True) + return cache_dir + + +def files_dir() -> str: + """文件缓存目录(cache_dir/files),确保目录存在""" + directory = os.path.join(get_cache_dir(), "files") + os.makedirs(directory, exist_ok=True) + return directory + + +def file_list_path() -> str: + """文件列表索引路径""" + return os.path.join(get_cache_dir(), "file_list.json") + + +def cached_url_path() -> str: + """URL 缓存索引路径""" + return os.path.join(get_cache_dir(), "cached_url.json") + + +def get_file_path(file_name: str) -> str: + """返回缓存文件的绝对路径""" + return os.path.abspath(os.path.join(files_dir(), file_name)) + + +def cache_config() -> dict: + """system.cache 配置段(不存在时返回空 dict)""" + return config["system"].get("cache", {}) or {} + + +async def clean_expired() -> None: + """按 system.cache 的 TTL 清理过期缓存;TTL=0(默认)表示不过期不清理""" + cache = cache_config() + now = int(time.time()) + + # 1) 文件缓存(files/ 目录下全部文件 + cache_dir 根下的 node.* 合并转发节点文件,按 mtime) + files_ttl = int(cache.get("files_ttl", 0) or 0) + if files_ttl > 0: + try: + for path in ( + os.path.join(files_dir(), f) for f in os.listdir(files_dir()) + ): + if now - int(os.path.getmtime(path)) > files_ttl: + os.remove(path) + logger.info(f"已清理过期缓存文件:{path}") + for name in os.listdir(get_cache_dir()): + if not name.startswith("node."): + continue + path = os.path.join(get_cache_dir(), name) + if now - int(os.path.getmtime(path)) > files_ttl: + os.remove(path) + logger.info(f"已清理过期节点缓存:{path}") + except Exception as e: + logger.warning(f"清理过期文件缓存失败:{e}") + + # 2) URL 缓存索引(条目含 time 字段才参与过期判断,旧条目视为不过期) + url_ttl = int(cache.get("url_cache_ttl", 0) or 0) + if url_ttl > 0: + try: + with open(cached_url_path(), "r", encoding="utf-8") as f: + data = json.load(f) + keep = {} + for file_id, item in data.items(): + created = item.get("time") if isinstance(item, dict) else None + if created is None or now - int(created) <= url_ttl: + keep[file_id] = item + if len(keep) != len(data): + with open(cached_url_path(), "w", encoding="utf-8") as f: + json.dump(keep, f) + logger.info(f"已清理过期 URL 缓存索引({len(data) - len(keep)} 条)") + except Exception as e: + logger.warning(f"清理过期 URL 缓存失败:{e}") + + # 3) 本地消息数据库(延迟导入避免循环依赖) + db_ttl = int(cache.get("db_ttl", 0) or 0) + if db_ttl > 0: + from .db import cleanup_expired_messages + + try: + await cleanup_expired_messages(db_ttl) + except Exception as e: + logger.warning(f"清理过期消息记录失败:{e}") + + +async def start_cleanup_loop() -> None: + """后台自动清理循环:按 system.cache.cleanup_interval(秒,默认 3600)周期执行""" + interval = int(cache_config().get("cleanup_interval", 3600) or 3600) + if interval <= 0: + return + while True: + await asyncio.sleep(interval) + await clean_expired() diff --git a/utils/db.py b/utils/db.py index 24f0313..f69c65c 100644 --- a/utils/db.py +++ b/utils/db.py @@ -1,13 +1,27 @@ +import os.path +import time from traceback import format_exc from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession from sqlalchemy.orm import declarative_base -from sqlalchemy import Column, Integer +from sqlalchemy import Column, Integer, delete from .config import config +from .cache import get_cache_dir from .logger import get_logger logger = get_logger() Base = declarative_base() -db_url = config["system"].get("database", "sqlite+aiosqlite:///:memory:") + + +def _resolve_db_url() -> str: + """database 未配置(或为 null)时默认持久化到缓存目录,避免重启丢失""" + database = config["system"].get("database") + if not database: + path = os.path.join(get_cache_dir(), "onedisc.db") + return f"sqlite+aiosqlite:///{path}" + return database + + +db_url = _resolve_db_url() logger.debug(f"使用数据库: {db_url}") engine = create_async_engine(db_url) del db_url @@ -39,3 +53,19 @@ async def commit_message(id_: int, channel: int, time_: int) -> None: logger.warning(f"写入数据库失败: {format_exc()}") else: await session.commit() + + +async def cleanup_expired_messages(ttl: int) -> None: + """删除 time 早于 now - ttl 的消息记录(缓存过期清理用)""" + now = int(time.time()) + async with get_session() as session: + try: + result = await session.execute( + delete(Message).where(Message.time < now - ttl) + ) + await session.commit() + if result.rowcount: + logger.info(f"已清理过期消息记录 {result.rowcount} 条") + except Exception: + await session.rollback() + logger.warning(f"清理过期消息记录失败: {format_exc()}") diff --git a/utils/event/discord_event.py b/utils/event/discord_event.py index e71f405..f695e7b 100644 --- a/utils/event/discord_event.py +++ b/utils/event/discord_event.py @@ -8,6 +8,8 @@ from discord import Object import asyncio from ..db import commit_message, init_database +from utils.cache import start_cleanup_loop +import utils.forward_merge as forward_merge from actions.v12.basic import get_status from actions.v11.basic import get_role @@ -30,6 +32,7 @@ async def on_ready() -> None: logger.info(config["system"].get("started_text", "OneDisc 已成功启动")) event.new_event("meta", "status_update", status=(await get_status())["data"]) await init_database() + asyncio.create_task(start_cleanup_loop()) @client.event @@ -43,6 +46,12 @@ async def on_message(message: discord.Message) -> None: and config["system"].get("ignore_bot_event") ): return + if forward_merge.handle_forward_message(message): + # 已进入合并转发流程(第一条已上报 forward 段,后续消息被吸收) + await commit_message( + message.id, message.channel.id, int(message.created_at.timestamp()) + ) + return print_message_log(message) await commit_message( message.id, message.channel.id, int(message.created_at.timestamp()) diff --git a/utils/forward_merge.py b/utils/forward_merge.py new file mode 100644 index 0000000..5f4b854 --- /dev/null +++ b/utils/forward_merge.py @@ -0,0 +1,270 @@ +"""Discord 转发消息自动合并为 OneBot V11 合并转发(forward) + +机制:同一频道同一发送者的「转发」消息(message_reference.type == forward), +在 merge_forward_interval(默认 500ms)窗口内自动合并: +- 第一条到达时立即上报一条 forward 段消息,id 编码为「频道id_消息id」 +- 窗口内后续满足条件的转发消息被吸收(不单独上报) +- 框架收到 forward 段后调用 get_forward_msg 动作,用 id 解码出第一条消息, + 通过 Discord REST API(GET /channels/{cid}/messages?around=)拉取上下文, + 按「同作者 + forward 类型 + 窗口时间范围」重建完整合并列表 + 实现完全无状态:不依赖内存暂存,进程重启后依然可用 + +可通过 system.merge_forward(默认 true)开关,system.merge_forward_interval +(毫秒,默认 500)调节窗口时长。 +""" +from datetime import datetime, timedelta, timezone +import discord +from utils.config import config +from utils.logger import get_logger +from utils import discord_api +import utils.message.v11.parser as parser + +logger = get_logger() + +# (channel_id, author_id) → {"first_id": int, "first_created_at": datetime} +# 仅用于实时拦截窗口,条目随新第一条覆盖,不会无限增长 +_first_records: dict[tuple[int, int], dict] = {} + +# Discord message_reference.type:1 = 转发,0 = 引用回复 +FORWARD_REFERENCE_TYPE = 1 + + +def _interval_ms() -> int: + return int(config["system"].get("merge_forward_interval", 500)) + + +def is_forward_message(message: discord.Message) -> bool: + """是否为 Discord 的「转发」消息(区别于引用回复)""" + return bool(message.reference) and ( + message.reference.type == discord.MessageReferenceType.forward + ) + + +def build_forward_id(channel_id: int, message_id: int) -> str: + """第一条消息的转发 id:编码频道与消息,可无状态重建""" + return f"{channel_id}_{message_id}" + + +def parse_forward_id(forward_id: str) -> tuple[int, int] | None: + """解码转发 id 为(频道 id,消息 id),非法返回 None""" + try: + channel_id, message_id = forward_id.split("_", 1) + return int(channel_id), int(message_id) + except (ValueError, AttributeError): + return None + + +def handle_forward_message(message: discord.Message) -> bool: + """ + 处理转发消息。返回 True 表示已被合并转发流程接管 + (第一条已上报 forward 段消息,后续消息被吸收),调用方不应再单独上报 + """ + if not is_forward_message(message): + return False + if not config["system"].get("merge_forward", True): + return False + + key = (message.channel.id, message.author.id) + record = _first_records.get(key) + if record is not None and ( + message.created_at - record["first_created_at"] + ).total_seconds() * 1000 <= _interval_ms(): + # 窗口内:吸收进合并列表,不单独上报 + logger.debug(f"转发消息 {message.id} 已并入合并转发 {record['first_id']}") + return True + + # 新第一条:记录窗口基准,上报 forward 段 + _first_records[key] = { + "first_id": message.id, + "first_created_at": message.created_at, + } + _report_forward_message( + message, build_forward_id(message.channel.id, message.id) + ) + logger.info( + f"开启合并转发 {build_forward_id(message.channel.id, message.id)}" + f"(频道 {message.channel.id},发送者 {message.author.id})" + ) + return True + + +def _report_forward_message(message: discord.Message, forward_id: str) -> None: + """上报一条 forward 段消息事件(与普通消息上报相同的频道类型分支)""" + from utils import event + + message_array = [{"type": "forward", "data": {"id": forward_id}}] + common = { + "_type": "message", + "_time": message.created_at.timestamp(), + "message_id": str(message.id), + "message": message_array, + "alt_message": forward_id, + "user_id": str(message.author.id), + } + if message.guild and config["system"].get("enable_channel_event"): + event.new_event( + detail_type="channel", + guild_id=str(message.guild.id), + channel_id=str(message.channel.id), + **common, + ) + elif message.guild: + event.new_event( + detail_type="group", + group_id=str(message.channel.id), + **common, + ) + else: + event.new_event(detail_type="private", **common) + + +def _is_forward_data(data: dict) -> bool: + """Discord 消息 JSON 是否为转发消息""" + reference = data.get("message_reference") + return bool(reference) and reference.get("type") == FORWARD_REFERENCE_TYPE + + +def _parse_timestamp(timestamp: str | None) -> datetime | None: + if not timestamp: + return None + try: + return datetime.fromisoformat(timestamp.replace("Z", "+00:00")) + except ValueError: + return None + + +# 翻页上限,防止异常情况下无限拉取 +MAX_PAGES = 10 + + +async def _fetch_merge_messages( + channel_id: int, message_id: int +) -> list[dict] | None: + """ + 以第一条消息为起点拉取上下文:第一页用 around(含第一条), + 若最新一条仍在窗口内且窗口内消息全部符合条件,则用 after 继续向前翻页, + 直到窗口结束、序列中断或拉不到更多。失败返回 None + """ + collected: list[dict] = [] + cursor = str(message_id) + first_page = True + for _ in range(MAX_PAGES): + params = { + "around" if first_page else "after": cursor, + "limit": "100", + } + try: + data = await discord_api.call( + "GET", f"/channels/{channel_id}/messages", params=params + ) + except Exception: + logger.warning(f"获取合并转发上下文失败(频道 {channel_id},消息 {cursor})") + return None + if not isinstance(data, list): + logger.warning(f"获取合并转发上下文返回异常(频道 {channel_id})") + return None + collected.extend(data) + first_page = False + + first = next( + (m for m in collected if str(m.get("id")) == str(message_id)), None + ) + if first is None: + logger.warning(f"合并转发第一条消息 {message_id} 不在返回上下文中") + return None + first_time = _parse_timestamp(first.get("timestamp")) + if first_time is None: + return None + author_id = str(first.get("author", {}).get("id")) + window_end = first_time + timedelta(milliseconds=_interval_ms()) + + # 拉取结果中最新的一条(Discord 消息 API 按时间倒序,即列表首条) + newest = max( + collected, + key=lambda m: _parse_timestamp(m.get("timestamp")) or first_time, + ) + newest_time = _parse_timestamp(newest.get("timestamp")) + if newest_time is None or newest_time > window_end: + # 最后一条已超出窗口:窗口内消息已完整,停止 + break + # 窗口内所有消息都符合条件(同作者转发)才继续:序列中断则停止 + in_window = [ + m + for m in collected + if first_time <= _parse_timestamp(m.get("timestamp")) <= window_end + ] + if not in_window or not all( + _is_forward_data(m) + and str(m.get("author", {}).get("id")) == author_id + for m in in_window + ): + break + if str(newest.get("id")) == cursor: + # 翻页没有拉到更新的消息,停止 + break + # 最后一条仍在窗口内且全部符合:从它继续向前拉取 + cursor = str(newest.get("id")) + return collected + + +def _data_to_node(data: dict) -> dict: + """Discord 消息 JSON → v11 node 节点(显示快照中的原作者与内容)""" + snapshot = None + snapshots = data.get("message_snapshots") + if snapshots: + snapshot = snapshots[0].get("message") + source = snapshot or data + author = source.get("author", {}) + segments = parser.parse_string_to_array(source.get("content") or "") + for attachment in source.get("attachments", []): + content_type = attachment.get("content_type", "") or "" + if content_type.startswith("image"): + segments.append( + {"type": "image", "data": {"file": attachment.get("url", "")}} + ) + elif content_type.startswith("video"): + segments.append( + {"type": "video", "data": {"file": attachment.get("url", "")}} + ) + return { + "user_id": author.get("id"), + "nickname": author.get("username") or str(author.get("id")), + "content": segments, + } + + +async def get_forward(forward_id: str) -> list[dict] | None: + """ + 按 id 取合并转发节点列表(get_forward_msg 动作使用)。 + 通过第一条消息调 Discord REST API 拉取上下文重建,失败或不存在返回 None + """ + parsed = parse_forward_id(forward_id) + if parsed is None: + return None + channel_id, message_id = parsed + collected = await _fetch_merge_messages(channel_id, message_id) + if collected is None: + return None + + first = next( + (message for message in collected if str(message.get("id")) == str(message_id)), + None, + ) + if first is None: + return None + first_time = _parse_timestamp(first.get("timestamp")) + if first_time is None: + return None + author_id = str(first.get("author", {}).get("id")) + window = timedelta(milliseconds=_interval_ms()) + merged = [ + message + for message in collected + if _is_forward_data(message) + and str(message.get("author", {}).get("id")) == author_id + and first_time + <= _parse_timestamp(message.get("timestamp")) + <= first_time + window + ] + merged.sort(key=lambda message: message["timestamp"]) + return [_data_to_node(message) for message in merged] diff --git a/utils/message/v12/parser.py b/utils/message/v12/parser.py index 19f2918..d3727eb 100644 --- a/utils/message/v12/parser.py +++ b/utils/message/v12/parser.py @@ -8,6 +8,8 @@ import utils.message.v12.tokenizer as tokenizer import re import discord.file +from utils import discord_api +from utils.db import get_session, Message class BadSegmentData(Exception): @@ -48,7 +50,35 @@ def escape_mentions(text): return text -async def parse_message(message: list) -> dict: +async def _resolve_reply( + message_id: int, channel_id: int | None +) -> discord.Message | discord.MessageReference | None: + """解析 reply 段引用的消息:内存缓存 → 本地数据库 → 同频道 REST 预检""" + for msg in client.cached_messages: + if msg.id == message_id: + return msg + async with get_session() as session: + record = await session.get(Message, message_id) + if record is not None: + return discord.MessageReference( + message_id=message_id, channel_id=record.channel + ) + if channel_id is not None: + try: + await discord_api.call( + "GET", f"/channels/{channel_id}/messages/{message_id}" + ) + return discord.MessageReference( + message_id=message_id, channel_id=channel_id + ) + except Exception: + logger.debug( + f"REST 预检失败:消息 {message_id}(频道 {channel_id})不存在或不可读" + ) + return None + + +async def parse_message(message: list, channel_id: int | None = None) -> dict: logger.debug(config) message_data = {"content": "", "files": []} for segment in message: @@ -106,10 +136,11 @@ async def parse_message(message: list) -> dict: case "discord.navigation": message_data["content"] += f"" case "reply": - for msg in client.cached_messages: - if msg.id == int(segment["data"]["message_id"]): - message_data["reference"] = msg - break + reference = await _resolve_reply( + int(segment["data"]["message_id"]), channel_id + ) + if reference is not None: + message_data["reference"] = reference else: logger.warning( f"解析消息段 {segment} 时出现错误:找不到指定消息,已忽略" From 7f56bc80cee1ca04bf1d43d2f4f5898ca86e8773 Mon Sep 17 00:00:00 2001 From: This-is-XiaoDeng <1744793737@qq.com> Date: Fri, 14 Aug 2026 08:21:32 +0800 Subject: [PATCH 2/5] =?UTF-8?q?test:=20=E6=B5=8B=E8=AF=95=E7=A7=BB?= =?UTF-8?q?=E5=85=A5=20tests/=20=E7=9B=AE=E5=BD=95=E5=B9=B6=E6=8E=A5?= =?UTF-8?q?=E5=85=A5=20pytest=20=E4=B8=8E=20CI?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 根目录 test_*.py 移入 tests/,改写为 pytest 风格(async def + assert),新增 conftest.py(CI 环境自动创建 dummy config.json、清理 .cache 副作用) - pyproject.toml 新增 [tool.pytest.ini_options](testpaths/pythonpath/asyncio_mode)与 test 依赖组(pytest、pytest-asyncio) - ci.yml 新增 test job:poetry install --only main,test + pytest(不装 nuitka 等构建依赖) - 本地验证 40 例全部通过 --- .github/workflows/ci.yml | 26 ++++ pyproject.toml | 9 ++ test_cache_and_reply.py | 189 ------------------------ test_forward_merge.py | 266 ---------------------------------- test_native_forward.py | 138 ------------------ tests/conftest.py | 29 ++++ tests/test_cache_and_reply.py | 148 +++++++++++++++++++ tests/test_forward_merge.py | 263 +++++++++++++++++++++++++++++++++ tests/test_native_forward.py | 129 +++++++++++++++++ 9 files changed, 604 insertions(+), 593 deletions(-) delete mode 100644 test_cache_and_reply.py delete mode 100644 test_forward_merge.py delete mode 100644 test_native_forward.py create mode 100644 tests/conftest.py create mode 100644 tests/test_cache_and_reply.py create mode 100644 tests/test_forward_merge.py create mode 100644 tests/test_native_forward.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 929b9c6..3ea782c 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -7,6 +7,32 @@ on: workflow_dispatch: jobs: + test: + runs-on: ubuntu-latest + steps: + - name: "Checkout" + uses: actions/checkout@v4 + + - name: "Setup Python" + uses: actions/setup-python@v5 + with: + python-version: '3.12' + + - name: "Setup Poetry" + uses: snok/install-poetry@v1 + with: + version: latest + virtualenvs-create: true + virtualenvs-in-project: true + + - name: "Install dependencies" + run: | + poetry install --only main,test + + - name: "Run tests" + run: | + poetry run pytest + get-version-number: runs-on: ubuntu-latest outputs: diff --git a/pyproject.toml b/pyproject.toml index 5bb4457..9f71b2c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -32,3 +32,12 @@ package-mode = false nuitka = "^4.1.3" imageio = "^2.37.3" +[tool.poetry.group.test.dependencies] +pytest = "^8.3.0" +pytest-asyncio = "^0.24.0" + +[tool.pytest.ini_options] +testpaths = ["tests"] +pythonpath = ["."] +asyncio_mode = "auto" + diff --git a/test_cache_and_reply.py b/test_cache_and_reply.py deleted file mode 100644 index aea2909..0000000 --- a/test_cache_and_reply.py +++ /dev/null @@ -1,189 +0,0 @@ -"""缓存目录 / DB 默认路径 / reply 三层解析 / 过期清理 单元测试(无需真实 Discord)""" -import asyncio -import json -import os -import shutil -import sys -import tempfile -import time -from types import SimpleNamespace -from unittest.mock import patch, AsyncMock - -sys.path.insert(0, "/vol2/@apphome/trim.openclaw/data/workspace/onedisc-dev") - -import discord - -from utils import cache -from utils import db -from utils.message.v12 import parser - - -def run(coro): - return asyncio.get_event_loop().run_until_complete(coro) - - -results = [] - - -def check(name, cond): - results.append((name, bool(cond))) - print(f"{'✅' if cond else '❌'} {name}") - - -class FakeSession: - def __init__(self, record): - self.record = record - - async def __aenter__(self): - return self - - async def __aexit__(self, *a): - pass - - async def get(self, model, id_): - return self.record - - -def make_msg(mid, cid): - return SimpleNamespace(id=mid, channel=SimpleNamespace(id=cid)) - - -# ---------- reply 三层解析 ---------- - -# 1. 第一层:cached_messages 命中 → 返回原消息对象 -with patch.object(parser, "client", SimpleNamespace(cached_messages=[make_msg(1001, 2001)])), patch.object( - parser, "get_session", lambda: FakeSession(None) -), patch.object(parser.discord_api, "call", AsyncMock()): - r = run(parser._resolve_reply(1001, 3001)) - check("reply: cached 命中返回原消息对象", getattr(r, "id", None) == 1001) - -# 2. 第二层:cached 空 + DB 命中 → MessageReference(channel=record.channel) -with patch.object(parser, "client", SimpleNamespace(cached_messages=[])), patch.object( - parser, "get_session", lambda: FakeSession(SimpleNamespace(channel=2001)) -), patch.object(parser.discord_api, "call", AsyncMock()): - r = run(parser._resolve_reply(1001, None)) - check( - "reply: DB 命中返回 MessageReference(channel=2001)", - isinstance(r, discord.MessageReference) and r.channel_id == 2001, - ) - -# 3. 第三层:cached/DB 空 + REST 预检成功 → MessageReference(channel=传入) -with patch.object(parser, "client", SimpleNamespace(cached_messages=[])), patch.object( - parser, "get_session", lambda: FakeSession(None) -), patch.object(parser.discord_api, "call", AsyncMock(return_value={})): - r = run(parser._resolve_reply(1001, 3001)) - check( - "reply: REST 预检成功返回 MessageReference(channel=3001)", - isinstance(r, discord.MessageReference) and r.channel_id == 3001, - ) - -# 4. 全 miss(REST 抛异常)→ None -async def boom(*a, **k): - raise RuntimeError("404") - - -with patch.object(parser, "client", SimpleNamespace(cached_messages=[])), patch.object( - parser, "get_session", lambda: FakeSession(None) -), patch.object(parser.discord_api, "call", boom): - r = run(parser._resolve_reply(1001, 3001)) - check("reply: 全 miss 返回 None", r is None) - -# 5. parse_message 集成:reply 段 + DB 命中 → message_data["reference"] -with patch.object(parser, "client", SimpleNamespace(cached_messages=[])), patch.object( - parser, "get_session", lambda: FakeSession(SimpleNamespace(channel=2001)) -), patch.object(parser.discord_api, "call", AsyncMock()): - data = run(parser.parse_message([{"type": "reply", "data": {"message_id": "1001"}}])) - check( - "parse_message: reply 段生成 reference(channel=2001)", - data.get("reference") is not None and data["reference"].channel_id == 2001, - ) - -# ---------- DB 默认路径 ---------- - -# 6. database 未配置 → 拼接缓存目录 + onedisc.db -with patch.dict(db.config, {"system": {"cache_dir": "/tmp/onedisc_test_cache", "database": None}}): - url = db._resolve_db_url() - check( - "db_url: 默认拼接缓存目录/onedisc.db", - url.endswith("onedisc.db") and "onedisc_test_cache" in url, - ) - -# 7. database 显式配置 → 原样使用 -with patch.dict(db.config, {"system": {"cache_dir": ".cache", "database": "sqlite+aiosqlite:///:memory:"}}): - check( - "db_url: 显式配置原样使用", - db._resolve_db_url() == "sqlite+aiosqlite:///:memory:", - ) - -# ---------- 过期清理 ---------- - -# 8. files_ttl=1:过期文件删除,新文件保留(files/ 目录 + node.* 节点文件) -tmp = tempfile.mkdtemp() -try: - files = os.path.join(tmp, "files") - os.makedirs(files) - old = os.path.join(files, "old.bin") - new = os.path.join(files, "new.bin") - open(old, "w").write("x") - open(new, "w").write("x") - os.utime(old, (time.time() - 100, time.time() - 100)) - old_node = os.path.join(tmp, "node.123") - open(old_node, "w").write("x") - os.utime(old_node, (time.time() - 100, time.time() - 100)) - with patch.object(cache, "get_cache_dir", lambda: tmp), patch.object( - cache, "cache_config", lambda: {"files_ttl": 1, "url_cache_ttl": 0, "db_ttl": 0} - ): - run(cache.clean_expired()) - check( - "clean: files_ttl=1 删旧留新", - not os.path.exists(old) and os.path.exists(new) and not os.path.exists(old_node), - ) -finally: - shutil.rmtree(tmp, ignore_errors=True) - -# 9. TTL=0(默认)→ 不过期不清理 -tmp = tempfile.mkdtemp() -try: - files = os.path.join(tmp, "files") - os.makedirs(files) - old = os.path.join(files, "old.bin") - open(old, "w").write("x") - os.utime(old, (time.time() - 10000, time.time() - 10000)) - with patch.object(cache, "get_cache_dir", lambda: tmp), patch.object( - cache, "cache_config", lambda: {"files_ttl": 0, "url_cache_ttl": 0, "db_ttl": 0} - ): - run(cache.clean_expired()) - check("clean: TTL=0 不过期不清理", os.path.exists(old)) -finally: - shutil.rmtree(tmp, ignore_errors=True) - -# 10. url_cache_ttl:过期条目删除、未过期保留、无 time 字段的旧条目视为不过期 -tmp = tempfile.mkdtemp() -try: - url_path = os.path.join(tmp, "cached_url.json") - json.dump( - { - "old": {"name": "a", "url": "u", "time": int(time.time()) - 100}, - "new": {"name": "b", "url": "u", "time": int(time.time())}, - "legacy": {"name": "c", "url": "u"}, - }, - open(url_path, "w"), - ) - with patch.object(cache, "cached_url_path", lambda: url_path), patch.object( - cache, "cache_config", lambda: {"files_ttl": 0, "url_cache_ttl": 10, "db_ttl": 0} - ): - run(cache.clean_expired()) - data = json.load(open(url_path)) - check( - "clean: url_cache_ttl 删过期/留新/留旧格式条目", - "old" not in data and "new" in data and "legacy" in data, - ) -finally: - shutil.rmtree(tmp, ignore_errors=True) - -# ---------- 清理 import 副作用(模块级创建的真实 .cache 目录) ---------- -shutil.rmtree(".cache", ignore_errors=True) - -failed = [n for n, c in results if not c] -print(f"\n共 {len(results)} 例,失败 {len(failed)} 例") -sys.exit(1 if failed else 0) diff --git a/test_forward_merge.py b/test_forward_merge.py deleted file mode 100644 index 23de33c..0000000 --- a/test_forward_merge.py +++ /dev/null @@ -1,266 +0,0 @@ -"""转发消息自动合并(forward merge)单元测试:无状态 + REST 上下文重建""" -import asyncio -import sys -from datetime import datetime, timedelta, timezone -from types import SimpleNamespace -from unittest.mock import AsyncMock, patch - -sys.path.insert(0, "/vol2/@apphome/trim.openclaw/data/workspace/onedisc-dev") - -import discord - -from utils import forward_merge as fm - - -def run(coro): - return asyncio.get_event_loop().run_until_complete(coro) - - -results = [] - - -def check(name, cond): - results.append((name, bool(cond))) - print(f"{'✅' if cond else '❌'} {name}") - - -NOW = datetime(2026, 8, 14, 7, 0, 0, tzinfo=timezone.utc) - - -def make_msg( - mid, - cid=2001, - author=("1001", "Alice"), - ref_type=discord.MessageReferenceType.forward, - content="", - snapshot=None, - guild=True, - created_at=NOW, -): - ref = ( - SimpleNamespace(type=ref_type, message_id=123, channel_id=cid) - if ref_type is not None - else None - ) - snapshots = [] - if snapshot is not None: - snapshots = [ - SimpleNamespace( - message=SimpleNamespace( - content=snapshot.get("content", "快照内容"), - attachments=snapshot.get("attachments", []), - author=SimpleNamespace( - id=snapshot.get("user_id", 9999), - name=snapshot.get("nickname", "原作者"), - ), - ) - ) - ] - return SimpleNamespace( - id=mid, - channel=SimpleNamespace(id=cid), - author=SimpleNamespace(id=author[0], name=author[1]), - content=content, - attachments=[], - reference=ref, - message_snapshots=snapshots, - guild=SimpleNamespace(id=111) if guild else None, - created_at=created_at, - ) - - -def msg_json(mid, ts, author="1001", ref_type=1, content="", snapshot=None): - """REST 返回的消息 JSON""" - ref = {"type": ref_type, "message_id": 123, "channel_id": 2001} if ref_type is not None else None - snapshots = [] - if snapshot is not None: - snapshots = [ - { - "message": { - "content": snapshot.get("content", "快照内容"), - "attachments": snapshot.get("attachments", []), - "author": { - "id": str(snapshot.get("user_id", 9999)), - "username": snapshot.get("nickname", "原作者"), - }, - } - } - ] - return { - "id": str(mid), - "channel_id": "2001", - "author": {"id": str(author), "username": f"用户{author}"}, - "content": content, - "attachments": [], - "message_reference": ref, - "message_snapshots": snapshots, - "timestamp": ts.isoformat(), - } - - -def patch_config(**system): - merged = { - **fm.config["system"], - "merge_forward": True, - "merge_forward_interval": 500, - **system, - } - return patch.dict(fm.config, {"system": merged}) - - -def reset(): - fm._first_records.clear() - - -# ---------- is_forward_message ---------- -check( - "is_forward: forward 类型 → True", - fm.is_forward_message(make_msg(1, ref_type=discord.MessageReferenceType.forward)), -) -check( - "is_forward: reply 类型 → False", - not fm.is_forward_message(make_msg(2, ref_type=discord.MessageReferenceType.default)), -) -check("is_forward: 无 reference → False", not fm.is_forward_message(make_msg(3, ref_type=None))) - -# ---------- build/parse forward id ---------- -fid = fm.build_forward_id(2001, 12345) -check("build_forward_id 编码频道+消息", fid == "2001_12345") -check("parse_forward_id 解码正常", fm.parse_forward_id(fid) == (2001, 12345)) -check("parse_forward_id 非法返回 None", fm.parse_forward_id("no_such_id") is None) -check("parse_forward_id 空串返回 None", fm.parse_forward_id("") is None) - -# ---------- handle_forward_message ---------- -reset() -with patch_config(), patch("utils.event.new_event") as new_event: - r = fm.handle_forward_message(make_msg(10)) - check("handle: 第一条返回 True", r is True) - check("handle: 第一条上报 forward 段", new_event.called) - if new_event.called: - kwargs = new_event.call_args.kwargs - seg = kwargs["message"][0] - check( - "handle: 上报内容为 forward 段且 id=频道_消息", - seg["type"] == "forward" and seg["data"]["id"] == "2001_10", - ) - - # 窗口内第二条(created_at +300ms):吸收、不重复上报 - r2 = fm.handle_forward_message(make_msg(11, created_at=NOW + timedelta(milliseconds=300))) - check("handle: 窗口内第二条被吸收返回 True", r2 is True) - check("handle: 窗口内第二条不重复上报", new_event.call_count == 1) - - # 超过窗口(+600ms):开新窗口并上报 - r3 = fm.handle_forward_message(make_msg(12, created_at=NOW + timedelta(milliseconds=600))) - check("handle: 超窗口消息开新窗口", r3 is True and new_event.call_count == 2) - check( - "handle: 新窗口 id 用第二条消息", - new_event.call_args.kwargs["message"][0]["data"]["id"] == "2001_12", - ) - - # 同频道不同发送者:各自独立窗口 - r4 = fm.handle_forward_message(make_msg(13, author=("2002", "Bob"))) - check("handle: 不同发送者开新窗口", r4 is True and new_event.call_count == 3) - -reset() -with patch_config(merge_forward=False), patch("utils.event.new_event") as new_event: - r5 = fm.handle_forward_message(make_msg(20)) - check("handle: merge_forward=false 不处理", r5 is False and not new_event.called) - -with patch_config(), patch("utils.event.new_event") as new_event: - r6 = fm.handle_forward_message(make_msg(21, ref_type=None)) - check("handle: 非转发消息不处理", r6 is False and not new_event.called) - -# ---------- get_forward(REST 上下文重建) ---------- -# 场景:第一条 100(t0),窗口内 101(+300ms)同作者转发;窗口外 102(+1s)同作者转发 -# 应被排除:103(+200ms,reply 类型)、104(+200ms,不同作者转发)、105(+200ms,普通消息) -messages = [ - msg_json(100, NOW, snapshot={"content": "原始内容", "user_id": 9999, "nickname": "原作者"}), - msg_json(101, NOW + timedelta(milliseconds=300), snapshot={"content": "第二条", "user_id": 8888, "nickname": "二号"}), - msg_json(102, NOW + timedelta(seconds=1), snapshot={"content": "超窗", "user_id": 7777}), - msg_json(103, NOW + timedelta(milliseconds=200), ref_type=0, content="回复"), - msg_json(104, NOW + timedelta(milliseconds=200), author="2002", snapshot={"content": "别人"}), - msg_json(105, NOW + timedelta(milliseconds=200), ref_type=None, content="普通"), -] -with patch_config(), patch.object( - fm.discord_api, "call", new=AsyncMock(return_value=messages) -) as call_mock: - nodes = run(fm.get_forward("2001_100")) - check("get_forward: 重建出 2 个节点", nodes is not None and len(nodes) == 2) - check("get_forward: 按时间升序(第一条在前)", nodes[0]["user_id"] == "9999" and nodes[1]["user_id"] == "8888") - check( - "get_forward: 节点用快照原作者/内容", - nodes[0]["nickname"] == "原作者" and nodes[0]["content"][0]["data"]["text"] == "原始内容", - ) - check( - "get_forward: 调用 REST around 接口", - call_mock.await_args.args[1] == "/channels/2001/messages" - and call_mock.await_args.kwargs["params"]["around"] == "100", - ) - -# REST 失败 → None -with patch_config(), patch.object( - fm.discord_api, "call", new=AsyncMock(side_effect=Exception("网络错误")) -): - check("get_forward: REST 异常返回 None", run(fm.get_forward("2001_100")) is None) - -# 第一条不在上下文中 → None -with patch_config(), patch.object( - fm.discord_api, "call", new=AsyncMock(return_value=[msg_json(999, NOW)]) -): - check("get_forward: 第一条缺失返回 None", run(fm.get_forward("2001_100")) is None) - -# 非法 id → None(不调 REST) -with patch_config(), patch.object( - fm.discord_api, "call", new=AsyncMock() -) as call_mock: - check("get_forward: 非法 id 返回 None", run(fm.get_forward("no_such_id")) is None) - check("get_forward: 非法 id 不调 REST", not call_mock.called) - -# ---------- 翻页:最后一条仍在窗口内且全部符合 → after 继续拉 ---------- -# 第一页(倒序):101(+300ms) 100(t0),最后一条 101 在窗口内且全部符合 → 翻页 -# 第二页(after=101,倒序):103(+600ms 超窗) 102(+400ms) → 停止 -with patch_config(), patch.object( - fm.discord_api, "call", new=AsyncMock(side_effect=[[ - msg_json(101, NOW + timedelta(milliseconds=300), snapshot={"content": "第二", "user_id": 8888}), - msg_json(100, NOW, snapshot={"content": "原始内容", "user_id": 9999, "nickname": "原作者"}), - ], [ - msg_json(103, NOW + timedelta(milliseconds=600), snapshot={"content": "超窗", "user_id": 9999}), - msg_json(102, NOW + timedelta(milliseconds=400), snapshot={"content": "第三", "user_id": 7777}), - ]])) as call_mock: - - nodes = run(fm.get_forward("2001_100")) - check("翻页: 合并 3 条(100/101/102)", nodes is not None and len(nodes) == 3) - check("翻页: 按时间升序", [n["user_id"] for n in nodes] == ["9999", "8888", "7777"]) - check("翻页: 共调用 2 次 REST", call_mock.await_count == 2) - check( - "翻页: 第二次用 after=最后一条(101)", - call_mock.await_args.args[1] == "/channels/2001/messages" - and call_mock.await_args.kwargs["params"] == {"after": "101", "limit": "100"}, - ) - -# 第一页最后一条已超窗 → 不翻页 -with patch_config(), patch.object( - fm.discord_api, "call", new=AsyncMock(return_value=[ - msg_json(102, NOW + timedelta(milliseconds=600), snapshot={"content": "超窗"}), - msg_json(101, NOW + timedelta(milliseconds=300), snapshot={"content": "第二"}), - msg_json(100, NOW, snapshot={"content": "原始内容", "user_id": 9999}), - ]) -) as call_mock: - nodes = run(fm.get_forward("2001_100")) - check("超窗即停: 合并 2 条", nodes is not None and len(nodes) == 2) - check("超窗即停: 只调用 1 次 REST", call_mock.await_count == 1) - -# 窗口内出现不符合条件的消息(不同作者转发)→ 不翻页,序列中断 -with patch_config(), patch.object( - fm.discord_api, "call", new=AsyncMock(return_value=[ - msg_json(101, NOW + timedelta(milliseconds=300), author="2002", snapshot={"content": "别人"}), - msg_json(100, NOW, snapshot={"content": "原始内容", "user_id": 9999}), - ]) -) as call_mock: - nodes = run(fm.get_forward("2001_100")) - check("序列中断: 只合并第一条", nodes is not None and len(nodes) == 1) - check("序列中断: 不翻页", call_mock.await_count == 1) - -failed = [n for n, c in results if not c] -print(f"\n共 {len(results)} 例,失败 {len(failed)} 例") -sys.exit(1 if failed else 0) diff --git a/test_native_forward.py b/test_native_forward.py deleted file mode 100644 index 0931c1d..0000000 --- a/test_native_forward.py +++ /dev/null @@ -1,138 +0,0 @@ -"""can_native_forward 判定逻辑单元测试(无需真实 Discord)""" -import asyncio -import sys -from types import SimpleNamespace -from unittest.mock import patch, AsyncMock - -import discord - -sys.path.insert(0, "/vol2/@apphome/trim.openclaw/data/workspace/onedisc-dev") - -from utils import native_forward as nf - - -def msg(**kw): - base = dict( - id=1001, - type=discord.MessageType.default, - poll=None, - reference=None, - channel=SimpleNamespace(id=2001), - guild=SimpleNamespace(id=111), - ) - base.update(kw) - return SimpleNamespace(**base) - - -class FakeSession: - def __init__(self, record): - self.record = record - - async def __aenter__(self): - return self - - async def __aexit__(self, *a): - pass - - async def get(self, model, id_): - return self.record - - -def run(coro): - return asyncio.get_event_loop().run_until_complete(coro) - - -def patch_env(cached=None, db_record=None, target_guild=SimpleNamespace(id=111), api_get=None): - from contextlib import ExitStack - - sess = FakeSession(db_record) - stack = ExitStack() - stack.enter_context(patch.object(nf, "client", SimpleNamespace( - cached_messages=cached or [], - get_channel=lambda cid: SimpleNamespace(guild=target_guild), - ))) - stack.enter_context(patch.object(nf, "get_session", lambda: sess)) - stack.enter_context(patch.object(nf.discord_api, "call", AsyncMock(return_value=api_get))) - return stack - - -results = [] - - -def check(name, cond): - results.append((name, cond)) - print(f"{'✅' if cond else '❌'} {name}") - - -# --- 1. 内联内容节点 → 整体回退 --- -with patch_env(): - r = run(nf.can_native_forward( - [{"type": "node", "data": {"user_id": 1, "nickname": "x", "content": "hi"}}], 3001)) - check("内联节点 → None", r is None) - -# --- 2. 引用节点 cache/DB 均无记录 → 回退 --- -with patch_env(cached=[], db_record=None): - r = run(nf.can_native_forward([{"type": "node", "data": {"message_id": 999}}], 3001)) - check("channel 无法解析 → None", r is None) - -# --- 3. cache 命中但类型不可转发(pins_add)→ 回退 --- -with patch_env(cached=[msg(id=1001, type=discord.MessageType.pins_add)]): - r = run(nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001)) - check("不可转发类型 → None", r is None) - -# --- 4. cache 命中、类型可转发 → 返回 refs --- -with patch_env(cached=[msg(id=1001)]): - r = run(nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001)) - check("正常引用 → refs", r == [{"message_id": 1001, "channel_id": 2001}]) - -# --- 5. 转发消息本身(reference.type==forward)→ 回退 --- -with patch_env(cached=[msg(id=1001, reference=SimpleNamespace(type=discord.MessageReferenceType.forward))]): - r = run(nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001)) - check("转发消息不能再转发 → None", r is None) - -# --- 6. 数量超阈值 → 回退 --- -with patch.object(nf, "config", {**nf.config, "system": {**nf.config["system"], "native_forward_max_nodes": 1}}): - with patch_env(cached=[msg(id=1001), msg(id=1002)]): - r = run(nf.can_native_forward( - [{"type": "node", "data": {"message_id": 1001}}, - {"type": "node", "data": {"message_id": 1002}}], 3001)) - check("超阈值 → None", r is None) - # 阈值 1 但只有 1 条 → 放行 - r = run(nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001)) - check("未超阈值 → refs", r == [{"message_id": 1001, "channel_id": 2001}]) - -# --- 7. 跨服务器(cache 命中 guild 不同)→ 回退 --- -with patch_env(cached=[msg(id=1001, guild=SimpleNamespace(id=999))], target_guild=SimpleNamespace(id=111)): - r = run(nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001)) - check("跨服务器 → None", r is None) - -# --- 8. DB 命中 + REST 预检通过 → refs --- -with patch_env(db_record=SimpleNamespace(channel=2001), api_get={"type": 0, "guild_id": 111}): - r = run(nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001)) - check("DB+预检通过 → refs", r == [{"message_id": 1001, "channel_id": 2001}]) - -# --- 9. DB 命中 + 预检失败(不可转发类型 6=pins_add)→ 回退 --- -with patch_env(db_record=SimpleNamespace(channel=2001), api_get={"type": 6, "guild_id": 111}): - r = run(nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001)) - check("DB+预检类型不可转发 → None", r is None) - -# --- 10. DB 命中 + 预检错误响应(消息不存在)→ 回退 --- -with patch_env(db_record=SimpleNamespace(channel=2001), api_get={"code": 10008, "message": "Unknown Message"}): - r = run(nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001)) - check("DB+预检消息不存在 → None", r is None) - -# --- 11. 混合:一条正常 + 一条内联 → 整体回退 --- -with patch_env(cached=[msg(id=1001)]): - r = run(nf.can_native_forward( - [{"type": "node", "data": {"message_id": 1001}}, - {"type": "node", "data": {"user_id": 1, "nickname": "x", "content": "hi"}}], 3001)) - check("混合节点 → None", r is None) - -# --- 12. 非 node 结构 → 回退 --- -with patch_env(): - r = run(nf.can_native_forward([{"type": "text", "data": {"text": "hi"}}], 3001)) - check("非 node 结构 → None", r is None) - -failed = [n for n, c in results if not c] -print(f"\n共 {len(results)} 例,失败 {len(failed)} 例") -sys.exit(1 if failed else 0) diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..2dda1af --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,29 @@ +"""pytest 公共配置:保证 import utils 链的依赖环境可用""" +import json +import os +import shutil + +import pytest + + +@pytest.fixture(scope="session", autouse=True) +def ensure_config_json(): + """import utils 链需要读取 config.json(CI/干净环境不存在时创建 dummy,本地已有则不动)""" + if not os.path.exists("config.json"): + with open("config.json", "w", encoding="utf-8") as f: + json.dump( + { + "account_token": "dummy_token", + "system": {"proxy": None, "logger": {"level": 20}}, + "servers": [], + }, + f, + ) + yield + + +@pytest.fixture(autouse=True) +def cleanup_cache_dir(): + """清理测试 import 副作用创建的 .cache 目录""" + yield + shutil.rmtree(".cache", ignore_errors=True) diff --git a/tests/test_cache_and_reply.py b/tests/test_cache_and_reply.py new file mode 100644 index 0000000..74787d5 --- /dev/null +++ b/tests/test_cache_and_reply.py @@ -0,0 +1,148 @@ +"""缓存目录 / DB 默认路径 / reply 三层解析 / 过期清理 单元测试(无需真实 Discord)""" +import json +import os +import time +from types import SimpleNamespace +from unittest.mock import patch, AsyncMock + +import discord + +from utils import cache +from utils import db +from utils.message.v12 import parser + + +class FakeSession: + def __init__(self, record): + self.record = record + + async def __aenter__(self): + return self + + async def __aexit__(self, *a): + pass + + async def get(self, model, id_): + return self.record + + +def make_msg(mid, cid): + return SimpleNamespace(id=mid, channel=SimpleNamespace(id=cid)) + + +def patch_parser_env(cached, db_record, api_get=None): + from contextlib import ExitStack + + stack = ExitStack() + stack.enter_context(patch.object(parser, "client", SimpleNamespace(cached_messages=cached))) + stack.enter_context(patch.object(parser, "get_session", lambda: FakeSession(db_record))) + stack.enter_context(patch.object(parser.discord_api, "call", AsyncMock(return_value=api_get))) + return stack + + +# ---------- reply 三层解析 ---------- + +async def test_reply_cached命中返回原消息对象(): + with patch_parser_env([make_msg(1001, 2001)], None): + r = await parser._resolve_reply(1001, 3001) + assert getattr(r, "id", None) == 1001 + + +async def test_reply_db命中返回MessageReference(): + with patch_parser_env([], SimpleNamespace(channel=2001)): + r = await parser._resolve_reply(1001, None) + assert isinstance(r, discord.MessageReference) and r.channel_id == 2001 + + +async def test_reply_rest预检成功返回MessageReference(): + with patch_parser_env([], None, api_get={}): + r = await parser._resolve_reply(1001, 3001) + assert isinstance(r, discord.MessageReference) and r.channel_id == 3001 + + +async def test_reply_全miss返回None(): + async def boom(*a, **k): + raise RuntimeError("404") + + with patch_parser_env([], None), patch.object(parser.discord_api, "call", boom): + r = await parser._resolve_reply(1001, 3001) + assert r is None + + +async def test_parse_message_reply段生成reference(): + with patch_parser_env([], SimpleNamespace(channel=2001)): + data = await parser.parse_message([{"type": "reply", "data": {"message_id": "1001"}}]) + assert data.get("reference") is not None and data["reference"].channel_id == 2001 + + +# ---------- DB 默认路径 ---------- + +def test_db_url_未配置拼接缓存目录(): + with patch.dict(db.config, {"system": {"cache_dir": "/tmp/onedisc_test_cache", "database": None}}): + url = db._resolve_db_url() + assert url.endswith("onedisc.db") and "onedisc_test_cache" in url + + +def test_db_url_显式配置原样使用(): + with patch.dict(db.config, {"system": {"cache_dir": ".cache", "database": "sqlite+aiosqlite:///:memory:"}}): + assert db._resolve_db_url() == "sqlite+aiosqlite:///:memory:" + + +# ---------- 过期清理 ---------- + +async def test_clean_files_ttl删除过期保留新文件(tmp_path): + files = tmp_path / "files" + files.mkdir() + old = files / "old.bin" + new = files / "new.bin" + old.write_bytes(b"x") + new.write_bytes(b"x") + os_old = str(old) + os_new = str(new) + os_utime_old = os_old + os.utime(os_old, (time.time() - 100, time.time() - 100)) + node = tmp_path / "node.123" + node.write_bytes(b"x") + os.utime(str(node), (time.time() - 100, time.time() - 100)) + with patch.object(cache, "get_cache_dir", lambda: str(tmp_path)), patch.object( + cache, "cache_config", lambda: {"files_ttl": 1, "url_cache_ttl": 0, "db_ttl": 0} + ): + await cache.clean_expired() + assert not os.path.exists(os_old) + assert os.path.exists(os_new) + assert not os.path.exists(str(node)) + + +async def test_clean_ttl为0不过期不清理(tmp_path): + files = tmp_path / "files" + files.mkdir() + old = files / "old.bin" + old.write_bytes(b"x") + os.utime(str(old), (time.time() - 10000, time.time() - 10000)) + with patch.object(cache, "get_cache_dir", lambda: str(tmp_path)), patch.object( + cache, "cache_config", lambda: {"files_ttl": 0, "url_cache_ttl": 0, "db_ttl": 0} + ): + await cache.clean_expired() + assert os.path.exists(str(old)) + + +async def test_clean_url_cache_ttl删过期留新留旧格式(tmp_path): + url_path = tmp_path / "cached_url.json" + url_path.write_text( + json.dumps( + { + "old": {"name": "a", "url": "u", "time": int(time.time()) - 100}, + "new": {"name": "b", "url": "u", "time": int(time.time())}, + "legacy": {"name": "c", "url": "u"}, + } + ), + encoding="utf-8", + ) + with patch.object(cache, "cached_url_path", lambda: str(url_path)), patch.object( + cache, "cache_config", lambda: {"files_ttl": 0, "url_cache_ttl": 10, "db_ttl": 0} + ): + await cache.clean_expired() + data = json.loads(url_path.read_text(encoding="utf-8")) + assert "old" not in data + assert "new" in data + assert "legacy" in data diff --git a/tests/test_forward_merge.py b/tests/test_forward_merge.py new file mode 100644 index 0000000..fec19a8 --- /dev/null +++ b/tests/test_forward_merge.py @@ -0,0 +1,263 @@ +"""转发消息自动合并(forward merge)单元测试:无状态 + REST 上下文重建""" +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +import pytest + +import discord + +from utils import forward_merge as fm + +NOW = datetime(2026, 8, 14, 7, 0, 0, tzinfo=timezone.utc) + + +@pytest.fixture(autouse=True) +def reset_state(): + fm._first_records.clear() + yield + fm._first_records.clear() + + +def make_msg( + mid, + cid=2001, + author=("1001", "Alice"), + ref_type=discord.MessageReferenceType.forward, + content="", + snapshot=None, + guild=True, + created_at=NOW, +): + ref = ( + SimpleNamespace(type=ref_type, message_id=123, channel_id=cid) + if ref_type is not None + else None + ) + snapshots = [] + if snapshot is not None: + snapshots = [ + SimpleNamespace( + message=SimpleNamespace( + content=snapshot.get("content", "快照内容"), + attachments=snapshot.get("attachments", []), + author=SimpleNamespace( + id=snapshot.get("user_id", 9999), + name=snapshot.get("nickname", "原作者"), + ), + ) + ) + ] + return SimpleNamespace( + id=mid, + channel=SimpleNamespace(id=cid), + author=SimpleNamespace(id=author[0], name=author[1]), + content=content, + attachments=[], + reference=ref, + message_snapshots=snapshots, + guild=SimpleNamespace(id=111) if guild else None, + created_at=created_at, + ) + + +def msg_json(mid, ts, author="1001", ref_type=1, content="", snapshot=None): + """REST 返回的消息 JSON""" + ref = {"type": ref_type, "message_id": 123, "channel_id": 2001} if ref_type is not None else None + snapshots = [] + if snapshot is not None: + snapshots = [ + { + "message": { + "content": snapshot.get("content", "快照内容"), + "attachments": snapshot.get("attachments", []), + "author": { + "id": str(snapshot.get("user_id", 9999)), + "username": snapshot.get("nickname", "原作者"), + }, + } + } + ] + return { + "id": str(mid), + "channel_id": "2001", + "author": {"id": str(author), "username": f"用户{author}"}, + "content": content, + "attachments": [], + "message_reference": ref, + "message_snapshots": snapshots, + "timestamp": ts.isoformat(), + } + + +def patch_config(**system): + merged = { + **fm.config["system"], + "merge_forward": True, + "merge_forward_interval": 500, + **system, + } + return patch.dict(fm.config, {"system": merged}) + + +# ---------- is_forward_message ---------- + +def test_is_forward_forward类型(): + assert fm.is_forward_message(make_msg(1, ref_type=discord.MessageReferenceType.forward)) + + +def test_is_forward_reply类型(): + assert not fm.is_forward_message(make_msg(2, ref_type=discord.MessageReferenceType.default)) + + +def test_is_forward_无reference(): + assert not fm.is_forward_message(make_msg(3, ref_type=None)) + + +# ---------- build/parse forward id ---------- + +def test_forward_id_编码解码(): + fid = fm.build_forward_id(2001, 12345) + assert fid == "2001_12345" + assert fm.parse_forward_id(fid) == (2001, 12345) + + +def test_forward_id_非法返回None(): + assert fm.parse_forward_id("no_such_id") is None + assert fm.parse_forward_id("") is None + + +# ---------- handle_forward_message ---------- + +def test_handle_第一条上报forward段(): + with patch_config(), patch("utils.event.new_event") as new_event: + assert fm.handle_forward_message(make_msg(10)) is True + assert new_event.called + seg = new_event.call_args.kwargs["message"][0] + assert seg["type"] == "forward" and seg["data"]["id"] == "2001_10" + + +def test_handle_窗口内第二条被吸收(): + with patch_config(), patch("utils.event.new_event") as new_event: + fm.handle_forward_message(make_msg(10)) + r = fm.handle_forward_message(make_msg(11, created_at=NOW + timedelta(milliseconds=300))) + assert r is True + assert new_event.call_count == 1 + + +def test_handle_超窗口消息开新窗口(): + with patch_config(), patch("utils.event.new_event") as new_event: + fm.handle_forward_message(make_msg(10)) + r = fm.handle_forward_message(make_msg(12, created_at=NOW + timedelta(milliseconds=600))) + assert r is True + assert new_event.call_count == 2 + assert new_event.call_args.kwargs["message"][0]["data"]["id"] == "2001_12" + + +def test_handle_不同发送者独立窗口(): + with patch_config(), patch("utils.event.new_event") as new_event: + fm.handle_forward_message(make_msg(10)) + r = fm.handle_forward_message(make_msg(13, author=("2002", "Bob"))) + assert r is True + assert new_event.call_count == 2 + + +def test_handle_开关关闭不处理(): + with patch_config(merge_forward=False), patch("utils.event.new_event") as new_event: + assert fm.handle_forward_message(make_msg(20)) is False + assert not new_event.called + + +def test_handle_非转发消息不处理(): + with patch_config(), patch("utils.event.new_event") as new_event: + assert fm.handle_forward_message(make_msg(21, ref_type=None)) is False + assert not new_event.called + + +# ---------- get_forward(REST 上下文重建) ---------- + +async def test_get_forward_窗口筛选重建(): + messages = [ + msg_json(100, NOW, snapshot={"content": "原始内容", "user_id": 9999, "nickname": "原作者"}), + msg_json(101, NOW + timedelta(milliseconds=300), snapshot={"content": "第二条", "user_id": 8888, "nickname": "二号"}), + msg_json(102, NOW + timedelta(seconds=1), snapshot={"content": "超窗", "user_id": 7777}), + msg_json(103, NOW + timedelta(milliseconds=200), ref_type=0, content="回复"), + msg_json(104, NOW + timedelta(milliseconds=200), author="2002", snapshot={"content": "别人"}), + msg_json(105, NOW + timedelta(milliseconds=200), ref_type=None, content="普通"), + ] + with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock(return_value=messages) + ) as call_mock: + nodes = await fm.get_forward("2001_100") + assert nodes is not None and len(nodes) == 2 + assert nodes[0]["user_id"] == "9999" and nodes[1]["user_id"] == "8888" + assert nodes[0]["nickname"] == "原作者" + assert nodes[0]["content"][0]["data"]["text"] == "原始内容" + assert call_mock.await_args.args[1] == "/channels/2001/messages" + assert call_mock.await_args.kwargs["params"]["around"] == "100" + + +async def test_get_forward_rest异常返回None(): + with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock(side_effect=Exception("网络错误")) + ): + assert await fm.get_forward("2001_100") is None + + +async def test_get_forward_第一条缺失返回None(): + with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock(return_value=[msg_json(999, NOW)]) + ): + assert await fm.get_forward("2001_100") is None + + +async def test_get_forward_非法id不调REST(): + with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock() + ) as call_mock: + assert await fm.get_forward("no_such_id") is None + assert not call_mock.called + + +# ---------- 翻页:最后一条仍在窗口内且全部符合 → after 继续拉 ---------- + +async def test_get_forward_翻页合并(): + with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock(side_effect=[[ + msg_json(101, NOW + timedelta(milliseconds=300), snapshot={"content": "第二", "user_id": 8888}), + msg_json(100, NOW, snapshot={"content": "原始内容", "user_id": 9999, "nickname": "原作者"}), + ], [ + msg_json(103, NOW + timedelta(milliseconds=600), snapshot={"content": "超窗", "user_id": 9999}), + msg_json(102, NOW + timedelta(milliseconds=400), snapshot={"content": "第三", "user_id": 7777}), + ]]) + ) as call_mock: + nodes = await fm.get_forward("2001_100") + assert nodes is not None and len(nodes) == 3 + assert [n["user_id"] for n in nodes] == ["9999", "8888", "7777"] + assert call_mock.await_count == 2 + assert call_mock.await_args.kwargs["params"] == {"after": "101", "limit": "100"} + + +async def test_get_forward_超窗即停不翻页(): + with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock(return_value=[ + msg_json(102, NOW + timedelta(milliseconds=600), snapshot={"content": "超窗"}), + msg_json(101, NOW + timedelta(milliseconds=300), snapshot={"content": "第二"}), + msg_json(100, NOW, snapshot={"content": "原始内容", "user_id": 9999}), + ]) + ) as call_mock: + nodes = await fm.get_forward("2001_100") + assert nodes is not None and len(nodes) == 2 + assert call_mock.await_count == 1 + + +async def test_get_forward_序列中断不翻页(): + with patch_config(), patch.object( + fm.discord_api, "call", new=AsyncMock(return_value=[ + msg_json(101, NOW + timedelta(milliseconds=300), author="2002", snapshot={"content": "别人"}), + msg_json(100, NOW, snapshot={"content": "原始内容", "user_id": 9999}), + ]) + ) as call_mock: + nodes = await fm.get_forward("2001_100") + assert nodes is not None and len(nodes) == 1 + assert call_mock.await_count == 1 diff --git a/tests/test_native_forward.py b/tests/test_native_forward.py new file mode 100644 index 0000000..8332888 --- /dev/null +++ b/tests/test_native_forward.py @@ -0,0 +1,129 @@ +"""can_native_forward 判定逻辑单元测试(无需真实 Discord)""" +from types import SimpleNamespace +from unittest.mock import patch, AsyncMock + +import discord + +from utils import native_forward as nf + + +def msg(**kw): + base = dict( + id=1001, + type=discord.MessageType.default, + poll=None, + reference=None, + channel=SimpleNamespace(id=2001), + guild=SimpleNamespace(id=111), + ) + base.update(kw) + return SimpleNamespace(**base) + + +class FakeSession: + def __init__(self, record): + self.record = record + + async def __aenter__(self): + return self + + async def __aexit__(self, *a): + pass + + async def get(self, model, id_): + return self.record + + +def patch_env(cached=None, db_record=None, target_guild=SimpleNamespace(id=111), api_get=None): + from contextlib import ExitStack + + sess = FakeSession(db_record) + stack = ExitStack() + stack.enter_context(patch.object(nf, "client", SimpleNamespace( + cached_messages=cached or [], + get_channel=lambda cid: SimpleNamespace(guild=target_guild), + ))) + stack.enter_context(patch.object(nf, "get_session", lambda: sess)) + stack.enter_context(patch.object(nf.discord_api, "call", AsyncMock(return_value=api_get))) + return stack + + +async def test_内联内容节点整体回退(): + with patch_env(): + r = await nf.can_native_forward( + [{"type": "node", "data": {"user_id": 1, "nickname": "x", "content": "hi"}}], 3001) + assert r is None + + +async def test_引用节点cache与db均无记录回退(): + with patch_env(cached=[], db_record=None): + r = await nf.can_native_forward([{"type": "node", "data": {"message_id": 999}}], 3001) + assert r is None + + +async def test_cache命中但类型不可转发回退(): + with patch_env(cached=[msg(id=1001, type=discord.MessageType.pins_add)]): + r = await nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001) + assert r is None + + +async def test_cache命中类型可转发返回refs(): + with patch_env(cached=[msg(id=1001)]): + r = await nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001) + assert r == [{"message_id": 1001, "channel_id": 2001}] + + +async def test_转发消息本身不能再转发(): + with patch_env(cached=[msg(id=1001, reference=SimpleNamespace(type=discord.MessageReferenceType.forward))]): + r = await nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001) + assert r is None + + +async def test_数量超阈值回退(): + with patch.object(nf, "config", {**nf.config, "system": {**nf.config["system"], "native_forward_max_nodes": 1}}): + with patch_env(cached=[msg(id=1001), msg(id=1002)]): + r = await nf.can_native_forward( + [{"type": "node", "data": {"message_id": 1001}}, + {"type": "node", "data": {"message_id": 1002}}], 3001) + assert r is None + # 阈值 1 但只有 1 条 → 放行 + r = await nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001) + assert r == [{"message_id": 1001, "channel_id": 2001}] + + +async def test_跨服务器回退(): + with patch_env(cached=[msg(id=1001, guild=SimpleNamespace(id=999))], target_guild=SimpleNamespace(id=111)): + r = await nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001) + assert r is None + + +async def test_db命中且rest预检通过返回refs(): + with patch_env(db_record=SimpleNamespace(channel=2001), api_get={"type": 0, "guild_id": 111}): + r = await nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001) + assert r == [{"message_id": 1001, "channel_id": 2001}] + + +async def test_db命中但预检类型不可转发回退(): + with patch_env(db_record=SimpleNamespace(channel=2001), api_get={"type": 6, "guild_id": 111}): + r = await nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001) + assert r is None + + +async def test_db命中但预检消息不存在回退(): + with patch_env(db_record=SimpleNamespace(channel=2001), api_get={"code": 10008, "message": "Unknown Message"}): + r = await nf.can_native_forward([{"type": "node", "data": {"message_id": 1001}}], 3001) + assert r is None + + +async def test_混合节点整体回退(): + with patch_env(cached=[msg(id=1001)]): + r = await nf.can_native_forward( + [{"type": "node", "data": {"message_id": 1001}}, + {"type": "node", "data": {"user_id": 1, "nickname": "x", "content": "hi"}}], 3001) + assert r is None + + +async def test_非node结构回退(): + with patch_env(): + r = await nf.can_native_forward([{"type": "text", "data": {"text": "hi"}}], 3001) + assert r is None From 0d0d6153a59f55518ca2885a16dff487cc3bde5b Mon Sep 17 00:00:00 2001 From: This-is-XiaoDeng <1744793737@qq.com> Date: Fri, 14 Aug 2026 08:22:04 +0800 Subject: [PATCH 3/5] =?UTF-8?q?ci:=20=E6=96=B0=E5=A2=9E=E7=8B=AC=E7=AB=8B?= =?UTF-8?q?=20test=20workflow=EF=BC=88PR=20+=20master=20=E8=A7=A6=E5=8F=91?= =?UTF-8?q?=EF=BC=8C=E4=BB=85=E8=B7=91=20pytest=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .github/workflows/ci.yml | 26 -------------------------- .github/workflows/test.yml | 34 ++++++++++++++++++++++++++++++++++ 2 files changed, 34 insertions(+), 26 deletions(-) create mode 100644 .github/workflows/test.yml diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 3ea782c..929b9c6 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -7,32 +7,6 @@ on: workflow_dispatch: jobs: - test: - runs-on: ubuntu-latest - steps: - - name: "Checkout" - uses: actions/checkout@v4 - - - name: "Setup Python" - uses: actions/setup-python@v5 - with: - python-version: '3.12' - - - name: "Setup Poetry" - uses: snok/install-poetry@v1 - with: - version: latest - virtualenvs-create: true - virtualenvs-in-project: true - - - name: "Install dependencies" - run: | - poetry install --only main,test - - - name: "Run tests" - run: | - poetry run pytest - get-version-number: runs-on: ubuntu-latest outputs: diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml new file mode 100644 index 0000000..85fe8e5 --- /dev/null +++ b/.github/workflows/test.yml @@ -0,0 +1,34 @@ +name: 'Test' + +on: + pull_request: + push: + branches: + - master + +jobs: + test: + runs-on: ubuntu-latest + steps: + - name: "Checkout" + uses: actions/checkout@v4 + + - name: "Setup Python" + uses: actions/setup-python@v5 + with: + python-version: '3.12' + + - name: "Setup Poetry" + uses: snok/install-poetry@v1 + with: + version: latest + virtualenvs-create: true + virtualenvs-in-project: true + + - name: "Install dependencies" + run: | + poetry install --only main,test + + - name: "Run tests" + run: | + poetry run pytest From bee48bd3b1dfde7ad4ebb2150909af2a66b9d313 Mon Sep 17 00:00:00 2001 From: This-is-XiaoDeng <1744793737@qq.com> Date: Fri, 14 Aug 2026 08:25:26 +0800 Subject: [PATCH 4/5] =?UTF-8?q?chore:=20poetry=20lock=20=E6=9B=B4=E6=96=B0?= =?UTF-8?q?=EF=BC=88=E6=96=B0=E5=A2=9E=20test=20=E4=BE=9D=E8=B5=96?= =?UTF-8?q?=E7=BB=84=20pytest/pytest-asyncio=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- poetry.lock | 104 ++++++++++++++++++++++++++++++++++++++++++++++++++-- 1 file changed, 100 insertions(+), 4 deletions(-) diff --git a/poetry.lock b/poetry.lock index f0c7960..293b212 100644 --- a/poetry.lock +++ b/poetry.lock @@ -1,4 +1,4 @@ -# This file is automatically @generated by Poetry 2.4.1 and should not be changed by hand. +# This file is automatically @generated by Poetry 2.3.2 and should not be changed by hand. [[package]] name = "aiohappyeyeballs" @@ -300,12 +300,12 @@ version = "0.4.6" description = "Cross-platform colored terminal text." optional = false python-versions = "!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,!=3.6.*,>=2.7" -groups = ["main"] -markers = "platform_system == \"Windows\"" +groups = ["main", "test"] files = [ {file = "colorama-0.4.6-py2.py3-none-any.whl", hash = "sha256:4f1d9991f5acc0ca119f9d443620b77f9d6b33703e51011c16baf57afb285fc6"}, {file = "colorama-0.4.6.tar.gz", hash = "sha256:08695f5cb7ed6e0531a20572697297273c47b8cae5a63ffc6d6ed5c201be6e44"}, ] +markers = {main = "platform_system == \"Windows\"", test = "sys_platform == \"win32\""} [[package]] name = "discord-py" @@ -670,6 +670,18 @@ files = [ [package.dependencies] six = "*" +[[package]] +name = "iniconfig" +version = "2.3.0" +description = "brain-dead simple config-ini parsing" +optional = false +python-versions = ">=3.10" +groups = ["test"] +files = [ + {file = "iniconfig-2.3.0-py3-none-any.whl", hash = "sha256:f631c04d2c48c52b84d0d0549c99ff3859c98df65b3101406327ecc7d53fbf12"}, + {file = "iniconfig-2.3.0.tar.gz", hash = "sha256:c76315c77db068650d49c5b56314774a7804df16fee4402c1f19d6d15d8c4730"}, +] + [[package]] name = "multidict" version = "6.6.4" @@ -892,6 +904,18 @@ files = [ {file = "numpy-2.3.2.tar.gz", hash = "sha256:e0486a11ec30cdecb53f184d496d1c6a20786c81e55e41640270130056f8ee48"}, ] +[[package]] +name = "packaging" +version = "26.3" +description = "Core utilities for Python packages" +optional = false +python-versions = ">=3.9" +groups = ["test"] +files = [ + {file = "packaging-26.3-py3-none-any.whl", hash = "sha256:d7193f7c8e4e93f444fde0262bf90af30e16fa0ad0ad44cb553c87339b23cd1c"}, + {file = "packaging-26.3.tar.gz", hash = "sha256:94edc256424af38762eb31306eed28beb9f0efc50a8837492c9d6fd6004aed79"}, +] + [[package]] name = "pillow" version = "11.3.0" @@ -1017,6 +1041,22 @@ tests = ["check-manifest", "coverage (>=7.4.2)", "defusedxml", "markdown2", "ole typing = ["typing-extensions ; python_version < \"3.10\""] xmp = ["defusedxml"] +[[package]] +name = "pluggy" +version = "1.6.0" +description = "plugin and hook calling mechanisms for python" +optional = false +python-versions = ">=3.9" +groups = ["test"] +files = [ + {file = "pluggy-1.6.0-py3-none-any.whl", hash = "sha256:e920276dd6813095e9377c0bc5566d94c932c33b27a3e3945d8389c374dd4746"}, + {file = "pluggy-1.6.0.tar.gz", hash = "sha256:7dcc130b76258d33b90f61b658791dede3486c3e6bfb003ee5c9bfb396dd22f3"}, +] + +[package.extras] +dev = ["pre-commit", "tox"] +testing = ["coverage", "pytest", "pytest-benchmark"] + [[package]] name = "propcache" version = "0.3.2" @@ -1259,6 +1299,62 @@ files = [ [package.dependencies] typing-extensions = ">=4.6.0,<4.7.0 || >4.7.0" +[[package]] +name = "pygments" +version = "2.20.0" +description = "Pygments is a syntax highlighting package written in Python." +optional = false +python-versions = ">=3.9" +groups = ["test"] +files = [ + {file = "pygments-2.20.0-py3-none-any.whl", hash = "sha256:81a9e26dd42fd28a23a2d169d86d7ac03b46e2f8b59ed4698fb4785f946d0176"}, + {file = "pygments-2.20.0.tar.gz", hash = "sha256:6757cd03768053ff99f3039c1a36d6c0aa0b263438fcab17520b30a303a82b5f"}, +] + +[package.extras] +windows-terminal = ["colorama (>=0.4.6)"] + +[[package]] +name = "pytest" +version = "8.4.2" +description = "pytest: simple powerful testing with Python" +optional = false +python-versions = ">=3.9" +groups = ["test"] +files = [ + {file = "pytest-8.4.2-py3-none-any.whl", hash = "sha256:872f880de3fc3a5bdc88a11b39c9710c3497a547cfa9320bc3c5e62fbf272e79"}, + {file = "pytest-8.4.2.tar.gz", hash = "sha256:86c0d0b93306b961d58d62a4db4879f27fe25513d4b969df351abdddb3c30e01"}, +] + +[package.dependencies] +colorama = {version = ">=0.4", markers = "sys_platform == \"win32\""} +iniconfig = ">=1" +packaging = ">=20" +pluggy = ">=1.5,<2" +pygments = ">=2.7.2" + +[package.extras] +dev = ["argcomplete", "attrs (>=19.2)", "hypothesis (>=3.56)", "mock", "requests", "setuptools", "xmlschema"] + +[[package]] +name = "pytest-asyncio" +version = "0.24.0" +description = "Pytest support for asyncio" +optional = false +python-versions = ">=3.8" +groups = ["test"] +files = [ + {file = "pytest_asyncio-0.24.0-py3-none-any.whl", hash = "sha256:a811296ed596b69bf0b6f3dc40f83bcaf341b155a269052d82efa2b25ac7037b"}, + {file = "pytest_asyncio-0.24.0.tar.gz", hash = "sha256:d081d828e576d85f875399194281e92bf8a68d60d72d1a2faf2feddb6c46b276"}, +] + +[package.dependencies] +pytest = ">=8.2,<9" + +[package.extras] +docs = ["sphinx (>=5.3)", "sphinx-rtd-theme (>=1.0)"] +testing = ["coverage (>=6.2)", "hypothesis (>=5.7.1)"] + [[package]] name = "six" version = "1.17.0" @@ -1658,4 +1754,4 @@ propcache = ">=0.2.1" [metadata] lock-version = "2.1" python-versions = ">=3.12" -content-hash = "909c78fe0ef4860e1d6b4eebd9c20055c2c8457128a078b4465e71fad75af95d" +content-hash = "560cfc72c024b1aa1d1b406097ad03bbc4a054950c842d6a6bc07be45e7fa330" From e63fc0d3ac48be44e685af286985af1e55a2b3de Mon Sep 17 00:00:00 2001 From: This-is-XiaoDeng <1744793737@qq.com> Date: Fri, 14 Aug 2026 08:26:44 +0800 Subject: [PATCH 5/5] =?UTF-8?q?fix:=20conftest=20=E5=9C=A8=E6=94=B6?= =?UTF-8?q?=E9=9B=86=E9=98=B6=E6=AE=B5=E5=88=9B=E5=BB=BA=20dummy=20config.?= =?UTF-8?q?json=EF=BC=88CI=20=E6=97=A0=E9=85=8D=E7=BD=AE=E6=96=87=E4=BB=B6?= =?UTF-8?q?=E6=97=B6=20import=20utils=20=E5=A4=B1=E8=B4=A5=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tests/conftest.py | 27 ++++++++++++--------------- 1 file changed, 12 insertions(+), 15 deletions(-) diff --git a/tests/conftest.py b/tests/conftest.py index 2dda1af..0026228 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -5,21 +5,18 @@ import pytest - -@pytest.fixture(scope="session", autouse=True) -def ensure_config_json(): - """import utils 链需要读取 config.json(CI/干净环境不存在时创建 dummy,本地已有则不动)""" - if not os.path.exists("config.json"): - with open("config.json", "w", encoding="utf-8") as f: - json.dump( - { - "account_token": "dummy_token", - "system": {"proxy": None, "logger": {"level": 20}}, - "servers": [], - }, - f, - ) - yield +# 模块级执行(收集阶段即生效):import utils 链需要 config.json, +# CI/干净环境不存在时创建 dummy,本地已有则不动 +if not os.path.exists("config.json"): + with open("config.json", "w", encoding="utf-8") as f: + json.dump( + { + "account_token": "dummy_token", + "system": {"proxy": None, "logger": {"level": 20}}, + "servers": [], + }, + f, + ) @pytest.fixture(autouse=True)