Files
PolyWeather/src/bot/command_guard.py
T

102 lines
4.0 KiB
Python

from __future__ import annotations
import os
from typing import Any
from loguru import logger
from src.bot.io_layer import BotIOLayer
from src.utils.telegram_chat_ids import get_telegram_chat_ids_from_env, parse_telegram_chat_ids
class CommandGuard:
"""Unified command gate: group membership + points charge."""
_ALLOWED_MEMBER_STATUSES = {"creator", "administrator", "member", "restricted"}
def __init__(self, io_layer: BotIOLayer, group_chat_id: str | None = None):
self.io_layer = io_layer
if group_chat_id is None:
self.group_chat_ids = get_telegram_chat_ids_from_env()
else:
self.group_chat_ids = parse_telegram_chat_ids(group_chat_id)
self.group_invite_url = str(os.getenv("POLYWEATHER_BOT_GROUP_INVITE_URL") or "").strip()
def _reply_group_required(self, message: Any) -> None:
lines = ["🔒 该指令仅限群成员使用,请先加入官方群。"]
if self.group_invite_url:
lines.append(f"👉 入群链接: {self.group_invite_url}")
self.io_layer.bot.reply_to(message, "\n".join(lines))
def ensure_group_member(self, message: Any, command_label: str) -> bool:
user = getattr(message, "from_user", None)
user_id = getattr(user, "id", None)
if user_id is None:
return False
if not self.group_chat_ids:
logger.error(
"group member blocked command={} user_id={} reason=missing_TELEGRAM_CHAT_IDS",
command_label,
user_id,
)
self.io_layer.bot.reply_to(
message,
"⚠️ 机器人未配置群组准入(TELEGRAM_CHAT_IDS / TELEGRAM_CHAT_ID),请联系管理员。",
)
return False
blocked_statuses: list[str] = []
lookup_failures: list[str] = []
for chat_id in self.group_chat_ids:
try:
member = self.io_layer.bot.get_chat_member(chat_id, int(user_id))
status = str(getattr(member, "status", "") or "").strip().lower()
except Exception as exc:
lookup_failures.append(f"{chat_id}:{type(exc).__name__}")
continue
if status in self._ALLOWED_MEMBER_STATUSES:
return True
blocked_statuses.append(f"{chat_id}:{status or 'unknown'}")
if blocked_statuses:
logger.info(
"group member blocked command={} user_id={} chat_ids={} statuses={}",
command_label,
user_id,
",".join(self.group_chat_ids),
"|".join(blocked_statuses),
)
else:
logger.info(
"group member blocked command={} user_id={} chat_ids={} reason=lookup_failed_all errors={}",
command_label,
user_id,
",".join(self.group_chat_ids),
"|".join(lookup_failures) or "unknown",
)
self._reply_group_required(message)
return False
def check_daily_query_limit(self, message: Any, query_type: str, limit: int) -> bool:
user = message.from_user
result = self.io_layer.db.track_query_usage(user.id, query_type)
if result.get("allowed"):
return True
if result.get("reason") == "user_missing":
self.io_layer.bot.reply_to(message, "⚠️ 用户未注册,请先在群内发言。")
return False
if result.get("reason") == "daily_limit":
used = int(result.get("used") or 0)
self.io_layer.send_query_message(
message,
f"❌ 今日 <b>/{query_type}</b> 免费次数已用完 ({used}/{limit})\n请明天再来,或通过群内发言赚取积分。",
parse_mode="HTML",
)
return False
return True
def ensure_access_and_points(self, message: Any, cost: int, command_label: str) -> bool:
return self.io_layer.ensure_query_points(message, cost, command_label)