feat(core): 出站媒体内联化全局钩子 / 本地文件改写为协议端可达 URL
协议端不在本机,而 NoneBot 的 f2s() 会把 Path 转成 file:///D:/...(file:// 的
语义是「协议端那台机器上的路径」),必然 ENOENT。受影响的不止本仓调用点,还有
装不了改不动的 pip 社区插件(本次报错来自 nonebot_plugin_doroending),只能全局拦截。
- 新增 hexi/core/outbound_media.py:挂 Bot.on_calling_api(公开钩子、适配器不覆盖),
在真正调用适配器前原地改写发送类 API 的 data:默认改写成
http://<本机可达IP>:<端口>/media/<token> 由协议端主动回拉(地址与 /hub 首页
「协议端配对」同源,零新配置),拿不到可达地址时回退 base64://
- 文件上传类(upload_group_file / upload_private_file)单独处理:file 在顶层且
只做 URL 不做 base64(几十上百 MB 不适合内联),覆盖 video_analysis 群文件
「本地直传」兜底与 alconna 的 $onebot11:file
- 媒体路由 get_app().mount("/media", sub) 复用现有端口不新开服务;免鉴权是必须的
(协议端登录不了 hub),安全模型改为「路径由 bot 改写那一刻自己登记、请求方只能
出示不可猜 token」,配短 TTL;_mounted 为 False 时绝不发 URL
- 钩子内每个段独立 try/except,任何异常只告警不抛,绝不反过来把发送搞挂
- message_utils.common_proc_reply 占位图不再自己转 base64,改传 path= 交给钩子,
避免两处各有一套策略
- tests/test_outbound_media.py:用 tmp_path 造真实文件而非 mock 文件系统,重点覆盖
file:// 解析(空格/中文编码、手工未转义的字面 %XX)与 walker 各消息形态
Co-Authored-By: Claude Code <noreply@anthropic.com>
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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/<token>`**,由协议端主动来拉。
|
||||
地址取自 `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/<token>")
|
||||
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}")
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user