from __future__ import annotations from dataclasses import dataclass, field from enum import StrEnum from typing import Any class DmPolicy(StrEnum): OPEN = "open" PAIRING = "pairing" ALLOWLIST = "allowlist" DISABLED = "disabled" class GroupPolicy(StrEnum): OPEN = "open" ALLOWLIST = "allowlist" DISABLED = "disabled" DM_POLICY_VALUES = frozenset(p.value for p in DmPolicy) GROUP_POLICY_VALUES = frozenset(p.value for p in GroupPolicy) DEFAULT_DM_POLICY = DmPolicy.ALLOWLIST DEFAULT_GROUP_POLICY = GroupPolicy.ALLOWLIST @dataclass class SecurityConfig: dm_policy: DmPolicy = DEFAULT_DM_POLICY group_policy: GroupPolicy = DEFAULT_GROUP_POLICY require_mention: bool = False allow_from: set[str] = field(default_factory=set) allow_from_wildcard: bool = False @classmethod def from_config(cls, config: dict[str, Any] | None) -> SecurityConfig: if not config: return cls() dm_raw = str(config.get("dm_policy", config.get("DM_POLICY", ""))).strip().lower() dm_policy = DmPolicy(dm_raw) if dm_raw in DM_POLICY_VALUES else DEFAULT_DM_POLICY group_raw = str(config.get("group_policy", config.get("GROUP_POLICY", ""))).strip().lower() group_policy = GroupPolicy(group_raw) if group_raw in GROUP_POLICY_VALUES else DEFAULT_GROUP_POLICY require_mention = bool(config.get("require_mention", config.get("REQUIRE_MENTION", False))) allow_from_raw = config.get("allow_from", config.get("ALLOW_FROM", [])) if isinstance(allow_from_raw, str): allow_from_raw = [x.strip() for x in allow_from_raw.split(",") if x.strip()] elif not isinstance(allow_from_raw, (list, tuple)): allow_from_raw = [] allow_from = set() allow_from_wildcard = False for entry in allow_from_raw: entry = str(entry).strip() if entry == "*": allow_from_wildcard = True elif entry: allow_from.add(entry) return cls( dm_policy=dm_policy, group_policy=group_policy, require_mention=require_mention, allow_from=allow_from, allow_from_wildcard=allow_from_wildcard, ) def is_allowed_user(self, user_id: str) -> bool: if self.allow_from_wildcard: return True if not self.allow_from: return False return user_id in self.allow_from def is_allowed_channel(self, channel_id: str) -> bool: if self.allow_from_wildcard: return True if not self.allow_from: return False return channel_id in self.allow_from def is_allowed_dm(self, user_id: str) -> bool: if self.dm_policy == DmPolicy.OPEN: return True if self.dm_policy == DmPolicy.DISABLED: return False if self.dm_policy == DmPolicy.PAIRING: return True if self.dm_policy == DmPolicy.ALLOWLIST: return self.is_allowed_user(user_id) return False def is_allowed_group(self, channel_id: str) -> bool: if self.group_policy == GroupPolicy.OPEN: return True if self.group_policy == GroupPolicy.DISABLED: return False if self.group_policy == GroupPolicy.ALLOWLIST: return self.is_allowed_channel(channel_id) return False def should_require_mention(self, channel_id: str) -> bool: return self.require_mention def to_dict(self) -> dict[str, Any]: return { "dm_policy": self.dm_policy.value, "group_policy": self.group_policy.value, "require_mention": self.require_mention, "allow_from": sorted(self.allow_from), "allow_from_wildcard": self.allow_from_wildcard, } @dataclass class SecurityDecision: allowed: bool reason: str = "" requires_pairing: bool = False pairing_code: str | None = None