ForcePilot/backend/package/yuxi/channels/adapters/msteams/security.py
Kris 939f1ba82a refactor(msteams): 整理代码结构并新增多项功能
本次提交对Microsoft Teams适配器代码进行了多维度优化与新增:
1.  调整多处导入顺序,优化代码可读性
2.  新增media_tools工具模块,提供媒体相关辅助函数
3.  新增thread_history模块,实现对话历史拉取与缓存功能
4.  新增connection_modes模块,支持webhook/websocket/polling三种连接模式
5.  扩展security.py与tool_policy.py,新增通配符配置校验与三级策略解析
6.  新增feedback会话记录功能
7.  为sent_message_cache添加自动清理任务
8.  优化normalizer模块,新增引用、编辑消息解析与线程上下文注入
9.  重构file_upload的SSRF防护逻辑,复用公共校验工具
10. 修复多处导入顺序与代码排版问题
11. 为消息发送添加断路器保护与异步去重锁
2026-05-13 16:12:31 +08:00

291 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Microsoft Teams 安全策略模块。
DM Policy (open/pairing/allowlist/disabled) 和 Group Policy (open/allowlist/disabled)
支持通配符白名单、名称模糊匹配、Access Groups (Azure AD 组授权)。
"""
from __future__ import annotations
import fnmatch
import os
from typing import TYPE_CHECKING, Any
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 = "*"
_ALLOW_WILDCARD_ENABLED = os.environ.get("MSTEAMS_ALLOW_WILDCARD", "").lower() in ("1", "true", "yes")
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)
if self.dm_policy == DM_POLICY_ALLOWLIST and _ALLOW_WILDCARD in self.allow_from:
logger.warning(
"MSTeams SecurityPolicy: wildcard '*' detected in DM allowlist. "
"All users will bypass DM policy. Consider using dm_policy='open' for explicit intent."
)
if self.group_policy == GROUP_POLICY_ALLOWLIST and _ALLOW_WILDCARD in self.group_allow_from:
logger.warning(
"MSTeams SecurityPolicy: wildcard '*' detected in Group allowlist. "
"All groups/users will bypass Group policy."
)
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:
if _ALLOW_WILDCARD_ENABLED:
logger.warning("MSTeams SecurityPolicy: wildcard matched — bypassing allowlist check")
return True
return False
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)