from __future__ import annotations import re from dataclasses import dataclass, field DmPolicy = str DM_POLICY_PAIRING: DmPolicy = "pairing" DM_POLICY_OPEN: DmPolicy = "open" DM_POLICY_ALLOWLIST: DmPolicy = "allowlist" DM_POLICY_DISABLED: DmPolicy = "disabled" DM_POLICIES = frozenset({DM_POLICY_PAIRING, DM_POLICY_OPEN, DM_POLICY_ALLOWLIST, DM_POLICY_DISABLED}) GroupPolicy = str GROUP_POLICY_OPEN: GroupPolicy = "open" GROUP_POLICY_ALLOWLIST: GroupPolicy = "allowlist" GROUP_POLICY_DISABLED: GroupPolicy = "disabled" GROUP_POLICIES = frozenset({GROUP_POLICY_OPEN, GROUP_POLICY_ALLOWLIST, GROUP_POLICY_DISABLED}) _PREFIX_RE = re.compile(r"^(?:mattermost:|user:|channel:)", re.IGNORECASE) _AT_USERNAME_RE = re.compile(r"^@(\S+)$") def normalize_allow_entry(raw: str) -> str: """规范化白名单条目为纯 Mattermost User ID。 支持格式: - 纯 User ID(如 "abc123") - @username 格式(保留为 username,后续需要查表转换) - user: 或 mattermost: 前缀 """ if not raw or not raw.strip(): return "" entry = raw.strip() entry = _PREFIX_RE.sub("", entry).strip() return entry def is_allow_entry_match(entry: str, user_id: str, username: str = "") -> bool: """检查白名单条目是否匹配给定的用户 ID 或 username。""" if not entry: return False m = _AT_USERNAME_RE.match(entry) if m: return m.group(1).casefold() == username.casefold() return entry.casefold() == user_id.casefold() @dataclass class MattermostSecurityConfig: dm_policy: DmPolicy = DM_POLICY_OPEN group_policy: GroupPolicy = GROUP_POLICY_OPEN allow_from: set[str] = field(default_factory=set) group_allow_from: set[str] = field(default_factory=set) dangerously_allow_name_matching: bool = False @classmethod def from_config(cls, config: dict) -> MattermostSecurityConfig: dm_policy = config.get("dm_policy", DM_POLICY_OPEN) if dm_policy not in DM_POLICIES: dm_policy = DM_POLICY_OPEN group_policy = config.get("group_policy", GROUP_POLICY_OPEN) if group_policy not in GROUP_POLICIES: group_policy = GROUP_POLICY_OPEN allow_from = {normalize_allow_entry(e) for e in config.get("allow_from", []) if normalize_allow_entry(e)} group_allow_from = { normalize_allow_entry(e) for e in config.get("group_allow_from", []) if normalize_allow_entry(e) } dangerously_allow_name_matching = bool(config.get("dangerously_allow_name_matching", False)) return cls( dm_policy=dm_policy, group_policy=group_policy, allow_from=frozenset(allow_from), group_allow_from=frozenset(group_allow_from), dangerously_allow_name_matching=dangerously_allow_name_matching, ) @dataclass class SecurityCheckResult: allowed: bool reason: str = "" class MattermostSecurity: def __init__(self, config: dict | None = None): self._config = MattermostSecurityConfig.from_config(config or {}) self._group_configs: dict[str, dict] = (config or {}).get("groups", {}) @property def dm_policy(self) -> str: return self._config.dm_policy @property def group_policy(self) -> str: return self._config.group_policy @property def allow_from_count(self) -> int: return len(self._config.allow_from) def resolve_group_require_mention(self, channel_id: str) -> bool | None: group_cfg = self._group_configs.get(channel_id, {}) if "requireMention" in group_cfg: return bool(group_cfg["requireMention"]) return None def check_dm(self, user_id: str, username: str = "") -> SecurityCheckResult: """检查 DM 消息是否允许通过。""" policy = self._config.dm_policy if policy == DM_POLICY_DISABLED: return SecurityCheckResult(False, "DM is disabled") if policy == DM_POLICY_OPEN: return SecurityCheckResult(True) if policy == DM_POLICY_ALLOWLIST: return self._check_allowlist(user_id, username, self._config.allow_from) if policy == DM_POLICY_PAIRING: return SecurityCheckResult(True) return SecurityCheckResult(False, f"Unknown DM policy: {policy}") def check_group(self, user_id: str, username: str = "") -> SecurityCheckResult: """检查组消息是否允许通过。""" policy = self._config.group_policy if policy == GROUP_POLICY_DISABLED: return SecurityCheckResult(False, "Group is disabled") if policy == GROUP_POLICY_OPEN: return SecurityCheckResult(True) if policy == GROUP_POLICY_ALLOWLIST: return self._check_allowlist(user_id, username, self._config.group_allow_from) return SecurityCheckResult(False, f"Unknown Group policy: {policy}") def check_inbound(self, chat_type: str, user_id: str, username: str = "") -> SecurityCheckResult: """统一的入站安全检查入口。""" if chat_type == "direct": return self.check_dm(user_id, username) if chat_type in ("group", "channel"): return self.check_group(user_id, username) return SecurityCheckResult(False, f"Unknown chat_type: {chat_type}") def _check_allowlist(self, user_id: str, username: str, allowlist: frozenset[str]) -> SecurityCheckResult: if not allowlist: return SecurityCheckResult(False, "Allowlist is empty — no senders allowed") allow_name_match = self._config.dangerously_allow_name_matching for entry in allowlist: match_username = username if allow_name_match else "" if is_allow_entry_match(entry, user_id, match_username): return SecurityCheckResult(True) return SecurityCheckResult(False, f"User {user_id} not in allowlist") def authorize_command_invocation( self, user_id: str, command: str, channel_id: str = "", ) -> bool: """授权用户调用特定命令。 当前仅允许白名单用户调用管理命令。 """ managed_commands = {"restart", "sudo", "exec", "delete", "config"} if command not in managed_commands: return True allowlist = self._config.allow_from if not allowlist: return False return any(is_allow_entry_match(entry, user_id) for entry in allowlist) def audit_event( self, event_type: str, user_id: str = "", detail: str = "", channel_id: str = "", ) -> None: from yuxi.utils.logging_config import logger logger.info(f"[Mattermost][AUDIT] {event_type} user={user_id} channel={channel_id} detail={detail}")