Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions astrbot/core/star/star_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -964,7 +964,7 @@ async def reload_failed_plugin(self, dir_name):
else:
return False, error

async def reload(self, specified_plugin_name=None):
async def reload(self, specified_plugin_name=None, *, force_unbind=False):
"""重新加载插件

Args:
Expand Down Expand Up @@ -1013,7 +1013,7 @@ async def reload(self, specified_plugin_name=None):
logger.warning(
f"插件 {smd.name} 未被正常终止: {e!s}, 可能会导致该插件运行不正常。",
)
if smd.name and smd.activated:
if smd.name and (smd.activated or force_unbind):
await self._unbind_plugin(smd.name, specified_module_path)

result = await self.load(specified_module_path)
Expand Down Expand Up @@ -1933,7 +1933,7 @@ async def turn_on_plugin(self, plugin_name: str) -> None:
for func_tool in self._iter_plugin_llm_tools(plugin.module_path):
func_tool.active = func_tool.name not in inactivated_llm_tools

success, error = await self.reload(plugin_name)
success, error = await self.reload(plugin_name, force_unbind=True)
if not success:
raise Exception(error or f"插件 {plugin_name} 启用失败。")
current_plugin = self.context.get_registered_star(plugin_name)
Expand Down
83 changes: 80 additions & 3 deletions tests/test_plugin_manager.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import asyncio
import functools
import json
import os
from pathlib import Path
Expand All @@ -9,6 +10,7 @@
import yaml

from astrbot.core.star import star_manager as star_manager_module
from astrbot.core.star.star_handler import EventType, StarHandlerMetadata
from astrbot.core.star.star_manager import PluginDependencyInstallError, PluginManager
from astrbot.core.utils.pip_installer import PipInstallError
from astrbot.core.utils.requirements_utils import MissingRequirementsPlan
Expand Down Expand Up @@ -627,7 +629,7 @@ async def mock_global_put(key, value):
async def mock_terminate(star_metadata):
assert star_metadata is plugin

async def mock_reload(plugin_name):
async def mock_reload(plugin_name, **_):
assert plugin_name == plugin.root_dir_name
return True, None

Expand Down Expand Up @@ -721,7 +723,7 @@ async def mock_global_put(key, value):
async def mock_terminate(star_metadata):
assert star_metadata is plugin

async def mock_reload(plugin_name):
async def mock_reload(plugin_name, **_):
assert plugin_name == plugin.root_dir_name
return True, None

Expand Down Expand Up @@ -2109,6 +2111,81 @@ async def mock_load(specified_module_path=None, **kwargs):
_clear_star_runtime_state()


@pytest.mark.asyncio
async def test_turn_on_plugin_rebinds_handlers_after_deactivated_reload(
plugin_manager_pm: PluginManager, monkeypatch
):
_clear_star_runtime_state()
plugin_name = "demo_plugin"
module_path = f"data.plugins.{plugin_name}.main"
plugin = star_manager_module.StarMetadata(
name=plugin_name,
root_dir_name=plugin_name,
module_path=module_path,
activated=False,
)
cast(Any, plugin_manager_pm.context).stars.append(plugin)
star_manager_module.star_map[module_path] = plugin
star_manager_module.star_registry.append(plugin)

async def raw_handler(plugin_instance, event):
return plugin_instance, event

handler = StarHandlerMetadata(
event_type=EventType.OnLLMRequestEvent,
handler_full_name=f"{module_path}.handler",
handler_name="handler",
handler_module_path=module_path,
handler=functools.partial(raw_handler, None),
event_filters=[],
)
star_manager_module.star_handlers_registry.append(handler)

async def mock_global_get(key, default=None):
return [module_path] if key == "inactivated_plugins" else default

async def mock_global_put(*_):
pass

async def mock_load(specified_module_path=None, **kwargs):
assert specified_module_path == module_path

related_handlers = (
star_manager_module.star_handlers_registry.get_handlers_by_module_name(
module_path
)
)
if not related_handlers:
handler.handler = raw_handler
star_manager_module.star_handlers_registry.append(handler)
related_handlers = [handler]

for registered_handler in related_handlers:
registered_handler.handler = functools.partial(
registered_handler.handler, object()
)
return True, None

monkeypatch.setattr(star_manager_module.sp, "global_get", mock_global_get)
monkeypatch.setattr(star_manager_module.sp, "global_put", mock_global_put)
monkeypatch.setattr(plugin_manager_pm, "load", mock_load)

try:
await plugin_manager_pm.turn_on_plugin(plugin_name)
handlers = (
star_manager_module.star_handlers_registry.get_handlers_by_module_name(
module_path
)
)
assert len(handlers) == 1
result = await handlers[0].handler("event")
assert result[0] is not None
assert result[1] == "event"
finally:
cast(Any, plugin_manager_pm.context).stars.remove(plugin)
_clear_star_runtime_state()


@pytest.mark.asyncio
async def test_reload_activated_plugin_still_unbinds(
plugin_manager_pm: PluginManager, monkeypatch
Expand Down Expand Up @@ -2238,7 +2315,7 @@ async def mock_global_put(key, value):
async def mock_terminate(smd):
pass

async def mock_reload(plugin_name_arg):
async def mock_reload(plugin_name_arg, **_):
assert plugin_name_arg == plugin_name
return True, None

Expand Down
Loading