from __future__ import annotations import fnmatch from dataclasses import dataclass, field from typing import Any, Literal from yuxi.channels.models import ChatType @dataclass class FeishuPolicy: dm_policy: Literal["open", "pairing", "allowlist"] = "pairing" group_policy: Literal["open", "allowlist", "disabled"] = "allowlist" allowlist: list[str] = field(default_factory=list) group_overrides: dict[str, dict[str, Any]] = field(default_factory=dict) dm_overrides: dict[str, dict[str, Any]] = field(default_factory=dict) _dm_prefixes = {"user:", "dm:", "open_id:"} _group_prefixes = {"chat:", "group:", "channel:"} def __post_init__(self): self._dm_exact: set[str] = set() self._dm_patterns: list[str] = [] self._group_exact: set[str] = set() self._group_patterns: list[str] = [] self._has_wildcard = False for entry in self.allowlist: entry = entry.strip() if not entry: continue if entry == "*": self._has_wildcard = True continue if "*" in entry or "?" in entry: self._group_patterns.append(entry) continue prefix, _, bare = _parse_prefix(entry) if prefix in self._group_prefixes: self._group_exact.add(bare) elif prefix in self._dm_prefixes: self._dm_exact.add(bare) else: self._group_exact.add(bare) self._dm_exact.add(bare) def check_dm_access(self, open_id: str) -> bool: if self.dm_policy == "open": return True if self.dm_policy == "pairing": return True if self._has_wildcard: return True return open_id in self._dm_exact def check_group_access(self, chat_id: str) -> bool: if self.group_policy == "open": return True if self.group_policy == "disabled": return False if self._has_wildcard: return True if chat_id in self._group_exact: return True for pattern in self._group_patterns: if fnmatch.fnmatch(chat_id, pattern): return True return False def should_require_mention(self, chat_type: ChatType) -> bool: if chat_type == ChatType.DIRECT: return False if self.group_policy == "open": return False return True def get_group_override(self, chat_id: str) -> dict[str, Any]: override = self.group_overrides.get(chat_id, {}) if override: return override for pattern, cfg in self.group_overrides.items(): if fnmatch.fnmatch(chat_id, pattern): return cfg return {} def get_dm_override(self, open_id: str) -> dict[str, Any]: return self.dm_overrides.get(open_id, {}) def get_effective_group_config(self, chat_id: str) -> dict[str, Any]: override = self.get_group_override(chat_id) if not override: return {} return { "requireMention": override.get("requireMention"), "systemPrompt": override.get("systemPrompt", ""), "tools": override.get("tools", {}), "skills": override.get("skills", {}), "groupPolicy": override.get("groupPolicy", ""), } def get_effective_dm_config(self, open_id: str) -> dict[str, Any]: override = self.get_dm_override(open_id) if not override: return {} return { "enabled": override.get("enabled", True), "systemPrompt": override.get("systemPrompt", ""), } @classmethod def from_config(cls, config: dict[str, Any]) -> FeishuPolicy: group_overrides = {} for entry in config.get("groupConfigOverrides", []): chat_id = entry.get("chatId", entry.get("chat_id", "")) if chat_id: group_overrides[chat_id] = entry dm_overrides = {} for entry in config.get("dmConfigOverrides", []): open_id = entry.get("openId", entry.get("open_id", "")) if open_id: dm_overrides[open_id] = entry group_policy = config.get("group_policy", config.get("groupPolicy", "allowlist")) if group_policy == "allowall": group_policy = "open" return cls( dm_policy=config.get("dm_policy", config.get("dmPolicy", "pairing")), group_policy=group_policy, allowlist=config.get("allowlist", config.get("allowFrom", [])), group_overrides=group_overrides, dm_overrides=dm_overrides, ) def _parse_prefix(entry: str) -> tuple[str, str, str]: for sep in (":",): if sep in entry: prefix, bare = entry.split(sep, 1) return f"{prefix}:", prefix, bare return "", "", entry