mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-08-14 14:31:00 +00:00
fix(aiocqhttp): correct listener lifecycle handling (#2336)
* fix(aiocqhttp): correct listener lifecycle handling * docs(aiocqhttp): add listener lifecycle evidence
This commit is contained in:
Binary file not shown.
|
After Width: | Height: | Size: 33 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 20 KiB |
@@ -491,7 +491,11 @@ class AiocqhttpAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
|||||||
message_converter: AiocqhttpMessageConverter = AiocqhttpMessageConverter()
|
message_converter: AiocqhttpMessageConverter = AiocqhttpMessageConverter()
|
||||||
event_converter: AiocqhttpEventConverter = pydantic.Field(default_factory=AiocqhttpEventConverter)
|
event_converter: AiocqhttpEventConverter = pydantic.Field(default_factory=AiocqhttpEventConverter)
|
||||||
|
|
||||||
on_websocket_connection_event_cache: typing.List[typing.Callable[[aiocqhttp.Event], None]] = []
|
on_websocket_connection_event_cache: list[aiocqhttp.Event] = []
|
||||||
|
_listener_wrappers: dict[
|
||||||
|
tuple[typing.Type[platform_events.Event], typing.Callable],
|
||||||
|
tuple[str, typing.Callable],
|
||||||
|
] = {}
|
||||||
|
|
||||||
def __init__(self, config: dict, logger: abstract_platform_logger.AbstractEventLogger):
|
def __init__(self, config: dict, logger: abstract_platform_logger.AbstractEventLogger):
|
||||||
super().__init__(
|
super().__init__(
|
||||||
@@ -506,6 +510,7 @@ class AiocqhttpAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
|||||||
self.config['shutdown_trigger'] = shutdown_trigger_placeholder
|
self.config['shutdown_trigger'] = shutdown_trigger_placeholder
|
||||||
|
|
||||||
self.on_websocket_connection_event_cache = []
|
self.on_websocket_connection_event_cache = []
|
||||||
|
self._listener_wrappers = {}
|
||||||
|
|
||||||
if 'access-token' in config:
|
if 'access-token' in config:
|
||||||
self.bot = aiocqhttp.CQHttp(access_token=config['access-token'])
|
self.bot = aiocqhttp.CQHttp(access_token=config['access-token'])
|
||||||
@@ -513,6 +518,16 @@ class AiocqhttpAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
|||||||
else:
|
else:
|
||||||
self.bot = aiocqhttp.CQHttp()
|
self.bot = aiocqhttp.CQHttp()
|
||||||
|
|
||||||
|
self.bot.on_websocket_connection(self._on_websocket_connection)
|
||||||
|
|
||||||
|
async def _on_websocket_connection(self, event: aiocqhttp.Event):
|
||||||
|
for cached_event in self.on_websocket_connection_event_cache:
|
||||||
|
if cached_event.self_id == event.self_id and cached_event.time == event.time:
|
||||||
|
return
|
||||||
|
|
||||||
|
self.on_websocket_connection_event_cache.append(event)
|
||||||
|
await self.logger.info(f'WebSocket connection established, bot id: {event.self_id}')
|
||||||
|
|
||||||
async def send_message(self, target_type: str, target_id: str, message: platform_message.MessageChain):
|
async def send_message(self, target_type: str, target_id: str, message: platform_message.MessageChain):
|
||||||
# Check if message contains a Forward component
|
# Check if message contains a Forward component
|
||||||
forward_msg = message.get_first(platform_message.Forward)
|
forward_msg = message.get_first(platform_message.Forward)
|
||||||
@@ -648,22 +663,14 @@ class AiocqhttpAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
|||||||
|
|
||||||
if event_type == platform_events.GroupMessage:
|
if event_type == platform_events.GroupMessage:
|
||||||
self.bot.on_message('group')(on_message)
|
self.bot.on_message('group')(on_message)
|
||||||
|
self._listener_wrappers[(event_type, callback)] = ('message.group', on_message)
|
||||||
# self.bot.on_notice()(on_message)
|
# self.bot.on_notice()(on_message)
|
||||||
elif event_type == platform_events.FriendMessage:
|
elif event_type == platform_events.FriendMessage:
|
||||||
self.bot.on_message('private')(on_message)
|
self.bot.on_message('private')(on_message)
|
||||||
|
self._listener_wrappers[(event_type, callback)] = ('message.private', on_message)
|
||||||
# self.bot.on_notice()(on_message)
|
# self.bot.on_notice()(on_message)
|
||||||
# print(event_type)
|
# print(event_type)
|
||||||
|
|
||||||
async def on_websocket_connection(event: aiocqhttp.Event):
|
|
||||||
for event in self.on_websocket_connection_event_cache:
|
|
||||||
if event.self_id == event.self_id and event.time == event.time:
|
|
||||||
return
|
|
||||||
|
|
||||||
self.on_websocket_connection_event_cache.append(event)
|
|
||||||
await self.logger.info(f'WebSocket connection established, bot id: {event.self_id}')
|
|
||||||
|
|
||||||
self.bot.on_websocket_connection(on_websocket_connection)
|
|
||||||
|
|
||||||
def unregister_listener(
|
def unregister_listener(
|
||||||
self,
|
self,
|
||||||
event_type: typing.Type[platform_events.Event],
|
event_type: typing.Type[platform_events.Event],
|
||||||
@@ -671,7 +678,12 @@ class AiocqhttpAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
|
|||||||
[platform_events.Event, abstract_platform_adapter.AbstractMessagePlatformAdapter], None
|
[platform_events.Event, abstract_platform_adapter.AbstractMessagePlatformAdapter], None
|
||||||
],
|
],
|
||||||
):
|
):
|
||||||
return super().unregister_listener(event_type, callback)
|
listener = self._listener_wrappers.pop((event_type, callback), None)
|
||||||
|
if listener is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
event_name, wrapper = listener
|
||||||
|
self.bot._bus.unsubscribe(event_name, wrapper)
|
||||||
|
|
||||||
async def run_async(self):
|
async def run_async(self):
|
||||||
await self.bot._server_app.run_task(**self.config)
|
await self.bot._server_app.run_task(**self.config)
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ import pytest
|
|||||||
import aiocqhttp
|
import aiocqhttp
|
||||||
|
|
||||||
import langbot_plugin.api.entities.builtin.platform.message as platform_message
|
import langbot_plugin.api.entities.builtin.platform.message as platform_message
|
||||||
|
import langbot_plugin.api.entities.builtin.platform.events as platform_events
|
||||||
from langbot.pkg.platform.sources.aiocqhttp import (
|
from langbot.pkg.platform.sources.aiocqhttp import (
|
||||||
AiocqhttpAdapter,
|
AiocqhttpAdapter,
|
||||||
AiocqhttpEventConverter,
|
AiocqhttpEventConverter,
|
||||||
@@ -15,6 +16,72 @@ async def _convert_single(component: platform_message.MessageComponent):
|
|||||||
return message[0]
|
return message[0]
|
||||||
|
|
||||||
|
|
||||||
|
class _TestLogger:
|
||||||
|
def __init__(self):
|
||||||
|
self.messages = []
|
||||||
|
|
||||||
|
async def info(self, message):
|
||||||
|
self.messages.append(message)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_adapter():
|
||||||
|
logger = _TestLogger()
|
||||||
|
adapter = AiocqhttpAdapter.model_construct(
|
||||||
|
config={},
|
||||||
|
logger=logger,
|
||||||
|
bot=aiocqhttp.CQHttp(),
|
||||||
|
on_websocket_connection_event_cache=[],
|
||||||
|
_listener_wrappers={},
|
||||||
|
)
|
||||||
|
adapter.bot.on_websocket_connection(adapter._on_websocket_connection)
|
||||||
|
return adapter, logger
|
||||||
|
|
||||||
|
|
||||||
|
def test_connection_listener_is_registered_once_for_multiple_message_listeners():
|
||||||
|
adapter, _ = _make_adapter()
|
||||||
|
|
||||||
|
async def callback(event, source_adapter):
|
||||||
|
return None
|
||||||
|
|
||||||
|
adapter.register_listener(platform_events.FriendMessage, callback)
|
||||||
|
adapter.register_listener(platform_events.GroupMessage, callback)
|
||||||
|
adapter.register_listener(platform_events.FeedbackEvent, callback)
|
||||||
|
|
||||||
|
assert len(adapter.bot._bus._subscribers['meta_event.lifecycle.connect']) == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_connection_listener_only_suppresses_exact_duplicates():
|
||||||
|
adapter, logger = _make_adapter()
|
||||||
|
first = aiocqhttp.Event({'self_id': 1001, 'time': 10})
|
||||||
|
duplicate = aiocqhttp.Event({'self_id': 1001, 'time': 10})
|
||||||
|
second = aiocqhttp.Event({'self_id': 2002, 'time': 20})
|
||||||
|
|
||||||
|
await adapter._on_websocket_connection(first)
|
||||||
|
await adapter._on_websocket_connection(duplicate)
|
||||||
|
await adapter._on_websocket_connection(second)
|
||||||
|
|
||||||
|
assert adapter.on_websocket_connection_event_cache == [first, second]
|
||||||
|
assert logger.messages == [
|
||||||
|
'WebSocket connection established, bot id: 1001',
|
||||||
|
'WebSocket connection established, bot id: 2002',
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_unregister_listener_removes_registered_wrapper():
|
||||||
|
adapter, _ = _make_adapter()
|
||||||
|
|
||||||
|
async def callback(event, source_adapter):
|
||||||
|
return None
|
||||||
|
|
||||||
|
adapter.register_listener(platform_events.GroupMessage, callback)
|
||||||
|
assert len(adapter.bot._bus._subscribers['message.group']) == 1
|
||||||
|
|
||||||
|
adapter.unregister_listener(platform_events.GroupMessage, callback)
|
||||||
|
|
||||||
|
assert not adapter.bot._bus._subscribers['message.group']
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
('payload', 'expected'),
|
('payload', 'expected'),
|
||||||
|
|||||||
Reference in New Issue
Block a user