diff --git a/hexi/core/__init__.py b/hexi/core/__init__.py index d4bf708..159778a 100644 --- a/hexi/core/__init__.py +++ b/hexi/core/__init__.py @@ -5,6 +5,7 @@ from . import ( # noqa: E402,F401 cooldown, custom_utils, message_utils, + outbound_media, plugin_control, plugin_manager, rate_limit, @@ -16,3 +17,9 @@ async def _startup_plugin_gate() -> None: """等所有插件 matcher 注册完成后,把统一 filter 规则注入到 application 插件。""" n = plugin_control.instrument_plugin_gate() logger.info(f"插件控制面: 已注入 {n} 条 gateway 规则") + + +@get_driver().on_startup +async def _startup_outbound_media() -> None: + """挂载本机媒体服务:协议端据此回拉 bot 生成的本地文件(内部已处理降级)。""" + outbound_media.mount_media_endpoint() diff --git a/hexi/core/message_utils.py b/hexi/core/message_utils.py index d9dd539..49e428d 100644 --- a/hexi/core/message_utils.py +++ b/hexi/core/message_utils.py @@ -1,4 +1,3 @@ -import mimetypes from pathlib import Path from nonebot import get_bot @@ -148,15 +147,14 @@ async def common_proc_reply( def build() -> UniMessage: # 每次重建:alconna 的 reply_to 会把引用段插入消息自身,降级重发需干净副本 - # 本地文件转 base64 内联发送:后端(NapCat)读的是它自己 CWD 下的路径, - # file:/// 形式不可靠;base64:// 由 bot 直接携带字节,OneBot 实现通用支持 + # 本地文件统一交给 hexi.core.outbound_media 的 call_api 钩子改写 + # (URL 优先/base64 兜底),这里不再自己转 base64——占位图本就是发图, + # 走同一套策略才不会出现「一处能发一处不能发」。关掉 HEXI_INLINE_MEDIA + # 的语义是「协议端就在本机」,此时 file:// 同样是有效的。 if isinstance(image, str) and image.startswith(("http://", "https://")): img = UniMessage.image(url=image) else: - img = UniMessage.image( - raw=Path(image).read_bytes(), - mimetype=mimetypes.guess_type(str(image))[0], - ) + img = UniMessage.image(path=Path(image)) return UniMessage.text(f"{text} ") + img if text else img if message_id: diff --git a/hexi/core/outbound_media.py b/hexi/core/outbound_media.py new file mode 100644 index 0000000..842c0d1 --- /dev/null +++ b/hexi/core/outbound_media.py @@ -0,0 +1,487 @@ +"""出站媒体内联化:把本地文件改写为协议端可达的 URL(或 base64 兜底)。 + +生产环境的 OneBot 协议端不在本机:给 API 传本地路径时,NoneBot 的 `f2s()` +会把 `Path` 转成 `file:///D:/...`(`str` 则原样透传)——而 `file://` 的语义是 +「**协议端那台机器上的路径**」,协议端 stat 必然 ENOENT。受影响的不止本仓插件, +还有 pip 装的社区插件(如 nonebot_plugin_doroending),它们的代码改不了, +只能全局拦截。 + +本模块挂在 `Bot.on_calling_api`(NoneBot 公开钩子,`call_api` 的必经之路, +`bot.send` / alconna `finish` / 直接 `send_group_msg` 全从这里过)上, +在真正调用适配器前改写发送类 API 里的本地文件: + +1. **默认改写为 `http://<本机可达IP>:<端口>/media/`**,由协议端主动来拉。 + 地址取自 `driver.config` 的 host/port + 本机网卡,与 /hub 首页「协议端配对」 + 展示的是同一来源;反向 WS 部署下协议端本就连得到这个地址(它连的就是它), + 所以不需要任何新配置。体积不受 base64 膨胀影响,视频/多媒体转发尤其受益。 +2. **拿不到可达地址时回退 `base64://`**:字节由 bot 直接携带,不依赖协议端 + 反过来拨通本机的任何假设,跨 NAT/隧道也成立。 + +不改写 `http(s)://` / `base64://` / 不存在的路径——收到的图片 id(如 `{abc}.jpg`) +和协议端自己的缓存路径都靠这条自然放行。任何异常只告警不抛: +拦截层绝不能反过来把发送搞挂。 + +两类 API 的处理方式不同: + +- **消息发送**(`send_*_msg` / `send_*_forward_msg`):媒体在消息段里,递归遍历 + 后按上面两级改写(URL 优先,base64 兜底)。 +- **文件上传**(`upload_group_file` / `upload_private_file`):`file` 在顶层, + 且**只做 URL 不做 base64**——文件体积不适合内联。覆盖这条让 + `video_analysis` 的群文件「本地直传」兜底和 alconna 的 `$onebot11:file` + 一起受益(后者传的是 `Path.as_posix()` 裸路径)。 +""" + +from __future__ import annotations + +import asyncio +import os +import secrets +import socket +import time +from base64 import b64encode +from pathlib import Path +from typing import Any +from urllib.parse import urlparse +from urllib.request import url2pathname + +from nonebot.adapters import Bot as BaseBot + +# 用 nonebot.log 而非 `from nonebot import logger`:与 rate_limit 保持一致, +# 模块级 import 保持干净,测试才能 importlib 按路径裸加载(不经包 __init__)。 +from nonebot.log import logger + +try: # NoneBot 内部路径:V11/V12 的 MessageSegment 都继承自它 + from nonebot.internal.adapter.message import MessageSegment as BaseMessageSegment +except Exception: # pragma: no cover - 内部路径变更时的兜底,见 _as_segment + BaseMessageSegment = None # type: ignore[assignment,misc] + +# 消息段里代表「本地/远端媒体文件」的 type(file 字段语义) +MEDIA_SEG_TYPES = frozenset({"image", "record", "video"}) + +# OneBot V11 发送类 API(与 picstatus misc_statistics 的清单一致,另加合并转发)。 +_SEND_APIS = frozenset( + { + "send_private_msg", + "send_group_msg", + "send_msg", + "send_private_forward_msg", + "send_group_forward_msg", + "send_forward_msg", + } +) + +# 文件上传类 API:file 字段在**顶层**而不是消息段里,所以走单独一小段逻辑, +# 且**只做 URL、不做 base64 兜底**——文件动辄几十上百 MB,内联会撑爆内存和 +# WS 帧,OneBot 的文件上传语义本来也不是为内联设计的。 +_FILE_UPLOAD_APIS = frozenset({"upload_group_file", "upload_private_file"}) + +# 只有 OneBot 的 file 字段认 file:// / base64:// 语义;bot.py 当前只注册 V11 +# (pyproject.toml 虽列了 V12,但代码里没装适配器)。 +_ONEBOT_NAMES = frozenset({"OneBot V11"}) + +# 通配绑定:值本身不可拨号,需要换成具体网卡地址 +_WILDCARD_HOSTS = frozenset({"", "0.0.0.0", "::", "*"}) + +_MAX_DEPTH = 8 +_MAX_TOKENS = 256 +_DEFAULT_TTL = 300.0 +_DEFAULT_MAX_MB = 4.0 + +# token -> (文件路径, 过期时刻 monotonic) +_tokens: dict[str, tuple[Path, float]] = {} + +# 媒体服务是否真的挂上了:没挂上就不能发 URL(否则协议端必然 404) +_mounted = False + + +# ---------------------------------------------------------------- 环境开关 + + +def _env_flag(name: str, default: bool) -> bool: + raw = os.environ.get(name) + if raw is None: + return default + return raw.strip().lower() not in {"", "0", "false", "no", "off"} + + +def _enabled() -> bool: + """总开关(默认开)。关闭的语义是「协议端就在本机」,此时 file:// 是有效的。""" + return _env_flag("HEXI_INLINE_MEDIA", True) + + +def _url_enabled() -> bool: + """URL 优先模式(默认开);关掉则一律走 base64。""" + return _env_flag("HEXI_INLINE_MEDIA_URL", True) + + +def _ttl() -> float: + try: + return float(os.environ.get("HEXI_MEDIA_URL_TTL", "") or _DEFAULT_TTL) + except ValueError: + return _DEFAULT_TTL + + +def _max_inline_mb() -> float: + try: + return float(os.environ.get("HEXI_INLINE_MEDIA_MAX_MB", "") or _DEFAULT_MAX_MB) + except ValueError: + return _DEFAULT_MAX_MB + + +# ---------------------------------------------------------------- 路径解析 + + +def local_file_of(value: Any) -> Path | None: + """把 file 字段值解析成本机真实存在的文件路径;不是本地文件则 None。 + + - `http(s)://` / `base64://` → None(原样放行) + - `file://` → `url2pathname` 还原 + - 其余按原始路径候选(`f2s()` 对 `str` 是原样透传的) + - 不存在 / 非法 → None(收到的图片 id、协议端缓存路径靠这条放行) + """ + if not isinstance(value, str) or not value: + return None + if value.startswith(("http://", "https://", "base64://")): + return None + + candidates: list[str] = [] + if value.startswith("file://"): + # 关键陷阱:不能直接 unquote(urlparse(uri).path)——Windows 上会得到 + # 前导斜杠的 `/D:/...`,is_file() 恒为 False 而**静默不转换**, + # 只在跨机时才暴露。url2pathname 才能正确还原成 `D:\\...`。 + try: + uri_path = urlparse(value).path + candidates.append(url2pathname(uri_path)) + except Exception: + return None + # 有人手工拼 file:// 且不做百分号转义(video_analysis._file_uri 即如此: + # `"file:///" + path.replace("\\", "/")`)。对已经未转义的串再 unquote, + # 文件名含字面 %XX 时会认错,所以把原始路径也列为候选兜底。 + if len(uri_path) > 2 and uri_path[0] == "/" and uri_path[2] == ":": + candidates.append(uri_path[1:]) # /D:/x → D:/x + else: + # 裸路径:f2s() 对 str 原样透传;alconna 的 $onebot11:file 走 as_posix() + candidates.append(value) + + for raw in candidates: + try: + path = Path(raw) + if path.is_file(): + return path + except (OSError, ValueError): + continue + return None + + +# ---------------------------------------------------------------- 媒体服务 + + +def _is_private(ip: str) -> bool: + parts = ip.split(".") + if len(parts) != 4: + return False + try: + a, b = int(parts[0]), int(parts[1]) + except ValueError: + return False + return a == 10 or (a == 172 and 16 <= b <= 31) or (a == 192 and b == 168) + + +def _local_ips() -> list[str]: + """本机可被协议端访问的 IPv4 列表(私有网段优先,排除环回/链路本地)。 + + 与 hexi/web_hub/dashboard.py 的 `_local_ips` 同源逻辑(那边供 /hub 首页 + 「协议端配对」展示,即用户已经验证过能用的地址);此处刻意不 import, + 避免 hexi.core → hexi.web_hub 的反向依赖。 + """ + ips: list[str] = [] + try: + import psutil + + for addrs in psutil.net_if_addrs().values(): + for addr in addrs: + if addr.family == socket.AF_INET: + ips.append(str(addr.address or "")) + except Exception: # noqa: BLE001 - psutil 缺失/异常时退回 getaddrinfo + pass + + if not ips: + try: + infos = socket.getaddrinfo(socket.gethostname(), None, socket.AF_INET) + ips.extend(info[4][0] for info in infos) + except OSError: + pass + + out: list[str] = [] + for ip in ips: + ip = ip.strip() + if not ip or ":" in ip or ip.count(".") < 3: + continue + # 环回与链路本地(169.254.x)协议端访问不到,排除 + if ip.startswith(("127.", "169.254.")): + continue + if ip not in out: + out.append(ip) + out.sort(key=lambda ip: (not _is_private(ip), ip)) + return out + + +def _listen_addr() -> tuple[str, int] | None: + """本机媒体服务的对外可达地址(host, port)。""" + try: + from nonebot import get_driver + + config = get_driver().config + port = int(getattr(config, "port", 0) or 0) + host = str(getattr(config, "host", "") or "").strip() + except Exception: # noqa: BLE001 + return None + + if not port: + return None + if host not in _WILDCARD_HOSTS: + # 显式绑定(如 dev 的 127.0.0.1)直接用 + return host, port + for ip in _local_ips(): + return ip, port + return None + + +def _base_url() -> str | None: + """媒体服务基地址;未挂载/无可达地址时 None(调用方据此回退 base64)。""" + if not _mounted: + return None + override = os.environ.get("HEXI_MEDIA_BASE_URL", "").strip() + if override: + return override.rstrip("/") + addr = _listen_addr() + if addr is None: + return None + host, port = addr + return f"http://{host}:{port}" + + +def _gc() -> None: + now = time.monotonic() + for token, (_, expires) in list(_tokens.items()): + if expires <= now: + del _tokens[token] + while len(_tokens) >= _MAX_TOKENS: + _tokens.pop(next(iter(_tokens))) + + +def _publish(path: Path) -> str: + """登记一个一次性 token(值随机不可猜),返回它。""" + _gc() + token = secrets.token_urlsafe(16) + _tokens[token] = (path, time.monotonic() + _ttl()) + return token + + +def mount_media_endpoint() -> None: + """把媒体路由挂到现有 ASGI app 上(同一个端口,不新开服务)。 + + 免鉴权是**必须**的——协议端登录不了 hub——所以安全模型改为 + 「路径由 bot 在改写那一刻自己登记,请求方只能出示不可猜 token、 + 根本无法表达路径」,天然没有目录穿越面,配合短 TTL。 + """ + global _mounted + + if not _enabled(): + logger.info("出站媒体内联化已禁用 (HEXI_INLINE_MEDIA=false),按原值发送") + return + if _mounted: + return + + try: + from fastapi import FastAPI + from nonebot import get_app + + sub = FastAPI( + title="HeXi Outbound Media", + docs_url=None, + redoc_url=None, + openapi_url=None, + ) + + @sub.get("/{token}") + async def media(token: str): + from fastapi import HTTPException + from fastapi.responses import FileResponse + + entry = _tokens.get(token) + if entry is None: + raise HTTPException(status_code=404, detail="媒体不存在或已过期") + path, expires = entry + if expires <= time.monotonic(): + _tokens.pop(token, None) + raise HTTPException(status_code=404, detail="媒体链接已过期") + if not path.is_file(): + raise HTTPException(status_code=404, detail="文件已不存在") + # no-store:过期媒体不该被中间层缓存住 + return FileResponse(path, headers={"Cache-Control": "no-store"}) + + get_app().mount("/media", sub) + _mounted = True + logger.info("出站媒体服务已挂载: /media/") + except Exception as e: # noqa: BLE001 - 挂载失败只降级,不影响启动 + logger.warning( + f"出站媒体服务挂载失败,将回退 base64: {type(e).__name__}: {e}" + ) + + +# ---------------------------------------------------------------- 改写 + + +async def _to_base64(path: Path) -> str: + # 同步 IO 挪到线程,别在事件循环里读文件 + data = await asyncio.to_thread(path.read_bytes) + size_mb = len(data) / 1048576 + limit = _max_inline_mb() + if size_mb > limit: + # 超阈值仍要发:不转是 100% 必失败,转了才可能成功 + logger.warning( + f"出站媒体内联: {path.name} 体积 {size_mb:.1f}MB 超过 " + f"HEXI_INLINE_MEDIA_MAX_MB={limit:g},仍以 base64 发送" + ) + return f"base64://{b64encode(data).decode()}" + + +async def rewrite_file(value: Any) -> str | None: + """改写单个 file 字段值;无需改写(或改写失败)返回 None。""" + path = local_file_of(value) + if path is None: + return None + try: + if _url_enabled(): + base = _base_url() + if base: + return f"{base}/media/{_publish(path)}" + return await _to_base64(path) + except Exception as e: # noqa: BLE001 - 单个段失败只告警,按原值发送 + logger.warning(f"出站媒体改写失败,按原值发送 {path}: {type(e).__name__}: {e}") + return None + + +async def rewrite_upload_file(value: Any) -> str | None: + """文件上传类 API 的 file 字段改写;无需改写返回 None。 + + 只做 URL,**不做 base64 兜底**:拿不到可达地址就原样放行,行为与改动前 + 一致(不保证成功,但绝不会比原来更糟)。`HEXI_INLINE_MEDIA_URL=false` + 同样会让这里放行——那个开关的语义是「媒体服务这条路不可用」,文件上传 + 依赖同一条路,理应一起关掉。 + """ + if not _url_enabled(): + return None + path = local_file_of(value) + if path is None: + return None + try: + base = _base_url() + if not base: + return None + return f"{base}/media/{_publish(path)}" + except Exception as e: # noqa: BLE001 + logger.warning(f"群/私聊文件改写失败,按原值发送 {path}: {type(e).__name__}: {e}") + return None + + +def _as_segment(node: Any) -> tuple[str, dict] | None: + """识别消息段对象 → (type, data);不是消息段返回 None。 + + MessageSegment 不是 dict 子类(但实现了 keys/get),所以单靠 isinstance(dict) + 认不出来;这里以 NoneBot 的公共基类为准,并留一条结构兜底以防内部路径变更。 + """ + if BaseMessageSegment is not None and isinstance(node, BaseMessageSegment): + seg_type, seg_data = node.type, node.data + else: + seg_type = getattr(node, "type", None) + seg_data = getattr(node, "data", None) + if isinstance(seg_type, str) and isinstance(seg_data, dict): + return seg_type, seg_data + return None + + +async def _walk_segment(seg_type: str, seg_data: dict, depth: int) -> None: + if seg_type in MEDIA_SEG_TYPES: + rewritten = await rewrite_file(seg_data.get("file")) + if rewritten is not None: + seg_data["file"] = rewritten + # 合并转发节点:content 可能是 str / list[dict] / Message,里面还可能嵌媒体 + # (MessageSegment.node_custom 把 Message 原样塞进 data["content"]) + for key in ("content", "messages"): + if key in seg_data: + await _walk(seg_data[key], depth + 1) + + +async def _walk(node: Any, depth: int = 0) -> None: + """深度受限的通用遍历,原地改写。 + + 形态是散的(Message / 纯 dict 段 / list[dict] / 转发节点的三种 content), + 写死形状必漏,所以按结构特征递归而不是枚举。 + """ + if node is None or depth > _MAX_DEPTH: + return + # str/bytes 是叶子;数字等标量也直接跳过 + if isinstance(node, (str, bytes, int, float, bool)): + return + + seg = _as_segment(node) + if seg is not None: + await _walk_segment(seg[0], seg[1], depth) + return + + if isinstance(node, (list, tuple)): # 含 Message(list 子类) + for item in node: + await _walk(item, depth + 1) + return + + if isinstance(node, dict): + seg_type, seg_data = node.get("type"), node.get("data") + if isinstance(seg_type, str) and isinstance(seg_data, dict): + # 段字典({"type": "image", "data": {...}} / node 字典) + await _walk_segment(seg_type, seg_data, depth) + return + for value in node.values(): + await _walk(value, depth + 1) + + +# ---------------------------------------------------------------- 钩子 + + +@BaseBot.on_calling_api +async def _inline_outbound_media(bot: BaseBot, api: str, data: dict[str, Any]) -> None: + """call_api 前置钩子:把出站 API 里的本地文件改写成可跨机访问的形式。 + + 钩子拿到的是 `_call_api` 之前**同一个可变 data dict**,且 hook 的 task group + 在 `_call_api` 之前 await 完成,所以原地改生效。 + """ + if not _enabled(): + return + is_send = api in _SEND_APIS + is_upload = api in _FILE_UPLOAD_APIS + if not (is_send or is_upload) or not isinstance(data, dict): + return + try: + if bot.adapter.get_name() not in _ONEBOT_NAMES: + return + except Exception: # noqa: BLE001 + return + + if is_upload: + # 文件上传的 file 在顶层(不是消息段);name/folder 等参数一律不碰 + try: + rewritten = await rewrite_upload_file(data.get("file")) + except Exception as e: # noqa: BLE001 - 绝不反过来搞挂发送 + logger.warning(f"出站媒体改写异常(upload.file),按原值发送: {e}") + return + if rewritten is not None: + data["file"] = rewritten + return + + for key in ("message", "messages"): + if key in data: + try: + await _walk(data[key]) + except Exception as e: # noqa: BLE001 - 绝不反过来搞挂发送 + logger.warning(f"出站媒体改写异常({key}),按原值发送: {e}") diff --git a/tests/test_outbound_media.py b/tests/test_outbound_media.py new file mode 100644 index 0000000..dd48ff2 --- /dev/null +++ b/tests/test_outbound_media.py @@ -0,0 +1,433 @@ +"""出站媒体内联化单元测试 + +hexi/core 包 __init__ 会触发 NoneBot 运行时初始化(get_driver), +无法裸导入,故用 importlib 以真实模块名加载 outbound_media 并注册进 +sys.modules(game: 与 test_rate_limit / test_cooldown 同一套路)。 + +文件一律用 tmp_path 造真实文件,**不 mock 文件系统**——这类转换的坑 +恰恰在真实路径语义上(file:// 的百分号编码、Windows 盘符、相对路径), +mock 掉就等于把要验的东西验没了。 +""" + +import importlib.util +import sys +from pathlib import Path + +import pytest +from nonebot.adapters.onebot.v11 import Message, MessageSegment + +_BASE = Path(__file__).resolve().parents[1] / "hexi" / "core" +_MOD_NAME = "hexi.core.outbound_media" + + +def _load_module(name: str, path: Path): + spec = importlib.util.spec_from_file_location(name, path) + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + return module + + +@pytest.fixture(scope="module") +def om(): + module = _load_module(_MOD_NAME, _BASE / "outbound_media.py") + yield module + sys.modules.pop(_MOD_NAME, None) + + +@pytest.fixture +def media_file(tmp_path: Path) -> Path: + """带空格与中文的真实文件:同时覆盖 URL 编码与多字节路径两种坑。""" + fp = tmp_path / "有 空格 的文件.jpg" + fp.write_bytes(b"hello") + return fp + + +class _FakeAdapter: + def __init__(self, name: str = "OneBot V11"): + self._name = name + + def get_name(self) -> str: + return self._name + + +class _FakeBot: + def __init__(self, name: str = "OneBot V11"): + self.adapter = _FakeAdapter(name) + + +# ------------------------------------------------------------ local_file_of + + +def test_file_uri_with_spaces_and_unicode(om, media_file: Path): + """file:// URI(带 %20 与中文百分号编码)必须能还原成本机真实路径。 + + 这是整个模块最容易静默失效的一处:天真的 + `Path(unquote(urlparse(uri).path))` 在 Windows 上会得到前导斜杠的 + `/D:/...`,is_file() 恒为 False 而**不报错**,只在跨机时才暴露。 + """ + uri = media_file.as_uri() + assert "%" in uri # 确认确实带百分号编码,否则这条测试没验到东西 + assert om.local_file_of(uri) == media_file + + +def test_plain_absolute_path(om, media_file: Path): + """f2s() 对 str 是原样透传的,所以裸路径也得认。""" + assert om.local_file_of(str(media_file)) == media_file + + +def test_relative_path(om, media_file: Path, monkeypatch): + monkeypatch.chdir(media_file.parent) + assert om.local_file_of(media_file.name) == Path(media_file.name) + + +@pytest.mark.parametrize( + "value", + [ + "http://example.com/a.jpg", + "https://example.com/a.jpg", + "base64://aGVsbG8=", + ], +) +def test_remote_and_inline_values_pass_through(om, value: str): + assert om.local_file_of(value) is None + + +@pytest.mark.parametrize( + "value", + [ + "D:/nope/does_not_exist.jpg", + "file:///D:/nope/does_not_exist.jpg", + # 协议端回传的图片 id / 它自己的缓存路径:不是本机文件,必须自然放行 + "{12345678-1234-1234-1234-123456789abc}.jpg", + "C:/NapCat/cache/abcdef.image", + "", + ], +) +def test_non_local_values_are_left_alone(om, value: str): + assert om.local_file_of(value) is None + + +@pytest.mark.parametrize("value", [None, 123, b"bytes", ["a.jpg"]]) +def test_non_string_values_are_left_alone(om, value): + assert om.local_file_of(value) is None + + +# ------------------------------------------------------------ 改写目标 + + +async def test_falls_back_to_base64_when_no_url(om, media_file: Path, monkeypatch): + """拿不到可达地址(未挂载)→ base64 兜底,字节由 bot 自己携带。""" + monkeypatch.setattr(om, "_mounted", False) + assert await om.rewrite_file(str(media_file)) == "base64://aGVsbG8=" + + +async def test_url_mode_when_mounted(om, media_file: Path, monkeypatch): + monkeypatch.setattr(om, "_mounted", True) + monkeypatch.setenv("HEXI_MEDIA_BASE_URL", "http://192.168.2.15:39697/") + url = await om.rewrite_file(str(media_file)) + assert url.startswith("http://192.168.2.15:39697/media/") + # 基地址结尾斜杠不该被拼成 //media + assert "//media" not in url + + token = url.rsplit("/", 1)[1] + assert om._tokens[token][0] == media_file + + +async def test_url_mode_can_be_disabled(om, media_file: Path, monkeypatch): + """URL 模式关掉后一律 base64,即使媒体服务挂着。""" + monkeypatch.setattr(om, "_mounted", True) + monkeypatch.setenv("HEXI_MEDIA_BASE_URL", "http://192.168.2.15:39697") + monkeypatch.setenv("HEXI_INLINE_MEDIA_URL", "false") + assert await om.rewrite_file(str(media_file)) == "base64://aGVsbG8=" + + +async def test_rewrite_is_idempotent(om, media_file: Path, monkeypatch): + """重复执行必须无副作用:第二次看到 base64:// 应直接放行。""" + monkeypatch.setattr(om, "_mounted", False) + once = await om.rewrite_file(str(media_file)) + assert await om.rewrite_file(once) is None + + +async def test_rewrite_returns_none_for_remote(om): + assert await om.rewrite_file("http://example.com/a.jpg") is None + + +# ------------------------------------------------------------ 遍历形态 + + +async def test_walk_rewrites_message_segments(om, media_file: Path, monkeypatch): + monkeypatch.setattr(om, "_mounted", False) + msg = Message( + [ + MessageSegment.text("看图"), + MessageSegment.image(media_file), + MessageSegment.face(1), + ] + ) + await om._walk(msg) + assert msg[0].data["text"] == "看图" + assert msg[1].data["file"] == "base64://aGVsbG8=" + assert msg[2].type == "face" # 非媒体段不动 + + +@pytest.mark.parametrize("seg_type", ["image", "record", "video"]) +async def test_walk_rewrites_all_media_types( + om, media_file: Path, monkeypatch, seg_type +): + monkeypatch.setattr(om, "_mounted", False) + seg = {"type": seg_type, "data": {"file": str(media_file)}} + await om._walk(seg) + assert seg["data"]["file"] == "base64://aGVsbG8=" + + +async def test_walk_handles_bare_dict_segment(om, media_file: Path, monkeypatch): + """group_daily_analysis 那种纯 dict 段(不是 Message 也不是 MessageSegment)。""" + monkeypatch.setattr(om, "_mounted", False) + payload = {"type": "image", "data": {"file": str(media_file)}} + await om._walk(payload) + assert payload["data"]["file"] == "base64://aGVsbG8=" + + +async def test_walk_handles_list_of_dict_segments(om, media_file: Path, monkeypatch): + monkeypatch.setattr(om, "_mounted", False) + payload = [ + {"type": "text", "data": {"text": "hi"}}, + {"type": "image", "data": {"file": str(media_file)}}, + ] + await om._walk(payload) + assert payload[1]["data"]["file"] == "base64://aGVsbG8=" + + +async def test_walk_forward_node_with_str_content(om, media_file: Path, monkeypatch): + """message_utils.build_forward_nodes: content 是 str(不含媒体,应保持原样)。""" + monkeypatch.setattr(om, "_mounted", False) + nodes = [ + {"type": "node", "data": {"name": "x", "uin": "1", "content": "纯文本"}}, + ] + await om._walk(nodes) + assert nodes[0]["data"]["content"] == "纯文本" + + +async def test_walk_forward_node_with_dict_list_content( + om, media_file: Path, monkeypatch +): + """video_analysis._build_forward_nodes: content 是段字典数组。""" + monkeypatch.setattr(om, "_mounted", False) + nodes = [ + { + "type": "node", + "data": { + "name": "x", + "uin": "1", + "content": [{"type": "video", "data": {"file": str(media_file)}}], + }, + } + ] + await om._walk(nodes) + assert nodes[0]["data"]["content"][0]["data"]["file"] == "base64://aGVsbG8=" + + +async def test_walk_forward_node_with_message_content( + om, media_file: Path, monkeypatch +): + """MessageSegment.node_custom 会把 Message 原样塞进 data["content"]。""" + monkeypatch.setattr(om, "_mounted", False) + node = MessageSegment.node_custom( + 1, "n", Message([MessageSegment.image(media_file)]) + ) + await om._walk([node]) + inner = node.data["content"][0] + assert isinstance(inner, MessageSegment) + assert inner.data["file"] == "base64://aGVsbG8=" + + +async def test_walk_recurses_messages_key(om, media_file: Path, monkeypatch): + """转发的 messages 还可能是 Message 对象(alconna 的 include() 返回同类型)。""" + monkeypatch.setattr(om, "_mounted", False) + data = { + "messages": Message( + [ + MessageSegment.node_custom( + 1, "n", Message([MessageSegment.image(media_file)]) + ) + ] + ) + } + await om._walk(data["messages"]) + assert data["messages"][0].data["content"][0].data["file"] == "base64://aGVsbG8=" + + +async def test_walk_leaves_remote_media_alone(om, media_file: Path, monkeypatch): + monkeypatch.setattr(om, "_mounted", False) + seg = {"type": "image", "data": {"file": "https://example.com/a.jpg"}} + await om._walk(seg) + assert seg["data"]["file"] == "https://example.com/a.jpg" + + +async def test_walk_respects_depth_limit(om, media_file: Path, monkeypatch): + """超深嵌套不再下钻:拦截层不该在畸形结构上无界递归。""" + monkeypatch.setattr(om, "_mounted", False) + node: dict = {"type": "image", "data": {"file": str(media_file)}} + for _ in range(20): + node = {"type": "node", "data": {"content": [node]}} + await om._walk(node) + deepest = node + while isinstance(deepest, dict) and "content" in deepest.get("data", {}): + deepest = deepest["data"]["content"][0] + assert deepest["data"]["file"] == str(media_file) # 没被改写 + + +# ------------------------------------------------------------ 钩子 + + +async def test_hook_rewrites_send_api(om, media_file: Path, monkeypatch): + monkeypatch.setattr(om, "_mounted", False) + data = {"group_id": 1, "message": Message([MessageSegment.image(media_file)])} + await om._inline_outbound_media(_FakeBot(), "send_group_msg", data) + assert data["message"][0].data["file"] == "base64://aGVsbG8=" + + +async def test_hook_ignores_unrelated_api(om, media_file: Path, monkeypatch): + """既非发送也非文件上传的 API 一律不碰。""" + monkeypatch.setattr(om, "_mounted", True) + monkeypatch.setenv("HEXI_MEDIA_BASE_URL", "http://192.168.2.186:39697") + data = {"group_id": 1, "file": str(media_file)} + await om._inline_outbound_media(_FakeBot(), "get_group_info", data) + assert data["file"] == str(media_file) + + +async def test_hook_ignores_other_adapters(om, media_file: Path, monkeypatch): + monkeypatch.setattr(om, "_mounted", False) + data = {"message": Message([MessageSegment.image(media_file)])} + before = data["message"][0].data["file"] + await om._inline_outbound_media(_FakeBot("Telegram"), "send_group_msg", data) + assert data["message"][0].data["file"] == before + + +async def test_hook_disabled_by_env(om, media_file: Path, monkeypatch): + """关掉开关 = 声明「协议端就在本机」,此时原样发送才是对的。""" + monkeypatch.setenv("HEXI_INLINE_MEDIA", "false") + data = {"message": Message([MessageSegment.image(media_file)])} + before = data["message"][0].data["file"] + await om._inline_outbound_media(_FakeBot(), "send_group_msg", data) + assert data["message"][0].data["file"] == before + # f2s() 对 Path 产出 file:// URI,正是本模块要拦的那种值 + assert before.startswith("file://") + + +async def test_hook_handles_forward_messages_key(om, media_file: Path, monkeypatch): + monkeypatch.setattr(om, "_mounted", False) + data = { + "group_id": 1, + "messages": [ + { + "type": "node", + "data": { + "name": "n", + "uin": "1", + "content": [{"type": "image", "data": {"file": str(media_file)}}], + }, + } + ], + } + await om._inline_outbound_media(_FakeBot(), "send_group_forward_msg", data) + seg = data["messages"][0]["data"]["content"][0] + assert seg["data"]["file"] == "base64://aGVsbG8=" + + +async def test_hook_never_raises_on_broken_segment(om, monkeypatch): + """拦截层绝不能反过来把发送搞挂:畸形段只告警。""" + monkeypatch.setattr(om, "_mounted", False) + data = {"message": [{"type": "image", "data": {"file": object()}}]} + await om._inline_outbound_media(_FakeBot(), "send_group_msg", data) + + +# ------------------------------------------------------------ 文件上传 + + +def _unescaped_uri(path: Path) -> str: + """video_analysis._file_uri 的手工拼法:正斜杠、不转义。""" + return "file:///" + str(path.resolve()).replace("\\", "/") + + +async def test_upload_file_url_when_mounted(om, media_file: Path, monkeypatch): + monkeypatch.setattr(om, "_mounted", True) + monkeypatch.setenv("HEXI_MEDIA_BASE_URL", "http://192.168.2.186:39697") + data = {"group_id": 1, "file": str(media_file), "name": "显示名.zip"} + await om._inline_outbound_media(_FakeBot(), "upload_group_file", data) + assert data["file"].startswith("http://192.168.2.186:39697/media/") + assert data["name"] == "显示名.zip" # 其它参数一律不碰 + assert om._tokens[data["file"].rsplit("/", 1)[1]][0] == media_file + + +@pytest.mark.parametrize( + "builder", + [ + str, # 裸路径(alconna $onebot11:file 的 as_posix() 形态) + lambda p: p.as_uri(), # Path.as_uri():带百分号转义 + _unescaped_uri, # 手工拼:不转义 + ], + ids=["plain", "as_uri", "unescaped"], +) +async def test_upload_file_accepts_all_uri_forms( + om, media_file: Path, monkeypatch, builder +): + monkeypatch.setattr(om, "_mounted", True) + monkeypatch.setenv("HEXI_MEDIA_BASE_URL", "http://192.168.2.186:39697") + data = {"file": builder(media_file)} + await om._inline_outbound_media(_FakeBot(), "upload_private_file", data) + assert data["file"].startswith("http://192.168.2.186:39697/media/") + + +async def test_upload_file_has_no_base64_fallback(om, media_file: Path, monkeypatch): + """文件上传不做 base64 兜底:没 URL 就原样放行(不能比改动前更糟)。""" + monkeypatch.setattr(om, "_mounted", False) + data = {"group_id": 1, "file": str(media_file)} + await om._inline_outbound_media(_FakeBot(), "upload_group_file", data) + assert data["file"] == str(media_file) + + +async def test_upload_file_respects_url_toggle(om, media_file: Path, monkeypatch): + monkeypatch.setattr(om, "_mounted", True) + monkeypatch.setenv("HEXI_MEDIA_BASE_URL", "http://192.168.2.186:39697") + monkeypatch.setenv("HEXI_INLINE_MEDIA_URL", "false") + data = {"group_id": 1, "file": str(media_file)} + await om._inline_outbound_media(_FakeBot(), "upload_group_file", data) + assert data["file"] == str(media_file) + + +async def test_upload_file_ignores_s3_url(om, media_file: Path, monkeypatch): + """S3 预签名链接是群文件的首选路径,不能被碰。""" + monkeypatch.setattr(om, "_mounted", True) + monkeypatch.setenv("HEXI_MEDIA_BASE_URL", "http://192.168.2.186:39697") + url = "http://192.168.2.15:5246/PLANA/x.zip?X-Amz-Signature=abc" + data = {"group_id": 1, "file": url} + await om._inline_outbound_media(_FakeBot(), "upload_group_file", data) + assert data["file"] == url + + +async def test_local_file_of_handles_literal_percent(tmp_path: Path, om): + """文件名含字面 %XX 时,未转义的手拼 URI 不能被 unquote 带偏。 + + 先按解码结果找(会解析成错误路径),再退回原始路径 —— 这条兜底的意义就在此。 + """ + fp = tmp_path / "100%20off.zip" + fp.write_bytes(b"zip") + assert om.local_file_of(_unescaped_uri(fp)) == fp + + +async def test_media_endpoint_roundtrip(om, media_file: Path, monkeypatch): + """token → 路径 的登记与过期语义(不依赖 FastAPI,直接验注册表)。""" + monkeypatch.setattr(om, "_mounted", True) + monkeypatch.setenv("HEXI_MEDIA_BASE_URL", "http://127.0.0.1:11011") + monkeypatch.setenv("HEXI_MEDIA_URL_TTL", "60") + url = await om.rewrite_file(str(media_file)) + token = url.rsplit("/", 1)[1] + + path, expires = om._tokens[token] + assert path == media_file + assert expires > om.time.monotonic() + # token 不可猜:长度足够且不含路径信息 + assert len(token) >= 20 + assert "空格" not in token