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

161 lines
5.6 KiB
Python
Raw Normal View History

from __future__ import annotations
from typing import Any
def check_dm_policy(config: dict[str, Any], open_id: str) -> bool:
dm_policy = config.get("dm_policy", "open")
if dm_policy == "open":
return True
elif dm_policy == "disabled":
return False
elif dm_policy == "allowlist":
allow_from = config.get("allow_from", [])
return f"wx:{open_id}" in allow_from
return False
def check_group_policy(config: dict[str, Any], chat_id: str, open_id: str) -> bool:
group_policy = config.get("group_policy", "allowlist")
if group_policy == "open":
return True
elif group_policy == "disabled":
return False
elif group_policy == "allowlist":
global_allow = config.get("group_allow_from", [])
if f"wx:{open_id}" in global_allow:
return True
groups_config = config.get("groups", {})
chat_config = groups_config.get(chat_id, {})
per_group_allow = chat_config.get("allow_from", [])
return f"wx:{open_id}" in per_group_allow
return False
def is_bot_mentioned(config: dict[str, Any], payload: dict, mode: str) -> bool:
if mode == "wecom":
content = payload.get("Content", "")
agent_name = config.get("agent_name", "")
if agent_name and ("@" in content and agent_name in content):
return True
return "@" in content
if mode == "mp":
content = payload.get("Content", "")
return "@" in content
if mode == "personal":
at_list = payload.get("at_list", [])
return len(at_list) > 0
return False
def should_require_mention(config: dict[str, Any], chat_id: str) -> bool:
groups_config = config.get("groups", {})
chat_config = groups_config.get(chat_id, {})
require_mention = chat_config.get("require_mention", True)
return require_mention
WECHAT_SAFETY_STRATEGIES: dict[str, dict[str, Any]] = {
"rateLimit": {
"description": "API 调用频率控制",
"implementation": "TokenBucketRateLimiter(rate=20, per=60)",
},
"banned": {
"description": "封禁检测与处理",
"implementation": "检测微信错误码 48001API未授权/ 违规记录 → 自动禁用适配器",
},
"autoJoin": {
"description": "自动加入群组策略",
"implementation": "config.auto_join_groups 白名单 → connect() 时自动加入",
},
"autoDerive": {
"description": "自动派生会话策略",
"implementation": "群聊消息按 sender 自动派生子会话(可配置开关)",
},
}
class WeChatSecurityAdapter:
@staticmethod
def resolve_dm_policy(config: dict[str, Any], account_id: str = "default") -> str:
accounts = config.get("accounts", {})
if account_id in accounts:
return accounts[account_id].get("dm_policy", "pairing")
return config.get("dm_policy", "pairing")
@staticmethod
def resolve_dm_allow_from(config: dict[str, Any], account_id: str = "default") -> list[str]:
accounts = config.get("accounts", {})
if account_id in accounts:
return accounts[account_id].get("allow_from", [])
return config.get("allow_from", [])
@staticmethod
def resolve_group_policy(config: dict[str, Any], account_id: str = "default") -> str:
accounts = config.get("accounts", {})
if account_id in accounts:
return accounts[account_id].get("group_policy", "allowlist")
return config.get("group_policy", "allowlist")
@staticmethod
def resolve_group_allow_from(config: dict[str, Any], chat_id: str) -> list[str]:
global_allow = config.get("group_allow_from", [])
groups_config = config.get("groups", {})
chat_config = groups_config.get(chat_id, {})
per_group_allow = chat_config.get("allow_from", [])
result = list(global_allow)
for entry in per_group_allow:
if entry not in result:
result.append(entry)
return result
@staticmethod
def collect_warnings(config: dict[str, Any]) -> list[str]:
warnings: list[str] = []
dm_policy = config.get("dm_policy", "pairing")
if dm_policy == "open":
warnings.append("dm_policy='open' 存在高危安全风险,建议启用配对机制")
allow_from = config.get("allow_from", [])
group_allow = config.get("group_allow_from", [])
if dm_policy == "allowlist" and not allow_from:
warnings.append("allowlist 为空,无人可访问")
if not group_allow:
groups_config = config.get("groups", {})
has_per_group = any(chat_config.get("allow_from") for chat_config in groups_config.values())
if not has_per_group and config.get("group_policy") == "allowlist":
warnings.append("群组 allowlist 为空,所有群组将被拒绝访问")
has_pairing = any(acc.get("dm_policy") == "pairing" for acc in config.get("accounts", {}).values())
if dm_policy != "pairing" and not has_pairing:
warnings.append("未启用配对机制 (pairing),建议开启以增强安全性")
return warnings
@staticmethod
def collect_audit_findings(config: dict[str, Any]) -> list[dict[str, Any]]:
findings: list[dict[str, Any]] = []
allow_from = config.get("allow_from", [])
if len(allow_from) > 1000:
findings.append(
{
"severity": "warning",
"category": "allowlist_size",
"message": f"allowlist 规模异常({len(allow_from)} 条目 > 1000",
}
)
return findings