Files

487 lines
17 KiB
Python
Raw Permalink Normal View History

"""群策略模型与存储(list.json v3)。
策略以「群」为单位,字段全部可选,缺省值见 `Policy`:
auto 自动解析 默认 False
auto_link 自动策略(命中平台即解析) 默认 []
ban_link 禁用策略(命中平台不解析) 默认 []
plan 存储策略 A/B/C 默认 C
upload_public 上传公网 默认 False
send_link 发送下载链接 默认 False(依赖 upload_public)
upload_group_file 上传群文件(并行通道) 默认 False
group_file_platforms 群文件限定平台 默认 [](空 = 全部平台)
文件结构(v3)::
{
"groups": {"<群号>": {<Policy 字段>}},
"default": {<Policy 字段>}, # 私聊 + 群条目缺字段时的兜底
"blacklist": ["<QQ>", ...] # 全局用户黑名单
}
读路径全内存:`verify_user` 在每条消息的热路径上,首次 `load()` 之后不再读盘。
写路径走 `asyncio.Lock` + `.tmp` 原子替换;Web 子应用与 bot 同进程同事件循环,
asyncio 锁即可覆盖两边并发写(Web 不在独立线程里跑写操作)。
换 ORM 时只需再实现一个同样接口的 Store,调用方零改动。
"""
from __future__ import annotations
import asyncio
import json
import os
import re
from collections.abc import Iterable
from dataclasses import dataclass, field, replace
from pathlib import Path
from typing import Any
from nonebot import logger
# ─────────────────────────── 平台 ───────────────────────────
XHS = "小红书"
BILIBILI = "哔哩哔哩"
DOUYIN = "抖音"
YOUTUBE = "Youtube"
X = "X"
#: 规范标签(存储值 = 展示值,命令/Web 输入同一个词)
PLATFORMS: tuple[str, ...] = (XHS, BILIBILI, DOUYIN, YOUTUBE, X)
#: 平台 → 域名片段(判定 auto_link / ban_link 命中用)
PLATFORM_DOMAINS: dict[str, tuple[str, ...]] = {
XHS: ("xiaohongshu.com", "xhslink.com", "xhslink.cn"),
BILIBILI: ("bilibili.com", "b23.tv", "bili2233.cn"),
DOUYIN: (
"douyin.com",
"iesdouyin.com",
"m.douyin.com",
"jingxuan.douyin.com",
),
YOUTUBE: ("youtube.com", "youtu.be"),
X: ("x.com", "twitter.com"),
}
#: 输入别名 → 规范标签(统一小写比较)
_PLATFORM_ALIASES: dict[str, str] = {
"xhs": XHS,
"小红书": XHS,
"bilibili": BILIBILI,
"b23": BILIBILI,
"b站": BILIBILI,
"哔哩哔哩": BILIBILI,
"douyin": DOUYIN,
"抖音": DOUYIN,
"yt": YOUTUBE,
"youtube": YOUTUBE,
"油管": YOUTUBE,
"x": X,
"twitter": X,
"推特": X,
}
#: 分隔符:兼容手写 JSON 里的 "小红书,xhs 抖音" 这类写法
_SPLIT_RE = re.compile(r"[,,、/\s]+")
def normalize_platform(raw: Any) -> str | None:
"""用户输入 → 规范标签,未知返回 None。"""
if not isinstance(raw, str):
return None
return _PLATFORM_ALIASES.get(raw.strip().lower())
def normalize_platforms(values: Any) -> list[str]:
"""归一 + 去重 + 丢弃非法值,结果按 PLATFORMS 顺序排列。
接受字符串(按分隔符拆)或列表。
"""
if values is None:
items: Iterable[Any] = ()
elif isinstance(values, str):
items = _SPLIT_RE.split(values)
elif isinstance(values, (list, tuple, set)):
items = list(values)
else:
return []
picked = {p for p in (normalize_platform(v) for v in items) if p}
return [p for p in PLATFORMS if p in picked]
def apply_platform_diff(current: Iterable[str], tokens: Iterable[str]) -> list[str] | None:
"""按 `+平台` / `-平台` 增删;出现裸平台名时整体覆盖。
全部 token 都带 +/- 时做增量,否则视为覆盖(命令 `视频策略 自动策略 …` 用)。
含未知平台(或覆盖时没有合法平台)返回 None,由调用方提示用法。
"""
tokens = [t for t in tokens if str(t).strip()]
if not tokens:
return None
if all(str(t)[:1] in "+-" for t in tokens):
result = normalize_platforms(current)
for token in tokens:
picked = normalize_platforms([str(token)[1:]])
if not picked:
return None
if str(token)[0] == "+" and picked[0] not in result:
result.append(picked[0])
elif str(token)[0] == "-" and picked[0] in result:
result.remove(picked[0])
return normalize_platforms(result)
picked = normalize_platforms(tokens)
return picked or None
def match_platform(url: str) -> str | None:
"""URL → 平台规范标签;不属于任何已支持平台时返回 None。"""
for name, domains in PLATFORM_DOMAINS.items():
if any(domain in url for domain in domains):
return name
return None
# ─────────────────────────── 策略 ───────────────────────────
PLANS: tuple[str, ...] = ("A", "B", "C")
DEFAULT_PLAN = "C"
POLICY_FIELDS: tuple[str, ...] = (
"auto",
"auto_link",
"ban_link",
"plan",
"upload_public",
"send_link",
"upload_group_file",
"group_file_platforms",
)
@dataclass
class Policy:
"""单个群的策略(字段缺省即默认值,见类文档)。"""
auto: bool = False
auto_link: list[str] = field(default_factory=list)
ban_link: list[str] = field(default_factory=list)
plan: str = DEFAULT_PLAN
upload_public: bool = False
send_link: bool = False
upload_group_file: bool = False
#: 群文件限定平台;空列表 = 所有平台都传群文件
group_file_platforms: list[str] = field(default_factory=list)
def __post_init__(self) -> None:
"""构造即归一:直接 Policy(...) 传入平台别名/非法 plan 也不至于静默失效。"""
self.auto_link = normalize_platforms(self.auto_link)
self.ban_link = normalize_platforms(self.ban_link)
self.group_file_platforms = normalize_platforms(self.group_file_platforms)
if self.plan not in PLANS:
self.plan = DEFAULT_PLAN
@property
def sends_link(self) -> bool:
"""是否真的发下载链接:没有公网链接可发时恒为 False。"""
return self.send_link and self.upload_public
def allows_group_file(self, platform: str | None) -> bool:
"""这个平台的作品要不要传群文件。
开关关着 → 不传;清单为空 → 全部平台都传;否则只认清单里的平台
(平台识别不出来时按"不在清单里"处理)。
"""
if not self.upload_group_file:
return False
if not self.group_file_platforms:
return True
return platform is not None and platform in self.group_file_platforms
def banned(self, url: str) -> str | None:
"""URL 命中的禁用平台(未命中返回 None)。"""
platform = match_platform(url)
return platform if platform and platform in self.ban_link else None
def auto_matched(self, url: str) -> str | None:
"""URL 命中的自动策略平台(未命中返回 None)。"""
platform = match_platform(url)
return platform if platform and platform in self.auto_link else None
def to_dict(self) -> dict[str, Any]:
return {
"auto": self.auto,
"auto_link": list(self.auto_link),
"ban_link": list(self.ban_link),
"plan": self.plan,
"upload_public": self.upload_public,
"send_link": self.send_link,
"upload_group_file": self.upload_group_file,
"group_file_platforms": list(self.group_file_platforms),
}
@classmethod
def from_dict(cls, raw: Any) -> Policy:
"""脏数据归一:非法 plan 回退 C,平台别名转规范标签,多余键丢弃。"""
if not isinstance(raw, dict):
return cls()
plan = str(raw.get("plan", DEFAULT_PLAN)).strip().upper()
return cls(
auto=bool(raw.get("auto", False)),
auto_link=normalize_platforms(raw.get("auto_link")),
ban_link=normalize_platforms(raw.get("ban_link")),
plan=plan if plan in PLANS else DEFAULT_PLAN,
upload_public=bool(raw.get("upload_public", False)),
send_link=bool(raw.get("send_link", False)),
upload_group_file=bool(raw.get("upload_group_file", False)),
group_file_platforms=normalize_platforms(raw.get("group_file_platforms")),
)
# ───────────────────────── 迁移 ─────────────────────────
def _v1_to_v2(raw: dict) -> dict:
"""旧格式(平铺数组)→ group-centric。"""
white = raw.get("WHITE_LIST", [])
auto_list = raw.get("AUTO_ANALYSIS", [])
pa = raw.get("PLANA", [])
pb = raw.get("PLANB", [])
groups: dict[str, dict] = {}
for gid in white:
entry: dict = {}
if gid in auto_list:
entry["auto"] = True
if gid in pa:
entry["plan"] = "A"
elif gid in pb:
entry["plan"] = "B"
groups[str(gid)] = entry
return {"groups": groups, "blacklist": raw.get("BLACK_LIST", [])}
def migrate(raw: Any) -> tuple[dict, bool]:
"""v1 / v2 → v3,返回 (数据, 是否发生迁移)。
v2 → v3 的关键一步:v2 的 `plan=B` 隐含"上传公网",解耦后给它显式补上
`upload_public`,保证已有群行为不变。
"""
if not isinstance(raw, dict):
return _empty_data(), False
changed = False
if "groups" not in raw:
raw = _v1_to_v2(raw)
changed = True
groups: dict[str, dict] = {}
for gid, entry in (raw.get("groups") or {}).items():
if not isinstance(entry, dict):
changed = True
continue
policy = Policy.from_dict(entry)
if policy.plan == "B" and "upload_public" not in entry:
policy.upload_public = True
if entry != policy.to_dict():
changed = True
groups[str(gid)] = policy.to_dict()
if "default" not in raw:
changed = True
data = {
"groups": groups,
"default": Policy.from_dict(raw.get("default")).to_dict(),
"blacklist": [str(x) for x in (raw.get("blacklist") or [])],
}
return data, changed
def _empty_data() -> dict:
return {
"groups": {},
"default": Policy().to_dict(),
"blacklist": [],
}
# ─────────────────────────── 存储 ───────────────────────────
class PolicyStore:
"""list.json 读写:读全内存,写加锁 + 原子替换。"""
def __init__(self, path: Path) -> None:
self.path = Path(path)
self._groups: dict[str, Policy] = {}
self._default = Policy()
self._blacklist: list[str] = []
self._loaded = False
self._lock = asyncio.Lock()
# ── 读 ──────────────────────────────────────────────
def load(self) -> None:
"""首次调用读盘 + 迁移(幂等,之后不再读盘)。"""
if self._loaded:
return
raw: Any = {}
if self.path.exists():
try:
raw = json.loads(self.path.read_text(encoding="utf-8"))
except Exception:
logger.exception(f"读取 {self.path.name} 失败,改用空配置")
data, migrated = migrate(raw)
self._apply(data)
self._loaded = True
if migrated:
self._write_sync(self._snapshot(data))
logger.info(f"{self.path.name} 已迁移为 v3 格式(群 {len(self._groups)} 个)")
def _apply(self, data: dict) -> None:
self._groups = {
str(gid): Policy.from_dict(entry)
for gid, entry in (data.get("groups") or {}).items()
}
self._default = Policy.from_dict(data.get("default"))
self._blacklist = [str(x) for x in (data.get("blacklist") or [])]
def is_whitelisted(self, group_id: Any) -> bool:
"""群是否在白名单里(白名单即 groups 的键,仍是准入门槛)。"""
self.load()
return str(group_id) in self._groups
def get(self, group_id: Any) -> Policy:
"""群策略:未配置的群回落到 default 节。"""
self.load()
return self._groups.get(str(group_id), self._default)
def all_groups(self) -> dict[str, Policy]:
self.load()
return dict(self._groups)
def default_policy(self) -> Policy:
self.load()
return self._default
def blacklist(self) -> list[str]:
self.load()
return list(self._blacklist)
def is_blacklisted(self, user_id: Any) -> bool:
self.load()
return str(user_id) in self._blacklist
# ── 写 ──────────────────────────────────────────────
def _snapshot(self, data: dict | None = None) -> dict:
"""在事件循环线程内构造完整快照,写线程不再触碰共享状态。"""
if data is not None:
return data
return {
"groups": {gid: p.to_dict() for gid, p in self._groups.items()},
"default": self._default.to_dict(),
"blacklist": list(self._blacklist),
}
def _write_sync(self, data: dict) -> None:
tmp = self.path.with_name(self.path.name + ".tmp")
tmp.parent.mkdir(parents=True, exist_ok=True)
tmp.write_text(
json.dumps(data, ensure_ascii=False, indent=2), encoding="utf-8"
)
os.replace(tmp, self.path)
async def _save(self) -> None:
snapshot = self._snapshot()
async with self._lock:
await asyncio.to_thread(self._write_sync, snapshot)
async def update_group(
self, group_id: Any, *, create: bool = False, **changes: Any
) -> Policy | None:
"""更新群策略的部分字段;群不存在且 create=False 时返回 None。
changes 传最终值(列表类字段由调用方算好新列表)。
"""
self.load()
gid = str(group_id)
current = self._groups.get(gid)
if current is None:
if not create:
return None
current = Policy()
valid = {k: v for k, v in changes.items() if k in POLICY_FIELDS}
policy = replace(current, **valid) if valid else current
policy = Policy.from_dict(policy.to_dict()) # 归一化(平台别名/非法 plan)
self._groups[gid] = policy
await self._save()
return policy
async def set_group(self, group_id: Any, policy: Policy) -> Policy:
"""整体覆盖群策略(Web 编辑用)。"""
self.load()
gid = str(group_id)
policy = Policy.from_dict(policy.to_dict())
self._groups[gid] = policy
await self._save()
return policy
async def remove_group(self, group_id: Any) -> bool:
self.load()
gid = str(group_id)
if gid not in self._groups:
return False
del self._groups[gid]
await self._save()
return True
async def set_default(self, policy: Policy) -> Policy:
self.load()
self._default = Policy.from_dict(policy.to_dict())
await self._save()
return self._default
async def add_blacklist(self, user_id: Any) -> bool:
self.load()
uid = str(user_id)
if uid in self._blacklist:
return False
self._blacklist.append(uid)
await self._save()
return True
async def set_blacklist(self, values: Iterable[Any]) -> list[str]:
"""整体替换黑名单(去重保序,一次落盘)。"""
self.load()
cleaned: list[str] = []
for value in values:
uid = str(value).strip()
if uid and uid not in cleaned:
cleaned.append(uid)
self._blacklist = cleaned
await self._save()
return list(self._blacklist)
async def remove_blacklist(self, user_id: Any) -> bool:
self.load()
uid = str(user_id)
if uid not in self._blacklist:
return False
self._blacklist.remove(uid)
await self._save()
return True
#: 单例(插件自己的 data/ 目录,与历史 list.json 同路径,原地迁移)
STORE = PolicyStore(Path(__file__).resolve().parent / "data" / "list.json")