This commit is contained in:
qsc
2026-09-05 02:36:59 +08:00
parent 9c9808adf7
commit 694f6a9604
17 changed files with 1300 additions and 466 deletions
+4
View File
@@ -15,6 +15,7 @@ from astrbot.api.star import Context
from .sqlite import AsyncSQLiteDB
from .server_binding import ServerBindingService
from .session_control import SessionControlService
DEFAULT_WSS_URL = "wss://socket.nicemoe.cn"
@@ -294,11 +295,13 @@ class EventPushService:
config: AstrBotConfig,
sqlite: AsyncSQLiteDB,
server_binding: ServerBindingService,
session_control: SessionControlService,
):
self.context = context
self.config = config
self.sql = sqlite
self.server_binding = server_binding
self.session_control = session_control
self.url = str(config.get("jx3api_wss", "") or DEFAULT_WSS_URL).strip()
self.token = str(config.get("jx3api_wss_token", "") or "").strip()
self._runner: Optional[asyncio.Task] = None
@@ -520,6 +523,7 @@ class EventPushService:
return [
(str(row["session_id"]), str(row.get("server") or "").strip())
for row in rows
if self.session_control.is_allowed(row["session_id"])
]
async def list_subscription_statuses(self) -> list[dict[str, Any]]:
+34 -3
View File
@@ -31,6 +31,18 @@ class KungfuAliasService:
await self._seed_defaults()
async def _seed_defaults(self):
records = self._load_seed_defaults()
for values in records:
await self.sql.execute(
"""
INSERT OR IGNORE INTO kungfu
(pzid, name, name1, name2, name3, name4, name5)
VALUES (?, ?, ?, ?, ?, ?, ?)
""",
values,
)
def _load_seed_defaults(self) -> list[tuple[Any, ...]]:
try:
records = json.loads(self.seed_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
@@ -39,21 +51,40 @@ class KungfuAliasService:
if not isinstance(records, list):
raise RuntimeError("心法种子数据必须是数组")
normalized_records: list[tuple[Any, ...]] = []
seen_pzids: set[int] = set()
for record in records:
if not isinstance(record, dict):
continue
pzid = self._parse_pzid(record.get("pzid"))
if pzid in seen_pzids:
continue
seen_pzids.add(pzid)
name = self._clean(record.get("name"))
aliases = self._normalize_aliases(name, record.get("aliases") or [])
values = [*aliases, *([None] * (self.MAX_ALIASES - len(aliases)))]
await self.sql.execute(
normalized_records.append((pzid, name, *values))
return normalized_records
async def restore_defaults(self) -> int:
"""使用随插件分发的 JSON 种子完整重写心法表。"""
records = self._load_seed_defaults()
statements: list[tuple[str, tuple[Any, ...]]] = [
("DELETE FROM kungfu", ()),
]
statements.extend(
(
"""
INSERT OR IGNORE INTO kungfu
INSERT INTO kungfu
(pzid, name, name1, name2, name3, name4, name5)
VALUES (?, ?, ?, ?, ?, ?, ?)
""",
(pzid, name, *values),
values,
)
for values in records
)
await self.sql.execute_transaction(statements)
return len(records)
async def list_kungfu(self) -> list[dict[str, Any]]:
rows = await self.sql.fetch_all(
+1 -1
View File
@@ -486,7 +486,7 @@ class MessageBuilder:
return await self.T2I_image_msg(event, lambda: self.jx3api.qixue(name,0))
async def liaotian(self, event: AstrMessageEvent, server:str, name: str, limit:int = 20, page:int = 1):
""" 聊天 服务器 角色 条数 页数"""
""" 发言 服务器 角色 条数 页数"""
return await self.T2I_image_msg(event, lambda: self.jx3api.juesheliaotian(server,name,limit,page))
async def tongzhanyy(self, event: AstrMessageEvent, server: str = ""):
+55 -9
View File
@@ -14,7 +14,9 @@ class ServerBindingService:
self.sql = sqlite
self.seed_path = seed_path
self._remote_servers: set[str] = set()
self._standard_servers: set[str] = set()
self._known_servers: set[str] = set()
self._standard_lookup: dict[str, str] = {}
self._server_lookup: dict[str, str] = {}
async def initialize(self):
@@ -38,6 +40,17 @@ class ServerBindingService:
await self._reload_cache()
async def _seed_aliases(self):
records = self._load_seed_aliases()
for server, aliases_json in records:
await self.sql.execute(
"""
INSERT OR IGNORE INTO server_aliases (server, aliases)
VALUES (?, ?)
""",
(server, aliases_json),
)
def _load_seed_aliases(self) -> list[tuple[str, str]]:
try:
records = json.loads(self.seed_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
@@ -46,13 +59,21 @@ class ServerBindingService:
if not isinstance(records, list):
raise RuntimeError("区服别名种子数据必须是数组")
normalized_records: list[tuple[str, str]] = []
seen_servers: set[str] = set()
for record in records:
if not isinstance(record, dict):
continue
server = self._clean(record.get("server"))
raw_aliases = record.get("aliases") or []
if not server or not isinstance(raw_aliases, list):
server_key = self._key(server)
if (
not server
or not isinstance(raw_aliases, list)
or server_key in seen_servers
):
continue
seen_servers.add(server_key)
aliases = []
seen: set[str] = set()
for value in raw_aliases:
@@ -61,20 +82,35 @@ class ServerBindingService:
if alias and key != self._key(server) and key not in seen:
seen.add(key)
aliases.append(alias)
await self.sql.execute(
"""
INSERT OR IGNORE INTO server_aliases (server, aliases)
VALUES (?, ?)
""",
(server, json.dumps(aliases, ensure_ascii=False)),
normalized_records.append(
(server, json.dumps(aliases, ensure_ascii=False))
)
return normalized_records
async def restore_default_aliases(self) -> int:
"""使用随插件分发的 JSON 种子完整重写区服别名表。"""
records = self._load_seed_aliases()
statements: list[tuple[str, tuple[Any, ...]]] = [
("DELETE FROM server_aliases", ()),
]
statements.extend(
(
"INSERT INTO server_aliases (server, aliases) VALUES (?, ?)",
(server, aliases_json),
)
for server, aliases_json in records
)
await self.sql.execute_transaction(statements)
await self._reload_cache()
return len(records)
async def _reload_cache(self):
bindings = await self.list_bindings()
alias_rows = await self.list_aliases()
known = set(self._remote_servers)
standard = set(self._remote_servers)
standard.update(row["server"] for row in alias_rows)
known = set(standard)
known.update(row["server"] for row in bindings)
known.update(row["server"] for row in alias_rows)
lookup = {self._key(server): server for server in known}
for row in alias_rows:
@@ -84,7 +120,11 @@ class ServerBindingService:
# 新增的官方区服名优先于历史别名,避免目录更新后误解析。
lookup.setdefault(self._key(alias), server)
self._standard_servers = standard
self._known_servers = known
self._standard_lookup = {
self._key(server): server for server in standard
}
self._server_lookup = lookup
async def update_server_catalog(self, servers: Iterable[str]):
@@ -98,6 +138,12 @@ class ServerBindingService:
def known_servers(self) -> list[str]:
return sorted(self._known_servers)
def standard_servers(self) -> list[str]:
return sorted(self._standard_servers)
def resolve_standard_server(self, value: Any) -> str:
return self._standard_lookup.get(self._key(value), "")
def is_known_server(self, value: Any) -> bool:
return self._key(value) in self._server_lookup
+154
View File
@@ -0,0 +1,154 @@
from typing import Any
from .sqlite import AsyncSQLiteDB
class SessionControlService:
"""使用 SQLite 保存并缓存插件的会话访问策略。"""
MODE_ALL = "all"
MODE_WHITELIST = "whitelist"
MODE_BLACKLIST = "blacklist"
MODES = {MODE_ALL, MODE_WHITELIST, MODE_BLACKLIST}
def __init__(self, sqlite: AsyncSQLiteDB):
self.sql = sqlite
self._mode = self.MODE_ALL
self._entries: dict[str, dict[str, str]] = {}
async def initialize(self):
await self.sql.execute(
"""
CREATE TABLE IF NOT EXISTS session_control_settings (
key TEXT PRIMARY KEY,
value TEXT NOT NULL,
updated_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP
)
"""
)
await self.sql.execute(
"""
CREATE TABLE IF NOT EXISTS session_control_entries (
session_id TEXT PRIMARY KEY,
list_type TEXT NOT NULL CHECK(list_type IN ('whitelist', 'blacklist')),
remark TEXT NOT NULL DEFAULT '',
updated_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP
)
"""
)
await self.sql.execute(
"""
INSERT OR IGNORE INTO session_control_settings (key, value)
VALUES ('mode', 'all')
"""
)
await self._reload()
async def _reload(self):
setting = await self.sql.fetch_one(
"SELECT value FROM session_control_settings WHERE key='mode'"
)
mode = str(setting.get("value") or "") if setting else ""
self._mode = mode if mode in self.MODES else self.MODE_ALL
rows = await self.sql.fetch_all(
"""
SELECT session_id, list_type, remark
FROM session_control_entries
ORDER BY session_id
"""
)
self._entries = {
str(row["session_id"]): {
"session_id": str(row["session_id"]),
"list_type": str(row["list_type"]),
"remark": str(row.get("remark") or ""),
}
for row in rows
}
@staticmethod
def _normalize_session_id(session_id: Any) -> str:
value = str(session_id or "").strip()
if not value:
raise ValueError("会话 ID 不能为空")
if len(value) > 512:
raise ValueError("会话 ID 不能超过 512 个字符")
return value
def is_allowed(self, session_id: Any) -> bool:
"""同步检查缓存策略,供高频消息和推送分发入口调用。"""
value = str(session_id or "").strip()
if self._mode == self.MODE_ALL:
return True
entry = self._entries.get(value)
if self._mode == self.MODE_WHITELIST:
return bool(entry and entry["list_type"] == self.MODE_WHITELIST)
return not (entry and entry["list_type"] == self.MODE_BLACKLIST)
async def set_mode(self, mode: Any):
value = str(mode or "").strip().lower()
if value not in self.MODES:
raise ValueError("会话控制模式无效")
await self.sql.execute(
"""
INSERT INTO session_control_settings (key, value, updated_at)
VALUES ('mode', ?, CURRENT_TIMESTAMP)
ON CONFLICT(key) DO UPDATE SET
value=excluded.value,
updated_at=CURRENT_TIMESTAMP
""",
(value,),
)
self._mode = value
async def save_entry(
self,
session_id: Any,
list_type: Any,
remark: Any = "",
):
normalized_session_id = self._normalize_session_id(session_id)
normalized_type = str(list_type or "").strip().lower()
if normalized_type not in {self.MODE_WHITELIST, self.MODE_BLACKLIST}:
raise ValueError("名单类型必须是白名单或黑名单")
normalized_remark = str(remark or "").strip()
if len(normalized_remark) > 200:
raise ValueError("备注不能超过 200 个字符")
await self.sql.execute(
"""
INSERT INTO session_control_entries (
session_id, list_type, remark, updated_at
) VALUES (?, ?, ?, CURRENT_TIMESTAMP)
ON CONFLICT(session_id) DO UPDATE SET
list_type=excluded.list_type,
remark=excluded.remark,
updated_at=CURRENT_TIMESTAMP
""",
(normalized_session_id, normalized_type, normalized_remark),
)
self._entries[normalized_session_id] = {
"session_id": normalized_session_id,
"list_type": normalized_type,
"remark": normalized_remark,
}
async def delete_entry(self, session_id: Any):
normalized_session_id = self._normalize_session_id(session_id)
await self.sql.delete(
"session_control_entries",
"session_id=?",
(normalized_session_id,),
)
self._entries.pop(normalized_session_id, None)
async def get_state(self) -> dict[str, Any]:
return {
"mode": self._mode,
"entries": sorted(
(dict(entry) for entry in self._entries.values()),
key=lambda entry: entry["session_id"],
),
}
+14
View File
@@ -37,6 +37,20 @@ class AsyncSQLiteDB:
async with self.conn.execute(sql, params):
await self.conn.commit()
async def execute_transaction(
self,
statements: List[Tuple[str, Tuple[Any, ...]]],
):
"""在同一事务内顺序执行多条参数化 SQL。"""
await self.conn.execute("BEGIN IMMEDIATE")
try:
for sql, params in statements:
await self.conn.execute(sql, params)
await self.conn.commit()
except Exception:
await self.conn.rollback()
raise
async def fetch_one(self, sql: str, params: Tuple = ()) -> Optional[Dict[str, Any]]:
async with self.conn.execute(sql, params) as cursor:
row = await cursor.fetchone()
+216
View File
@@ -0,0 +1,216 @@
from __future__ import annotations
import re
from typing import TYPE_CHECKING, Any
from astrbot.api.star import Context
from astrbot.api.web import error_response, json_response, request
from .event_push import EVENT_NAMES
if TYPE_CHECKING:
from .event_push import EventPushService
from .jx3api_data import JX3APIService
from .kungfu_alias import KungfuAliasService
from .server_binding import ServerBindingService
from .session_control import SessionControlService
class WebUIService:
"""注册插件管理页接口,并处理 WebUI 的数据读写。"""
def __init__(
self,
jx3api: JX3APIService,
event_push: EventPushService,
server_binding: ServerBindingService,
kungfu_alias: KungfuAliasService,
session_control: SessionControlService,
):
self.jx3api = jx3api
self.event_push = event_push
self.server_binding = server_binding
self.kungfu_alias = kungfu_alias
self.session_control = session_control
def register(self, context: Context, plugin_name: str):
routes = (
("dashboard", self.dashboard, ["GET"], "读取会话管理数据"),
("bindings/save", self.save_binding, ["POST"], "保存会话区服绑定"),
("bindings/delete", self.delete_binding, ["POST"], "删除会话区服绑定"),
("aliases/save", self.save_aliases, ["POST"], "保存区服别名"),
("aliases/delete", self.delete_aliases, ["POST"], "删除区服别名"),
("aliases/restore", self.restore_aliases, ["POST"], "恢复默认区服别名"),
("kungfu/save", self.save_kungfu, ["POST"], "保存心法别名"),
("kungfu/restore", self.restore_kungfu, ["POST"], "恢复默认心法别名"),
("servers/refresh", self.refresh_servers, ["POST"], "刷新区服目录"),
(
"session-control/mode",
self.save_session_control_mode,
["POST"],
"保存会话控制模式",
),
(
"session-control/save",
self.save_session_control_entry,
["POST"],
"保存会话控制名单",
),
(
"session-control/delete",
self.delete_session_control_entry,
["POST"],
"删除会话控制名单",
),
)
for path, handler, methods, description in routes:
context.register_web_api(
f"/{plugin_name}/{path}",
handler,
methods,
description,
)
@staticmethod
async def _json_payload() -> dict[str, Any]:
payload = await request.json(default={})
if not isinstance(payload, dict):
raise ValueError("请求正文必须是 JSON 对象")
return payload
@staticmethod
def _parse_aliases(raw_aliases: Any) -> list[str]:
if isinstance(raw_aliases, str):
return re.split(r"[,;\n]+", raw_aliases)
if isinstance(raw_aliases, list):
return [str(value) for value in raw_aliases]
raise ValueError("别名必须是字符串或数组")
async def dashboard(self):
bindings = await self.server_binding.list_bindings()
subscriptions = await self.event_push.list_subscription_statuses()
aliases = await self.server_binding.list_aliases()
kungfu = await self.kungfu_alias.list_kungfu()
session_control = await self.session_control.get_state()
return json_response(
{
"bindings": bindings,
"subscriptions": subscriptions,
"aliases": aliases,
"kungfu": kungfu,
"servers": self.server_binding.standard_servers(),
"events": {
str(action): name for action, name in EVENT_NAMES.items()
},
"session_control": session_control,
}
)
async def save_binding(self):
try:
payload = await self._json_payload()
server = self.server_binding.resolve_standard_server(
payload.get("server")
)
if not server:
raise ValueError("绑定区服必须选择标准区服")
await self.server_binding.set_binding(
str(payload.get("session_id") or ""),
server,
)
except ValueError as exc:
return error_response(str(exc), status_code=400)
return json_response({"saved": True})
async def delete_binding(self):
try:
payload = await self._json_payload()
session_id = str(payload.get("session_id") or "")
if not session_id.strip():
raise ValueError("会话 ID 不能为空")
await self.server_binding.delete_binding(session_id)
except ValueError as exc:
return error_response(str(exc), status_code=400)
return json_response({"deleted": True})
async def save_aliases(self):
try:
payload = await self._json_payload()
await self.server_binding.set_aliases(
str(payload.get("server") or ""),
self._parse_aliases(payload.get("aliases", [])),
)
except ValueError as exc:
return error_response(str(exc), status_code=400)
return json_response({"saved": True})
async def delete_aliases(self):
try:
payload = await self._json_payload()
server = str(payload.get("server") or "")
if not server.strip():
raise ValueError("标准区服名不能为空")
await self.server_binding.delete_aliases(server)
except ValueError as exc:
return error_response(str(exc), status_code=400)
return json_response({"deleted": True})
async def restore_aliases(self):
try:
restored = await self.server_binding.restore_default_aliases()
except (RuntimeError, ValueError) as exc:
return error_response(str(exc), status_code=500)
return json_response({"restored": restored})
async def save_kungfu(self):
try:
payload = await self._json_payload()
await self.kungfu_alias.save_aliases(
payload.get("pzid"),
self._parse_aliases(payload.get("aliases", [])),
)
except ValueError as exc:
return error_response(str(exc), status_code=400)
return json_response({"saved": True})
async def restore_kungfu(self):
try:
restored = await self.kungfu_alias.restore_defaults()
except (RuntimeError, ValueError) as exc:
return error_response(str(exc), status_code=500)
return json_response({"restored": restored})
async def refresh_servers(self):
servers = await self.jx3api.server_list()
if not servers:
return error_response("区服目录刷新失败", status_code=502)
await self.server_binding.update_server_catalog(servers)
return json_response({"servers": self.server_binding.standard_servers()})
async def save_session_control_mode(self):
try:
payload = await self._json_payload()
await self.session_control.set_mode(payload.get("mode"))
except ValueError as exc:
return error_response(str(exc), status_code=400)
return json_response({"saved": True})
async def save_session_control_entry(self):
try:
payload = await self._json_payload()
await self.session_control.save_entry(
payload.get("session_id"),
payload.get("list_type"),
payload.get("remark"),
)
except ValueError as exc:
return error_response(str(exc), status_code=400)
return json_response({"saved": True})
async def delete_session_control_entry(self):
try:
payload = await self._json_payload()
await self.session_control.delete_entry(payload.get("session_id"))
except ValueError as exc:
return error_response(str(exc), status_code=400)
return json_response({"deleted": True})