From 9fa0a198a4b726b15b1e7d5d7c1f4ba664e8b109 Mon Sep 17 00:00:00 2001 From: Ruochen Pan Date: Fri, 7 Aug 2026 12:20:18 +0800 Subject: [PATCH] refactor: remove sync shared preferences (#9582) * refactor: migrate async session preference calls * refactor: migrate internal preference calls to async * perf: reuse tool permissions when listing tools --- astrbot/core/provider/func_tool_manager.py | 8 +- astrbot/dashboard/api/sessions.py | 8 +- .../services/session_management_service.py | 46 ++++---- astrbot/dashboard/services/tools_service.py | 40 +++---- tests/unit/test_session_management_service.py | 40 +++++++ tests/unit/test_tool_permission.py | 108 ++++++++++-------- 6 files changed, 152 insertions(+), 98 deletions(-) create mode 100644 tests/unit/test_session_management_service.py diff --git a/astrbot/core/provider/func_tool_manager.py b/astrbot/core/provider/func_tool_manager.py index 216aa084c..b79b54e9c 100644 --- a/astrbot/core/provider/func_tool_manager.py +++ b/astrbot/core/provider/func_tool_manager.py @@ -246,7 +246,7 @@ class _PermissionGuardedTool(FunctionTool): async def call(self, context: Any, **kwargs: Any) -> Any: import inspect as _inspect - error = self._mgr._check_tool_permission(self.name, context) + error = await self._mgr._check_tool_permission(self.name, context) if error is not None: return error @@ -457,7 +457,7 @@ class FunctionToolManager: Builtin tools are never routed through this method.""" return "member" - def _check_tool_permission( + async def _check_tool_permission( self, tool_name: str, context: Any, @@ -469,9 +469,7 @@ class FunctionToolManager: no explicit entry exists the tool inherits the fallback ``_default_permission``.""" try: - perms_raw = sp.get( - "tool_permissions", {}, scope="global", scope_id="global" - ) + perms_raw = await sp.global_get("tool_permissions", {}) except Exception: perms_raw = {} defaults = perms_raw.get("_default", {}) if isinstance(perms_raw, dict) else {} diff --git a/astrbot/dashboard/api/sessions.py b/astrbot/dashboard/api/sessions.py index ac1152bdd..68ae27ef1 100644 --- a/astrbot/dashboard/api/sessions.py +++ b/astrbot/dashboard/api/sessions.py @@ -202,7 +202,7 @@ async def list_session_groups( service: SessionManagementService = Depends(get_service), ): try: - return ok(service.list_groups()) + return ok(await service.list_groups()) except SessionManagementServiceError as exc: return _service_error(exc) except Exception as exc: @@ -216,7 +216,7 @@ async def create_session_group( service: SessionManagementService = Depends(get_service), ): try: - return ok(service.create_group(payload.model_dump(exclude_none=True))) + return ok(await service.create_group(payload.model_dump(exclude_none=True))) except SessionManagementServiceError as exc: return _service_error(exc) except Exception as exc: @@ -232,7 +232,7 @@ async def update_session_group( ): try: body = payload.model_dump(exclude_none=True) - return ok(service.update_group({"group_id": group_id, **body})) + return ok(await service.update_group({"group_id": group_id, **body})) except SessionManagementServiceError as exc: return _service_error(exc) except Exception as exc: @@ -246,7 +246,7 @@ async def delete_session_group( service: SessionManagementService = Depends(get_service), ): try: - return ok(service.delete_group({"group_id": group_id})) + return ok(await service.delete_group({"group_id": group_id})) except SessionManagementServiceError as exc: return _service_error(exc) except Exception as exc: diff --git a/astrbot/dashboard/services/session_management_service.py b/astrbot/dashboard/services/session_management_service.py index a66f4e8e9..2307ae1ad 100644 --- a/astrbot/dashboard/services/session_management_service.py +++ b/astrbot/dashboard/services/session_management_service.py @@ -91,7 +91,7 @@ class SessionManagementService: if scope == "custom_group": if not group_id: raise SessionManagementServiceError("请指定分组 ID") - groups = self.get_groups() + groups = await self.get_groups() if group_id not in groups: raise SessionManagementServiceError(f"分组 '{group_id}' 不存在") return groups[group_id].get("umos", []) @@ -459,8 +459,7 @@ class SessionManagementService: for umo in umos: try: session_config = ( - sp.get("session_service_config", {}, scope="umo", scope_id=umo) - or {} + await sp.session_get(umo, "session_service_config", {}) or {} ) if llm_enabled is not None: @@ -470,11 +469,10 @@ class SessionManagementService: if session_enabled is not None: session_config["session_enabled"] = session_enabled - sp.put( + await sp.session_put( + umo, "session_service_config", session_config, - scope="umo", - scope_id=umo, ) success_count += 1 except Exception as exc: @@ -548,14 +546,20 @@ class SessionManagementService: "failed_umos": failed_umos, } - def get_groups(self) -> dict: - return sp.get("session_groups", {}) + async def get_groups(self) -> dict: + groups = await sp.get_async( + "unknown", + "unknown", + "session_groups", + {}, + ) + return groups if groups is not None else {} - def save_groups(self, groups: dict) -> None: - sp.put("session_groups", groups) + async def save_groups(self, groups: dict) -> None: + await sp.put_async("unknown", "unknown", "session_groups", groups) - def list_groups(self) -> dict: - groups = self.get_groups() + async def list_groups(self) -> dict: + groups = await self.get_groups() return { "groups": [ { @@ -568,7 +572,7 @@ class SessionManagementService: ] } - def create_group(self, data: object) -> dict: + async def create_group(self, data: object) -> dict: payload = self._payload(data) name = str(payload.get("name", "")).strip() umos = payload.get("umos", []) @@ -576,13 +580,13 @@ class SessionManagementService: if not name: raise SessionManagementServiceError("分组名称不能为空") - groups = self.get_groups() + groups = await self.get_groups() group_id = str(uuid.uuid4())[:8] groups[group_id] = { "name": name, "umos": umos, } - self.save_groups(groups) + await self.save_groups(groups) return { "message": f"分组 '{name}' 创建成功", @@ -594,7 +598,7 @@ class SessionManagementService: }, } - def update_group(self, data: object) -> dict: + async def update_group(self, data: object) -> dict: payload = self._payload(data) group_id = payload.get("id") or payload.get("group_id") name = payload.get("name") @@ -605,7 +609,7 @@ class SessionManagementService: if not group_id: raise SessionManagementServiceError("分组 ID 不能为空") - groups = self.get_groups() + groups = await self.get_groups() if group_id not in groups: raise SessionManagementServiceError(f"分组 '{group_id}' 不存在") @@ -623,7 +627,7 @@ class SessionManagementService: current_umos.difference_update(remove_umos) group["umos"] = list(current_umos) - self.save_groups(groups) + await self.save_groups(groups) return { "message": f"分组 '{group['name']}' 更新成功", @@ -635,20 +639,20 @@ class SessionManagementService: }, } - def delete_group(self, data: object) -> dict: + async def delete_group(self, data: object) -> dict: payload = self._payload(data) group_id = payload.get("id") or payload.get("group_id") if not group_id: raise SessionManagementServiceError("分组 ID 不能为空") - groups = self.get_groups() + groups = await self.get_groups() if group_id not in groups: raise SessionManagementServiceError(f"分组 '{group_id}' 不存在") group_name = groups[group_id].get("name", group_id) del groups[group_id] - self.save_groups(groups) + await self.save_groups(groups) return {"message": f"分组 '{group_name}' 已删除"} @staticmethod diff --git a/astrbot/dashboard/services/tools_service.py b/astrbot/dashboard/services/tools_service.py index a3288946e..7be88f75c 100644 --- a/astrbot/dashboard/services/tools_service.py +++ b/astrbot/dashboard/services/tools_service.py @@ -247,7 +247,7 @@ class ToolsService: logger.error(traceback.format_exc()) raise ToolsServiceError(f"Failed to test MCP connection: {exc!s}") from exc - def get_tool_list(self) -> list[dict]: + async def get_tool_list(self) -> list[dict]: try: tools = list(self.tool_mgr.func_list) existing_names = {tool.name for tool in tools} @@ -256,15 +256,22 @@ class ToolsService: tools.append(tool) config_entries = self._get_config_entries() + perms_store = ( + await sp.global_get("tool_permissions", {}) + if any(not self.tool_mgr.is_builtin_tool(tool.name) for tool in tools) + else {} + ) tools_dict = [] for tool in tools: - tools_dict.append(self._serialize_tool(tool, config_entries)) + tools_dict.append( + self._serialize_tool(tool, config_entries, perms_store) + ) return tools_dict except Exception as exc: logger.error(traceback.format_exc()) raise ToolsServiceError(f"Failed to get tool list: {exc!s}") from exc - def update_tool_permission(self, data: Any) -> str: + async def update_tool_permission(self, data: Any) -> str: """Set a tool permission level. Args: @@ -294,12 +301,7 @@ class ToolsService: if not any(t.name == tool_name for t in self.tool_mgr.func_list): raise ToolsServiceError(f"Tool '{tool_name}' not found") - perms_store = sp.get( - "tool_permissions", - {}, - scope="global", - scope_id="global", - ) + perms_store = await sp.global_get("tool_permissions", {}) if not isinstance(perms_store, dict): perms_store = {} defaults = perms_store.get("_default", {}) @@ -307,12 +309,7 @@ class ToolsService: defaults = {} defaults[tool_name] = permission perms_store["_default"] = defaults - sp.put( - "tool_permissions", - perms_store, - scope="global", - scope_id="global", - ) + await sp.global_put("tool_permissions", perms_store) return f"Tool '{tool_name}' permission set to {permission}" except ToolsServiceError: @@ -530,7 +527,12 @@ class ToolsService: ) return config_entries - def _serialize_tool(self, tool, config_entries: list[dict]) -> dict: + def _serialize_tool( + self, + tool, + config_entries: list[dict], + perms_store: object, + ) -> dict: readonly = False builtin_config_statuses = [] builtin_config_tags = [] @@ -573,12 +575,6 @@ class ToolsService: "builtin_config_tags": builtin_config_tags, } if not readonly: - perms_store = sp.get( - "tool_permissions", - {}, - scope="global", - scope_id="global", - ) defaults = ( perms_store.get("_default", {}) if isinstance(perms_store, dict) else {} ) diff --git a/tests/unit/test_session_management_service.py b/tests/unit/test_session_management_service.py new file mode 100644 index 000000000..f1bf6a4c4 --- /dev/null +++ b/tests/unit/test_session_management_service.py @@ -0,0 +1,40 @@ +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from astrbot.core import sp +from astrbot.dashboard.services.session_management_service import ( + SessionManagementService, +) + + +@pytest.mark.asyncio +async def test_get_groups_preserves_legacy_preference_scope(monkeypatch): + get_async = AsyncMock(return_value=None) + monkeypatch.setattr(sp, "get_async", get_async) + service = SessionManagementService(MagicMock(), MagicMock()) + + assert await service.get_groups() == {} + get_async.assert_awaited_once_with( + "unknown", + "unknown", + "session_groups", + {}, + ) + + +@pytest.mark.asyncio +async def test_save_groups_preserves_legacy_preference_scope(monkeypatch): + put_async = AsyncMock() + monkeypatch.setattr(sp, "put_async", put_async) + service = SessionManagementService(MagicMock(), MagicMock()) + groups = {"group-id": {"name": "Group", "umos": []}} + + await service.save_groups(groups) + + put_async.assert_awaited_once_with( + "unknown", + "unknown", + "session_groups", + groups, + ) diff --git a/tests/unit/test_tool_permission.py b/tests/unit/test_tool_permission.py index bb7626aaa..d0d883548 100644 --- a/tests/unit/test_tool_permission.py +++ b/tests/unit/test_tool_permission.py @@ -1,6 +1,6 @@ """Tests for per-tool permission management.""" -from unittest.mock import MagicMock +from unittest.mock import AsyncMock, MagicMock import pytest @@ -50,8 +50,8 @@ def _dummy_tool(name: str = "test_tool") -> FunctionTool: ) -def _clear_tool_permissions() -> None: - sp.put("tool_permissions", {}, scope="global", scope_id="global") +async def _clear_tool_permissions() -> None: + await sp.global_put("tool_permissions", {}) def _make_tools_service( @@ -87,58 +87,52 @@ def test_default_permission_is_member(): @pytest.mark.asyncio async def test_check_permission_passes_when_no_config(): - _clear_tool_permissions() + await _clear_tool_permissions() mgr = FunctionToolManager() context = _make_context(role="member") - error = mgr._check_tool_permission("no_such_tool", context) + error = await mgr._check_tool_permission("no_such_tool", context) assert error is None @pytest.mark.asyncio async def test_check_permission_passes_for_admin_with_admin_tool(): - sp.put( + await sp.global_put( "tool_permissions", {"_default": {"dangerous_tool": "admin"}}, - scope="global", - scope_id="global", ) try: mgr = FunctionToolManager() context = _make_context(role="admin", sender_id="admin_001") - error = mgr._check_tool_permission("dangerous_tool", context) + error = await mgr._check_tool_permission("dangerous_tool", context) assert error is None finally: - _clear_tool_permissions() + await _clear_tool_permissions() @pytest.mark.asyncio async def test_check_permission_denies_member_for_admin_tool(): - sp.put( + await sp.global_put( "tool_permissions", {"_default": {"dangerous_tool": "admin"}}, - scope="global", - scope_id="global", ) try: mgr = FunctionToolManager() context = _make_context(role="member", sender_id="user_999") - error = mgr._check_tool_permission("dangerous_tool", context) + error = await mgr._check_tool_permission("dangerous_tool", context) assert error is not None assert "dangerous_tool" in str(error) assert "admin" in str(error).lower() assert "user_999" in str(error) finally: - _clear_tool_permissions() + await _clear_tool_permissions() @pytest.mark.asyncio async def test_check_permission_denies_when_no_event(): - sp.put( + await sp.global_put( "tool_permissions", {"_default": {"dangerous_tool": "admin"}}, - scope="global", - scope_id="global", ) try: mgr = FunctionToolManager() @@ -146,28 +140,26 @@ async def test_check_permission_denies_when_no_event(): class FakeWrapper: pass # no .context.event - error = mgr._check_tool_permission("dangerous_tool", FakeWrapper()) + error = await mgr._check_tool_permission("dangerous_tool", FakeWrapper()) assert error is not None assert "admin" in str(error).lower() finally: - _clear_tool_permissions() + await _clear_tool_permissions() @pytest.mark.asyncio async def test_check_permission_passes_for_member_when_configured_member(): - sp.put( + await sp.global_put( "tool_permissions", {"_default": {"safe_tool": "member"}}, - scope="global", - scope_id="global", ) try: mgr = FunctionToolManager() context = _make_context(role="member") - error = mgr._check_tool_permission("safe_tool", context) + error = await mgr._check_tool_permission("safe_tool", context) assert error is None finally: - _clear_tool_permissions() + await _clear_tool_permissions() # ── _PermissionGuardedTool ─────────────────────────────────────────── @@ -175,7 +167,7 @@ async def test_check_permission_passes_for_member_when_configured_member(): @pytest.mark.asyncio async def test_guarded_tool_delegates_handler_with_event_when_permission_passes(): - _clear_tool_permissions() + await _clear_tool_permissions() mgr = FunctionToolManager() called = False @@ -205,11 +197,9 @@ async def test_guarded_tool_delegates_handler_with_event_when_permission_passes( @pytest.mark.asyncio async def test_guarded_tool_blocks_when_permission_denied(): - sp.put( + await sp.global_put( "tool_permissions", {"_default": {"blocked_tool": "admin"}}, - scope="global", - scope_id="global", ) try: mgr = FunctionToolManager() @@ -234,12 +224,12 @@ async def test_guarded_tool_blocks_when_permission_denied(): assert isinstance(result, str) assert "Permission denied" in result finally: - _clear_tool_permissions() + await _clear_tool_permissions() @pytest.mark.asyncio async def test_guarded_tool_delegates_to_wrapped_call(): - _clear_tool_permissions() + await _clear_tool_permissions() mgr = FunctionToolManager() class CallableTool(FunctionTool): @@ -260,7 +250,7 @@ async def test_guarded_tool_delegates_to_wrapped_call(): @pytest.mark.asyncio async def test_guarded_tool_delegates_to_wrapped_run(): - _clear_tool_permissions() + await _clear_tool_permissions() mgr = FunctionToolManager() class RunnableTool(FunctionTool): @@ -281,7 +271,7 @@ async def test_guarded_tool_delegates_to_wrapped_run(): @pytest.mark.asyncio async def test_guarded_tool_handles_async_generator_handler(): - _clear_tool_permissions() + await _clear_tool_permissions() mgr = FunctionToolManager() async def gen_handler(event, **kw): # type: ignore[misc] @@ -321,7 +311,6 @@ def test_get_full_tool_set_excludes_builtin_tools(): def test_get_full_tool_set_wraps_non_builtin(): mgr = FunctionToolManager() - _clear_tool_permissions() mgr.func_list.append(_dummy_tool("my_plugin_tool")) tool_set = mgr.get_full_tool_set() @@ -340,33 +329,60 @@ class TestGetToolListPermission: @pytest.mark.asyncio async def test_list_includes_permission_fields_for_non_builtin(self): service = _make_tools_service() - sp.put( + await sp.global_put( "tool_permissions", {"_default": {"my_plugin_tool": "admin"}}, - scope="global", - scope_id="global", ) try: service.tool_mgr.func_list.append(_dummy_tool("my_plugin_tool")) - tools = service.get_tool_list() + tools = await service.get_tool_list() target = next(t for t in tools if t["name"] == "my_plugin_tool") assert target["permission"] == "admin" assert target["permission_configured"] is True assert target["readonly"] is False finally: - _clear_tool_permissions() + await _clear_tool_permissions() @pytest.mark.asyncio async def test_list_no_permission_fields_for_builtin(self): service = _make_tools_service() - tools = service.get_tool_list() + tools = await service.get_tool_list() target = next(t for t in tools if t["name"] == "astrbot_execute_shell") assert "permission" not in target assert "permission_configured" not in target assert target["readonly"] is True + @pytest.mark.asyncio + async def test_list_reads_permission_store_once(self, monkeypatch): + global_get = AsyncMock( + return_value={ + "_default": { + "first_tool": "admin", + "second_tool": "member", + } + } + ) + monkeypatch.setattr(sp, "global_get", global_get) + service = _make_tools_service() + service.tool_mgr.func_list.extend( + [_dummy_tool("first_tool"), _dummy_tool("second_tool")] + ) + + tools = await service.get_tool_list() + + permissions = { + tool["name"]: tool["permission"] + for tool in tools + if tool["name"] in {"first_tool", "second_tool"} + } + assert permissions == { + "first_tool": "admin", + "second_tool": "member", + } + global_get.assert_awaited_once_with("tool_permissions", {}) + # ── API: update_tool_permission ────────────────────────────────────── @@ -376,14 +392,14 @@ class TestUpdateToolPermission: async def test_set_admin_permission(self): service = _make_tools_service() service.tool_mgr.func_list.append(_dummy_tool("target_tool")) - _clear_tool_permissions() + await _clear_tool_permissions() - message = service.update_tool_permission( + message = await service.update_tool_permission( {"name": "target_tool", "permission": "admin"} ) assert "target_tool" in message - stored = sp.get("tool_permissions", {}, scope="global", scope_id="global") + stored = await sp.global_get("tool_permissions", {}) assert stored["_default"]["target_tool"] == "admin" @pytest.mark.asyncio @@ -391,7 +407,7 @@ class TestUpdateToolPermission: service = _make_tools_service() with pytest.raises(ToolsServiceError, match="Builtin"): - service.update_tool_permission( + await service.update_tool_permission( {"name": "astrbot_execute_shell", "permission": "admin"} ) @@ -400,7 +416,7 @@ class TestUpdateToolPermission: service = _make_tools_service() with pytest.raises(ToolsServiceError, match="not found"): - service.update_tool_permission( + await service.update_tool_permission( {"name": "ghost_tool", "permission": "admin"} ) @@ -410,6 +426,6 @@ class TestUpdateToolPermission: service.tool_mgr.func_list.append(_dummy_tool("target_tool")) with pytest.raises(ToolsServiceError, match="admin or member"): - service.update_tool_permission( + await service.update_tool_permission( {"name": "target_tool", "permission": "everyone"} )