from __future__ import annotations import re import secrets import unicodedata from collections.abc import Iterable, Mapping, Sequence from dataclasses import dataclass from datetime import UTC, datetime, timedelta from typing import Any from pyrogram.enums import ChatMemberStatus, MessageEntityType from pyrogram.types import ChatPermissions from wbb import BOT_ID, SUDOERS, app, log from wbb.utils.dbadmin import ( get_managed_chat_settings, update_managed_chat_settings, ) from wbb.utils.dbfunctions import ( add_warn, delete_blacklist_filter, get_blacklisted_words, get_warn, int_to_alpha, remove_warns, save_blacklist_filter, ) RISK_ACTIONS = {"delete", "warn", "mute", "kick", "ban"} MEMBER_ACTIONS = {"warn", "mute", "kick", "ban"} MEMBER_ACTION_PRIORITY = ("ban", "kick", "mute", "warn") WARNING_LIMIT = 3 DEFAULT_BAN_DURATION_SECONDS = 3600 MIN_BAN_DURATION_SECONDS = 60 MAX_BAN_DURATION_SECONDS = 365 * 24 * 60 * 60 MAX_RISK_RULES = 50 LEGACY_POLICY_RULE_ID = "legacy-risk-control" LEGACY_KEYWORD_RULE_ID = "legacy-blacklist" COMMAND_KEYWORD_RULE_ID = "command-blacklist" LEGACY_RULE_IDS = { LEGACY_POLICY_RULE_ID, LEGACY_KEYWORD_RULE_ID, COMMAND_KEYWORD_RULE_ID, } RULE_ID_PATTERN = re.compile(r"^[A-Za-z0-9_-]{1,40}$") LINK_PATTERN = re.compile( r"(?:https?://|ftp://|www\.|t\.me/|telegram\.me/)\S+", flags=re.IGNORECASE, ) class RiskControlValidationError(ValueError): pass @dataclass(frozen=True) class RiskMatch: rule_id: str rule_name: str trigger_types: tuple[str, ...] keyword: str | None = None @dataclass(frozen=True) class RiskEnforcementResult: matches: tuple[RiskMatch, ...] deleted: bool | None requested_member_action: str applied_member_action: str warning_count: int | None = None duration_seconds: int | None = None @property def match(self) -> RiskMatch: return self.matches[0] def _normalize(value: str) -> str: return unicodedata.normalize("NFKC", value).casefold().strip() def _normalized_keywords(words: Iterable[Any]) -> list[str]: values = {_normalize(str(word))[:100] for word in words} return sorted(value for value in values if value)[:200] def _normalize_actions(value: Any, *, strict: bool) -> list[str]: raw_actions = value if value is not None else ["delete"] if isinstance(raw_actions, str): raw_actions = [raw_actions] if not isinstance(raw_actions, (list, tuple, set)): if strict: raise RiskControlValidationError("风控操作格式不正确。") raw_actions = ["delete"] requested_actions = [str(action).strip().lower() for action in raw_actions] invalid_actions = sorted(set(requested_actions) - RISK_ACTIONS) if invalid_actions and strict: raise RiskControlValidationError("包含不支持的风控操作。") member_actions = [ action for action in MEMBER_ACTION_PRIORITY if action in requested_actions ] if len(member_actions) > 1 and strict: raise RiskControlValidationError("警告、禁言、踢出和封禁只能选择一种。") actions: list[str] = [] if "delete" in requested_actions: actions.append("delete") if member_actions: actions.append(member_actions[0]) return actions def _normalize_action_duration(value: Any, *, strict: bool) -> int: if value is None or value == "": return DEFAULT_BAN_DURATION_SECONDS try: duration = int(value) except (TypeError, ValueError): if strict: raise RiskControlValidationError("处罚时长必须是整数秒。") from None return DEFAULT_BAN_DURATION_SECONDS if MIN_BAN_DURATION_SECONDS <= duration <= MAX_BAN_DURATION_SECONDS: return duration if strict: raise RiskControlValidationError("处罚时长需要在 1 分钟到 365 天之间。") return DEFAULT_BAN_DURATION_SECONDS def _normalize_rule( value: Mapping[str, Any], *, fallback_keywords: Iterable[Any] = (), strict: bool = False, default_rule_id: str | None = None, default_name: str | None = None, ) -> dict[str, Any]: raw_rule_id = str(value.get("rule_id") or default_rule_id or "").strip() if raw_rule_id and not RULE_ID_PATTERN.fullmatch(raw_rule_id): if strict: raise RiskControlValidationError("规则 ID 格式不正确。") raw_rule_id = "" rule_id = raw_rule_id or secrets.token_hex(8) name = str(value.get("name") or default_name or "").strip() if not name: if strict: raise RiskControlValidationError("规则名称不能为空。") name = "未命名规则" if len(name) > 60: if strict: raise RiskControlValidationError("规则名称不能超过 60 个字符。") name = name[:60] keywords = _normalized_keywords( value["keywords"] if "keywords" in value else fallback_keywords ) match_images = bool(value.get("match_images", False)) match_links = bool(value.get("match_links", False)) actions = _normalize_actions(value.get("actions"), strict=strict) enabled = bool( value.get("enabled", bool(keywords or match_images or match_links)) ) if strict and enabled and not (keywords or match_images or match_links): raise RiskControlValidationError("启用规则时至少配置一种触发条件。") if strict and enabled and not actions: raise RiskControlValidationError("启用规则时至少选择一种处置操作。") return { "rule_id": rule_id, "name": name, "enabled": enabled, "keywords": keywords, "match_images": match_images, "match_links": match_links, "actions": actions, "duration_seconds": _normalize_action_duration( value.get("duration_seconds", value.get("ban_duration_seconds")), strict=strict, ), } def normalize_risk_control( value: Mapping[str, Any] | None, *, fallback_keywords: Iterable[Any] = (), strict: bool = False, ) -> dict[str, Any]: raw = value or {} rule = _normalize_rule( raw, fallback_keywords=fallback_keywords, strict=strict, default_rule_id=LEGACY_POLICY_RULE_ID, default_name="原风控策略", ) return { key: rule[key] for key in ( "enabled", "keywords", "match_images", "match_links", "actions", "duration_seconds", ) } def normalize_risk_rules( value: Sequence[Mapping[str, Any]] | None, *, legacy_policy: Mapping[str, Any] | None = None, fallback_keywords: Iterable[Any] = (), strict: bool = False, ) -> list[dict[str, Any]]: fallback = _normalized_keywords(fallback_keywords) if value is None: if legacy_policy is not None: legacy = dict(legacy_policy) legacy["keywords"] = _normalized_keywords( [*(legacy.get("keywords") or []), *fallback] ) return [ _normalize_rule( legacy, strict=strict, default_rule_id=LEGACY_POLICY_RULE_ID, default_name="原风控策略", ) ] if fallback: return [ _normalize_rule( { "enabled": True, "keywords": fallback, "actions": ["delete"], }, strict=strict, default_rule_id=LEGACY_KEYWORD_RULE_ID, default_name="原关键词黑名单", ) ] return [] if isinstance(value, (str, bytes)) or not isinstance(value, Sequence): if strict: raise RiskControlValidationError("风控规则必须是列表。") return [] if len(value) > MAX_RISK_RULES and strict: raise RiskControlValidationError(f"每个群最多配置 {MAX_RISK_RULES} 条风控规则。") rules: list[dict[str, Any]] = [] rule_ids: set[str] = set() for index, item in enumerate(value[:MAX_RISK_RULES]): if not isinstance(item, Mapping): if strict: raise RiskControlValidationError("风控规则格式不正确。") continue rule = _normalize_rule( item, strict=strict, default_name=f"规则 {index + 1}", ) if rule["rule_id"] in rule_ids: if strict: raise RiskControlValidationError("规则 ID 不能重复。") rule["rule_id"] = f"{rule['rule_id'][:30]}-{index + 1}" rule_ids.add(rule["rule_id"]) rules.append(rule) return rules def risk_rule_keywords( rules: Iterable[Mapping[str, Any]], *, enabled_only: bool = False ) -> list[str]: return _normalized_keywords( keyword for rule in rules if not enabled_only or bool(rule.get("enabled")) for keyword in rule.get("keywords", []) ) def find_blacklist_match(content: str, words: Iterable[str]) -> str | None: normalized_content = _normalize(content) if not normalized_content: return None for word in words: normalized_word = _normalize(str(word)) if normalized_word and normalized_word in normalized_content: return str(word).strip() return None def _is_link_entity(entity: Any) -> bool: entity_type = getattr(entity, "type", None) if entity_type in {MessageEntityType.URL, MessageEntityType.TEXT_LINK}: return True label = str(entity_type).rsplit(".", 1)[-1].lower() return label in {"url", "text_link"} def _message_has_link(message: Any, content: str) -> bool: for name in ("entities", "caption_entities"): if any(_is_link_entity(entity) for entity in (getattr(message, name, None) or [])): return True return bool(LINK_PATTERN.search(content)) def _message_has_image(message: Any) -> bool: if getattr(message, "photo", None) or getattr(message, "animation", None): return True document = getattr(message, "document", None) return bool( document and str(getattr(document, "mime_type", "")).lower().startswith("image/") ) def detect_risk_match( message: Any, policy: Mapping[str, Any] ) -> RiskMatch | None: if not policy.get("enabled"): return None content = str( getattr(message, "text", None) or getattr(message, "caption", None) or "" ) keyword = find_blacklist_match(content, policy.get("keywords", [])) triggers: list[str] = [] if keyword: triggers.append("keyword") if policy.get("match_images") and _message_has_image(message): triggers.append("image") if policy.get("match_links") and _message_has_link(message, content): triggers.append("link") if not triggers: return None return RiskMatch( rule_id=str(policy.get("rule_id") or LEGACY_POLICY_RULE_ID), rule_name=str(policy.get("name") or "风控规则"), trigger_types=tuple(triggers), keyword=keyword, ) def detect_risk_matches( message: Any, rules: Iterable[Mapping[str, Any]] ) -> list[tuple[Mapping[str, Any], RiskMatch]]: matches: list[tuple[Mapping[str, Any], RiskMatch]] = [] for rule in rules: match = detect_risk_match(message, rule) if match: matches.append((rule, match)) return matches async def get_effective_risk_rules(chat_id: int) -> list[dict[str, Any]]: settings = await get_managed_chat_settings(chat_id) legacy_words = await get_blacklisted_words(chat_id) if "risk_rules" in settings: return normalize_risk_rules(settings.get("risk_rules")) return normalize_risk_rules( None, legacy_policy=settings.get("risk_control"), fallback_keywords=legacy_words or settings.get("blacklist_words", []), ) async def get_effective_risk_control(chat_id: int) -> dict[str, Any]: rules = await get_effective_risk_rules(chat_id) active = [rule for rule in rules if rule["enabled"]] actions = {action for rule in active for action in rule["actions"]} member_action = next( (action for action in MEMBER_ACTION_PRIORITY if action in actions), None ) normalized_actions = ["delete"] if "delete" in actions else [] if member_action: normalized_actions.append(member_action) return { "enabled": bool(active), "keywords": risk_rule_keywords(active), "match_images": any(rule["match_images"] for rule in active), "match_links": any(rule["match_links"] for rule in active), "actions": normalized_actions, "duration_seconds": max( (int(rule["duration_seconds"]) for rule in active), default=DEFAULT_BAN_DURATION_SECONDS, ), } async def sync_legacy_blacklist(chat_id: int, keywords: Iterable[Any]) -> list[str]: normalized = _normalized_keywords(keywords) current = await get_blacklisted_words(chat_id) current_by_normalized = {_normalize(str(word)): str(word) for word in current} target = set(normalized) for normalized_word, stored_word in current_by_normalized.items(): if normalized_word not in target: while await delete_blacklist_filter(chat_id, stored_word): pass current_normalized = { _normalize(str(word)) for word in await get_blacklisted_words(chat_id) } for word in normalized: if word not in current_normalized: await save_blacklist_filter(chat_id, word) return normalized def _keyword_rule_index(rules: list[dict[str, Any]]) -> int | None: for index, rule in enumerate(rules): if rule["rule_id"] in LEGACY_RULE_IDS: return index return None def replace_legacy_rule_keywords( rules: Sequence[Mapping[str, Any]], keywords: Iterable[Any] ) -> list[dict[str, Any]]: normalized_rules = normalize_risk_rules(rules) index = _keyword_rule_index(normalized_rules) if index is None: normalized_rules.append( _normalize_rule( { "enabled": True, "keywords": [], "actions": ["delete"], }, default_rule_id=COMMAND_KEYWORD_RULE_ID, default_name="命令关键词黑名单", ) ) index = len(normalized_rules) - 1 rule = normalized_rules[index] rule["keywords"] = _normalized_keywords(keywords) if rule["keywords"]: rule["enabled"] = True elif not rule["match_images"] and not rule["match_links"]: rule["enabled"] = False return normalized_rules async def _store_risk_rules(chat_id: int, rules: list[dict[str, Any]]) -> None: keywords = await sync_legacy_blacklist( chat_id, risk_rule_keywords(rules, enabled_only=True) ) await update_managed_chat_settings( chat_id, {"risk_rules": rules, "blacklist_words": keywords}, ) async def add_risk_keyword(chat_id: int, word: str) -> tuple[bool, dict[str, Any]]: normalized_word = _normalize(word)[:100] if not normalized_word: raise RiskControlValidationError("关键词不能为空。") rules = await get_effective_risk_rules(chat_id) index = _keyword_rule_index(rules) existing_keywords = rules[index]["keywords"] if index is not None else [] rules = replace_legacy_rule_keywords( rules, [*existing_keywords, normalized_word] ) index = _keyword_rule_index(rules) assert index is not None rule = rules[index] changed = normalized_word not in existing_keywords await _store_risk_rules(chat_id, rules) return changed, rule async def remove_risk_keyword(chat_id: int, word: str) -> tuple[bool, dict[str, Any]]: normalized_word = _normalize(word)[:100] rules = await get_effective_risk_rules(chat_id) index = _keyword_rule_index(rules) if index is None: return False, normalize_risk_control(None) existing_keywords = rules[index]["keywords"] filtered = [ item for item in existing_keywords if _normalize(item) != normalized_word ] changed = len(filtered) != len(existing_keywords) rules = replace_legacy_rule_keywords(rules, filtered) index = _keyword_rule_index(rules) assert index is not None rule = rules[index] await _store_risk_rules(chat_id, rules) return changed, rule async def _is_privileged_member(chat_id: int, user_id: int) -> bool: if user_id in SUDOERS or user_id == BOT_ID: return True try: member = await app.get_chat_member(chat_id, user_id) except Exception as exc: log.error( f"无法确认群 {chat_id} 成员 {user_id} 的管理员身份,已跳过自动处罚:{exc}" ) return True return member.status in { ChatMemberStatus.OWNER, ChatMemberStatus.ADMINISTRATOR, } async def _bot_capabilities(chat_id: int) -> dict[str, bool | None]: try: member = await app.get_chat_member(chat_id, BOT_ID) except Exception as exc: log.error(f"无法读取机器人在群 {chat_id} 的实时权限,将直接尝试执行:{exc}") return {"delete": None, "restrict": None} if member.status == ChatMemberStatus.OWNER: return {"delete": True, "restrict": True} if member.status != ChatMemberStatus.ADMINISTRATOR: return {"delete": False, "restrict": False} privileges = getattr(member, "privileges", None) return { "delete": bool(getattr(privileges, "can_delete_messages", False)), "restrict": bool(getattr(privileges, "can_restrict_members", False)), } def _ban_until(duration_seconds: int) -> datetime: return datetime.now(UTC) + timedelta(seconds=duration_seconds) async def _apply_warning( chat_id: int, user_id: int, *, can_restrict: bool | None, duration_seconds: int, ) -> tuple[str, int]: key = await int_to_alpha(user_id) current = await get_warn(chat_id, key) count = int(current.get("warns", 0)) + 1 if current else 1 if count < WARNING_LIMIT: await add_warn(chat_id, key, {"warns": count}) return "warn", count if can_restrict is False: await add_warn(chat_id, key, {"warns": count}) log.error( f"群 {chat_id} 成员 {user_id} 已达到警告上限,但机器人没有封禁权限" ) return "warn", count try: await app.ban_chat_member( chat_id, user_id, until_date=_ban_until(duration_seconds), ) except Exception as exc: await add_warn(chat_id, key, {"warns": count}) log.error(f"群 {chat_id} 成员 {user_id} 已达到警告上限,但封禁失败:{exc}") return "warn", count await remove_warns(chat_id, key) return "ban", count async def _apply_member_action( chat_id: int, user_id: int, action: str, *, can_restrict: bool | None, duration_seconds: int, ) -> tuple[str, int | None]: if action == "warn": return await _apply_warning( chat_id, user_id, can_restrict=can_restrict, duration_seconds=duration_seconds, ) if can_restrict is False: log.error(f"群 {chat_id} 无法执行风控{action}:机器人没有限制成员权限") return "none", None try: if action == "kick": await app.ban_chat_member(chat_id, user_id) await app.unban_chat_member(chat_id, user_id) elif action == "mute": await app.restrict_chat_member( chat_id, user_id, ChatPermissions(), until_date=_ban_until(duration_seconds), ) elif action == "ban": await app.ban_chat_member( chat_id, user_id, until_date=_ban_until(duration_seconds), ) else: return "none", None except Exception as exc: log.error(f"群 {chat_id} 风控处置成员 {user_id} 失败:{exc}") return "none", None return action, None def _trigger_description(match: RiskMatch) -> str: labels = {"keyword": "关键词", "image": "图片", "link": "链接"} values = [labels[item] for item in match.trigger_types] if match.keyword: values[values.index("关键词")] = f"关键词“{match.keyword}”" return "、".join(values) def _duration_description(duration_seconds: int) -> str: if duration_seconds % 86400 == 0: return f"{duration_seconds // 86400} 天" if duration_seconds % 3600 == 0: return f"{duration_seconds // 3600} 小时" return f"{max(1, duration_seconds // 60)} 分钟" async def _send_action_notice( chat_id: int, user: Any, matches: tuple[RiskMatch, ...], action: str, warning_count: int | None, duration_seconds: int, ) -> None: mention = getattr(user, "mention", None) or f"成员 {user.id}" if action == "warn": outcome = f"已警告({warning_count}/{WARNING_LIMIT})" elif action == "ban" and warning_count: outcome = ( f"累计 {WARNING_LIMIT} 次警告,已封禁 " f"{_duration_description(duration_seconds)}" ) elif action == "ban": outcome = f"已封禁 {_duration_description(duration_seconds)}" elif action == "mute": outcome = f"已禁言 {_duration_description(duration_seconds)}" else: outcome = {"kick": "已踢出"}.get(action, action) matched = ";".join( f"{match.rule_name}:{_trigger_description(match)}" for match in matches ) try: await app.send_message( chat_id, f"{mention} 触发风控规则({matched}),{outcome}。", ) except Exception as exc: log.error(f"群 {chat_id} 风控处置通知发送失败:{exc}") async def enforce_risk_message( message: Any, policy: Mapping[str, Any] | Sequence[Mapping[str, Any]] | None = None, ) -> RiskEnforcementResult | None: user = getattr(message, "from_user", None) if user is not None and getattr(user, "is_bot", False): return None chat_id = int(message.chat.id) if policy is None: rules = await get_effective_risk_rules(chat_id) elif isinstance(policy, Mapping): rules = normalize_risk_rules(None, legacy_policy=policy) else: rules = normalize_risk_rules(policy) matched_rules = detect_risk_matches(message, rules) if not matched_rules: return None matches = tuple(match for _, match in matched_rules) delete_requested = any("delete" in rule["actions"] for rule, _ in matched_rules) member_action = "none" action_rules: list[Mapping[str, Any]] = [] for candidate in MEMBER_ACTION_PRIORITY: action_rules = [ rule for rule, _ in matched_rules if candidate in rule["actions"] ] if action_rules: member_action = candidate break duration_seconds = max( (int(rule["duration_seconds"]) for rule in action_rules), default=DEFAULT_BAN_DURATION_SECONDS, ) capabilities = await _bot_capabilities(chat_id) deleted: bool | None = None if delete_requested: deleted = False if capabilities["delete"] is False: log.error(f"群 {chat_id} 命中风控规则,但机器人没有删除消息权限") else: try: await message.delete() deleted = True except Exception as exc: log.error(f"群 {chat_id} 命中风控规则但删除消息失败:{exc}") applied_action = "none" warning_count: int | None = None if ( member_action != "none" and user is not None and not await _is_privileged_member(chat_id, int(user.id)) ): applied_action, warning_count = await _apply_member_action( chat_id, int(user.id), member_action, can_restrict=capabilities["restrict"], duration_seconds=duration_seconds, ) if applied_action != "none": await _send_action_notice( chat_id, user, matches, applied_action, warning_count, duration_seconds, ) return RiskEnforcementResult( matches=matches, deleted=deleted, requested_member_action=member_action, applied_member_action=applied_action, warning_count=warning_count, duration_seconds=( duration_seconds if member_action in {"ban", "mute", "warn"} else None ), ) async def enforce_blacklist_message( message: Any, words: Iterable[str] ) -> RiskEnforcementResult | None: keyword_list = list(words) return await enforce_risk_message( message, [ { "rule_id": LEGACY_KEYWORD_RULE_ID, "name": "关键词黑名单", "enabled": bool(keyword_list), "keywords": keyword_list, "match_images": False, "match_links": False, "actions": ["delete"], "duration_seconds": DEFAULT_BAN_DURATION_SECONDS, } ], )