ForcePilot/backend/package/yuxi/channels/adapters/msteams/security.py

273 lines
11 KiB
Python
Raw Normal View History

"""Microsoft Teams 安全策略模块。
DM Policy (open/pairing/allowlist/disabled) Group Policy (open/allowlist/disabled)
支持通配符白名单名称模糊匹配Access Groups (Azure AD 组授权)
"""
from __future__ import annotations
import fnmatch
from typing import Any, TYPE_CHECKING
from yuxi.utils.logging_config import logger
if TYPE_CHECKING:
pass
DM_POLICY_OPEN = "open"
DM_POLICY_PAIRING = "pairing"
DM_POLICY_ALLOWLIST = "allowlist"
DM_POLICY_DISABLED = "disabled"
GROUP_POLICY_OPEN = "open"
GROUP_POLICY_ALLOWLIST = "allowlist"
GROUP_POLICY_DISABLED = "disabled"
_VALID_DM_POLICIES = {DM_POLICY_OPEN, DM_POLICY_PAIRING, DM_POLICY_ALLOWLIST, DM_POLICY_DISABLED}
_VALID_GROUP_POLICIES = {GROUP_POLICY_OPEN, GROUP_POLICY_ALLOWLIST, GROUP_POLICY_DISABLED}
_ALLOW_WILDCARD = "*"
AccessGroupResolver = "Callable[[str, list[str]], Awaitable[set[str]]]"
def clean_allow_entry(entry: str) -> str:
"""清洗白名单条目:去除空白与下行前缀 (e.g. '28:' or 'did:')。"""
cleaned = entry.strip()
colon_idx = cleaned.find(":")
if colon_idx > 0:
cleaned = cleaned[colon_idx + 1 :]
return cleaned
class SecurityPolicy:
"""MSTeams DM 与群组安全策略控制器。
配置驱动 handle_webhook 入口处进行鉴权过滤
支持 teams.{teamId}.channels.{channelId} 双层嵌套 allowlist
支持 useAccessGroups (Azure AD 组授权)
"""
def __init__(self, config: dict[str, Any] | None = None):
config = config or {}
self.dm_policy = config.get("dm_policy", DM_POLICY_OPEN)
if self.dm_policy not in _VALID_DM_POLICIES:
logger.warning(f"Invalid dm_policy '{self.dm_policy}', falling back to 'open'")
self.dm_policy = DM_POLICY_OPEN
self.group_policy = config.get("group_policy", GROUP_POLICY_OPEN)
if self.group_policy not in _VALID_GROUP_POLICIES:
logger.warning(f"Invalid group_policy '{self.group_policy}', falling back to 'open'")
self.group_policy = GROUP_POLICY_OPEN
self.allow_from = self._normalize_allow_entries(config.get("allow_from", []))
self.group_allow_from = self._normalize_allow_entries(config.get("group_allow_from", []))
self.allow_name_matching = config.get("allow_name_matching", False)
self._teams_config: dict[str, dict[str, Any]] = config.get("teams", {})
self.use_access_groups: bool = config.get("use_access_groups", False)
self.access_group_ids: set[str] = set(config.get("access_group_ids", []))
self._access_group_resolver: Any = None # AccessGroupResolver
self._paired_users: set[str] = set()
@staticmethod
def _normalize_allow_entries(entries: list[str] | str) -> list[str]:
if isinstance(entries, str):
entries = [e.strip() for e in entries.split(",") if e.strip()]
return [clean_allow_entry(e) for e in entries]
def check_dm(self, user_id: str, user_name: str = "") -> bool:
"""检查 DM 会话是否允许该用户。"""
if self.dm_policy == DM_POLICY_OPEN:
return True
if self.dm_policy == DM_POLICY_DISABLED:
logger.debug("MSTeams: DM disabled by policy")
return False
if self.dm_policy == DM_POLICY_PAIRING:
allowed = user_id in self._paired_users
if not allowed:
logger.debug(f"MSTeams: DM rejected (not paired): user={user_id}")
return allowed
if self.dm_policy == DM_POLICY_ALLOWLIST:
return self._check_allowlist(user_id, user_name, self.allow_from)
return False
def check_group(self, user_id: str, user_name: str = "", conversation_id: str = "") -> bool:
"""检查群组会话是否允许该来源。"""
if self.group_policy == GROUP_POLICY_OPEN:
return True
if self.group_policy == GROUP_POLICY_DISABLED:
logger.debug("MSTeams: Group messaging disabled by policy")
return False
if self.group_policy == GROUP_POLICY_ALLOWLIST:
if self._check_allowlist(user_id, user_name, self.group_allow_from):
return True
if self._check_allowlist(conversation_id, "", self.group_allow_from):
return True
logger.debug(f"MSTeams: Group rejected by allowlist: user={user_id}, conv={conversation_id}")
return False
return False
def _check_allowlist(self, user_id: str, user_name: str, allowlist: list[str]) -> bool:
if not allowlist:
return False
if _ALLOW_WILDCARD in allowlist:
return True
for entry in allowlist:
if self._match_entry(user_id, entry):
return True
if self.allow_name_matching and user_name and self._match_entry(user_name, entry):
return True
return False
@staticmethod
def _match_entry(value: str, pattern: str) -> bool:
if pattern == value:
return True
if fnmatch.fnmatch(value, pattern):
return True
return False
def check_require_mention(self, is_mentioned: bool, chat_type: str, config: dict[str, Any] | None = None) -> bool:
"""检查群组消息是否需要 @提及 Bot。
config require_mention 配置优先默认 group 类型需要提及
"""
require = (config or {}).get("require_mention")
if require is not None:
return not require or is_mentioned
if chat_type in ("group",):
return is_mentioned
return True
def resolve_nested_allowlist(
self, team_id: str = "", channel_id: str = "",
) -> tuple[str, str, list[str], list[str]]:
effective_dm = self.dm_policy
effective_group = self.group_policy
effective_allow_from = list(self.allow_from)
effective_group_allow_from = list(self.group_allow_from)
if not self._teams_config:
return effective_dm, effective_group, effective_allow_from, effective_group_allow_from
team_key = team_id or "*"
team_config: dict[str, Any] | None = self._teams_config.get(team_key, self._teams_config.get("*"))
if team_config:
if "dm_policy" in team_config:
effective_dm = team_config["dm_policy"]
if "group_policy" in team_config:
effective_group = team_config["group_policy"]
if "allow_from" in team_config:
effective_allow_from = self._normalize_allow_entries(team_config["allow_from"])
if "group_allow_from" in team_config:
effective_group_allow_from = self._normalize_allow_entries(team_config["group_allow_from"])
channels_config = team_config.get("channels", {})
ch_key = channel_id or "*"
ch_config: dict[str, Any] | None = channels_config.get(ch_key, channels_config.get("*"))
if ch_config:
if "allow_from" in ch_config:
effective_allow_from = self._normalize_allow_entries(ch_config["allow_from"])
if "group_allow_from" in ch_config:
effective_group_allow_from = self._normalize_allow_entries(ch_config["group_allow_from"])
if "group_policy" in ch_config:
effective_group = ch_config["group_policy"]
if "require_mention" in ch_config:
if "config" not in ch_config:
ch_config["config"] = {}
ch_config["config"]["require_mention"] = ch_config["require_mention"]
return effective_dm, effective_group, effective_allow_from, effective_group_allow_from
def set_access_group_resolver(self, resolver: Any) -> None:
"""设置 Azure AD 组成员关系解析器。
resolver(user_id, group_ids) -> set[str]: 返回 user 所属的 group_ids 集合
"""
self._access_group_resolver = resolver
async def check_access_group(self, user_id: str) -> bool:
"""通过 Azure AD 组检查用户授权。
如果 use_access_groups 未启用或无 access_group_ids返回 True (跳过)
否则通过 access_group_resolver 查询用户所属组
仅当用户属于至少一个 access_group_ids 中的组时返回 True
"""
if not self.use_access_groups or not self.access_group_ids:
return True
if not self._access_group_resolver or not user_id:
logger.warning("MSTeams: use_access_groups enabled but no resolver configured")
return False
try:
member_groups = await self._access_group_resolver(user_id, list(self.access_group_ids))
if not member_groups:
logger.debug(f"MSTeams: user {user_id[:12]}... not in any access group")
return False
common = member_groups & self.access_group_ids
if common:
return True
logger.debug(f"MSTeams: user {user_id[:12]}... not in required access groups")
return False
except Exception as e:
logger.error(f"MSTeams: access group check failed for user {user_id[:12]}...: {e}")
return False
def resolve_nested_access_groups(self, team_id: str = "", channel_id: str = "") -> tuple[bool, set[str]]:
"""解析嵌套的 access groups 配置。
支持 teams.{teamId}.channels.{channelId} 层级覆盖
"""
effective_use = self.use_access_groups
effective_ids = set(self.access_group_ids)
if not self._teams_config:
return effective_use, effective_ids
team_key = team_id or "*"
team_config: dict[str, Any] | None = self._teams_config.get(team_key, self._teams_config.get("*"))
if team_config:
if "use_access_groups" in team_config:
effective_use = bool(team_config["use_access_groups"])
if "access_group_ids" in team_config:
effective_ids = set(team_config["access_group_ids"])
channels_config = team_config.get("channels", {})
ch_key = channel_id or "*"
ch_config: dict[str, Any] | None = channels_config.get(ch_key, channels_config.get("*"))
if ch_config:
if "use_access_groups" in ch_config:
effective_use = bool(ch_config["use_access_groups"])
if "access_group_ids" in ch_config:
effective_ids = set(ch_config["access_group_ids"])
return effective_use, effective_ids
def add_paired_user(self, user_id: str) -> None:
self._paired_users.add(user_id)
def remove_paired_user(self, user_id: str) -> None:
self._paired_users.discard(user_id)
@property
def paired_count(self) -> int:
return len(self._paired_users)
@property
def allow_from_count(self) -> int:
return len(self.allow_from)
@property
def group_allow_from_count(self) -> int:
return len(self.group_allow_from)