Add HeXi bot codebase: custom plugins, web frontends, tests
- hexi core: message handling, rate limiting, cooldown, plugin manager - Custom plugins: BF stats, daily check-in, quotes, persona cards, etc. - Community plugins vendored under hexi/plugins with local fixes - Web admin frontends (learning-chat, persona-admin), unified hexi/web - Tests for rate_limit/cooldown/memes/persona; poetry.lock Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,285 @@
|
||||
"""命令冷却单元测试
|
||||
|
||||
hexi_core 包 __init__ 会触发 NoneBot 运行时初始化(get_plugin_config),
|
||||
无法裸导入,故用 importlib 以真实模块名加载 rate_limit / cooldown 并注册进
|
||||
sys.modules —— cooldown 内部对 rate_limit 的 import 会直接命中缓存,
|
||||
不触发包 __init__ 的导入链。
|
||||
|
||||
NoneBot 的依赖注入会对参数做类型校验,故 guard 测试使用真实
|
||||
GroupMessageEvent / Matcher 实例,并通过 ContextVar 提供 finish 所需环境。
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
import re
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from nonebot.adapters.onebot.v11 import (
|
||||
Bot,
|
||||
GroupMessageEvent,
|
||||
Message,
|
||||
PrivateMessageEvent,
|
||||
)
|
||||
from nonebot.adapters.onebot.v11.event import Sender
|
||||
from nonebot.dependencies import Dependent
|
||||
from nonebot.exception import FinishedException
|
||||
from nonebot.internal.matcher.matcher import current_bot, current_event
|
||||
from nonebot.matcher import Matcher
|
||||
|
||||
_BASE = (
|
||||
Path(__file__).resolve().parents[1]
|
||||
/ "hexi"
|
||||
/ "plugins"
|
||||
/ "nonebot_plugin_hexi_core"
|
||||
)
|
||||
_RL_NAME = "hexi.plugins.nonebot_plugin_hexi_core.rate_limit"
|
||||
_CD_NAME = "hexi.plugins.nonebot_plugin_hexi_core.cooldown"
|
||||
|
||||
|
||||
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 cd():
|
||||
_load_module(_RL_NAME, _BASE / "rate_limit.py")
|
||||
module = _load_module(_CD_NAME, _BASE / "cooldown.py")
|
||||
yield module
|
||||
sys.modules.pop(_CD_NAME, None)
|
||||
sys.modules.pop(_RL_NAME, None)
|
||||
|
||||
|
||||
def _event(uid: int = 123456, group_id: int = 2) -> GroupMessageEvent:
|
||||
return GroupMessageEvent(
|
||||
time=0,
|
||||
self_id=1,
|
||||
post_type="message",
|
||||
message_type="group",
|
||||
sub_type="normal",
|
||||
user_id=uid,
|
||||
message_id=1,
|
||||
message=Message("hi"),
|
||||
raw_message="hi",
|
||||
font=0,
|
||||
sender=Sender(user_id=uid),
|
||||
group_id=group_id,
|
||||
)
|
||||
|
||||
|
||||
def _private_event(uid: int = 123456) -> PrivateMessageEvent:
|
||||
return PrivateMessageEvent(
|
||||
time=0,
|
||||
self_id=1,
|
||||
post_type="message",
|
||||
message_type="private",
|
||||
sub_type="friend",
|
||||
user_id=uid,
|
||||
message_id=1,
|
||||
message=Message("hi"),
|
||||
raw_message="hi",
|
||||
font=0,
|
||||
sender=Sender(user_id=uid),
|
||||
)
|
||||
|
||||
|
||||
class FakeBot:
|
||||
"""极简 bot:send 记录消息"""
|
||||
|
||||
def __init__(self):
|
||||
self.sent = []
|
||||
|
||||
async def send(self, event, message, **kwargs):
|
||||
self.sent.append(str(message))
|
||||
return 0
|
||||
|
||||
|
||||
def _make_matcher(cd, seconds=10, **kwargs):
|
||||
mcls = Matcher.new("message")
|
||||
return cd.cooldown(seconds, **kwargs)(mcls)
|
||||
|
||||
|
||||
async def _call_guard(mcls, bot, event):
|
||||
"""模拟 NoneBot 调用 guard handler(Dependent 全关键字调用)"""
|
||||
m = mcls()
|
||||
dep = mcls.handlers[0]
|
||||
tok_bot = current_bot.set(bot)
|
||||
tok_event = current_event.set(event)
|
||||
try:
|
||||
return await dep(
|
||||
matcher=m,
|
||||
bot=bot,
|
||||
event=event,
|
||||
state={},
|
||||
stack=None,
|
||||
dependency_cache=None,
|
||||
)
|
||||
finally:
|
||||
current_bot.reset(tok_bot)
|
||||
current_event.reset(tok_event)
|
||||
|
||||
|
||||
# ---------- Cooldown 类 ----------
|
||||
|
||||
|
||||
def test_trigger_and_query(cd):
|
||||
c = cd.Cooldown(3)
|
||||
assert c.try_trigger("u1") is True
|
||||
assert c.in_cd("u1") is True
|
||||
assert 0 < c.remaining("u1") <= 3.01 # 浮点容差
|
||||
assert c.try_trigger("u1") is False # CD 中再次触发失败
|
||||
|
||||
|
||||
def test_instances_isolated(cd):
|
||||
c1 = cd.Cooldown(3)
|
||||
c2 = cd.Cooldown(3)
|
||||
assert c1.try_trigger("u1") is True
|
||||
assert c2.try_trigger("u1") is True # 无 name 的不同实例互不影响
|
||||
|
||||
|
||||
def test_remaining_after_window(cd):
|
||||
c = cd.Cooldown(3)
|
||||
assert c.try_trigger("u1") is True
|
||||
limiter = c._limiter("u1")
|
||||
limiter._last_refill = time.monotonic() - 3.1 # 模拟 3.1 秒后
|
||||
assert c.in_cd("u1") is False
|
||||
assert c.remaining("u1") == 0
|
||||
assert c.try_trigger("u1") is True # 冷却结束可再次触发
|
||||
|
||||
|
||||
def test_independent_per_key(cd):
|
||||
c = cd.Cooldown(3)
|
||||
assert c.try_trigger("u1") is True
|
||||
assert c.in_cd("u2") is False # 其他用户不受影响
|
||||
assert c.try_trigger("u2") is True
|
||||
|
||||
|
||||
def test_invalid_seconds(cd):
|
||||
with pytest.raises(ValueError):
|
||||
cd.Cooldown(0)
|
||||
|
||||
|
||||
def test_rate_limiter_remaining_secs(cd):
|
||||
rl = sys.modules[_RL_NAME]
|
||||
limiter = rl.RateLimiter(rate=1, capacity=1)
|
||||
assert limiter.remaining_secs() == 0 # 满桶无等待
|
||||
assert limiter.try_acquire() is True
|
||||
assert 0 < limiter.remaining_secs() <= 1
|
||||
limiter._last_refill = time.monotonic() - 2.0 # 2 秒后已补满
|
||||
assert limiter.remaining_secs() == 0
|
||||
|
||||
|
||||
# ---------- cooldown 装饰器 ----------
|
||||
|
||||
|
||||
def test_guard_registered_first(cd):
|
||||
mcls = _make_matcher(cd, 10)
|
||||
assert len(mcls.handlers) == 1
|
||||
|
||||
async def real_handler(bot: Bot, ev: GroupMessageEvent):
|
||||
pass
|
||||
|
||||
mcls.handlers.append(
|
||||
Dependent.parse(call=real_handler, allow_types=Matcher.HANDLER_PARAM_TYPES)
|
||||
) # 模拟 @matcher.handle() 追加真 handler
|
||||
assert mcls.handlers[0].call is not real_handler # guard 仍在最前
|
||||
|
||||
|
||||
async def test_guard_passes_and_blocks_with_hint(cd):
|
||||
mcls = _make_matcher(cd, 10, hint="慢点,{secs} 秒后")
|
||||
bot = FakeBot()
|
||||
ev = _event(1001)
|
||||
|
||||
# 首次触发:放行,不发提示
|
||||
assert await _call_guard(mcls, bot, ev) is None
|
||||
assert bot.sent == []
|
||||
|
||||
# 冷却中触发:finish 提示并终止
|
||||
with pytest.raises(FinishedException):
|
||||
await _call_guard(mcls, bot, ev)
|
||||
assert len(bot.sent) == 1
|
||||
msg = bot.sent[0]
|
||||
assert "慢点," in msg and "秒后" in msg
|
||||
secs = int(re.search(r"\d+", msg).group())
|
||||
assert 1 <= secs <= 10 # 剩余秒数向上取整
|
||||
|
||||
|
||||
async def test_guard_blocks_only_same_user(cd):
|
||||
mcls = _make_matcher(cd, 10)
|
||||
bot = FakeBot()
|
||||
await _call_guard(mcls, bot, _event(1001))
|
||||
await _call_guard(mcls, bot, _event(1002)) # 其他用户不受影响
|
||||
assert bot.sent == []
|
||||
with pytest.raises(FinishedException):
|
||||
await _call_guard(mcls, bot, _event(1001))
|
||||
assert len(bot.sent) == 1
|
||||
|
||||
|
||||
async def test_guard_custom_key_shares_cd(cd):
|
||||
mcls = _make_matcher(cd, 10, key=lambda e: "global")
|
||||
bot = FakeBot()
|
||||
await _call_guard(mcls, bot, _event(1001))
|
||||
# 自定义 key 下不同用户共享同一冷却
|
||||
with pytest.raises(FinishedException):
|
||||
await _call_guard(mcls, bot, _event(1002))
|
||||
assert len(bot.sent) == 1
|
||||
|
||||
|
||||
# ---------- scope 作用域 ----------
|
||||
|
||||
|
||||
async def test_scope_group_shared_by_group(cd):
|
||||
mcls = _make_matcher(cd, 10, scope="group")
|
||||
bot = FakeBot()
|
||||
await _call_guard(mcls, bot, _event(1001, group_id=2))
|
||||
# 同群其他用户也处于 CD 中
|
||||
with pytest.raises(FinishedException):
|
||||
await _call_guard(mcls, bot, _event(1002, group_id=2))
|
||||
assert len(bot.sent) == 1
|
||||
|
||||
|
||||
async def test_scope_group_isolated_between_groups(cd):
|
||||
mcls = _make_matcher(cd, 10, scope="group")
|
||||
bot = FakeBot()
|
||||
await _call_guard(mcls, bot, _event(1001, group_id=2))
|
||||
# 不同群的用户不受影响
|
||||
await _call_guard(mcls, bot, _event(1001, group_id=3))
|
||||
assert bot.sent == []
|
||||
|
||||
|
||||
async def test_scope_group_private_falls_back_to_user(cd):
|
||||
mcls = _make_matcher(cd, 10, scope="group")
|
||||
bot = FakeBot()
|
||||
await _call_guard(mcls, bot, _event(1001, group_id=2))
|
||||
# 私聊没有群,退化为按用户:不同用户私聊互不影响
|
||||
await _call_guard(mcls, bot, _private_event(1002))
|
||||
# 同一用户私聊与群聊各自独立
|
||||
await _call_guard(mcls, bot, _private_event(1001))
|
||||
assert bot.sent == []
|
||||
with pytest.raises(FinishedException):
|
||||
await _call_guard(mcls, bot, _private_event(1001))
|
||||
assert len(bot.sent) == 1
|
||||
|
||||
|
||||
async def test_scope_global_shared(cd):
|
||||
mcls = _make_matcher(cd, 10, scope="global")
|
||||
bot = FakeBot()
|
||||
await _call_guard(mcls, bot, _event(1001, group_id=2))
|
||||
# 任何用户、任何群都共享
|
||||
with pytest.raises(FinishedException):
|
||||
await _call_guard(mcls, bot, _event(1002, group_id=3))
|
||||
assert len(bot.sent) == 1
|
||||
|
||||
|
||||
async def test_key_takes_priority_over_scope(cd):
|
||||
mcls = _make_matcher(cd, 10, scope="global", key=lambda e: e.get_user_id())
|
||||
bot = FakeBot()
|
||||
await _call_guard(mcls, bot, _event(1001))
|
||||
# 传了 key 则忽略 scope:不同用户互不影响
|
||||
await _call_guard(mcls, bot, _event(1002))
|
||||
assert bot.sent == []
|
||||
@@ -0,0 +1,104 @@
|
||||
"""表情裸词操作补丁(hexi/plugins/memes_ops)的回归测试。
|
||||
|
||||
验证 `对称 右` 这类“表情 + 裸词操作”能被新 matcher 正确解析为选项,
|
||||
同时 `-右`/`--right` 旧形式仍然可用。
|
||||
|
||||
若 nonebot_plugin_memes 缺失或补丁因版本不兼容而失效,本模块整体跳过。
|
||||
"""
|
||||
|
||||
import nonebot
|
||||
import pytest
|
||||
|
||||
try:
|
||||
nonebot.init()
|
||||
from nonebot import get_loaded_plugins, load_plugin
|
||||
|
||||
load_plugin("hexi.plugins.memes_ops")
|
||||
_memes_loaded = True
|
||||
except Exception:
|
||||
_memes_loaded = False
|
||||
|
||||
if not _memes_loaded:
|
||||
pytest.skip(
|
||||
"无法加载 nonebot_plugin_memes / memes_ops,跳过裸词操作补丁测试",
|
||||
allow_module_level=True,
|
||||
)
|
||||
|
||||
_patchers = [
|
||||
matcher
|
||||
for plugin in get_loaded_plugins()
|
||||
for matcher in plugin.matcher
|
||||
if matcher.priority == 11 and matcher.module_name == "hexi.plugins.memes_ops"
|
||||
]
|
||||
|
||||
if not _patchers:
|
||||
pytest.skip(
|
||||
"memes_ops 补丁 matcher 未注册(版本不兼容?),跳过裸词操作补丁测试",
|
||||
allow_module_level=True,
|
||||
)
|
||||
|
||||
_cmd_start = list(nonebot.get_driver().config.command_start)[0]
|
||||
|
||||
|
||||
def _parse(msg: str):
|
||||
"""在所有新 matcher 中解析消息,返回 (选项, 文字参数) 命中集合。"""
|
||||
hits = []
|
||||
for matcher in _patchers:
|
||||
try:
|
||||
result = matcher.command().parse(msg)
|
||||
except Exception:
|
||||
continue
|
||||
if result.matched:
|
||||
options = {
|
||||
key: value.value
|
||||
for key, value in result.options.items()
|
||||
if value.value not in (None, False)
|
||||
}
|
||||
params = [
|
||||
seg.text for seg in result.query("meme_params", ()) if seg.text
|
||||
]
|
||||
hits.append((options, params))
|
||||
return hits
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("msg", "expected_options"),
|
||||
[
|
||||
(f"{_cmd_start}对称 右", {"right"}),
|
||||
(f"{_cmd_start}对称 左", {"left"}),
|
||||
(f"{_cmd_start}对称 上", {"top"}),
|
||||
(f"{_cmd_start}对称 下", {"bottom"}),
|
||||
(f"{_cmd_start}鬼畜 右", {"right"}),
|
||||
(f"{_cmd_start}循环 下", {"bottom"}),
|
||||
(f"{_cmd_start}摸 圆", {"circle"}),
|
||||
(f"{_cmd_start}小丑 爷", {"person"}),
|
||||
(f"{_cmd_start}小丑面具 前", {"front"}),
|
||||
(f"{_cmd_start}ba说 左", {"left"}),
|
||||
],
|
||||
)
|
||||
def test_bare_operation_maps_to_option(msg, expected_options):
|
||||
hits = _parse(msg)
|
||||
assert hits, f"{msg!r} 未被任何补丁 matcher 命中"
|
||||
# 新 matcher 优先(priority 11),取第一个命中
|
||||
options, params = hits[0]
|
||||
assert set(options) == expected_options
|
||||
assert params == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"msg",
|
||||
[f"{_cmd_start}对称 -右", f"{_cmd_start}对称 --right", f"{_cmd_start}鬼畜 --bottom"],
|
||||
)
|
||||
def test_dash_syntax_still_works(msg):
|
||||
options, params = _parse(msg)[0]
|
||||
assert options, f"{msg!r} 应解析出选项"
|
||||
assert params == []
|
||||
|
||||
|
||||
def test_non_alias_word_stays_text():
|
||||
msg = f"{_cmd_start}对称 任意文本"
|
||||
hits = _parse(msg)
|
||||
assert hits
|
||||
options, params = hits[0]
|
||||
assert options == {}
|
||||
assert params == ["任意文本"]
|
||||
@@ -0,0 +1,135 @@
|
||||
"""人设卡治理层纯函数测试
|
||||
|
||||
galgame_card 包 __init__ 会触发 NoneBot 运行时初始化(require orm),
|
||||
无法裸导入,故用 importlib 按文件路径直接加载 processor 模块
|
||||
(与 test_rate_limit.py 同款做法)。
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
_MODULE_PATH = (
|
||||
Path(__file__).resolve().parents[1]
|
||||
/ "hexi"
|
||||
/ "plugins"
|
||||
/ "nonebot_plugin_galgame_card"
|
||||
/ "processor.py"
|
||||
)
|
||||
|
||||
|
||||
def _load_processor():
|
||||
spec = importlib.util.spec_from_file_location("persona_processor", _MODULE_PATH)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def processor():
|
||||
return _load_processor()
|
||||
|
||||
|
||||
COMMAND_STARTS = {"/", ""}
|
||||
|
||||
|
||||
def test_filter_content_normal(processor):
|
||||
assert processor.filter_content("草,又被加班干碎了", COMMAND_STARTS) == "草,又被加班干碎了"
|
||||
|
||||
|
||||
def test_filter_content_command(processor):
|
||||
assert processor.filter_content("/打卡", COMMAND_STARTS) is None
|
||||
|
||||
|
||||
def test_filter_content_command_empty_prefix_not_fatal(processor):
|
||||
# command_start 里空串不应导致全量误杀
|
||||
assert processor.filter_content("正常说话", {""}) == "正常说话"
|
||||
assert processor.filter_content("/打卡", {""}) == "/打卡"
|
||||
|
||||
|
||||
def test_desensitize_clear_patterns(processor):
|
||||
assert processor.desensitize("我电话是13800138000") == "我电话是[手机号]"
|
||||
assert processor.desensitize("手机号138-0013-8000") == "手机号[手机号]" # 带分隔符变体
|
||||
assert processor.desensitize("身份证110101199003070012") == "身份证[身份证]"
|
||||
assert processor.desensitize("卡号6222021234567890") == "卡号[银行卡]"
|
||||
assert processor.desensitize("联系我 a@b.com") == "联系我 [邮箱]"
|
||||
assert processor.desensitize("服务器192.168.1.1") == "服务器[IP]"
|
||||
assert processor.desensitize("加我 wxid_abc123xyz") == "加我 [微信号]"
|
||||
assert processor.desensitize("车牌川A12345") == "车牌[车牌]"
|
||||
assert processor.desensitize("定位30.5,104.0") == "定位[坐标]"
|
||||
|
||||
|
||||
def test_desensitize_key_value(processor):
|
||||
# 定位式:关键词 + 值,只替换值、保留关键词
|
||||
assert processor.desensitize("密码是 xyz789") == "密码是 [密码]"
|
||||
assert processor.desensitize("账号 abc123") == "账号 [账号]"
|
||||
assert processor.desensitize("账号是 1122334455") == "账号是 [账号]"
|
||||
assert processor.desensitize("激活码 ABCDE-12345") == "激活码 [兑换码]"
|
||||
assert processor.desensitize("验证码是 123456") == "验证码是 [验证码]"
|
||||
assert processor.desensitize("月薪 25000") == "月薪 [收入]"
|
||||
assert processor.desensitize("VX: xiaoming123") == "VX: [微信号]"
|
||||
|
||||
|
||||
def test_desensitize_fallback(processor):
|
||||
# 兜底:裸长数字串 → [账号],字母数字混合 → [密码]
|
||||
assert processor.desensitize("1122334455") == "[账号]"
|
||||
assert processor.desensitize("abc1234") == "[密码]"
|
||||
|
||||
|
||||
def test_desensitize_context_hint(processor):
|
||||
# 语境强化:附近出现"验证码"关键词时,6-8 位数字 → [验证码]
|
||||
assert processor.desensitize("123456", context_text="发下验证码") == "[验证码]"
|
||||
assert processor.desensitize("验证码是 123456") == "验证码是 [验证码]"
|
||||
# 无关键词时 6-8 位数字走兜底 → [账号]
|
||||
assert processor.desensitize("123456") == "[账号]"
|
||||
|
||||
|
||||
def test_desensitize_keeps_normal_text(processor):
|
||||
# 正常文本不受影响(短数字、中文、英文单词)
|
||||
assert processor.desensitize("今晚八点开黑,打了2800分") == "今晚八点开黑,打了2800分"
|
||||
assert processor.desensitize("有一说一确实wwww") == "有一说一确实wwww"
|
||||
assert processor.desensitize("房间码12345") == "房间码12345"
|
||||
assert processor.desensitize("体重120斤") == "体重120斤"
|
||||
|
||||
|
||||
def test_filter_content_url_flood(processor):
|
||||
text = "看看这个 http://a.com http://b.com http://c.com"
|
||||
assert processor.filter_content(text, COMMAND_STARTS) is None
|
||||
|
||||
|
||||
def test_filter_content_truncate(processor):
|
||||
long_text = "长" * 500
|
||||
assert len(processor.filter_content(long_text, COMMAND_STARTS)) == processor.MAX_CONTENT_LEN
|
||||
|
||||
|
||||
def test_is_command(processor):
|
||||
assert processor.is_command("/查询", {"/"})
|
||||
assert not processor.is_command("你好", {"/"})
|
||||
|
||||
|
||||
def test_truncate(processor):
|
||||
assert processor.truncate("abc", 2) == "ab"
|
||||
assert processor.truncate("abc", 10) == "abc"
|
||||
|
||||
|
||||
def test_image_hash_from_file(processor):
|
||||
HASH32 = "0123456789abcdef0123456789abcdef"
|
||||
# 优先 file 里的 32 位 hash
|
||||
assert processor.image_hash_from_file(f"{HASH32}.image") == HASH32
|
||||
# 回退 url 文件名里的 hash
|
||||
assert (
|
||||
processor.image_hash_from_file("", f"https://example.com/pics/{HASH32}.png")
|
||||
== HASH32
|
||||
)
|
||||
# 都没有则截断文件名
|
||||
assert processor.image_hash_from_file("somefile.png") == "somefile.png"
|
||||
|
||||
|
||||
def test_build_content(processor):
|
||||
assert processor.build_content("") is None
|
||||
assert processor.build_content("草") == "草"
|
||||
assert processor.build_content("", 2) == "[图片×2]"
|
||||
assert processor.build_content("", 0, 3) == "[表情×3]"
|
||||
assert processor.build_content("", 2, 3) == "[图片×2][表情×3]"
|
||||
assert processor.build_content("草", 1) == "草 [图片×1]"
|
||||
@@ -0,0 +1,219 @@
|
||||
"""限频器单元测试
|
||||
|
||||
hexi_core 包 __init__ 会触发 NoneBot 运行时初始化(get_plugin_config),
|
||||
无法裸导入,故用 importlib 按文件路径直接加载 rate_limit 模块
|
||||
(该模块只依赖 nonebot.log,可独立于 NoneBot 运行)。
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
_MODULE_PATH = (
|
||||
Path(__file__).resolve().parents[1]
|
||||
/ "hexi"
|
||||
/ "plugins"
|
||||
/ "nonebot_plugin_hexi_core"
|
||||
/ "rate_limit.py"
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def rl():
|
||||
"""以独立模块名加载 rate_limit.py,避免与包 __init__ 的依赖链冲突"""
|
||||
spec = importlib.util.spec_from_file_location("hexi_rate_limit_test", _MODULE_PATH)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
sys.modules[spec.name] = module
|
||||
spec.loader.exec_module(module)
|
||||
yield module
|
||||
sys.modules.pop(spec.name, None)
|
||||
|
||||
|
||||
def test_capacity_limits_burst(rl):
|
||||
limiter = rl.RateLimiter(rate=10, capacity=3)
|
||||
assert limiter.try_acquire() is True
|
||||
assert limiter.try_acquire() is True
|
||||
assert limiter.try_acquire() is True
|
||||
assert limiter.try_acquire() is False # 突发额度耗尽
|
||||
|
||||
|
||||
def test_default_capacity_at_least_one(rl):
|
||||
assert rl.RateLimiter(rate=5)._capacity == 5 # 默认容量 = rate
|
||||
assert rl.RateLimiter(rate=0.5)._capacity == 1 # rate < 1 时保底 1
|
||||
|
||||
|
||||
def test_invalid_params(rl):
|
||||
with pytest.raises(ValueError):
|
||||
rl.RateLimiter(rate=0)
|
||||
with pytest.raises(ValueError):
|
||||
rl.RateLimiter(rate=1, capacity=0)
|
||||
|
||||
|
||||
def test_lazy_refill_by_backdating(rl):
|
||||
"""把 last_refill 回拨模拟时间流逝,验证惰性补充与容量上限"""
|
||||
limiter = rl.RateLimiter(rate=2, capacity=2)
|
||||
assert limiter.try_acquire() is True
|
||||
assert limiter.try_acquire() is True
|
||||
assert limiter.try_acquire() is False # 桶空
|
||||
limiter._last_refill = time.monotonic() - 1.0 # 假装 1 秒前补充过
|
||||
assert limiter.try_acquire() is True # 1 秒补 2 个
|
||||
assert limiter.try_acquire() is True # 桶已补满,再取 1 个
|
||||
assert limiter.try_acquire() is False # 再次耗尽
|
||||
|
||||
|
||||
def test_capacity_caps_refill(rl):
|
||||
limiter = rl.RateLimiter(rate=10, capacity=2)
|
||||
limiter._last_refill = time.monotonic() - 10.0 # 10 秒应补 100 个,但桶上限 2
|
||||
assert limiter.try_acquire() is True
|
||||
assert limiter.try_acquire() is True
|
||||
assert limiter.try_acquire() is False
|
||||
|
||||
|
||||
def test_acquire_sync_waits_for_token(rl):
|
||||
limiter = rl.RateLimiter(rate=1, capacity=1)
|
||||
assert limiter.try_acquire() is True # 消耗唯一令牌
|
||||
start = time.monotonic()
|
||||
assert limiter.acquire_sync(timeout=5) is True # 约 1 秒后补出令牌
|
||||
assert 0.6 <= time.monotonic() - start < 3
|
||||
|
||||
|
||||
def test_acquire_sync_timeout(rl):
|
||||
limiter = rl.RateLimiter(rate=0.1, capacity=1) # 10 秒才补 1 个
|
||||
assert limiter.try_acquire() is True
|
||||
start = time.monotonic()
|
||||
assert limiter.acquire_sync(timeout=0.2) is False
|
||||
assert time.monotonic() - start < 2
|
||||
|
||||
|
||||
def test_acquire_sync_zero_timeout_fails_fast(rl):
|
||||
limiter = rl.RateLimiter(rate=0.001, capacity=1)
|
||||
assert limiter.try_acquire() is True
|
||||
assert limiter.acquire_sync(timeout=0) is False # 立即失败,不等
|
||||
|
||||
|
||||
async def test_acquire_async_waits_for_token(rl):
|
||||
limiter = rl.RateLimiter(rate=1, capacity=1)
|
||||
assert limiter.try_acquire() is True
|
||||
start = time.monotonic()
|
||||
assert await limiter.acquire(timeout=5) is True
|
||||
assert 0.6 <= time.monotonic() - start < 3
|
||||
|
||||
|
||||
async def test_acquire_async_timeout(rl):
|
||||
limiter = rl.RateLimiter(rate=0.1, capacity=1)
|
||||
assert limiter.try_acquire() is True
|
||||
assert await limiter.acquire(timeout=0.2) is False
|
||||
|
||||
|
||||
def test_get_limiter_shared_by_name(rl):
|
||||
a = rl.get_limiter("shared_test", rate=1)
|
||||
b = rl.get_limiter("shared_test", rate=1)
|
||||
assert a is b # 同名复用同一实例
|
||||
assert rl.get_limiter("other_test", rate=1) is not a
|
||||
|
||||
|
||||
async def test_acquire_helper(rl):
|
||||
# 首次取令牌成功,立即再取失败(不等待)
|
||||
assert await rl.acquire("helper_failfast", rate=0.001, wait=False) is True
|
||||
assert await rl.acquire("helper_failfast", rate=0.001, wait=False) is False
|
||||
# timeout=0 等模式同样立即失败
|
||||
assert await rl.acquire("helper_failfast", rate=0.001, timeout=0) is False
|
||||
|
||||
|
||||
def test_acquire_sync_helper(rl):
|
||||
assert rl.acquire_sync("helper_sync", rate=0.001, wait=False) is True
|
||||
assert rl.acquire_sync("helper_sync", rate=0.001, wait=False) is False
|
||||
|
||||
|
||||
def test_is_ratelimited(rl):
|
||||
assert rl.is_ratelimited(403) is True
|
||||
assert rl.is_ratelimited(429) is True
|
||||
assert rl.is_ratelimited(200) is False
|
||||
assert rl.is_ratelimited(404) is False
|
||||
assert rl.is_ratelimited(None) is False
|
||||
|
||||
|
||||
# ---------- steam_info 场景模拟 ----------
|
||||
|
||||
|
||||
class _FakeTime:
|
||||
"""可手动推进的 time 替身:验证限频时序无需真实等待"""
|
||||
|
||||
def __init__(self):
|
||||
self._t = 1000.0
|
||||
|
||||
def monotonic(self) -> float:
|
||||
return self._t
|
||||
|
||||
def advance(self, secs: float) -> None:
|
||||
self._t += secs
|
||||
|
||||
|
||||
class _FakeAsyncio:
|
||||
"""asyncio 替身:sleep 立即完成并推进回拨时钟(模拟"等待配额"过程)"""
|
||||
|
||||
def __init__(self, fake_time: _FakeTime):
|
||||
self._ft = fake_time
|
||||
|
||||
async def sleep(self, secs: float) -> None:
|
||||
self._ft.advance(secs)
|
||||
|
||||
|
||||
async def test_steam_quota_simulation(rl):
|
||||
"""模拟 steam_info 场景:播报 1/min(等待)+ check 高频插入(快速失败)
|
||||
|
||||
参数与 steam.py 的 STEAM_API_RATE/STEAM_API_CAPACITY 保持一致
|
||||
(rate=1/2 → 相邻请求间隔 >= 2s,双 key 合计 150 req/5min)。
|
||||
回拨时钟跑 600 模拟秒,验证:
|
||||
- 相邻成功请求间隔 >= 2s(短窗口突发被杜绝)
|
||||
- 等待模式的播报在 check 抢配额下仍全部成功
|
||||
- 快速失败的 check 在配额耗尽时被拒绝
|
||||
"""
|
||||
fake = _FakeTime()
|
||||
real_time, real_asyncio = rl.time, rl.asyncio
|
||||
rl.time = fake
|
||||
rl.asyncio = _FakeAsyncio(fake)
|
||||
try:
|
||||
rate, capacity = 1 / 2, 1
|
||||
gaps: list[float] = []
|
||||
last_ok = fake.monotonic()
|
||||
broadcast_ok = 0
|
||||
check_ok = 0
|
||||
check_rejected = 0
|
||||
|
||||
next_broadcast, next_check = 60.0, 2.0
|
||||
t = 0.0
|
||||
while t <= 600:
|
||||
if t >= next_broadcast:
|
||||
next_broadcast += 60
|
||||
if await rl.acquire(
|
||||
"steam_api", rate=rate, capacity=capacity, timeout=60
|
||||
):
|
||||
gaps.append(fake.monotonic() - last_ok)
|
||||
last_ok = fake.monotonic()
|
||||
broadcast_ok += 1
|
||||
if t >= next_check:
|
||||
next_check += 2
|
||||
if await rl.acquire(
|
||||
"steam_api", rate=rate, capacity=capacity, wait=False
|
||||
):
|
||||
gaps.append(fake.monotonic() - last_ok)
|
||||
last_ok = fake.monotonic()
|
||||
check_ok += 1
|
||||
else:
|
||||
check_rejected += 1
|
||||
fake.advance(1.0)
|
||||
t += 1.0
|
||||
|
||||
# 等待模式的播报 10 次全部成功(最多等 60s,不会被 check 挤掉)
|
||||
assert broadcast_ok == 10
|
||||
# 快速失败的 check 在配额耗尽时被拒(check 300 次尝试 vs 配额 300)
|
||||
assert check_rejected > 0
|
||||
# 相邻成功请求间隔 >= 2s(首请求无前驱,不计入)
|
||||
assert check_ok + broadcast_ok <= 300
|
||||
assert min(gaps[1:]) >= 1.9
|
||||
finally:
|
||||
rl.time, rl.asyncio = real_time, real_asyncio
|
||||
Reference in New Issue
Block a user