99 lines
3.2 KiB
Python
99 lines
3.2 KiB
Python
"""多渠道网关白名单检查。"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from typing import TYPE_CHECKING
|
|
|
|
from yuxi.channel.constants import InboundRejectionReason
|
|
from yuxi.channel.security.models import DmPolicy, GroupPolicy
|
|
|
|
from .policy import SecurityCheckResult
|
|
|
|
if TYPE_CHECKING:
|
|
from .registry import SecurityContext
|
|
|
|
|
|
class AllowlistChecker:
|
|
name = "allowlist"
|
|
default_priority = 100
|
|
|
|
async def check(self, ctx: SecurityContext) -> SecurityCheckResult | None:
|
|
"""SecurityChecker 协议入口。"""
|
|
policy = self._resolve_policy(ctx)
|
|
allowed = self.is_allowed(policy, ctx.inbound, ctx.resolved_sender_id)
|
|
if allowed:
|
|
return None
|
|
|
|
reason = (
|
|
InboundRejectionReason.DM_NOT_ALLOWED
|
|
if ctx.chat_type == "private"
|
|
else InboundRejectionReason.GROUP_NOT_ALLOWED
|
|
)
|
|
return SecurityCheckResult(allowed=False, reason=reason)
|
|
|
|
def _resolve_policy(self, ctx: SecurityContext) -> DmPolicy | GroupPolicy | None:
|
|
if ctx.chat_type == "private":
|
|
resolve = getattr(ctx.plugin, "resolve_dm_policy", None)
|
|
else:
|
|
resolve = getattr(ctx.plugin, "resolve_group_policy", None)
|
|
if resolve is not None:
|
|
return resolve(ctx.config, ctx.inbound.account_id)
|
|
return None
|
|
|
|
def is_allowed(
|
|
self,
|
|
policy: DmPolicy | GroupPolicy | None,
|
|
inbound: object,
|
|
resolved_sender_id: str | None = None,
|
|
) -> bool:
|
|
"""原有 check(policy, inbound, resolved_sender_id) 逻辑,改名以避免覆盖协议方法。"""
|
|
if policy is None:
|
|
return True
|
|
|
|
if isinstance(policy, DmPolicy):
|
|
if policy.mode == "open":
|
|
return True
|
|
if policy.mode == "deny":
|
|
return False
|
|
if policy.mode == "allow_from":
|
|
sender_id = resolved_sender_id
|
|
if sender_id is None:
|
|
sender_id = getattr(inbound, "sender_id", None)
|
|
return self._match_allowlist(
|
|
policy.allow_list,
|
|
sender_id,
|
|
getattr(inbound, "channel_type", ""),
|
|
)
|
|
return True
|
|
|
|
if isinstance(policy, GroupPolicy):
|
|
chat_type = getattr(inbound, "chat_type", None)
|
|
if chat_type != "private":
|
|
session_key = getattr(inbound, "session_key", None) or ""
|
|
if policy.deny_groups and session_key in policy.deny_groups:
|
|
return False
|
|
if policy.require_mention and not getattr(inbound, "is_at_bot", False):
|
|
return False
|
|
if policy.allow_groups:
|
|
return session_key in policy.allow_groups
|
|
return True
|
|
|
|
return True
|
|
|
|
def _match_allowlist(
|
|
self,
|
|
allow_list: list[str],
|
|
sender_id: str | None,
|
|
channel_type: str,
|
|
) -> bool:
|
|
if not sender_id:
|
|
return False
|
|
for entry in allow_list:
|
|
if entry == "*":
|
|
return True
|
|
if entry == sender_id:
|
|
return True
|
|
if entry == f"{channel_type}:{sender_id}":
|
|
return True
|
|
return False
|