增加webui界面
This commit is contained in:
+42
-4
@@ -14,6 +14,7 @@ from astrbot.api.event import MessageChain
|
||||
from astrbot.api.star import Context
|
||||
|
||||
from .sqlite import AsyncSQLiteDB
|
||||
from .server_binding import ServerBindingService
|
||||
|
||||
|
||||
DEFAULT_WSS_URL = "wss://socket.nicemoe.cn"
|
||||
@@ -70,10 +71,12 @@ class EventPushService:
|
||||
context: Context,
|
||||
config: AstrBotConfig,
|
||||
sqlite: AsyncSQLiteDB,
|
||||
server_binding: ServerBindingService,
|
||||
):
|
||||
self.context = context
|
||||
self.config = config
|
||||
self.sql = sqlite
|
||||
self.server_binding = server_binding
|
||||
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
|
||||
@@ -252,6 +255,16 @@ class EventPushService:
|
||||
return
|
||||
|
||||
recipients = await self._enabled_sessions(action)
|
||||
if "server" in detail:
|
||||
event_server = self.server_binding.resolve_server(detail.get("server"))
|
||||
recipients = [
|
||||
session_id
|
||||
for session_id, bound_server in recipients
|
||||
if not bound_server
|
||||
or self.server_binding.resolve_server(bound_server) == event_server
|
||||
]
|
||||
else:
|
||||
recipients = [session_id for session_id, _ in recipients]
|
||||
if not recipients:
|
||||
return
|
||||
|
||||
@@ -271,13 +284,38 @@ class EventPushService:
|
||||
message_chain = MessageChain().message(text)
|
||||
await self.context.send_message(session_id, message_chain)
|
||||
|
||||
async def _enabled_sessions(self, action: int) -> list[str]:
|
||||
async def _enabled_sessions(self, action: int) -> list[tuple[str, str]]:
|
||||
column = self._action_column(action)
|
||||
rows = await self.sql.fetch_all(
|
||||
f"SELECT session_id FROM event_push_subscriptions "
|
||||
f"WHERE enabled=1 AND {column}=1"
|
||||
f"SELECT subscriptions.session_id, "
|
||||
f"COALESCE(bindings.server, '') AS server "
|
||||
f"FROM event_push_subscriptions AS subscriptions "
|
||||
f"LEFT JOIN session_server_bindings AS bindings "
|
||||
f"ON bindings.session_id=subscriptions.session_id "
|
||||
f"WHERE subscriptions.enabled=1 "
|
||||
f"AND subscriptions.{column}=1"
|
||||
)
|
||||
return [str(row["session_id"]) for row in rows]
|
||||
return [
|
||||
(str(row["session_id"]), str(row.get("server") or "").strip())
|
||||
for row in rows
|
||||
]
|
||||
|
||||
async def list_subscription_statuses(self) -> list[dict[str, Any]]:
|
||||
rows = await self.sql.fetch_all(
|
||||
"SELECT * FROM event_push_subscriptions ORDER BY session_id"
|
||||
)
|
||||
return [
|
||||
{
|
||||
"session_id": str(row["session_id"]),
|
||||
"enabled": row.get("enabled") == 1,
|
||||
"actions": [
|
||||
action
|
||||
for action in EVENT_ACTIONS
|
||||
if row.get(self._action_column(action)) == 1
|
||||
],
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
|
||||
async def configure(
|
||||
self,
|
||||
|
||||
+27
-2
@@ -10,7 +10,7 @@ from astrbot.api import logger
|
||||
from astrbot.api import AstrBotConfig
|
||||
import astrbot.api.message_components as Comp
|
||||
|
||||
from .request import APIClient
|
||||
from .request import APIClient, APIErrorResponse
|
||||
from .sqlite import AsyncSQLiteDB
|
||||
from .fun_basic import load_template,gold_to_parts,week_to_num,compare_date_str,format_time,format_remaining
|
||||
|
||||
@@ -47,6 +47,22 @@ class JX3APIService:
|
||||
if self._api:
|
||||
await self._api.close()
|
||||
|
||||
async def server_list(self) -> list[str]:
|
||||
"""获取当前有效区服名称,供会话绑定和参数消歧使用。"""
|
||||
data = await self._base_request(
|
||||
"/server/status/check",
|
||||
{"server": "", "type": "其他"},
|
||||
)
|
||||
if not isinstance(data, list):
|
||||
return []
|
||||
return sorted(
|
||||
{
|
||||
str(item.get("server") or "").strip()
|
||||
for item in data
|
||||
if isinstance(item, dict) and item.get("server")
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _init_return_data(self) -> Dict[str, Any]:
|
||||
"""初始化标准的返回数据结构"""
|
||||
@@ -75,7 +91,13 @@ class JX3APIService:
|
||||
|
||||
base_url = "https://www.jx3api.com"
|
||||
api_url = base_url + api_path
|
||||
data = await self._api.get(api_url, params=params, out_key=out)
|
||||
data = await self._api.get(
|
||||
api_url,
|
||||
params=params,
|
||||
out_key=out,
|
||||
success_codes=(200, "200"),
|
||||
return_error=True,
|
||||
)
|
||||
|
||||
if not data:
|
||||
logger.warning(f"获取接口信息失败或返回空数据: {api_url}")
|
||||
@@ -100,6 +122,9 @@ class JX3APIService:
|
||||
return_data = self._init_return_data()
|
||||
|
||||
data = await self._base_request(path, params)
|
||||
if isinstance(data, APIErrorResponse):
|
||||
return_data["msg"] = data.message or "获取接口信息失败"
|
||||
return return_data
|
||||
if data is None:
|
||||
return_data["msg"] = "获取接口信息失败"
|
||||
return return_data
|
||||
|
||||
@@ -0,0 +1,148 @@
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable
|
||||
|
||||
from .sqlite import AsyncSQLiteDB
|
||||
|
||||
|
||||
class KungfuAliasService:
|
||||
"""维护本地心法名称、JX3BOX 配装 ID 和别名。"""
|
||||
|
||||
MAX_ALIASES = 5
|
||||
|
||||
def __init__(self, sqlite: AsyncSQLiteDB, seed_path: Path):
|
||||
self.sql = sqlite
|
||||
self.seed_path = seed_path
|
||||
|
||||
async def initialize(self):
|
||||
await self.sql.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS kungfu (
|
||||
pzid INTEGER PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
name1 TEXT,
|
||||
name2 TEXT,
|
||||
name3 TEXT,
|
||||
name4 TEXT,
|
||||
name5 TEXT
|
||||
)
|
||||
"""
|
||||
)
|
||||
await self._seed_defaults()
|
||||
|
||||
async def _seed_defaults(self):
|
||||
try:
|
||||
records = json.loads(self.seed_path.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError) as exc:
|
||||
raise RuntimeError(f"读取心法种子数据失败:{exc}") from exc
|
||||
|
||||
if not isinstance(records, list):
|
||||
raise RuntimeError("心法种子数据必须是数组")
|
||||
|
||||
for record in records:
|
||||
if not isinstance(record, dict):
|
||||
continue
|
||||
pzid = self._parse_pzid(record.get("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(
|
||||
"""
|
||||
INSERT OR IGNORE INTO kungfu
|
||||
(pzid, name, name1, name2, name3, name4, name5)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(pzid, name, *values),
|
||||
)
|
||||
|
||||
async def list_kungfu(self) -> list[dict[str, Any]]:
|
||||
rows = await self.sql.fetch_all(
|
||||
"SELECT pzid, name, name1, name2, name3, name4, name5 "
|
||||
"FROM kungfu ORDER BY pzid"
|
||||
)
|
||||
return [
|
||||
{
|
||||
"pzid": int(row["pzid"]),
|
||||
"name": self._clean(row.get("name")),
|
||||
"aliases": [
|
||||
alias
|
||||
for key in ("name1", "name2", "name3", "name4", "name5")
|
||||
if (alias := self._clean(row.get(key)))
|
||||
],
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
|
||||
async def save(self, pzid: Any, name: Any, aliases: Iterable[Any]):
|
||||
normalized_pzid = self._parse_pzid(pzid)
|
||||
normalized_name = self._clean(name)
|
||||
if not normalized_name:
|
||||
raise ValueError("标准心法名不能为空")
|
||||
if len(normalized_name) > 64:
|
||||
raise ValueError("标准心法名过长")
|
||||
|
||||
normalized_aliases = self._normalize_aliases(normalized_name, aliases)
|
||||
occupied: dict[str, str] = {}
|
||||
for row in await self.list_kungfu():
|
||||
if row["pzid"] == normalized_pzid:
|
||||
continue
|
||||
for value in [row["name"], *row["aliases"]]:
|
||||
occupied[self._key(value)] = row["name"]
|
||||
|
||||
for value in [normalized_name, *normalized_aliases]:
|
||||
conflict = occupied.get(self._key(value))
|
||||
if conflict:
|
||||
raise ValueError(f"“{value}”已被心法“{conflict}”使用")
|
||||
|
||||
values = [
|
||||
*normalized_aliases,
|
||||
*([None] * (self.MAX_ALIASES - len(normalized_aliases))),
|
||||
]
|
||||
await self.sql.execute(
|
||||
"""
|
||||
INSERT INTO kungfu (pzid, name, name1, name2, name3, name4, name5)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(pzid) DO UPDATE SET
|
||||
name=excluded.name,
|
||||
name1=excluded.name1,
|
||||
name2=excluded.name2,
|
||||
name3=excluded.name3,
|
||||
name4=excluded.name4,
|
||||
name5=excluded.name5
|
||||
""",
|
||||
(normalized_pzid, normalized_name, *values),
|
||||
)
|
||||
|
||||
def _normalize_aliases(self, name: str, aliases: Iterable[Any]) -> list[str]:
|
||||
result: list[str] = []
|
||||
seen: set[str] = {self._key(name)}
|
||||
for value in aliases:
|
||||
alias = self._clean(value)
|
||||
key = self._key(alias)
|
||||
if not alias or key in seen:
|
||||
continue
|
||||
if len(alias) > 64:
|
||||
raise ValueError(f"心法别名过长:{alias}")
|
||||
seen.add(key)
|
||||
result.append(alias)
|
||||
if len(result) > self.MAX_ALIASES:
|
||||
raise ValueError(f"每个心法最多配置 {self.MAX_ALIASES} 个别名")
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _parse_pzid(value: Any) -> int:
|
||||
try:
|
||||
pzid = int(value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError("心法 ID 必须是整数") from exc
|
||||
if pzid <= 0:
|
||||
raise ValueError("心法 ID 必须大于 0")
|
||||
return pzid
|
||||
|
||||
@staticmethod
|
||||
def _clean(value: Any) -> str:
|
||||
return str(value or "").strip()
|
||||
|
||||
@classmethod
|
||||
def _key(cls, value: Any) -> str:
|
||||
return cls._clean(value).casefold()
|
||||
+18
-16
@@ -38,14 +38,12 @@ class MessageBuilder:
|
||||
)
|
||||
|
||||
def __init__(self,
|
||||
server: str,
|
||||
jx3api: JX3APIService,
|
||||
jx3box: JX3BOXService,
|
||||
bilei: BiLeidata,
|
||||
event_push: EventPushService,
|
||||
icons: dict[str, dict[str, str]]
|
||||
):
|
||||
self.server = server
|
||||
self.jx3api = jx3api
|
||||
self.jx3box = jx3box
|
||||
self.bilei = bilei
|
||||
@@ -69,13 +67,6 @@ class MessageBuilder:
|
||||
)
|
||||
|
||||
|
||||
def serverdefault(self,server) -> str:
|
||||
"""加载配置默认服务器"""
|
||||
if server == "":
|
||||
return self.server
|
||||
return server
|
||||
|
||||
|
||||
async def plain_msg(self, event: AstrMessageEvent, action):
|
||||
"""最终将数据整理成文本发送"""
|
||||
data= await action()
|
||||
@@ -89,7 +80,12 @@ class MessageBuilder:
|
||||
await event.send(event.plain_result("猪脑过载,请稍后再试"))
|
||||
|
||||
|
||||
async def T2I_image_msg(self, event: AstrMessageEvent, action):
|
||||
async def T2I_image_msg(
|
||||
self,
|
||||
event: AstrMessageEvent,
|
||||
action,
|
||||
render_options: dict | None = None,
|
||||
):
|
||||
"""最终将数据渲染成图片发送"""
|
||||
data = await action()
|
||||
try:
|
||||
@@ -101,6 +97,9 @@ class MessageBuilder:
|
||||
"omit_background": False,
|
||||
"type": "jpeg"
|
||||
}
|
||||
options.update(render_options or {})
|
||||
if options.get("type") == "png":
|
||||
options.pop("quality", None)
|
||||
data["data"]["icons"] = self.icons
|
||||
url = await self.html_render(data["temp"], data["data"], options=options)
|
||||
await event.send(event.image_result(url))
|
||||
@@ -109,7 +108,7 @@ class MessageBuilder:
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"功能函数执行错误: {e}")
|
||||
await event.send(event.plain_result("猪脑过载,请稍后再试"))
|
||||
await event.send(event.plain_result("猪脑过载,请稍后再试"))
|
||||
|
||||
|
||||
async def image_msg(self, event: AstrMessageEvent, action):
|
||||
@@ -324,9 +323,9 @@ class MessageBuilder:
|
||||
""" 阵营事件 阵营"""
|
||||
return await self.T2I_image_msg(event, lambda: self.jx3api.zhenyingevent(name,50))
|
||||
|
||||
async def yanhuachaxun(self, event: AstrMessageEvent,server: str = "",name: str = "" ):
|
||||
async def yanhuachaxun(self, event: AstrMessageEvent,server: str,name: str = "" ):
|
||||
""" 烟花 服务器 角色"""
|
||||
return await self.T2I_image_msg(event, lambda: self.jx3api.yanhuachaxun( self.serverdefault(server),name))
|
||||
return await self.T2I_image_msg(event, lambda: self.jx3api.yanhuachaxun(server,name))
|
||||
|
||||
async def shuma(self, event: AstrMessageEvent,server: str ):
|
||||
""" 刷马 服务器"""
|
||||
@@ -448,10 +447,13 @@ class MessageBuilder:
|
||||
""" 帮战 服务器"""
|
||||
return await self.T2I_image_msg(event, lambda: self.jx3api.bangzhanjilu(server))
|
||||
|
||||
async def shapan(self, event: AstrMessageEvent,server: str = ""):
|
||||
async def shapan(self, event: AstrMessageEvent,server: str):
|
||||
""" 沙盘 服务器"""
|
||||
server = self.serverdefault(server)
|
||||
return await self.T2I_image_msg(event, lambda: self.jx3api.shapan(server))
|
||||
return await self.T2I_image_msg(
|
||||
event,
|
||||
lambda: self.jx3api.shapan(server),
|
||||
render_options={"omit_background": True, "type": "png"},
|
||||
)
|
||||
|
||||
async def zhueevent(self, event: AstrMessageEvent, server: str):
|
||||
""" 诛恶事件 服务器"""
|
||||
|
||||
+63
-16
@@ -2,11 +2,20 @@
|
||||
import json
|
||||
import aiohttp
|
||||
import asyncio
|
||||
from typing import Optional, Dict, Any, Union, List
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Dict, Any, Union, List, Collection
|
||||
from aiohttp import ClientTimeout, ClientSession
|
||||
|
||||
from astrbot.api import logger
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class APIErrorResponse:
|
||||
"""保留接口业务错误,避免提取 data 时丢失 code/msg。"""
|
||||
|
||||
code: Any
|
||||
message: str
|
||||
|
||||
class APIClient:
|
||||
"""
|
||||
API客户端类
|
||||
@@ -102,7 +111,7 @@ class APIClient:
|
||||
return None
|
||||
|
||||
def _validate_api_payload(self, data: Any) -> Any:
|
||||
"""校验业务层面的 JSON 数据结构"""
|
||||
"""校验并规范化 JSON 数据结构。"""
|
||||
if not data:
|
||||
logger.error("API返回空数据")
|
||||
return None
|
||||
@@ -114,32 +123,70 @@ class APIClient:
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
|
||||
if isinstance(data, dict) and 'code' in data:
|
||||
# 兼容多种成功状态码:200, "0", 0, 1
|
||||
code = data.get('code')
|
||||
if code not in [200, "0", 0, 1]:
|
||||
msg = data.get('msg') or data.get('message', '未知错误')
|
||||
logger.error(f"API业务报错: code={code}, msg={msg}")
|
||||
return None
|
||||
|
||||
return data
|
||||
|
||||
async def get(self, url: str, params: Optional[Dict] = None, out_key: Optional[str] = None) -> Any:
|
||||
async def get(
|
||||
self,
|
||||
url: str,
|
||||
params: Optional[Dict] = None,
|
||||
out_key: Optional[str] = None,
|
||||
success_codes: Optional[Collection[Any]] = None,
|
||||
return_error: bool = False,
|
||||
) -> Any:
|
||||
"""GET 请求封装"""
|
||||
data = await self._request('GET', url, params=params)
|
||||
return self._extract_data(data, out_key)
|
||||
return self._extract_data(
|
||||
data,
|
||||
out_key,
|
||||
success_codes=success_codes,
|
||||
return_error=return_error,
|
||||
)
|
||||
|
||||
async def post(self, url: str, data: Optional[Dict] = None, out_key: Optional[str] = None) -> Any:
|
||||
async def post(
|
||||
self,
|
||||
url: str,
|
||||
data: Optional[Dict] = None,
|
||||
out_key: Optional[str] = None,
|
||||
success_codes: Optional[Collection[Any]] = None,
|
||||
return_error: bool = False,
|
||||
) -> Any:
|
||||
"""POST 请求封装 (默认发送 JSON)"""
|
||||
data = await self._request('POST', url, json_data=data)
|
||||
return self._extract_data(data, out_key)
|
||||
return self._extract_data(
|
||||
data,
|
||||
out_key,
|
||||
success_codes=success_codes,
|
||||
return_error=return_error,
|
||||
)
|
||||
|
||||
def _extract_data(self, data: Any, key: Optional[str]) -> Any:
|
||||
def _extract_data(
|
||||
self,
|
||||
data: Any,
|
||||
key: Optional[str],
|
||||
success_codes: Optional[Collection[Any]] = None,
|
||||
return_error: bool = False,
|
||||
) -> Any:
|
||||
"""辅助方法:从结果中提取指定字段"""
|
||||
if data is None:
|
||||
return None
|
||||
if isinstance(data, bytes):
|
||||
return data
|
||||
if isinstance(data, dict) and 'code' in data:
|
||||
allowed_codes = (
|
||||
set(success_codes)
|
||||
if success_codes is not None
|
||||
else {200, "0", 0, 1}
|
||||
)
|
||||
code = data.get('code')
|
||||
if code not in allowed_codes:
|
||||
raw_message = data.get('msg') or data.get('message') or ""
|
||||
message = str(raw_message).strip()
|
||||
logger.error(
|
||||
f"API业务报错: code={code}, msg={message or '未知错误'}"
|
||||
)
|
||||
if return_error:
|
||||
return APIErrorResponse(code=code, message=message)
|
||||
return None
|
||||
if key and isinstance(data, dict):
|
||||
return data.get(key, {})
|
||||
return data
|
||||
@@ -188,4 +235,4 @@ class APIClient:
|
||||
current_page += 1
|
||||
logger.info(f"已获取第 {current_page} 页数据")
|
||||
|
||||
return all_data
|
||||
return all_data
|
||||
|
||||
@@ -0,0 +1,241 @@
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable
|
||||
|
||||
from .sqlite import AsyncSQLiteDB
|
||||
|
||||
|
||||
class ServerBindingService:
|
||||
"""维护会话区服绑定、区服别名与指令解析所需的区服目录。"""
|
||||
|
||||
def __init__(self, sqlite: AsyncSQLiteDB, seed_path: Path):
|
||||
self.sql = sqlite
|
||||
self.seed_path = seed_path
|
||||
self._remote_servers: set[str] = set()
|
||||
self._known_servers: set[str] = set()
|
||||
self._server_lookup: dict[str, str] = {}
|
||||
|
||||
async def initialize(self):
|
||||
await self.sql.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS session_server_bindings (
|
||||
session_id TEXT PRIMARY KEY,
|
||||
server TEXT NOT NULL
|
||||
)
|
||||
"""
|
||||
)
|
||||
await self.sql.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS server_aliases (
|
||||
server TEXT PRIMARY KEY,
|
||||
aliases TEXT NOT NULL DEFAULT '[]'
|
||||
)
|
||||
"""
|
||||
)
|
||||
await self._seed_aliases()
|
||||
await self._reload_cache()
|
||||
|
||||
async def _seed_aliases(self):
|
||||
try:
|
||||
records = json.loads(self.seed_path.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError) as exc:
|
||||
raise RuntimeError(f"读取区服别名种子数据失败:{exc}") from exc
|
||||
|
||||
if not isinstance(records, list):
|
||||
raise RuntimeError("区服别名种子数据必须是数组")
|
||||
|
||||
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):
|
||||
continue
|
||||
aliases = []
|
||||
seen: set[str] = set()
|
||||
for value in raw_aliases:
|
||||
alias = self._clean(value)
|
||||
key = self._key(alias)
|
||||
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)),
|
||||
)
|
||||
|
||||
async def _reload_cache(self):
|
||||
bindings = await self.list_bindings()
|
||||
alias_rows = await self.list_aliases()
|
||||
known = set(self._remote_servers)
|
||||
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:
|
||||
server = row["server"]
|
||||
lookup.setdefault(self._key(server), server)
|
||||
for alias in row["aliases"]:
|
||||
# 新增的官方区服名优先于历史别名,避免目录更新后误解析。
|
||||
lookup.setdefault(self._key(alias), server)
|
||||
|
||||
self._known_servers = known
|
||||
self._server_lookup = lookup
|
||||
|
||||
async def update_server_catalog(self, servers: Iterable[str]):
|
||||
self._remote_servers = {
|
||||
normalized
|
||||
for value in servers
|
||||
if (normalized := self._clean(value))
|
||||
}
|
||||
await self._reload_cache()
|
||||
|
||||
def known_servers(self) -> list[str]:
|
||||
return sorted(self._known_servers)
|
||||
|
||||
def is_known_server(self, value: Any) -> bool:
|
||||
return self._key(value) in self._server_lookup
|
||||
|
||||
def resolve_server(self, value: Any) -> str:
|
||||
server = self._clean(value)
|
||||
return self._server_lookup.get(self._key(server), server)
|
||||
|
||||
async def get_binding(self, session_id: str) -> str:
|
||||
row = await self.sql.select_one(
|
||||
"session_server_bindings",
|
||||
"session_id=?",
|
||||
(session_id,),
|
||||
)
|
||||
return self._clean(row.get("server")) if row else ""
|
||||
|
||||
async def set_binding(self, session_id: str, server: str):
|
||||
session_id = self._clean(session_id)
|
||||
server = self.resolve_server(server)
|
||||
if not session_id:
|
||||
raise ValueError("会话 ID 不能为空")
|
||||
if not server:
|
||||
raise ValueError("绑定区服不能为空")
|
||||
if len(session_id) > 512 or len(server) > 64:
|
||||
raise ValueError("会话 ID 或区服名称过长")
|
||||
|
||||
await self.sql.execute(
|
||||
"""
|
||||
INSERT INTO session_server_bindings (session_id, server)
|
||||
VALUES (?, ?)
|
||||
ON CONFLICT(session_id) DO UPDATE SET server=excluded.server
|
||||
""",
|
||||
(session_id, server),
|
||||
)
|
||||
await self._reload_cache()
|
||||
|
||||
async def delete_binding(self, session_id: str):
|
||||
await self.sql.delete(
|
||||
"session_server_bindings",
|
||||
"session_id=?",
|
||||
(self._clean(session_id),),
|
||||
)
|
||||
await self._reload_cache()
|
||||
|
||||
async def list_bindings(self) -> list[dict[str, str]]:
|
||||
rows = await self.sql.fetch_all(
|
||||
"SELECT session_id, server FROM session_server_bindings "
|
||||
"ORDER BY session_id"
|
||||
)
|
||||
return [
|
||||
{
|
||||
"session_id": self._clean(row.get("session_id")),
|
||||
"server": self._clean(row.get("server")),
|
||||
}
|
||||
for row in rows
|
||||
]
|
||||
|
||||
async def set_aliases(self, server: str, aliases: Iterable[str]):
|
||||
server = self._clean(server)
|
||||
if not server:
|
||||
raise ValueError("标准区服名不能为空")
|
||||
if len(server) > 64:
|
||||
raise ValueError("区服名称过长")
|
||||
|
||||
cleaned_aliases: list[str] = []
|
||||
seen: set[str] = set()
|
||||
for value in aliases:
|
||||
alias = self._clean(value)
|
||||
key = self._key(alias)
|
||||
if not alias or key == self._key(server) or key in seen:
|
||||
continue
|
||||
if len(alias) > 64:
|
||||
raise ValueError(f"区服别名过长:{alias}")
|
||||
seen.add(key)
|
||||
cleaned_aliases.append(alias)
|
||||
if len(cleaned_aliases) > 50:
|
||||
raise ValueError("每个区服最多配置 50 个别名")
|
||||
|
||||
alias_rows = await self.list_aliases()
|
||||
occupied: dict[str, str] = {
|
||||
self._key(known_server): known_server
|
||||
for known_server in self._known_servers
|
||||
if self._key(known_server) != self._key(server)
|
||||
}
|
||||
for row in alias_rows:
|
||||
existing_server = row["server"]
|
||||
if self._key(existing_server) == self._key(server):
|
||||
continue
|
||||
occupied[self._key(existing_server)] = existing_server
|
||||
for alias in row["aliases"]:
|
||||
occupied[self._key(alias)] = existing_server
|
||||
|
||||
for value in [server, *cleaned_aliases]:
|
||||
conflict = occupied.get(self._key(value))
|
||||
if conflict:
|
||||
raise ValueError(f"“{value}”已被区服“{conflict}”使用")
|
||||
|
||||
await self.sql.execute(
|
||||
"""
|
||||
INSERT INTO server_aliases (server, aliases)
|
||||
VALUES (?, ?)
|
||||
ON CONFLICT(server) DO UPDATE SET aliases=excluded.aliases
|
||||
""",
|
||||
(server, json.dumps(cleaned_aliases, ensure_ascii=False)),
|
||||
)
|
||||
await self._reload_cache()
|
||||
|
||||
async def delete_aliases(self, server: str):
|
||||
await self.sql.delete(
|
||||
"server_aliases",
|
||||
"server=?",
|
||||
(self._clean(server),),
|
||||
)
|
||||
await self._reload_cache()
|
||||
|
||||
async def list_aliases(self) -> list[dict[str, Any]]:
|
||||
rows = await self.sql.fetch_all(
|
||||
"SELECT server, aliases FROM server_aliases ORDER BY server"
|
||||
)
|
||||
result = []
|
||||
for row in rows:
|
||||
aliases: list[str] = []
|
||||
try:
|
||||
raw_aliases = json.loads(row.get("aliases") or "[]")
|
||||
if isinstance(raw_aliases, list):
|
||||
aliases = [
|
||||
self._clean(value)
|
||||
for value in raw_aliases
|
||||
if self._clean(value)
|
||||
]
|
||||
except (TypeError, json.JSONDecodeError):
|
||||
aliases = []
|
||||
result.append(
|
||||
{"server": self._clean(row.get("server")), "aliases": aliases}
|
||||
)
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _clean(value: Any) -> str:
|
||||
return str(value or "").strip()
|
||||
|
||||
@classmethod
|
||||
def _key(cls, value: Any) -> str:
|
||||
return cls._clean(value).casefold()
|
||||
Reference in New Issue
Block a user