124 lines
3.9 KiB
Python
124 lines
3.9 KiB
Python
|
|
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
|