fix: retain discord typing task owners

This commit is contained in:
jxxghp
2026-08-24 09:43:17 +08:00
parent ecbd39a0bb
commit 2274a354dc
4 changed files with 424 additions and 215 deletions
+163 -55
View File
@@ -47,6 +47,7 @@ class Discord:
DISCORD_CHANNEL_ID: Optional[Union[str, int]] = None, DISCORD_CHANNEL_ID: Optional[Union[str, int]] = None,
**kwargs, **kwargs,
): ):
"""初始化 Discord 客户端及其专用事件循环生命周期状态。"""
logger.debug( logger.debug(
f"[Discord] 初始化 Discord 实例: name={kwargs.get('name')}, " f"[Discord] 初始化 Discord 实例: name={kwargs.get('name')}, "
f"GUILD_ID={DISCORD_GUILD_ID}, CHANNEL_ID={DISCORD_CHANNEL_ID}, " f"GUILD_ID={DISCORD_GUILD_ID}, CHANNEL_ID={DISCORD_CHANNEL_ID}, "
@@ -84,9 +85,12 @@ class Discord:
self._bot_user_id: Optional[int] = None self._bot_user_id: Optional[int] = None
self._typing_tasks: Dict[str, asyncio.Task] = {} self._typing_tasks: Dict[str, asyncio.Task] = {}
self._typing_stop_events: Dict[str, asyncio.Event] = {} self._typing_stop_events: Dict[str, asyncio.Event] = {}
self._typing_lifecycle_lock = asyncio.Lock()
self._typing_accepting = True
self._typing_interval_seconds = 5 self._typing_interval_seconds = 5
self._typing_initial_delay_seconds = 1 self._typing_initial_delay_seconds = 1
self._typing_max_duration_seconds = 10 * 60 self._typing_max_duration_seconds = 10 * 60
self._typing_stop_timeout_seconds = 1
self._registered_commands: Optional[Dict[str, dict]] = None self._registered_commands: Optional[Dict[str, dict]] = None
self._register_events() self._register_events()
@@ -200,10 +204,12 @@ class Discord:
await self._post_to_ds(payload) await self._post_to_ds(payload)
def _start(self): def _start(self):
"""启动并持有 Discord 客户端专用事件循环线程。"""
if self._thread: if self._thread:
return return
def runner(): def runner():
"""在线程内运行客户端,并在退出前回收本实例拥有的异步任务。"""
loop = self._loop loop = self._loop
client = self._client client = self._client
asyncio.set_event_loop(loop) asyncio.set_event_loop(loop)
@@ -226,6 +232,10 @@ class Discord:
loop.run_until_complete(start_task) loop.run_until_complete(start_task)
except asyncio.CancelledError: except asyncio.CancelledError:
pass pass
try:
loop.run_until_complete(self._stop_all_typing_tasks())
except Exception as err:
logger.debug(f"Discord typing 任务收尾失败:{err}")
try: try:
loop.run_until_complete(client.close()) loop.run_until_complete(client.close())
except Exception as err: except Exception as err:
@@ -238,6 +248,7 @@ class Discord:
self._thread.start() self._thread.start()
def stop(self): def stop(self):
"""停止 Discord 客户端,并在关闭事件循环前收口 typing owner。"""
if not self._client or not self._loop or not self._thread: if not self._client or not self._loop or not self._thread:
return return
self._stop_requested.set() self._stop_requested.set()
@@ -644,75 +655,172 @@ class Discord:
max_duration_seconds: Optional[float] = None, max_duration_seconds: Optional[float] = None,
initial_delay_seconds: Optional[float] = None, initial_delay_seconds: Optional[float] = None,
) -> bool: ) -> bool:
await self._stop_typing_task(typing_key) """
为会话启动唯一的 Discord typing 续发任务。
:return: 是否取得该会话的唯一 typing owner
"""
if not self._typing_accepting:
logger.debug("Discord client 已停止,拒绝启动 typing 任务")
return False
channel = await self._resolve_channel(userid=userid, chat_id=chat_id) channel = await self._resolve_channel(userid=userid, chat_id=chat_id)
if not channel: if not channel:
return False return False
stop_event = asyncio.Event() async with self._typing_lifecycle_lock:
max_duration = max_duration_seconds or self._typing_max_duration_seconds if not self._typing_accepting:
initial_delay = ( logger.debug("Discord client 已停止,拒绝启动 typing 任务")
self._typing_initial_delay_seconds return False
if initial_delay_seconds is None found, terminal = await self._stop_typing_task_locked(typing_key)
else max(initial_delay_seconds, 0) if found and not terminal:
) logger.warning(
f"Discord typing 旧任务尚未结束,拒绝并行启动: key={typing_key}"
)
return False
async def _typing_worker() -> None: stop_event = asyncio.Event()
started_at = self._loop.time() max_duration = max_duration_seconds or self._typing_max_duration_seconds
try: initial_delay = (
# Discord typing 触发后也会在客户端自然保留一段时间, self._typing_initial_delay_seconds
# 先给短响应一个取消窗口,避免回复后残留输入状态。 if initial_delay_seconds is None
if initial_delay: else max(initial_delay_seconds, 0)
try: )
await asyncio.wait_for(
stop_event.wait(),
timeout=initial_delay,
)
return
except asyncio.TimeoutError:
pass
while not stop_event.is_set():
if self._loop.time() - started_at >= max_duration:
logger.warning(
"Discord typing状态超过最大续期,自动停止: key=%s",
typing_key,
)
break
try:
await channel.trigger_typing()
except Exception as err:
logger.debug(f"触发 Discord typing 状态失败:{err}")
try:
await asyncio.wait_for(
stop_event.wait(),
timeout=self._typing_interval_seconds,
)
except asyncio.TimeoutError:
pass
finally:
current_task = asyncio.current_task()
if self._typing_tasks.get(typing_key) is current_task:
self._typing_tasks.pop(typing_key, None)
self._typing_stop_events.pop(typing_key, None)
self._typing_stop_events[typing_key] = stop_event async def _typing_worker() -> None:
self._typing_tasks[typing_key] = asyncio.create_task(_typing_worker()) """延迟首发并定期续发当前会话的 typing 状态。"""
return True started_at = self._loop.time()
try:
# Discord typing 触发后也会在客户端自然保留一段时间,
# 先给短响应一个取消窗口,避免回复后残留输入状态。
if initial_delay:
try:
await asyncio.wait_for(
stop_event.wait(),
timeout=initial_delay,
)
return
except asyncio.TimeoutError:
pass
while not stop_event.is_set():
if self._loop.time() - started_at >= max_duration:
logger.warning(
f"Discord typing状态超过最大续期,自动停止: key={typing_key}"
)
break
try:
await channel.trigger_typing()
except Exception as err:
logger.debug(f"触发 Discord typing 状态失败:{err}")
try:
await asyncio.wait_for(
stop_event.wait(),
timeout=self._typing_interval_seconds,
)
except asyncio.TimeoutError:
pass
finally:
current_task = asyncio.current_task()
if self._typing_tasks.get(typing_key) is current_task:
self._typing_tasks.pop(typing_key, None)
self._typing_stop_events.pop(typing_key, None)
task = asyncio.create_task(
_typing_worker(),
name=f"discord.typing.{typing_key}",
)
self._typing_stop_events[typing_key] = stop_event
self._typing_tasks[typing_key] = task
return True
async def _stop_typing_task(self, typing_key: str) -> bool: async def _stop_typing_task(self, typing_key: str) -> bool:
stop_event = self._typing_stop_events.pop(typing_key, None) """
task = self._typing_tasks.pop(typing_key, None) 请求停止会话 typing 任务,并保留尚未进入终态的 owner。
:return: 是否找到并请求停止了既有 owner
"""
async with self._typing_lifecycle_lock:
found, _ = await self._stop_typing_task_locked(typing_key)
return found
async def _stop_typing_task_locked(self, typing_key: str) -> tuple[bool, bool]:
"""
在 lifecycle 锁内停止单个 owner。
:return: 是否找到 owner,以及 owner 是否已进入终态
"""
stop_event = self._typing_stop_events.get(typing_key)
task = self._typing_tasks.get(typing_key)
found = bool(stop_event or task)
if stop_event: if stop_event:
stop_event.set() stop_event.set()
if task and task is not asyncio.current_task() and not task.done(): if task and task is not asyncio.current_task() and not task.done():
try: try:
await asyncio.wait_for(asyncio.shield(task), timeout=1) await asyncio.wait_for(
asyncio.shield(task),
timeout=self._typing_stop_timeout_seconds,
)
except asyncio.TimeoutError: except asyncio.TimeoutError:
pass return found, False
return bool(stop_event or task) except asyncio.CancelledError:
if not task.cancelled():
raise
except Exception as err:
logger.debug(f"Discord typing 任务异常结束:{err}")
terminal = task is None or task.done()
if terminal:
if self._typing_tasks.get(typing_key) is task:
self._typing_tasks.pop(typing_key, None)
if self._typing_stop_events.get(typing_key) is stop_event:
self._typing_stop_events.pop(typing_key, None)
return found, terminal
async def _stop_all_typing_tasks(self) -> None: async def _stop_all_typing_tasks(self) -> None:
for typing_key in list(self._typing_tasks.keys()): """封住新增 typing,并在统一预算内停止、取消和回收现有 owner。"""
await self._stop_typing_task(typing_key) async with self._typing_lifecycle_lock:
self._typing_accepting = False
typing_keys = set(self._typing_tasks) | set(self._typing_stop_events)
owners = {
typing_key: self._typing_tasks.get(typing_key)
for typing_key in typing_keys
}
for typing_key in typing_keys:
stop_event = self._typing_stop_events.get(typing_key)
if stop_event:
stop_event.set()
current_task = asyncio.current_task()
active_tasks = {
task
for task in owners.values()
if task and task is not current_task and not task.done()
}
settled_tasks = {task for task in owners.values() if task and task.done()}
if active_tasks:
done, pending = await asyncio.wait(
active_tasks,
timeout=self._typing_stop_timeout_seconds,
)
settled_tasks.update(done)
for task in pending:
task.cancel()
if pending:
done, _ = await asyncio.wait(
pending,
timeout=self._typing_stop_timeout_seconds,
)
settled_tasks.update(done)
for task in settled_tasks:
if task.cancelled():
continue
error = task.exception()
if error:
logger.debug(f"Discord typing 任务异常结束:{error}")
for typing_key, task in owners.items():
if task is None or task.done():
if self._typing_tasks.get(typing_key) is task:
self._typing_tasks.pop(typing_key, None)
self._typing_stop_events.pop(typing_key, None)
def delete_msg( def delete_msg(
self, message_id: Union[str, int], chat_id: Optional[str] = None self, message_id: Union[str, int], chat_id: Optional[str] = None
@@ -6,7 +6,7 @@
> 审计范围:宿主后端;排除 `app/plugins/**` 运行时插件副本 > 审计范围:宿主后端;排除 `app/plugins/**` 运行时插件副本
> 规范优先级:`AGENTS.md` 与 `docs/rules/` 高于本文 > 规范优先级:`AGENTS.md` 与 `docs/rules/` 高于本文
> 相关文档:`docs/architecture-overview.md`、`docs/refactor/backend-architecture-governance.md`、`docs/refactor/backend-module-refactor-compatibility.md` > 相关文档:`docs/architecture-overview.md`、`docs/refactor/backend-architecture-governance.md`、`docs/refactor/backend-module-refactor-compatibility.md`
> 实施进度:阶段 0~6 的宿主架构能力已完成收口;API/Application 公共复杂度基线已清零,启动组合根的 SystemConfigOper 构造点已由 14 降至 1;API 进程内后台任务已完成首批统一登记,插件仓适配和 Outbox 外围扩展仍按风险切片推进。Model/Base 查询与写装饰器、legacy 隐式会话外壳均已清零,插件 SDK 也不再导出宿主 Model。2026-08-23 的长期整改阶段 0 已恢复宿主、启动性能、官方插件和 SDK 契约门禁的可信基线;阶段 1a 已补齐 TaskRegistry owner 零债务门禁和诚实的关停超时语义;阶段 1b1 已收口整理 worker、pending 回放、失败通知、进程内 AI 重试、插件监控与事件投递的生命周期所有权;2026-08-24 的阶段 2 已将 212 个已观察宿主模块方法的 legacy aggregation 清零,并补齐可执行 fanout 与下载器文件 DTO 边界;阶段 3 已将消息交互和远程命令的订阅删除统一到 Application/UoW/outbox,宿主不再调用裸线程统计入口;阶段 4 已统一七种消息渠道的宿主回环与后台执行边界;阶段 5 已补齐事件窗口聚合任务的生命周期所有权;阶段 6 已统一插件文件操作的取消完成语义;阶段 7 已统一插件协程补偿的终态等待;阶段 8 已统一宿主同步函数的异步线程池入口;阶段 9 已统一工作流运行时的宿主获取路径;阶段 10 已统一模块、插件与调度运行时的显式 getter 调用;阶段 11 已清除系统配置 getter 的 Oper 形别名;阶段 12 已完成工作流域的显式 Chain 数据端口迁移;阶段 13 已收口用户、交互与消息链的数据端口;阶段 14 已收口音乐订阅数据端口;阶段 15 已收口站点数据端口;阶段 16 已收口媒体服务器数据端口;阶段 17 已收口下载数据端口;阶段 18 已收口主订阅数据端口;阶段 19 已收口整理数据端口;阶段 20 已收口 Agent 数据端口;阶段 21 已收口监控历史端口;阶段 22 已统一服务配置应用边界;阶段 23 已补齐媒体服务器 API 遗留的类形配置读取路径;阶段 24 已清除 Scheduler 内部无 owner 的协程提交双轨;阶段 25 已补齐 TaskRegistry 跨线程 owner 并迁移整理 AI 接管;阶段 26 已统一 Agent 会话清理提交;阶段 27 已统一历史 AI 进度 owner;阶段 28 已托管旧插件订阅统计线程;阶段 29 已统一 Emby 系条目转换并清零重复代码白名单;阶段 30 已收口插件市场请求级子任务;阶段 31 已托管搜索 AI 推荐任务;阶段 32 已清除事件调度器绕过生命周期 owner 的投递回退;阶段 33 已统一宿主 Agent 运行时的获取路径;阶段 34 已统一 durable-required 事件与 Outbox topic 事实源;阶段 35 已统一 LLM provider 管理 API 的运行时解析路径;阶段 36 已统一 WebAgent 音频能力访问边界;阶段 37 已统一插件输入事件发布路径;阶段 38 已统一 WebAgent 通知事件监听与队列边界;阶段 39 已补齐搜索 SSE 断线时的上游任务清理;阶段 40 已补齐异步防抖取消的终态所有权;阶段 41 已统一优雅重启兜底线程的唯一所有权;阶段 42 已补齐 Telegram typing 的多实例隔离和终态 owner。 > 实施进度:阶段 0~6 的宿主架构能力已完成收口;API/Application 公共复杂度基线已清零,启动组合根的 SystemConfigOper 构造点已由 14 降至 1;API 进程内后台任务已完成首批统一登记,插件仓适配和 Outbox 外围扩展仍按风险切片推进。Model/Base 查询与写装饰器、legacy 隐式会话外壳均已清零,插件 SDK 也不再导出宿主 Model。2026-08-23 的长期整改阶段 0 已恢复宿主、启动性能、官方插件和 SDK 契约门禁的可信基线;阶段 1a 已补齐 TaskRegistry owner 零债务门禁和诚实的关停超时语义;阶段 1b1 已收口整理 worker、pending 回放、失败通知、进程内 AI 重试、插件监控与事件投递的生命周期所有权;2026-08-24 的阶段 2 已将 212 个已观察宿主模块方法的 legacy aggregation 清零,并补齐可执行 fanout 与下载器文件 DTO 边界;阶段 3 已将消息交互和远程命令的订阅删除统一到 Application/UoW/outbox,宿主不再调用裸线程统计入口;阶段 4 已统一七种消息渠道的宿主回环与后台执行边界;阶段 5 已补齐事件窗口聚合任务的生命周期所有权;阶段 6 已统一插件文件操作的取消完成语义;阶段 7 已统一插件协程补偿的终态等待;阶段 8 已统一宿主同步函数的异步线程池入口;阶段 9 已统一工作流运行时的宿主获取路径;阶段 10 已统一模块、插件与调度运行时的显式 getter 调用;阶段 11 已清除系统配置 getter 的 Oper 形别名;阶段 12 已完成工作流域的显式 Chain 数据端口迁移;阶段 13 已收口用户、交互与消息链的数据端口;阶段 14 已收口音乐订阅数据端口;阶段 15 已收口站点数据端口;阶段 16 已收口媒体服务器数据端口;阶段 17 已收口下载数据端口;阶段 18 已收口主订阅数据端口;阶段 19 已收口整理数据端口;阶段 20 已收口 Agent 数据端口;阶段 21 已收口监控历史端口;阶段 22 已统一服务配置应用边界;阶段 23 已补齐媒体服务器 API 遗留的类形配置读取路径;阶段 24 已清除 Scheduler 内部无 owner 的协程提交双轨;阶段 25 已补齐 TaskRegistry 跨线程 owner 并迁移整理 AI 接管;阶段 26 已统一 Agent 会话清理提交;阶段 27 已统一历史 AI 进度 owner;阶段 28 已托管旧插件订阅统计线程;阶段 29 已统一 Emby 系条目转换并清零重复代码白名单;阶段 30 已收口插件市场请求级子任务;阶段 31 已托管搜索 AI 推荐任务;阶段 32 已清除事件调度器绕过生命周期 owner 的投递回退;阶段 33 已统一宿主 Agent 运行时的获取路径;阶段 34 已统一 durable-required 事件与 Outbox topic 事实源;阶段 35 已统一 LLM provider 管理 API 的运行时解析路径;阶段 36 已统一 WebAgent 音频能力访问边界;阶段 37 已统一插件输入事件发布路径;阶段 38 已统一 WebAgent 通知事件监听与队列边界;阶段 39 已补齐搜索 SSE 断线时的上游任务清理;阶段 40 已补齐异步防抖取消的终态所有权;阶段 41 已统一优雅重启兜底线程的唯一所有权;阶段 42 已补齐 Telegram typing 的多实例隔离和终态 owner;阶段 43 已统一 Discord typing 的异步 owner 和 shutdown 收尾
## 当前复核结论(2026-08-24 ## 当前复核结论(2026-08-24
@@ -450,6 +450,17 @@
均保持不变;私有类级可变状态已清除,正常 V1/V2/V3 模块实例不再共享运行状态。本阶段未修改插件仓、 均保持不变;私有类级可变状态已清除,正常 V1/V2/V3 模块实例不再共享运行状态。本阶段未修改插件仓、
SDK 或 Compat 映射。 SDK 或 Compat 映射。
### 长期整改阶段 43Discord typing 异步 owner 与 shutdown 收口(2026-08-24
- Discord typing 原先在等待旧 task 进入终态前就删除字典 owner;`trigger_typing()` 阻塞超过一秒时,
新请求会覆盖仍运行的 task,模块停止也无法再取得它。现在同一实例通过异步 lifecycle 锁串行替换,
超时后保留 owner 并拒绝并行启动,task 只在自己的 `finally` 或已确认终态后释放登记。
- client shutdown 先封住新增 typing,再按统一预算通知全部 owner;未自然结束的 task 会被取消并再次等待,
已完成 task 的异常也会被读取。Discord 长连接意外退出时,线程 runner 同样执行这条收尾路径,事件循环
不再直接关闭仍登记的 typing task。
- `Discord` 类路径、构造参数、同步 `start_typing()`/`stop_typing()` 布尔合同、模块方法、消息格式和配置
字段均保持不变;V1/V2/V3 插件仍通过原模块能力调用。本阶段未修改插件仓、SDK 或 Compat 映射。
### 总体判断 ### 总体判断
当前架构总体合理,已经从跨层混合的遗留单体收敛为**边界清晰的模块化单体**: 当前架构总体合理,已经从跨层混合的遗留单体收敛为**边界清晰的模块化单体**:
+124
View File
@@ -1,6 +1,9 @@
import asyncio import asyncio
import threading import threading
from typing import Optional from typing import Optional
from unittest.mock import AsyncMock
import pytest
from app.modules.discord.discord import Discord from app.modules.discord.discord import Discord
@@ -40,6 +43,20 @@ class _YieldingCloseDiscordClientStub(_DiscordClientStub):
self.closed.set() self.closed.set()
class _BlockingTypingChannelStub:
"""模拟阻塞中的 Discord typing 请求。"""
def __init__(self) -> None:
"""初始化请求进入与释放屏障。"""
self.entered = asyncio.Event()
self.release = asyncio.Event()
async def trigger_typing(self) -> None:
"""阻塞请求,直到测试释放或 owner 被取消。"""
self.entered.set()
await self.release.wait()
def _discord(client: _DiscordClientStub) -> Discord: def _discord(client: _DiscordClientStub) -> Discord:
"""构造只包含线程与事件循环生命周期状态的 Discord 实例。""" """构造只包含线程与事件循环生命周期状态的 Discord 实例。"""
instance = Discord.__new__(Discord) instance = Discord.__new__(Discord)
@@ -51,6 +68,24 @@ def _discord(client: _DiscordClientStub) -> Discord:
instance._ready_event = threading.Event() instance._ready_event = threading.Event()
instance._typing_tasks = {} instance._typing_tasks = {}
instance._typing_stop_events = {} instance._typing_stop_events = {}
instance._typing_lifecycle_lock = asyncio.Lock()
instance._typing_accepting = True
instance._typing_stop_timeout_seconds = 0.01
return instance
def _typing_discord() -> Discord:
"""构造绑定当前测试循环且不连接外部服务的 typing client。"""
instance = Discord.__new__(Discord)
instance._loop = asyncio.get_running_loop()
instance._typing_tasks = {}
instance._typing_stop_events = {}
instance._typing_lifecycle_lock = asyncio.Lock()
instance._typing_accepting = True
instance._typing_interval_seconds = 0.01
instance._typing_initial_delay_seconds = 0
instance._typing_max_duration_seconds = 1
instance._typing_stop_timeout_seconds = 0.01
return instance return instance
@@ -129,3 +164,92 @@ def test_stop_during_thread_bootstrap_preserves_runner_cleanup(monkeypatch) -> N
finally: finally:
release_runner.set() release_runner.set()
_cleanup(instance) _cleanup(instance)
@pytest.mark.anyio
async def test_short_typing_task_can_stop_before_first_trigger(monkeypatch) -> None:
"""短响应在首发前结束时,不应留下 Discord 客户端 typing 状态。"""
discord_client = _typing_discord()
channel = AsyncMock()
channel.trigger_typing = AsyncMock()
monkeypatch.setattr(
discord_client,
"_resolve_channel",
AsyncMock(return_value=channel),
)
started = await discord_client._start_typing_task(
typing_key="chat:30003",
chat_id="30003",
max_duration_seconds=1,
initial_delay_seconds=0.05,
)
stopped = await discord_client._stop_typing_task("chat:30003")
await asyncio.sleep(0.08)
assert started
assert stopped
channel.trigger_typing.assert_not_called()
assert "chat:30003" not in discord_client._typing_tasks
@pytest.mark.anyio
async def test_typing_stop_retains_blocked_owner_until_terminal(monkeypatch) -> None:
"""typing 请求阻塞超过预算时,不得删除或覆盖仍运行的 task owner。"""
discord_client = _typing_discord()
channel = _BlockingTypingChannelStub()
resolve_channel = AsyncMock(return_value=channel)
monkeypatch.setattr(discord_client, "_resolve_channel", resolve_channel)
try:
assert await discord_client._start_typing_task(
typing_key="chat:blocked",
chat_id="blocked",
)
await asyncio.wait_for(channel.entered.wait(), timeout=1)
owner = discord_client._typing_tasks["chat:blocked"]
assert await discord_client._stop_typing_task("chat:blocked")
assert discord_client._typing_tasks["chat:blocked"] is owner
assert not owner.done()
assert not await discord_client._start_typing_task(
typing_key="chat:blocked",
chat_id="blocked",
)
assert discord_client._typing_tasks["chat:blocked"] is owner
channel.release.set()
await asyncio.wait_for(owner, timeout=1)
assert "chat:blocked" not in discord_client._typing_tasks
finally:
channel.release.set()
await discord_client._stop_all_typing_tasks()
@pytest.mark.anyio
async def test_stop_all_seals_and_drains_typing_owners(monkeypatch) -> None:
"""client shutdown 必须封住新增任务并取消、等待既有 owner 进入终态。"""
discord_client = _typing_discord()
channel = _BlockingTypingChannelStub()
resolve_channel = AsyncMock(return_value=channel)
monkeypatch.setattr(discord_client, "_resolve_channel", resolve_channel)
assert await discord_client._start_typing_task(
typing_key="chat:shutdown",
chat_id="shutdown",
)
await asyncio.wait_for(channel.entered.wait(), timeout=1)
owner = discord_client._typing_tasks["chat:shutdown"]
await discord_client._stop_all_typing_tasks()
assert owner.done()
assert owner.cancelled()
assert discord_client._typing_tasks == {}
assert discord_client._typing_stop_events == {}
resolve_channel.reset_mock()
assert not await discord_client._start_typing_task(
typing_key="chat:after-stop",
chat_id="after-stop",
)
resolve_channel.assert_not_awaited()
+125 -159
View File
@@ -1,179 +1,145 @@
import asyncio
import json import json
import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest import IsolatedAsyncioTestCase from unittest.mock import MagicMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
from app.agent import _finish_processing_status from app.agent.orchestrator import _finish_processing_status
from app.modules.discord import DiscordModule from app.modules.discord import DiscordModule
from app.modules.discord.discord import Discord
from app.modules.slack import SlackModule from app.modules.slack import SlackModule
from app.schemas.notification import ChannelCapability, ChannelCapabilityManager from app.schemas.notification import ChannelCapability, ChannelCapabilityManager
from app.schemas.types import NotificationChannel from app.schemas.types import NotificationChannel
class TestMessageProcessingStatus(unittest.TestCase): def test_processing_status_capability_only_enabled_for_supported_channels() -> None:
def test_processing_status_capability_only_enabled_for_supported_channels(self): supported = {
supported = { NotificationChannel.Telegram,
NotificationChannel.Telegram, NotificationChannel.Feishu,
NotificationChannel.Feishu, NotificationChannel.Slack,
NotificationChannel.Slack, NotificationChannel.Discord,
NotificationChannel.Discord, }
}
for channel in NotificationChannel: for channel in NotificationChannel:
self.assertEqual( assert ChannelCapabilityManager.supports_capability(
ChannelCapabilityManager.supports_capability( channel, ChannelCapability.PROCESSING_STATUS
channel, ChannelCapability.PROCESSING_STATUS ) is (channel in supported)
),
channel in supported,
)
def test_slack_processing_status_uses_reaction(self):
module = SlackModule()
module._channel = NotificationChannel.Slack
client = MagicMock()
client.add_reaction.return_value = True
client.remove_reaction.return_value = True
with ( def test_slack_processing_status_uses_reaction() -> None:
patch.object( module = SlackModule()
module, "get_config", return_value=SimpleNamespace(name="slack-main") module._channel = NotificationChannel.Slack
), client = MagicMock()
patch.object(module, "get_instance", return_value=client), client.add_reaction.return_value = True
): client.remove_reaction.return_value = True
status = module.mark_message_processing_started(
channel=NotificationChannel.Slack,
source="slack-main",
userid="U01",
message_id="1710000000.000100",
chat_id="C01",
text="hello",
)
removed = module.mark_message_processing_finished(
channel=NotificationChannel.Slack,
source="slack-main",
userid="U01",
status=status,
)
client.add_reaction.assert_called_once_with( with (
channel="C01", patch.object(
timestamp="1710000000.000100", module, "get_config", return_value=SimpleNamespace(name="slack-main")
emoji="eyes", ),
patch.object(module, "get_instance", return_value=client),
):
status = module.mark_message_processing_started(
channel=NotificationChannel.Slack,
source="slack-main",
userid="U01",
message_id="1710000000.000100",
chat_id="C01",
text="hello",
) )
client.remove_reaction.assert_called_once_with( removed = module.mark_message_processing_finished(
channel="C01", channel=NotificationChannel.Slack,
timestamp="1710000000.000100", source="slack-main",
emoji="eyes", userid="U01",
)
self.assertEqual(status["metadata"]["kind"], "reaction")
self.assertTrue(removed)
def test_slack_parser_exposes_message_location_for_reaction_status(self):
module = SlackModule()
with patch.object(
module,
"get_config",
return_value=SimpleNamespace(name="slack-main", config={}),
):
message = module.message_parser(
source="slack-main",
body=json.dumps(
{
"type": "message",
"user": "U01",
"text": "hello",
"ts": "1710000000.000100",
"channel": "C01",
}
),
form=None,
args=None,
)
self.assertEqual(message.message_id, "1710000000.000100")
self.assertEqual(message.chat_id, "C01")
def test_discord_processing_status_starts_and_stops_typing(self):
module = DiscordModule()
module._channel = NotificationChannel.Discord
client = MagicMock()
client.start_typing.return_value = True
client.stop_typing.return_value = True
with (
patch.object(
module, "get_config", return_value=SimpleNamespace(name="discord-main")
),
patch.object(module, "get_instance", return_value=client),
):
status = module.mark_message_processing_started(
channel=NotificationChannel.Discord,
source="discord-main",
userid="10001",
message_id="20002",
chat_id="30003",
text="hello",
)
finished = module.mark_message_processing_finished(
channel=NotificationChannel.Discord,
source="discord-main",
userid="10001",
status=status,
)
client.start_typing.assert_called_once_with(userid="10001", chat_id="30003")
client.stop_typing.assert_called_once_with(userid="10001", chat_id="30003")
self.assertEqual(status["metadata"]["kind"], "typing")
self.assertTrue(finished)
def test_agent_finish_processing_status_uses_module_interface(self):
status = {
"channel": NotificationChannel.Telegram.value,
"source": "telegram-main",
"userid": "10001",
"message_id": None,
"chat_id": "-100",
"metadata": {"kind": "typing"},
}
with patch("app.agent.orchestrator.AgentChain") as chain_cls:
_finish_processing_status(status, user_id="fallback")
chain_cls.return_value.finish_message_processing_status.assert_called_once_with(
status=status, status=status,
userid="fallback",
) )
client.add_reaction.assert_called_once_with(
channel="C01",
timestamp="1710000000.000100",
emoji="eyes",
)
client.remove_reaction.assert_called_once_with(
channel="C01",
timestamp="1710000000.000100",
emoji="eyes",
)
assert status["metadata"]["kind"] == "reaction"
assert removed
class TestDiscordTypingLifecycle(IsolatedAsyncioTestCase):
async def test_short_typing_task_can_stop_before_first_trigger(self):
"""
短响应在首次 Discord typing 触发前结束时,不应留下客户端自然保留的输入状态。
"""
discord_client = Discord.__new__(Discord)
discord_client._loop = asyncio.get_running_loop()
discord_client._typing_tasks = {}
discord_client._typing_stop_events = {}
discord_client._typing_interval_seconds = 0.01
discord_client._typing_max_duration_seconds = 1
channel = MagicMock()
channel.trigger_typing = AsyncMock()
with patch.object(discord_client, "_resolve_channel", return_value=channel): def test_slack_parser_exposes_message_location_for_reaction_status() -> None:
started = await discord_client._start_typing_task( module = SlackModule()
typing_key="chat:30003",
chat_id="30003",
max_duration_seconds=1,
initial_delay_seconds=0.05,
)
stopped = await discord_client._stop_typing_task("chat:30003")
await asyncio.sleep(0.08)
self.assertTrue(started) with patch.object(
self.assertTrue(stopped) module,
channel.trigger_typing.assert_not_called() "get_config",
self.assertNotIn("chat:30003", discord_client._typing_tasks) return_value=SimpleNamespace(name="slack-main", config={}),
):
message = module.message_parser(
source="slack-main",
body=json.dumps(
{
"type": "message",
"user": "U01",
"text": "hello",
"ts": "1710000000.000100",
"channel": "C01",
}
),
form=None,
args=None,
)
assert message.message_id == "1710000000.000100"
assert message.chat_id == "C01"
def test_discord_processing_status_starts_and_stops_typing() -> None:
module = DiscordModule()
module._channel = NotificationChannel.Discord
client = MagicMock()
client.start_typing.return_value = True
client.stop_typing.return_value = True
with (
patch.object(
module, "get_config", return_value=SimpleNamespace(name="discord-main")
),
patch.object(module, "get_instance", return_value=client),
):
status = module.mark_message_processing_started(
channel=NotificationChannel.Discord,
source="discord-main",
userid="10001",
message_id="20002",
chat_id="30003",
text="hello",
)
finished = module.mark_message_processing_finished(
channel=NotificationChannel.Discord,
source="discord-main",
userid="10001",
status=status,
)
client.start_typing.assert_called_once_with(userid="10001", chat_id="30003")
client.stop_typing.assert_called_once_with(userid="10001", chat_id="30003")
assert status["metadata"]["kind"] == "typing"
assert finished
def test_agent_finish_processing_status_uses_module_interface() -> None:
status = {
"channel": NotificationChannel.Telegram.value,
"source": "telegram-main",
"userid": "10001",
"message_id": None,
"chat_id": "-100",
"metadata": {"kind": "typing"},
}
with patch("app.agent.orchestrator.AgentChain") as chain_cls:
_finish_processing_status(status, user_id="fallback")
chain_cls.return_value.finish_message_processing_status.assert_called_once_with(
status=status,
userid="fallback",
)