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
This commit is contained in:
Ruochen Pan
2026-08-07 12:20:18 +08:00
committed by GitHub
parent 78214cac5a
commit 9fa0a198a4
6 changed files with 152 additions and 98 deletions
+3 -5
View File
@@ -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 {}
+4 -4
View File
@@ -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:
@@ -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
+18 -22
View File
@@ -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 {}
)
@@ -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,
)
+62 -46
View File
@@ -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"}
)