diff --git a/alembic/versions/limit_policy_global_bundle.py b/alembic/versions/limit_policy_global_bundle.py new file mode 100644 index 0000000..d90942b --- /dev/null +++ b/alembic/versions/limit_policy_global_bundle.py @@ -0,0 +1,193 @@ +"""store all global limit values in one complete JSON document + +Revision ID: limit_policy_global_bundle +Revises: limit_policy_whitelist +""" +from __future__ import annotations + +from collections.abc import Sequence +from typing import Any + +import sqlalchemy as sa +from sqlalchemy.dialects import postgresql + +from alembic import op + +revision: str = "limit_policy_global_bundle" +down_revision: str | Sequence[str] | None = "limit_policy_whitelist" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + +_JSON = sa.JSON().with_variant(postgresql.JSONB(), "postgresql") +_BUNDLE_KEY = "limit_policy_global" + +# rule_code, old sparse key, default, legacy structured key, legacy JSON field +_RULES: tuple[tuple[str, str, int, str | None, str | None], ...] = ( + ("compare.start.daily", "compare_daily_limit", 100, None, None), + ("sms.send.hourly", "sms_send_hourly_limit", 5, None, None), + ("sms.send.daily", "sms_send_daily_limit", 20, None, None), + ("sms.phone.cooldown", "sms_phone_cooldown_seconds", 60, None, None), + ("sms.code.failed_attempts", "sms_code_max_failed_attempts", 5, None, None), + ("sms.login.hourly", "sms_login_hourly_limit", 5, None, None), + ("wechat.bind.hourly", "wechat_bind_sms_hourly_limit", 5, None, None), + ("wechat.conflict.hourly", "wechat_conflict_hourly_limit", 5, None, None), + ( + "ad.reward_video.daily", + "ad_reward_video_daily_limit", + 500, + "ad_daily_limit", + None, + ), + ( + "ad.feed.daily", + "ad_feed_daily_limit", + 500, + "ad_daily_limit", + None, + ), + ("ad.reward_video.cooldown", "ad_cooldown_sec", 3, None, None), + ( + "guide.video.lifetime", + "guide_video_max_plays", + 3, + "coupon_guide_video", + "max_plays", + ), + ("phone.rebind.days", "phone_rebind_days", 30, None, None), + ("risk.sms.hourly", "risk_sms_hourly_threshold", 5, None, None), + ( + "risk.oneclick.daily", + "risk_oneclick_daily_threshold", + 20, + None, + None, + ), + ( + "risk.compare.daily", + "risk_compare_daily_threshold", + 100, + None, + None, + ), +) + + +def _table() -> sa.TableClause: + return sa.table( + "app_config", + sa.column("key", sa.String(64)), + sa.column("value", _JSON), + sa.column("updated_by_admin_id", sa.Integer), + sa.column("updated_at", sa.DateTime(timezone=True)), + ) + + +def _row(conn, table, key: str): + return conn.execute( + sa.select( + table.c.value, + table.c.updated_by_admin_id, + ).where(table.c.key == key) + ).mappings().first() + + +def _int_or_none(value: Any) -> int | None: + if isinstance(value, bool): + return None + try: + return int(value) + except (TypeError, ValueError): + return None + + +def upgrade() -> None: + conn = op.get_bind() + table = _table() + values = {rule_code: default for rule_code, _, default, _, _ in _RULES} + + # Old shared/structured values are the lowest-precedence compatibility + # source. Dedicated sparse keys override them. + for rule_code, _, _, legacy_key, legacy_field in _RULES: + if legacy_key is None: + continue + legacy = _row(conn, table, legacy_key) + if legacy is None: + continue + raw = legacy["value"] + if legacy_field is not None: + raw = raw.get(legacy_field) if isinstance(raw, dict) else None + parsed = _int_or_none(raw) + if parsed is not None: + values[rule_code] = parsed + + for rule_code, sparse_key, _, _, _ in _RULES: + sparse = _row(conn, table, sparse_key) + parsed = _int_or_none(sparse["value"]) if sparse is not None else None + if parsed is not None: + values[rule_code] = parsed + + # If a deployment already wrote the new key, preserve it over old keys. + bundle = _row(conn, table, _BUNDLE_KEY) + if bundle is not None and isinstance(bundle["value"], dict): + for rule_code, raw in bundle["value"].items(): + if rule_code not in values: + continue + parsed = _int_or_none(raw) + if parsed is not None: + values[rule_code] = parsed + + if bundle is None: + conn.execute( + table.insert().values( + key=_BUNDLE_KEY, + value=values, + updated_by_admin_id=None, + ) + ) + else: + conn.execute( + table.update() + .where(table.c.key == _BUNDLE_KEY) + .values(value=values, updated_at=sa.func.now()) + ) + + sparse_keys = [sparse_key for _, sparse_key, _, _, _ in _RULES] + conn.execute(table.delete().where(table.c.key.in_(sparse_keys))) + + +def downgrade() -> None: + conn = op.get_bind() + table = _table() + bundle = _row(conn, table, _BUNDLE_KEY) + values = ( + bundle["value"] + if bundle is not None and isinstance(bundle["value"], dict) + else {} + ) + admin_id = bundle["updated_by_admin_id"] if bundle is not None else None + + for rule_code, sparse_key, default, _, _ in _RULES: + value = _int_or_none(values.get(rule_code)) + if value is None: + value = default + existing = _row(conn, table, sparse_key) + if existing is None: + conn.execute( + table.insert().values( + key=sparse_key, + value=value, + updated_by_admin_id=admin_id, + ) + ) + else: + conn.execute( + table.update() + .where(table.c.key == sparse_key) + .values( + value=value, + updated_by_admin_id=admin_id, + updated_at=sa.func.now(), + ) + ) + + conn.execute(table.delete().where(table.c.key == _BUNDLE_KEY)) diff --git a/alembic/versions/limit_policy_whitelist.py b/alembic/versions/limit_policy_whitelist.py new file mode 100644 index 0000000..3a1c8b4 --- /dev/null +++ b/alembic/versions/limit_policy_whitelist.py @@ -0,0 +1,128 @@ +"""add per-subject limit policy whitelist + +Revision ID: limit_policy_whitelist +Revises: push_binding_isolation +""" +from __future__ import annotations + +from collections.abc import Sequence + +import sqlalchemy as sa +from sqlalchemy.dialects import postgresql + +from alembic import op + +revision: str = "limit_policy_whitelist" +down_revision: str | Sequence[str] | None = "push_binding_isolation" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + +_JSON = sa.JSON().with_variant(postgresql.JSONB(), "postgresql") +_PAGE = "limit-whitelist" +_DEFAULT_ROLES = ("operator", "tech") + + +def _role_table() -> sa.TableClause: + return sa.table( + "admin_role", + sa.column("name", sa.String), + sa.column("pages", _JSON), + ) + + +def _add_default_role_permissions() -> None: + """Grant the page without replacing existing role customisations.""" + role = _role_table() + conn = op.get_bind() + rows = conn.execute( + sa.select(role.c.name, role.c.pages).where( + role.c.name.in_(_DEFAULT_ROLES) + ) + ).all() + for name, pages in rows: + current_pages = list(pages or []) + if _PAGE not in current_pages: + conn.execute( + role.update() + .where(role.c.name == name) + .values(pages=[*current_pages, _PAGE]) + ) + + +def _remove_default_role_permissions() -> None: + role = _role_table() + conn = op.get_bind() + rows = conn.execute( + sa.select(role.c.name, role.c.pages).where( + role.c.name.in_(_DEFAULT_ROLES) + ) + ).all() + for name, pages in rows: + current_pages = list(pages or []) + if _PAGE in current_pages: + conn.execute( + role.update() + .where(role.c.name == name) + .values( + pages=[page for page in current_pages if page != _PAGE] + ) + ) + + +def upgrade() -> None: + op.create_table( + "limit_policy_override", + sa.Column("id", sa.Integer(), autoincrement=True, nullable=False), + sa.Column("subject_type", sa.String(length=16), nullable=False), + sa.Column("subject_value", sa.String(length=128), nullable=False), + sa.Column("rule_code", sa.String(length=64), nullable=False), + sa.Column("mode", sa.String(length=24), nullable=False), + sa.Column("limit_value", sa.Integer(), nullable=True), + sa.Column( + "enabled", sa.Boolean(), server_default=sa.true(), nullable=False + ), + sa.Column("starts_at", sa.DateTime(timezone=True), nullable=True), + sa.Column("expires_at", sa.DateTime(timezone=True), nullable=True), + sa.Column("reset_at", sa.DateTime(timezone=True), nullable=True), + sa.Column("reason", sa.String(length=256), nullable=True), + sa.Column("created_by_admin_id", sa.Integer(), nullable=True), + sa.Column( + "created_at", + sa.DateTime(timezone=True), + server_default=sa.func.now(), + nullable=False, + ), + sa.Column( + "updated_at", + sa.DateTime(timezone=True), + server_default=sa.func.now(), + nullable=False, + ), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint( + "subject_type", + "subject_value", + "rule_code", + name="uq_limit_policy_subject_rule", + ), + ) + op.create_index( + "ix_limit_policy_lookup", + "limit_policy_override", + ["subject_type", "subject_value", "rule_code", "enabled"], + unique=False, + ) + op.create_index( + "ix_limit_policy_expires", + "limit_policy_override", + ["expires_at"], + unique=False, + ) + _add_default_role_permissions() + + +def downgrade() -> None: + _remove_default_role_permissions() + op.drop_index("ix_limit_policy_expires", table_name="limit_policy_override") + op.drop_index("ix_limit_policy_lookup", table_name="limit_policy_override") + op.drop_table("limit_policy_override") diff --git a/app/admin/main.py b/app/admin/main.py index 862e1a4..6ecbb44 100644 --- a/app/admin/main.py +++ b/app/admin/main.py @@ -31,6 +31,7 @@ from app.admin.routers.feedback import router as feedback_router from app.admin.routers.feedback_qr import router as feedback_qr_router from app.admin.routers.guide_video import router as guide_video_router from app.admin.routers.huawei_review import router as huawei_review_router +from app.admin.routers.limit_whitelist import router as limit_whitelist_router from app.admin.routers.onboarding import router as onboarding_router from app.admin.routers.ops_marquee_seed import router as ops_marquee_seed_router from app.admin.routers.ops_stat_config import router as ops_stat_config_router @@ -103,6 +104,7 @@ admin_app.include_router(wallet_router) admin_app.include_router(withdraw_router) admin_app.include_router(price_report_router) admin_app.include_router(risk_monitor_router) +admin_app.include_router(limit_whitelist_router) admin_app.include_router(feedback_router) admin_app.include_router(event_logs_router) admin_app.include_router(analytics_health_router) diff --git a/app/admin/permissions.py b/app/admin/permissions.py index 6ac62b4..b2ee2af 100644 --- a/app/admin/permissions.py +++ b/app/admin/permissions.py @@ -40,6 +40,7 @@ PERMISSION_CATALOG: list[dict] = [ {"key": "analytics-health", "label": "埋点成功率"}, {"key": "event-logs", "label": "埋点日志"}, {"key": "audit-logs", "label": "审计日志"}, + {"key": "limit-whitelist", "label": "白名单"}, ]}, {"group": "其他", "pages": [ {"key": "admins", "label": "权限管理"}, @@ -58,13 +59,14 @@ BUILTIN_ROLES: list[dict] = [ {"name": "operator", "label": "运营", "pages": [ "dashboard", "coupon-data", "ad-revenue-report", "comparison-records", "cps", "risk-monitor", "device-liveness", "price-reports", "feedbacks", "huawei-review", + "limit-whitelist", ]}, {"name": "finance", "label": "财务", "pages": [ "dashboard", "ad-revenue-report", "cps", "invite-withdraws", "withdraws", ]}, {"name": "tech", "label": "技术", "pages": [ "dashboard", "risk-monitor", "device-liveness", "analytics-health", "config", "ad-revenue", "huawei-review", - "event-logs", "audit-logs", + "event-logs", "audit-logs", "limit-whitelist", ]}, ] diff --git a/app/admin/repositories/limit_whitelist.py b/app/admin/repositories/limit_whitelist.py new file mode 100644 index 0000000..9585689 --- /dev/null +++ b/app/admin/repositories/limit_whitelist.py @@ -0,0 +1,583 @@ +"""CRUD and presentation helpers for limit policy overrides.""" + +from __future__ import annotations + +from datetime import UTC, datetime +from hashlib import blake2b + +from sqlalchemy import String, cast, func, or_, select, text +from sqlalchemy.exc import IntegrityError +from sqlalchemy.orm import Session + +from app.core import limit_policy +from app.models.app_config import AppConfig +from app.models.comparison import ComparisonRecord +from app.models.limit_policy import LimitPolicyOverride +from app.models.risk import BehaviorEvent +from app.models.user import User +from app.repositories import risk as risk_repo + + +class DuplicateOverrideError(Exception): + pass + + +def lock_subject(db: Session, *, subject_type: str, subject_value: str) -> None: + """Serialize writes for one logical whitelist subject on PostgreSQL. + + Row locks cannot protect a brand-new subject because there is no row to + lock yet. A transaction-scoped advisory lock closes that gap and also + serializes concurrent appends that touch different rule rows. + """ + + bind = db.get_bind() + if bind.dialect.name != "postgresql": + return + identity = f"limit-whitelist:{subject_type}:{subject_value}".encode() + lock_id = int.from_bytes( + blake2b(identity, digest_size=8).digest(), + byteorder="big", + signed=True, + ) + db.execute( + text("SELECT pg_advisory_xact_lock(:lock_id)"), + {"lock_id": lock_id}, + ) + + +_EVENT_SOURCE_LABELS = { + risk_repo.EVENT_SMS_SEND: "短信验证码", + risk_repo.EVENT_SMS_LOGIN: "短信登录", + risk_repo.EVENT_ONECLICK_LOGIN: "一键登录", +} + + +def _matching_users(keyword: str): + pattern = f"%{keyword}%" + return select(User.id).where( + or_( + cast(User.id, String).ilike(pattern), + User.username.ilike(pattern), + User.phone.ilike(pattern), + User.nickname.ilike(pattern), + ) + ) + + +def _user_maps( + db: Session, *, user_ids: set[int], phones: set[str] +) -> tuple[dict[int, User], dict[str, User]]: + conditions = [] + if user_ids: + conditions.append(User.id.in_(user_ids)) + if phones: + conditions.append(User.phone.in_(phones)) + if not conditions: + return {}, {} + users = list(db.scalars(select(User).where(or_(*conditions))).all()) + return {user.id: user for user in users}, {user.phone: user for user in users} + + +def _behavior_device_candidates(db: Session, *, keyword: str | None, limit: int) -> list[dict]: + conditions = [ + BehaviorEvent.subject_type == limit_policy.SUBJECT_DEVICE, + BehaviorEvent.subject_id != "", + BehaviorEvent.subject_id.not_like(f"{limit_policy.LEGACY_IP_DEVICE_PREFIX}%"), + ] + if keyword: + value = keyword.strip() + pattern = f"%{value}%" + matched_users = _matching_users(value) + conditions.append( + or_( + BehaviorEvent.subject_id.ilike(pattern), + BehaviorEvent.device_id.ilike(pattern), + BehaviorEvent.device_model.ilike(pattern), + BehaviorEvent.phone.ilike(pattern), + cast(BehaviorEvent.user_id, String).ilike(pattern), + BehaviorEvent.user_id.in_(matched_users), + BehaviorEvent.phone.in_( + select(User.phone).where(User.id.in_(_matching_users(value))) + ), + ) + ) + + ranked = ( + select( + BehaviorEvent.subject_id.label("device_id"), + BehaviorEvent.event_type, + BehaviorEvent.user_id, + BehaviorEvent.phone, + BehaviorEvent.device_model, + BehaviorEvent.occurred_at.label("last_active_at"), + func.row_number() + .over( + partition_by=BehaviorEvent.subject_id, + order_by=( + BehaviorEvent.occurred_at.desc(), + BehaviorEvent.id.desc(), + ), + ) + .label("rank"), + ) + .where(*conditions) + .subquery() + ) + rows = db.execute( + select(ranked) + .where(ranked.c.rank == 1) + .order_by(ranked.c.last_active_at.desc(), ranked.c.device_id.asc()) + .limit(limit) + ).all() + by_id, by_phone = _user_maps( + db, + user_ids={row.user_id for row in rows if row.user_id is not None}, + phones={row.phone for row in rows if row.phone}, + ) + result = [] + for row in rows: + user = by_id.get(row.user_id) if row.user_id is not None else None + user = user or (by_phone.get(row.phone) if row.phone else None) + result.append( + { + "device_id": row.device_id, + "source": row.event_type, + "source_label": _EVENT_SOURCE_LABELS.get(row.event_type, "登录/鉴权流水"), + "user_id": user.id if user else row.user_id, + "username": user.username if user else None, + "phone": user.phone if user else row.phone, + "nickname": user.nickname if user else None, + "device_model": row.device_model, + "last_active_at": row.last_active_at, + } + ) + return result + + +def _comparison_device_candidates(db: Session, *, keyword: str | None, limit: int) -> list[dict]: + conditions = [ + ComparisonRecord.device_id.is_not(None), + ComparisonRecord.device_id != "", + ] + if keyword: + value = keyword.strip() + pattern = f"%{value}%" + conditions.append( + or_( + ComparisonRecord.device_id.ilike(pattern), + ComparisonRecord.device_model.ilike(pattern), + cast(ComparisonRecord.user_id, String).ilike(pattern), + ComparisonRecord.user_id.in_(_matching_users(value)), + ) + ) + ranked = ( + select( + ComparisonRecord.device_id, + ComparisonRecord.user_id, + ComparisonRecord.device_model, + ComparisonRecord.created_at.label("last_active_at"), + func.row_number() + .over( + partition_by=ComparisonRecord.device_id, + order_by=( + ComparisonRecord.created_at.desc(), + ComparisonRecord.id.desc(), + ), + ) + .label("rank"), + ) + .where(*conditions) + .subquery() + ) + rows = db.execute( + select(ranked) + .where(ranked.c.rank == 1) + .order_by(ranked.c.last_active_at.desc(), ranked.c.device_id.asc()) + .limit(limit) + ).all() + by_id, _ = _user_maps( + db, + user_ids={row.user_id for row in rows if row.user_id is not None}, + phones=set(), + ) + return [ + { + "device_id": row.device_id, + "source": "comparison_record", + "source_label": "比价记录", + "user_id": user.id if user else row.user_id, + "username": user.username if user else None, + "phone": user.phone if user else None, + "nickname": user.nickname if user else None, + "device_model": row.device_model, + "last_active_at": row.last_active_at, + } + for row in rows + for user in [by_id.get(row.user_id) if row.user_id is not None else None] + ] + + +def list_device_candidates( + db: Session, + *, + rule_code: str, + keyword: str | None = None, + limit: int = 30, +) -> list[dict]: + rule = limit_policy.get_rule(rule_code) + if limit_policy.SUBJECT_DEVICE not in rule.subject_types: + raise ValueError("当前限制项不支持设备白名单") + if limit_policy.device_source_scope(rule_code) == "comparison": + return _comparison_device_candidates(db, keyword=keyword, limit=limit) + return _behavior_device_candidates(db, keyword=keyword, limit=limit) + + +def _aware(value: datetime | None) -> datetime | None: + if value is None: + return None + return value.replace(tzinfo=UTC) if value.tzinfo is None else value + + +def status_of(row: LimitPolicyOverride, *, now: datetime | None = None) -> str: + now = _aware(now) or datetime.now(UTC) + if not row.enabled: + return "disabled" + if _aware(row.starts_at) and _aware(row.starts_at) > now: + return "scheduled" + if _aware(row.expires_at) and _aware(row.expires_at) <= now: + return "expired" + return "active" + + +def to_dict(db: Session, row: LimitPolicyOverride) -> dict: + rule = limit_policy.get_rule(row.rule_code) + effective = limit_policy.resolve( + db, + row.rule_code, + phone=row.subject_value if row.subject_type == "phone" else None, + device=row.subject_value if row.subject_type == "device" else None, + ) + return { + "id": row.id, + "subject_type": row.subject_type, + "subject_value": row.subject_value, + "rule_code": row.rule_code, + "rule_label": rule.label, + "rule_group": rule.group, + "mode": row.mode, + "limit_value": row.limit_value, + "global_limit": effective.global_limit, + "effective_limit": effective.limit, + "enabled": row.enabled, + "starts_at": row.starts_at, + "expires_at": row.expires_at, + "reset_at": row.reset_at, + "reason": row.reason, + "status": status_of(row), + "created_by_admin_id": row.created_by_admin_id, + "created_at": row.created_at, + "updated_at": row.updated_at, + } + + +def list_rows( + db: Session, + *, + subject_type: str | None = None, + keyword: str | None = None, + rule_code: str | None = None, + offset: int = 0, + limit: int = 100, +) -> tuple[list[LimitPolicyOverride], int]: + stmt = select(LimitPolicyOverride) + count_stmt = select(func.count(LimitPolicyOverride.id)) + conditions = [] + if subject_type: + conditions.append(LimitPolicyOverride.subject_type == subject_type) + if keyword: + conditions.append(LimitPolicyOverride.subject_value.ilike(f"%{keyword.strip()}%")) + if rule_code: + conditions.append(LimitPolicyOverride.rule_code == rule_code) + if conditions: + stmt = stmt.where(*conditions) + count_stmt = count_stmt.where(*conditions) + total = int(db.scalar(count_stmt) or 0) + rows = list( + db.execute( + stmt.order_by( + LimitPolicyOverride.created_at.desc(), + LimitPolicyOverride.id.desc(), + ) + .offset(offset) + .limit(limit) + ).scalars() + ) + return rows, total + + +def list_subject_rows( + db: Session, + *, + subject_type: str | None = None, + keyword: str | None = None, + rule_code: str | None = None, + offset: int = 0, + limit: int = 10, +) -> tuple[list[tuple[str, str, list[LimitPolicyOverride], datetime]], int]: + """List whitelist entries grouped and paginated by subject. + + The filter decides which subjects match. Once a subject matches, all of + its rules are returned so category counts and the edit form stay complete. + Ordering uses the original creation time, therefore editing a subject does + not unexpectedly move it to the first page. + """ + + conditions = [ + LimitPolicyOverride.mode.in_( + [limit_policy.MODE_UNLIMITED, limit_policy.MODE_SUPPRESS_ALERT] + ) + ] + if subject_type: + conditions.append(LimitPolicyOverride.subject_type == subject_type) + if keyword: + conditions.append(LimitPolicyOverride.subject_value.ilike(f"%{keyword.strip()}%")) + if rule_code: + conditions.append(LimitPolicyOverride.rule_code == rule_code) + + matched_subjects = ( + select( + LimitPolicyOverride.subject_type.label("subject_type"), + LimitPolicyOverride.subject_value.label("subject_value"), + func.min(LimitPolicyOverride.created_at).label("subject_created_at"), + ) + .where(*conditions) + .group_by( + LimitPolicyOverride.subject_type, + LimitPolicyOverride.subject_value, + ) + .subquery() + ) + total = int(db.scalar(select(func.count()).select_from(matched_subjects)) or 0) + subjects = list( + db.execute( + select( + matched_subjects.c.subject_type, + matched_subjects.c.subject_value, + matched_subjects.c.subject_created_at, + ) + .order_by( + matched_subjects.c.subject_created_at.desc(), + matched_subjects.c.subject_type, + matched_subjects.c.subject_value, + ) + .offset(offset) + .limit(limit) + ).all() + ) + if not subjects: + return [], total + + pair_conditions = [ + ( + (LimitPolicyOverride.subject_type == item.subject_type) + & (LimitPolicyOverride.subject_value == item.subject_value) + ) + for item in subjects + ] + rows = list( + db.scalars( + select(LimitPolicyOverride) + .where( + or_(*pair_conditions), + LimitPolicyOverride.mode.in_( + [ + limit_policy.MODE_UNLIMITED, + limit_policy.MODE_SUPPRESS_ALERT, + ] + ), + ) + .order_by( + LimitPolicyOverride.created_at.desc(), + LimitPolicyOverride.id.desc(), + ) + ).all() + ) + rows_by_subject: dict[tuple[str, str], list[LimitPolicyOverride]] = {} + for row in rows: + rows_by_subject.setdefault((row.subject_type, row.subject_value), []).append(row) + + return [ + ( + item.subject_type, + item.subject_value, + rows_by_subject.get((item.subject_type, item.subject_value), []), + item.subject_created_at, + ) + for item in subjects + ], total + + +def rows_for_subject( + db: Session, *, subject_type: str, subject_value: str +) -> list[LimitPolicyOverride]: + return list( + db.scalars( + select(LimitPolicyOverride) + .where( + LimitPolicyOverride.subject_type == subject_type, + LimitPolicyOverride.subject_value == subject_value, + ) + .order_by( + LimitPolicyOverride.created_at.desc(), + LimitPolicyOverride.id.desc(), + ) + ).all() + ) + + +def create( + db: Session, + *, + subject_type: str, + subject_value: str, + rule_code: str, + mode: str, + limit_value: int | None, + enabled: bool, + starts_at: datetime | None, + expires_at: datetime | None, + reason: str, + admin_id: int, + reactivate_inactive: bool = False, +) -> LimitPolicyOverride: + rule = limit_policy.get_rule(rule_code) + subject_value = limit_policy.validate_whitelist_subject(subject_type, subject_value) + limit_policy.validate_override( + rule, + subject_type=subject_type, + mode=mode, + limit_value=limit_value, + starts_at=starts_at, + expires_at=expires_at, + ) + existing = db.execute( + select(LimitPolicyOverride).where( + LimitPolicyOverride.subject_type == subject_type, + LimitPolicyOverride.subject_value == subject_value, + LimitPolicyOverride.rule_code == rule_code, + ) + ).scalar_one_or_none() + if existing is not None: + if not reactivate_inactive or status_of(existing) not in {"disabled", "expired"}: + raise DuplicateOverrideError + # “恢复全局”与自然过期都会保留原行供审计。再次加入白名单时 + # 复用该行,避免唯一键让批量新增永久失败;完整变更仍写 admin 审计日志。 + existing.mode = mode + existing.limit_value = limit_value if mode == limit_policy.MODE_OVERRIDE else None + existing.enabled = enabled + existing.starts_at = starts_at + existing.expires_at = expires_at + existing.reset_at = None + existing.reason = reason.strip() or None + db.flush() + return existing + + row = LimitPolicyOverride( + subject_type=subject_type, + subject_value=subject_value, + rule_code=rule_code, + mode=mode, + limit_value=limit_value if mode == limit_policy.MODE_OVERRIDE else None, + enabled=enabled, + starts_at=starts_at, + expires_at=expires_at, + reason=reason.strip() or None, + created_by_admin_id=admin_id, + ) + db.add(row) + try: + db.flush() + except IntegrityError as exc: + db.rollback() + raise DuplicateOverrideError from exc + return row + + +def update( + db: Session, + row: LimitPolicyOverride, + *, + enabled: bool | None, + starts_at: datetime | None, + expires_at: datetime | None, + reason: str | None, + fields_set: set[str], +) -> LimitPolicyOverride: + new_starts = starts_at if "starts_at" in fields_set else row.starts_at + new_expires = expires_at if "expires_at" in fields_set else row.expires_at + new_enabled = enabled if "enabled" in fields_set and enabled is not None else row.enabled + rule = limit_policy.get_rule(row.rule_code) + if row.mode == limit_policy.MODE_OVERRIDE and new_enabled: + raise ValueError("历史覆盖配置只能停用、恢复全局或删除") + if new_enabled: + limit_policy.validate_override( + rule, + subject_type=row.subject_type, + mode=row.mode, + limit_value=None, + starts_at=new_starts, + expires_at=new_expires, + ) + if "enabled" in fields_set and enabled is not None: + row.enabled = enabled + if "starts_at" in fields_set: + row.starts_at = starts_at + if "expires_at" in fields_set: + row.expires_at = expires_at + if "reason" in fields_set: + row.reason = reason.strip() if reason else None + db.flush() + return row + + +def update_global_limit( + db: Session, rule_code: str, value: int, *, admin_id: int +) -> tuple[int, int]: + rule = limit_policy.get_rule(rule_code) + if not rule.min_value <= value <= rule.max_value: + raise ValueError(f"限制值必须在 {rule.min_value} 到 {rule.max_value} 之间") + before = limit_policy.resolve(db, rule_code).global_limit + limit_policy.set_global_limits( + db, + {rule.code: value}, + admin_id=admin_id, + commit=False, + ) + # 引导视频仍有一个旧的专用配置页。同步其结构化配置,保证两个入口读取 + # 同一数值;业务判定仍统一走 limit_policy。 + if rule.legacy_config_key and rule.legacy_json_field: + legacy = db.get(AppConfig, rule.legacy_config_key) + legacy_value = ( + dict(legacy.value) if legacy is not None and isinstance(legacy.value, dict) else {} + ) + legacy_value[rule.legacy_json_field] = value + if legacy is None: + legacy = AppConfig( + key=rule.legacy_config_key, + value=legacy_value, + updated_by_admin_id=admin_id, + ) + db.add(legacy) + else: + legacy.value = legacy_value + legacy.updated_by_admin_id = admin_id + db.flush() + return before, value + + +def restore_global(row: LimitPolicyOverride) -> LimitPolicyOverride: + """Stop applying the exception while retaining its history for audit.""" + + row.enabled = False + row.reset_at = None + return row diff --git a/app/admin/routers/config.py b/app/admin/routers/config.py index 270b1b2..8a3d7f5 100644 --- a/app/admin/routers/config.py +++ b/app/admin/routers/config.py @@ -12,9 +12,11 @@ from fastapi import APIRouter, Depends, HTTPException, Request from app.admin.audit import write_audit from app.admin.deps import AdminDb, get_client_ip, get_current_admin, require_role from app.admin.schemas.config import ConfigItemOut, ConfigUpdateRequest +from app.core import limit_policy from app.core.config_schema import CONFIG_DEFS from app.core.rewards import SIGNIN_CYCLE_LEN from app.models.admin import AdminUser +from app.models.app_config import AppConfig from app.repositories import app_config router = APIRouter( @@ -54,7 +56,35 @@ def _validate(key: str, value: Any) -> None: raise ValueError("需为布尔值") +def _limit_item( + key: str, + values: dict[str, int], + bundle: AppConfig | None, +) -> ConfigItemOut: + definition = CONFIG_DEFS[key] + rule_code = limit_policy.RULE_CODE_BY_CONFIG_KEY[key] + return ConfigItemOut( + key=key, + value=values[rule_code], + label=definition["label"], + group=definition["group"], + type=definition["type"], + help=definition.get("help"), + default=definition["default"], + overridden=bundle is not None, + updated_at=( + bundle.updated_at.isoformat() if bundle is not None else None + ), + ) + + def _item(db, key: str) -> ConfigItemOut: + if key in limit_policy.RULE_CODE_BY_CONFIG_KEY: + return _limit_item( + key, + limit_policy.get_global_limits(db), + db.get(AppConfig, limit_policy.LIMIT_POLICY_GLOBAL_KEY), + ) for item in app_config.list_all(db): if item["key"] == key: return ConfigItemOut(**item) @@ -64,11 +94,20 @@ def _item(db, key: str) -> ConfigItemOut: @router.get("", response_model=list[ConfigItemOut], summary="所有可配项 + 当前值(不含 hidden)") def list_config(db: AdminDb) -> list[ConfigItemOut]: # hidden 项(已下线/由专用页管理,如福利页任务·里程碑·看广告调参、首页轮播数据源)不在本页渲染。 - return [ - ConfigItemOut(**item) - for item in app_config.list_all(db) - if not CONFIG_DEFS[item["key"]].get("hidden") - ] + legacy_items = { + item["key"]: item for item in app_config.list_all(db) + } + values = limit_policy.get_global_limits(db) + bundle = db.get(AppConfig, limit_policy.LIMIT_POLICY_GLOBAL_KEY) + out: list[ConfigItemOut] = [] + for key, definition in CONFIG_DEFS.items(): + if definition.get("hidden"): + continue + if key in limit_policy.RULE_CODE_BY_CONFIG_KEY: + out.append(_limit_item(key, values, bundle)) + else: + out.append(ConfigItemOut(**legacy_items[key])) + return out @router.patch("/{key}", response_model=ConfigItemOut, summary="改某项配置(带审计)") @@ -83,11 +122,39 @@ def update_config( raise HTTPException(status_code=404, detail="未知配置项") try: _validate(key, body.value) + rule_code = limit_policy.RULE_CODE_BY_CONFIG_KEY.get(key) + if rule_code is not None: + before = limit_policy.get_global_limits(db)[rule_code] + limit_policy.set_global_limits( + db, + {rule_code: body.value}, + admin_id=admin.id, + commit=False, + ) + else: + before = app_config.get_value(db, key) + app_config.set_value( + db, + key, + body.value, + admin_id=admin.id, + commit=False, + ) + if key == "ad_daily_limit": + # 旧系统配置接口过去只有一个广告日上限。仍有人直接调用时,同时同步 + # 新的激励视频/Draw 两项,避免旧入口写入后业务实际值不变。 + limit_policy.set_global_limits( + db, + { + "ad.reward_video.daily": body.value, + "ad.feed.daily": body.value, + }, + admin_id=admin.id, + commit=False, + ) except ValueError as e: + db.rollback() raise HTTPException(status_code=400, detail=str(e)) from e - - before = app_config.get_value(db, key) - app_config.set_value(db, key, body.value, admin_id=admin.id, commit=False) write_audit( db, admin, action="config.set", target_type="config", target_id=key, detail={"before": before, "after": body.value}, ip=get_client_ip(request), commit=False, diff --git a/app/admin/routers/limit_whitelist.py b/app/admin/routers/limit_whitelist.py new file mode 100644 index 0000000..e581ab2 --- /dev/null +++ b/app/admin/routers/limit_whitelist.py @@ -0,0 +1,838 @@ +"""Unified limit rules and per-phone/device whitelist overrides.""" + +from __future__ import annotations + +from fastapi import APIRouter, Depends, HTTPException, Query, Request, status + +from app.admin.audit import write_audit +from app.admin.deps import AdminDb, CurrentAdmin, get_client_ip, require_page +from app.admin.repositories import limit_whitelist as repo +from app.admin.schemas.limit_whitelist import ( + DeviceCandidateOut, + GlobalLimitUpdate, + LimitOverrideBulkWrite, + LimitOverrideList, + LimitOverrideOut, + LimitOverridePatch, + LimitOverrideWrite, + LimitRuleOut, + LimitSubjectEnabledPatch, + LimitSubjectList, + LimitSubjectOut, +) +from app.core import limit_policy +from app.models.limit_policy import LimitPolicyOverride +from app.repositories import risk as risk_repo + +router = APIRouter( + prefix="/admin/api/limit-whitelist", + tags=["admin-limit-whitelist"], + dependencies=[Depends(require_page("limit-whitelist"))], +) + +_WHITELIST_MODES = { + limit_policy.MODE_UNLIMITED, + limit_policy.MODE_SUPPRESS_ALERT, +} + + +def _subject_whitelist_rows( + db, + *, + subject_type: str, + subject_value: str, +) -> list[LimitPolicyOverride]: + return [ + row + for row in repo.rows_for_subject( + db, + subject_type=subject_type, + subject_value=subject_value, + ) + if row.mode in _WHITELIST_MODES + ] + + +def _reconcile_risk_rule(db, rule_code: str) -> None: + now = risk_repo.utcnow() + if rule_code == "risk.sms.hourly": + risk_repo.reconcile_behavior_rule( + db, + rule_code=risk_repo.RULE_SMS_HOURLY, + at=now, + commit=False, + ) + elif rule_code == "risk.oneclick.daily": + risk_repo.reconcile_behavior_rule( + db, + rule_code=risk_repo.RULE_ONECLICK_DAILY, + at=now, + commit=False, + ) + elif rule_code == "risk.compare.daily": + risk_repo.reconcile_compare_rule(db, at=now, commit=False) + + +def _row_or_404(db, override_id: int) -> LimitPolicyOverride: + row = db.get(LimitPolicyOverride, override_id, populate_existing=True) + if row is None: + raise HTTPException(status_code=404, detail="白名单配置不存在") + return row + + +def _out(db, row: LimitPolicyOverride) -> LimitOverrideOut: + return LimitOverrideOut(**repo.to_dict(db, row)) + + +def _audit_payload(value: LimitOverrideOut) -> dict: + return value.model_dump(mode="json") + + +def _subject_out( + db, + *, + subject_type: str, + subject_value: str, + rows: list[LimitPolicyOverride], + created_at, +) -> LimitSubjectOut: + items = [_out(db, row) for row in rows] + group_counts: dict[str, int] = {} + for item in items: + group_counts[item.rule_group] = group_counts.get(item.rule_group, 0) + 1 + return LimitSubjectOut( + subject_type=subject_type, + subject_value=subject_value, + group_counts=group_counts, + total_rules=len(items), + items=items, + created_at=created_at, + updated_at=max( + (item.updated_at for item in items), + default=created_at, + ), + ) + + +@router.get("/rules", response_model=list[LimitRuleOut], summary="读取所有限制规则") +def list_rules(db: AdminDb) -> list[LimitRuleOut]: + return [LimitRuleOut(**item) for item in limit_policy.rule_catalog(db)] + + +@router.patch( + "/rules/{rule_code}", + response_model=LimitRuleOut, + summary="修改规则的全局限制值", +) +def update_rule( + rule_code: str, + body: GlobalLimitUpdate, + request: Request, + admin: CurrentAdmin, + db: AdminDb, +) -> LimitRuleOut: + try: + before, after = repo.update_global_limit(db, rule_code, body.value, admin_id=admin.id) + except KeyError as exc: + raise HTTPException(status_code=404, detail="未知限制规则") from exc + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + _reconcile_risk_rule(db, rule_code) + write_audit( + db, + admin, + action="limit.rule.update", + target_type="limit_rule", + target_id=rule_code, + detail={"before": before, "after": after}, + ip=get_client_ip(request), + commit=False, + ) + db.commit() + item = next(item for item in limit_policy.rule_catalog(db) if item["code"] == rule_code) + return LimitRuleOut(**item) + + +@router.get( + "/device-candidates", + response_model=list[DeviceCandidateOut], + summary="按限制项搜索可加入白名单的设备", +) +def list_device_candidates( + db: AdminDb, + rule_code: str = Query(..., min_length=1, max_length=64), + keyword: str | None = Query(None, max_length=128), + limit: int = Query(30, ge=1, le=50), +) -> list[DeviceCandidateOut]: + try: + rows = repo.list_device_candidates( + db, + rule_code=rule_code, + keyword=keyword, + limit=limit, + ) + except KeyError as exc: + raise HTTPException(status_code=404, detail="未知限制规则") from exc + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + return [DeviceCandidateOut(**row) for row in rows] + + +@router.get("", response_model=LimitOverrideList, summary="读取白名单配置") +def list_overrides( + db: AdminDb, + subject_type: str | None = Query(None, pattern="^(phone|device)$"), + keyword: str | None = Query(None, max_length=128), + rule_code: str | None = Query(None, max_length=64), + offset: int = Query(0, ge=0), + limit: int = Query(100, ge=1, le=500), +) -> LimitOverrideList: + rows, total = repo.list_rows( + db, + subject_type=subject_type, + keyword=keyword, + rule_code=rule_code, + offset=offset, + limit=limit, + ) + return LimitOverrideList(items=[_out(db, row) for row in rows], total=total) + + +@router.get( + "/subjects", + response_model=LimitSubjectList, + summary="按手机号或设备聚合读取白名单配置", +) +def list_override_subjects( + db: AdminDb, + subject_type: str | None = Query(None, pattern="^(phone|device)$"), + keyword: str | None = Query(None, max_length=128), + rule_code: str | None = Query(None, max_length=64), + offset: int = Query(0, ge=0), + limit: int = Query(10, ge=1, le=100), +) -> LimitSubjectList: + subjects, total = repo.list_subject_rows( + db, + subject_type=subject_type, + keyword=keyword, + rule_code=rule_code, + offset=offset, + limit=limit, + ) + return LimitSubjectList( + items=[ + _subject_out( + db, + subject_type=item_subject_type, + subject_value=subject_value, + rows=rows, + created_at=created_at, + ) + for item_subject_type, subject_value, rows, created_at in subjects + ], + total=total, + ) + + +@router.patch( + "/subjects/enabled", + response_model=LimitSubjectOut, + summary="整体启用或停用一个手机号或设备的白名单", +) +def set_override_subject_enabled( + body: LimitSubjectEnabledPatch, + request: Request, + admin: CurrentAdmin, + db: AdminDb, +) -> LimitSubjectOut: + try: + subject_value = limit_policy.validate_whitelist_subject( + body.subject_type, body.subject_value + ) + repo.lock_subject( + db, + subject_type=body.subject_type, + subject_value=subject_value, + ) + rows = _subject_whitelist_rows( + db, + subject_type=body.subject_type, + subject_value=subject_value, + ) + if not rows: + raise HTTPException(status_code=404, detail="白名单主体不存在") + before = [_audit_payload(_out(db, row)) for row in rows] + for row in rows: + repo.update( + db, + row, + enabled=body.enabled, + starts_at=None, + expires_at=None, + reason=None, + fields_set={"enabled"}, + ) + for rule_code in {row.rule_code for row in rows}: + _reconcile_risk_rule(db, rule_code) + after = [_audit_payload(_out(db, row)) for row in rows] + write_audit( + db, + admin, + action="limit.override.subject_enabled", + target_type="limit_override_subject", + target_id=f"{body.subject_type}:{subject_value}", + detail={"before": before, "after": after}, + ip=get_client_ip(request), + commit=False, + ) + db.commit() + except HTTPException: + raise + except (KeyError, ValueError) as exc: + db.rollback() + raise HTTPException(status_code=400, detail=str(exc)) from exc + + rows = [ + row + for row in repo.rows_for_subject( + db, + subject_type=body.subject_type, + subject_value=subject_value, + ) + if row.mode + in { + limit_policy.MODE_UNLIMITED, + limit_policy.MODE_SUPPRESS_ALERT, + } + ] + return _subject_out( + db, + subject_type=body.subject_type, + subject_value=subject_value, + rows=rows, + created_at=min(row.created_at for row in rows), + ) + + +@router.put( + "/subjects", + response_model=LimitSubjectOut, + summary="整体更新一个手机号或设备的白名单限制项", +) +def replace_override_subject( + body: LimitOverrideBulkWrite, + request: Request, + admin: CurrentAdmin, + db: AdminDb, +) -> LimitSubjectOut: + current_rule_code = "" + try: + subject_value = limit_policy.validate_whitelist_subject( + body.subject_type, body.subject_value + ) + repo.lock_subject( + db, + subject_type=body.subject_type, + subject_value=subject_value, + ) + existing_rows = repo.rows_for_subject( + db, + subject_type=body.subject_type, + subject_value=subject_value, + ) + subject_created_at = min( + (row.created_at for row in existing_rows), + default=None, + ) + before = [_audit_payload(_out(db, row)) for row in existing_rows] + existing_by_rule = {row.rule_code: row for row in existing_rows} + selected_codes = set(body.rule_codes) + touched_rule_codes = set(selected_codes) + + for current_rule_code in body.rule_codes: + rule = limit_policy.get_rule(current_rule_code) + mode = ( + limit_policy.MODE_SUPPRESS_ALERT if rule.alert_only else limit_policy.MODE_UNLIMITED + ) + row = existing_by_rule.get(current_rule_code) + if row is None: + row = repo.create( + db, + subject_type=body.subject_type, + subject_value=subject_value, + rule_code=current_rule_code, + mode=mode, + limit_value=None, + enabled=body.enabled, + starts_at=body.starts_at, + expires_at=body.expires_at, + reason=body.reason, + admin_id=admin.id, + ) + # The table is ordered by the subject's original creation + # time. If an edit replaces every rule, carry that timestamp + # to the new rows so the subject does not jump to the top. + if subject_created_at is not None: + row.created_at = subject_created_at + continue + + # Historical custom-value rows are converted to the only supported + # product modes when the administrator selects that rule again. + row.mode = mode + row.limit_value = None + row.reset_at = None + repo.update( + db, + row, + enabled=body.enabled, + starts_at=body.starts_at, + expires_at=body.expires_at, + reason=body.reason, + fields_set={"enabled", "starts_at", "expires_at", "reason"}, + ) + + for row in existing_rows: + if row.rule_code not in selected_codes and row.mode in { + limit_policy.MODE_UNLIMITED, + limit_policy.MODE_SUPPRESS_ALERT, + }: + touched_rule_codes.add(row.rule_code) + db.delete(row) + db.flush() + except KeyError as exc: + raise HTTPException(status_code=404, detail="未知限制规则") from exc + except repo.DuplicateOverrideError as exc: + rule = limit_policy.get_rule(current_rule_code) + raise HTTPException( + status_code=409, + detail=f"“{rule.label}”已有白名单配置,请刷新后重试", + ) from exc + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + + for rule_code in touched_rule_codes: + _reconcile_risk_rule(db, rule_code) + rows = [ + row + for row in repo.rows_for_subject( + db, + subject_type=body.subject_type, + subject_value=subject_value, + ) + if row.mode + in { + limit_policy.MODE_UNLIMITED, + limit_policy.MODE_SUPPRESS_ALERT, + } + ] + after = [_audit_payload(_out(db, row)) for row in rows] + write_audit( + db, + admin, + action="limit.override.subject_replace", + target_type="limit_override_subject", + target_id=f"{body.subject_type}:{subject_value}", + detail={"before": before, "after": after}, + ip=get_client_ip(request), + commit=False, + ) + db.commit() + rows = [ + row + for row in repo.rows_for_subject( + db, + subject_type=body.subject_type, + subject_value=subject_value, + ) + if row.mode + in { + limit_policy.MODE_UNLIMITED, + limit_policy.MODE_SUPPRESS_ALERT, + } + ] + return _subject_out( + db, + subject_type=body.subject_type, + subject_value=subject_value, + rows=rows, + created_at=min(row.created_at for row in rows), + ) + + +@router.post( + "", + response_model=LimitOverrideOut, + status_code=status.HTTP_201_CREATED, + summary="新增白名单配置", +) +def create_override( + body: LimitOverrideWrite, + request: Request, + admin: CurrentAdmin, + db: AdminDb, +) -> LimitOverrideOut: + before: list[dict] = [] + touched_rows: list[LimitPolicyOverride] = [] + try: + subject_value = limit_policy.validate_whitelist_subject( + body.subject_type, body.subject_value + ) + repo.lock_subject( + db, + subject_type=body.subject_type, + subject_value=subject_value, + ) + existing_rows = _subject_whitelist_rows( + db, + subject_type=body.subject_type, + subject_value=subject_value, + ) + before = [_audit_payload(_out(db, item)) for item in existing_rows] + row = repo.create( + db, + subject_type=body.subject_type, + subject_value=subject_value, + rule_code=body.rule_code, + mode=body.mode, + enabled=body.enabled, + starts_at=body.starts_at, + expires_at=body.expires_at, + reason=body.reason, + limit_value=None, + admin_id=admin.id, + ) + touched_rows = [*existing_rows, row] + if row.mode in _WHITELIST_MODES: + # Keep the legacy single-rule endpoint compatible without letting + # it create a second validity period for the same logical subject. + for existing in existing_rows: + repo.update( + db, + existing, + enabled=body.enabled, + starts_at=body.starts_at, + expires_at=body.expires_at, + reason=None, + fields_set={"enabled", "starts_at", "expires_at"}, + ) + except KeyError as exc: + raise HTTPException(status_code=404, detail="未知限制规则") from exc + except repo.DuplicateOverrideError as exc: + raise HTTPException( + status_code=409, + detail="该手机号或设备已配置此规则,请编辑现有配置", + ) from exc + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + for rule_code in {item.rule_code for item in touched_rows}: + _reconcile_risk_rule(db, rule_code) + after_rows = _subject_whitelist_rows( + db, + subject_type=row.subject_type, + subject_value=row.subject_value, + ) + write_audit( + db, + admin, + action="limit.override.create", + target_type="limit_override_subject", + target_id=f"{row.subject_type}:{row.subject_value}", + detail={ + "before": before, + "after": [_audit_payload(_out(db, item)) for item in after_rows], + }, + ip=get_client_ip(request), + commit=False, + ) + db.commit() + return _out(db, row) + + +@router.post( + "/bulk", + response_model=list[LimitOverrideOut], + status_code=status.HTTP_201_CREATED, + summary="批量新增临时不限或免告警白名单", +) +def create_overrides_bulk( + body: LimitOverrideBulkWrite, + request: Request, + admin: CurrentAdmin, + db: AdminDb, +) -> list[LimitOverrideOut]: + rows: list[LimitPolicyOverride] = [] + before: list[dict] = [] + current_rule_code = "" + try: + subject_value = limit_policy.validate_whitelist_subject( + body.subject_type, body.subject_value + ) + repo.lock_subject( + db, + subject_type=body.subject_type, + subject_value=subject_value, + ) + existing_rows = repo.rows_for_subject( + db, + subject_type=body.subject_type, + subject_value=subject_value, + ) + whitelist_modes = { + limit_policy.MODE_UNLIMITED, + limit_policy.MODE_SUPPRESS_ALERT, + } + existing_whitelist_rows = [row for row in existing_rows if row.mode in whitelist_modes] + before = [_audit_payload(_out(db, row)) for row in existing_whitelist_rows] + existing_by_rule = {row.rule_code: row for row in existing_rows} + selected_codes = set(body.rule_codes) + for current_rule_code in body.rule_codes: + rule = limit_policy.get_rule(current_rule_code) + mode = ( + limit_policy.MODE_SUPPRESS_ALERT if rule.alert_only else limit_policy.MODE_UNLIMITED + ) + row = existing_by_rule.get(current_rule_code) + if row is None: + row = repo.create( + db, + subject_type=body.subject_type, + subject_value=subject_value, + rule_code=current_rule_code, + mode=mode, + limit_value=None, + enabled=body.enabled, + starts_at=body.starts_at, + expires_at=body.expires_at, + reason=body.reason, + admin_id=admin.id, + ) + else: + row.mode = mode + row.limit_value = None + row.reset_at = None + repo.update( + db, + row, + enabled=body.enabled, + starts_at=body.starts_at, + expires_at=body.expires_at, + reason=body.reason, + fields_set={ + "enabled", + "starts_at", + "expires_at", + "reason", + }, + ) + + # 同一手机号或设备在产品上是一条白名单。追加限制项时保留原规则, + # 但统一使用最后一次配置的启用状态与有效期,避免一个主体出现多套时间。 + rows = [ + row + for row in repo.rows_for_subject( + db, + subject_type=body.subject_type, + subject_value=subject_value, + ) + if row.mode in whitelist_modes + ] + for row in rows: + if row.rule_code in selected_codes: + continue + repo.update( + db, + row, + enabled=body.enabled, + starts_at=body.starts_at, + expires_at=body.expires_at, + reason=None, + fields_set={"enabled", "starts_at", "expires_at"}, + ) + except KeyError as exc: + raise HTTPException(status_code=404, detail="未知限制规则") from exc + except repo.DuplicateOverrideError as exc: + rule = limit_policy.get_rule(current_rule_code) + raise HTTPException( + status_code=409, + detail=f"“{rule.label}”白名单配置发生并发更新,请刷新后重试", + ) from exc + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + + for row in rows: + _reconcile_risk_rule(db, row.rule_code) + results = [_out(db, row) for row in rows] + write_audit( + db, + admin, + action="limit.override.bulk_create", + target_type="limit_override", + target_id=",".join(str(row.id) for row in rows), + detail={ + "before": before, + "after": [_audit_payload(item) for item in results], + }, + ip=get_client_ip(request), + commit=False, + ) + db.commit() + return [_out(db, row) for row in rows] + + +@router.patch( + "/{override_id}", + response_model=LimitOverrideOut, + summary="编辑白名单配置", +) +def update_override( + override_id: int, + body: LimitOverridePatch, + request: Request, + admin: CurrentAdmin, + db: AdminDb, +) -> LimitOverrideOut: + row = _row_or_404(db, override_id) + repo.lock_subject( + db, + subject_type=row.subject_type, + subject_value=row.subject_value, + ) + row = _row_or_404(db, override_id) + subject_rows = ( + _subject_whitelist_rows( + db, + subject_type=row.subject_type, + subject_value=row.subject_value, + ) + if row.mode in _WHITELIST_MODES + else [row] + ) + before = [_audit_payload(_out(db, item)) for item in subject_rows] + fields_set = set(body.model_fields_set) + try: + period_fields = {"enabled", "starts_at", "expires_at"} + if row.mode in _WHITELIST_MODES and fields_set & period_fields: + new_enabled = body.enabled if "enabled" in fields_set else row.enabled + new_starts = body.starts_at if "starts_at" in fields_set else row.starts_at + new_expires = body.expires_at if "expires_at" in fields_set else row.expires_at + for item in subject_rows: + item_fields = set(period_fields) + if item.id == row.id and "reason" in fields_set: + item_fields.add("reason") + repo.update( + db, + item, + enabled=new_enabled, + starts_at=new_starts, + expires_at=new_expires, + reason=body.reason if item.id == row.id else None, + fields_set=item_fields, + ) + else: + repo.update( + db, + row, + **body.model_dump(), + fields_set=fields_set, + ) + except (KeyError, ValueError) as exc: + db.rollback() + raise HTTPException(status_code=400, detail=str(exc)) from exc + for rule_code in {item.rule_code for item in subject_rows}: + _reconcile_risk_rule(db, rule_code) + after_rows = [_audit_payload(_out(db, item)) for item in subject_rows] + write_audit( + db, + admin, + action="limit.override.update", + target_type="limit_override_subject", + target_id=f"{row.subject_type}:{row.subject_value}", + detail={"before": before, "after": after_rows}, + ip=get_client_ip(request), + commit=False, + ) + db.commit() + return _out(db, row) + + +@router.post( + "/{override_id}/reset", + response_model=LimitOverrideOut, + summary="停用例外配置并恢复全局策略", +) +def restore_global_policy( + override_id: int, + request: Request, + admin: CurrentAdmin, + db: AdminDb, +) -> LimitOverrideOut: + row = _row_or_404(db, override_id) + repo.lock_subject( + db, + subject_type=row.subject_type, + subject_value=row.subject_value, + ) + row = _row_or_404(db, override_id) + subject_rows = ( + _subject_whitelist_rows( + db, + subject_type=row.subject_type, + subject_value=row.subject_value, + ) + if row.mode in _WHITELIST_MODES + else [row] + ) + before = [_audit_payload(_out(db, item)) for item in subject_rows] + for item in subject_rows: + repo.restore_global(item) + for rule_code in {item.rule_code for item in subject_rows}: + _reconcile_risk_rule(db, rule_code) + after = [_audit_payload(_out(db, item)) for item in subject_rows] + write_audit( + db, + admin, + action="limit.override.restore_global", + target_type="limit_override_subject", + target_id=f"{row.subject_type}:{row.subject_value}", + detail={"before": before, "after": after}, + ip=get_client_ip(request), + commit=False, + ) + db.commit() + return _out(db, row) + + +@router.delete( + "/{override_id}", + status_code=status.HTTP_204_NO_CONTENT, + summary="删除白名单配置", +) +def delete_override( + override_id: int, + request: Request, + admin: CurrentAdmin, + db: AdminDb, +) -> None: + row = _row_or_404(db, override_id) + repo.lock_subject( + db, + subject_type=row.subject_type, + subject_value=row.subject_value, + ) + row = _row_or_404(db, override_id) + before = _audit_payload(_out(db, row)) + target_id = str(row.id) + rule_code = row.rule_code + db.delete(row) + db.flush() + _reconcile_risk_rule(db, rule_code) + write_audit( + db, + admin, + action="limit.override.delete", + target_type="limit_override", + target_id=target_id, + detail={"before": before}, + ip=get_client_ip(request), + commit=False, + ) + db.commit() diff --git a/app/admin/routers/risk_monitor.py b/app/admin/routers/risk_monitor.py index 3c37d64..3e767d6 100644 --- a/app/admin/routers/risk_monitor.py +++ b/app/admin/routers/risk_monitor.py @@ -23,13 +23,8 @@ from app.admin.schemas.risk_monitor import ( RiskResetResponse, RiskRuleConfig, ) -from app.core.config_schema import ( - RISK_COMPARE_DAILY_THRESHOLD_KEY, - RISK_ONECLICK_DAILY_THRESHOLD_KEY, - RISK_SMS_HOURLY_THRESHOLD_KEY, -) +from app.core import limit_policy from app.models.risk import RiskIncident, SubjectRestriction -from app.repositories import app_config from app.repositories import risk as risk_repo router = APIRouter( @@ -83,15 +78,16 @@ def update_rules( ) -> RiskRuleConfig: before = _rule_config(db).model_dump() after = body.model_dump() - values = ( - (RISK_SMS_HOURLY_THRESHOLD_KEY, body.sms_hourly_threshold), - (RISK_ONECLICK_DAILY_THRESHOLD_KEY, body.oneclick_daily_threshold), - (RISK_COMPARE_DAILY_THRESHOLD_KEY, body.compare_daily_threshold), + limit_policy.set_global_limits( + db, + { + "risk.sms.hourly": body.sms_hourly_threshold, + "risk.oneclick.daily": body.oneclick_daily_threshold, + "risk.compare.daily": body.compare_daily_threshold, + }, + admin_id=admin.id, + commit=False, ) - for key, value in values: - app_config.set_value( - db, key, value, admin_id=admin.id, commit=False - ) now = risk_repo.utcnow() risk_repo.reconcile_behavior_rule( diff --git a/app/admin/schemas/limit_whitelist.py b/app/admin/schemas/limit_whitelist.py new file mode 100644 index 0000000..115c194 --- /dev/null +++ b/app/admin/schemas/limit_whitelist.py @@ -0,0 +1,150 @@ +"""Admin contracts for global limit rules and per-subject overrides.""" +from __future__ import annotations + +from datetime import UTC, datetime +from typing import Literal + +from pydantic import BaseModel, ConfigDict, Field, model_validator + +SubjectType = Literal["phone", "device"] +PolicyMode = Literal["unlimited", "suppress_alert"] + + +class LimitRuleOut(BaseModel): + code: str + label: str + group: str + global_limit: int + default_limit: int + window_label: str + subject_types: list[str] + allowed_modes: list[str] + min_value: int + max_value: int + supports_reset: bool + alert_only: bool + + +class GlobalLimitUpdate(BaseModel): + value: int = Field(ge=0, le=1_000_000) + + +class DeviceCandidateOut(BaseModel): + device_id: str + source: str + source_label: str + user_id: int | None = None + username: str | None = None + phone: str | None = None + nickname: str | None = None + device_model: str | None = None + last_active_at: datetime + + +class LimitOverrideWrite(BaseModel): + model_config = ConfigDict(extra="forbid") + + subject_type: SubjectType + subject_value: str = Field(min_length=1, max_length=128) + rule_code: str = Field(min_length=1, max_length=64) + mode: PolicyMode + enabled: bool = True + starts_at: datetime | None = None + expires_at: datetime | None = None + reason: str = Field("", max_length=256) + + @model_validator(mode="after") + def validate_time_range(self): + if self.starts_at and self.expires_at and self.expires_at <= self.starts_at: + raise ValueError("失效时间必须晚于生效时间") + return self + + +class LimitOverrideBulkWrite(BaseModel): + model_config = ConfigDict(extra="forbid") + + subject_type: SubjectType + subject_value: str = Field(min_length=1, max_length=128) + rule_codes: list[str] = Field(min_length=1, max_length=32) + enabled: bool = True + starts_at: datetime | None = None + expires_at: datetime + reason: str = Field("", max_length=256) + + @model_validator(mode="after") + def validate_bulk_request(self): + self.rule_codes = list(dict.fromkeys(self.rule_codes)) + starts_at = ( + self.starts_at.replace(tzinfo=UTC) + if self.starts_at and self.starts_at.tzinfo is None + else self.starts_at + ) + expires_at = ( + self.expires_at.replace(tzinfo=UTC) + if self.expires_at.tzinfo is None + else self.expires_at + ) + if starts_at and expires_at <= starts_at: + raise ValueError("失效时间必须晚于生效时间") + if expires_at <= datetime.now(UTC): + raise ValueError("失效时间必须晚于当前时间") + return self + + +class LimitOverridePatch(BaseModel): + model_config = ConfigDict(extra="forbid") + + enabled: bool | None = None + starts_at: datetime | None = None + expires_at: datetime | None = None + reason: str | None = Field(None, max_length=256) + + +class LimitSubjectEnabledPatch(BaseModel): + model_config = ConfigDict(extra="forbid") + + subject_type: SubjectType + subject_value: str = Field(min_length=1, max_length=128) + enabled: bool + + +class LimitOverrideOut(BaseModel): + id: int + subject_type: str + subject_value: str + rule_code: str + rule_label: str + rule_group: str + mode: str + limit_value: int | None + global_limit: int + effective_limit: int | None + enabled: bool + starts_at: datetime | None + expires_at: datetime | None + reset_at: datetime | None + reason: str | None + status: str + created_by_admin_id: int | None + created_at: datetime + updated_at: datetime + + +class LimitOverrideList(BaseModel): + items: list[LimitOverrideOut] + total: int + + +class LimitSubjectOut(BaseModel): + subject_type: str + subject_value: str + group_counts: dict[str, int] + total_rules: int + items: list[LimitOverrideOut] + created_at: datetime + updated_at: datetime + + +class LimitSubjectList(BaseModel): + items: list[LimitSubjectOut] + total: int diff --git a/app/admin/schemas/risk_monitor.py b/app/admin/schemas/risk_monitor.py index 826393e..3e514ca 100644 --- a/app/admin/schemas/risk_monitor.py +++ b/app/admin/schemas/risk_monitor.py @@ -21,9 +21,9 @@ class RiskMonitorSummary(BaseModel): class RiskRuleConfig(BaseModel): - sms_hourly_threshold: int = Field(ge=1, le=5) + sms_hourly_threshold: int = Field(ge=1, le=100_000) oneclick_daily_threshold: int = Field(ge=1, le=100_000) - compare_daily_threshold: int = Field(ge=1, le=100) + compare_daily_threshold: int = Field(ge=1, le=100_000) class RiskIncidentItem(BaseModel): diff --git a/app/api/v1/ad.py b/app/api/v1/ad.py index bf98e80..0368ff3 100644 --- a/app/api/v1/ad.py +++ b/app/api/v1/ad.py @@ -18,7 +18,7 @@ import uuid from fastapi import APIRouter, Depends, HTTPException, Path, Request, status from app.api.deps import CurrentUser, DbSession -from app.core import rewards +from app.core import limit_policy, rewards from app.core.config import settings from app.core.ratelimit import rate_limit from app.integrations import pangle @@ -413,12 +413,23 @@ def feed_reward(payload: FeedRewardIn, user: CurrentUser, db: DbSession) -> Feed "feed ad reward user_id=%d event=%s status=%s units=%d coin=%d", user.id, rec.client_event_id, rec.status, rec.unit_count, rec.coin, ) + feed_policy = limit_policy.resolve_for_user(db, "ad.feed.daily", user.id) + feed_limit = ( + rewards.get_ad_daily_limit(db) + if feed_policy.override_id is None + and feed_policy.bucket_version == "default" + else feed_policy.limit + ) return FeedRewardOut( granted=(rec.status == "granted"), status=rec.status, coin=rec.coin, unit_count=rec.unit_count, - daily_limit=rewards.get_ad_daily_limit(db), + daily_limit=( + feed_limit + if feed_limit is not None + else limit_policy.get_rule("ad.feed.daily").max_value + ), ) diff --git a/app/api/v1/auth.py b/app/api/v1/auth.py index cccddac..9995e34 100644 --- a/app/api/v1/auth.py +++ b/app/api/v1/auth.py @@ -16,7 +16,7 @@ from fastapi import APIRouter, HTTPException, Request from sqlalchemy.exc import IntegrityError from app.api.deps import CurrentUser, DbSession -from app.core import test_account +from app.core import limit_policy, test_account from app.core.ratelimit import ( RateLimitRule, check_rate_limits, @@ -69,6 +69,7 @@ SMS_LOGIN_MAX_PER_HOUR = 5 # 堵「换手机号绕开单号 60s 冷却」的洞 —— 冷却是单号维度,一机换号能绕开。 SMS_SEND_MAX_PER_HOUR_PER_DEVICE = 5 # 每小时上限 SMS_SEND_MAX_PER_DAY_PER_DEVICE = 20 # 每天上限(再叠一层日封顶,挡低频长时间轰炸) +UNLIMITED_VERIFY_ATTEMPTS = 2_147_483_647 def _client_ip(request: Request) -> str: @@ -198,16 +199,53 @@ def sms_send(req: SmsSendRequest, request: Request, db: DbSession) -> SmsSendRes # 补「换手机号绕开单号 60s 冷却」的洞(冷却是单号维度,一机换号能绕);设备维度按机器封顶,挡短信轰炸/烧钱。 # 关键:被单号 60s 冷却挡下的重发是「没真发、没烧钱」→ 不该占额度。故 check(先判)放在真发之前 # (超限直接 429、不真发),record(计数)只在 send_code 成功后调 —— 冷却/供应商失败抛 429 时直接返回、不计数。 + hourly_policy = limit_policy.resolve( + db, + "sms.send.hourly", + phone=req.phone, + device=subject_id, + ) + daily_policy = limit_policy.resolve( + db, + "sms.send.daily", + phone=req.phone, + device=subject_id, + ) + cooldown_policy = limit_policy.resolve( + db, + "sms.phone.cooldown", + phone=req.phone, + ) + hourly_limit = ( + SMS_SEND_MAX_PER_HOUR_PER_DEVICE + if hourly_policy.override_id is None + and hourly_policy.bucket_version == "default" + else hourly_policy.limit + ) + daily_limit = ( + SMS_SEND_MAX_PER_DAY_PER_DEVICE + if daily_policy.override_id is None + and daily_policy.bucket_version == "default" + else daily_policy.limit + ) send_rules = [ - RateLimitRule("sms-send-device", SMS_SEND_MAX_PER_HOUR_PER_DEVICE, 3600, - "操作过于频繁,请稍后再试"), - RateLimitRule("sms-send-device-daily", SMS_SEND_MAX_PER_DAY_PER_DEVICE, 86400, - "今日验证码发送次数过多,请明天再试"), + RateLimitRule("sms-send-device", hourly_limit, 3600, + "操作过于频繁,请稍后再试", hourly_policy.bucket_version), + RateLimitRule("sms-send-device-daily", daily_limit, 86400, + "今日验证码发送次数过多,请明天再试", daily_policy.bucket_version), ] - check_rate_limits(request, subject=req.device_id, rules=send_rules) + check_rate_limits(request, subject=subject_id, rules=send_rules) try: - send_result = send_code(req.phone) + send_result = send_code( + req.phone, + cooldown_sec=( + None + if cooldown_policy.override_id is None + and cooldown_policy.bucket_version == "default" + else (0 if cooldown_policy.limit is None else cooldown_policy.limit) + ), + ) cooldown = send_result.cooldown_sec except SmsError as e: risk_repo.record_behavior_event( @@ -225,7 +263,7 @@ def sms_send(req: SmsSendRequest, request: Request, db: DbSession) -> SmsSendRes raise HTTPException(status_code=e.status_code, detail=str(e)) from e # 发码成功 → 两道闸各 +1(被单号冷却挡下的重发走不到这里,故不占额度) - record_rate_limits(request, subject=req.device_id, rules=send_rules) + record_rate_limits(request, subject=subject_id, rules=send_rules) from app.core.config import settings # 局部 import 避免循环 @@ -286,17 +324,48 @@ def sms_login(req: SmsLoginRequest, request: Request, db: DbSession) -> TokenWit # **之前** → 输错验证码的失败尝试也计数,才挡得住撞库/爆破(另有单码失败 SMS_MAX_VERIFY_ATTEMPTS 次即作废兜底)。 # ⚠️ 按设备而非手机号 → 一台机器换不同手机号刷登录也受限(防一机狂登多号);device_id 空(老客户端)时 # 退化为该 IP 下所有空设备聚一桶,仍受限。 + login_policy = limit_policy.resolve( + db, + "sms.login.hourly", + phone=req.phone, + device=subject_id, + ) + verify_policy = limit_policy.resolve( + db, + "sms.code.failed_attempts", + phone=req.phone, + ) + login_limit = ( + SMS_LOGIN_MAX_PER_HOUR + if login_policy.override_id is None + and login_policy.bucket_version == "default" + else login_policy.limit + ) enforce_rate_limit( request, scope="sms-login-device", - subject=req.device_id, - limit=SMS_LOGIN_MAX_PER_HOUR, + subject=subject_id, + limit=login_limit, window_sec=3600, detail="登录尝试过于频繁,请稍后再试", + bucket_suffix=login_policy.bucket_version, ) try: - ok = verify_code(req.phone, req.code) + ok = verify_code( + req.phone, + req.code, + max_failed_attempts=( + None + if verify_policy.override_id is None + and verify_policy.bucket_version == "default" + else ( + UNLIMITED_VERIFY_ATTEMPTS + if verify_policy.limit is None + else verify_policy.limit + ) + ), + ) except SmsError as e: # provider 校验降级(如阿里云接口异常)→ 原样透出其状态码(503),别误报「验证码错误」 raise HTTPException(status_code=e.status_code, detail=str(e)) from e if not ok: @@ -397,15 +466,25 @@ def _finish_wechat_bind( 未占用 → 新建微信账号(channel=wechat,昵称头像取微信)→ 签 token 登入。""" existing = user_repo.get_user_by_phone(db, phone) if existing is not None: - from app.core.config import settings # 局部 import,避免循环 - ticket = create_conflict_ticket( openid=openid, wechat_nickname=wechat_nickname, wechat_avatar_url=wechat_avatar_url, phone=phone, ) - blocked = rebind_repo.rebound_within_days(db, phone, settings.PHONE_REBIND_LIMIT_DAYS) + rebind_policy = limit_policy.resolve( + db, + "phone.rebind.days", + phone=phone, + device=device_id, + ) + rebind_days = rebind_policy.limit or 0 + blocked = rebind_repo.rebound_within_days( + db, + phone, + rebind_days, + reset_at=rebind_policy.reset_at, + ) logger.info( "wechat bind phone occupied phone=%s by user_id=%d has_wechat=%s", mask_phone(phone), existing.id, bool(existing.wechat_openid), @@ -421,7 +500,12 @@ def _finish_wechat_bind( conflict_ticket=ticket, rebind_available=not blocked, rebind_blocked_days=( - rebind_repo.remaining_block_days(db, phone, settings.PHONE_REBIND_LIMIT_DAYS) + rebind_repo.remaining_block_days( + db, + phone, + rebind_days, + reset_at=rebind_policy.reset_at, + ) if blocked else 0 ), ) @@ -454,18 +538,50 @@ def wechat_bind_phone_sms( except TokenError as e: raise HTTPException(status_code=401, detail="授权已过期,请重新用微信登录") from e + subject_id = _device_subject(req.device_id, request) # 防刷:同 sms/login,按 设备+IP 每小时限流(放在验证码校验之前,失败也计数) + bind_policy = limit_policy.resolve( + db, + "wechat.bind.hourly", + phone=req.phone, + device=subject_id, + ) + verify_policy = limit_policy.resolve( + db, + "sms.code.failed_attempts", + phone=req.phone, + ) + bind_limit = ( + SMS_LOGIN_MAX_PER_HOUR + if bind_policy.override_id is None + and bind_policy.bucket_version == "default" + else bind_policy.limit + ) enforce_rate_limit( request, scope="wechat-bind-sms-device", - subject=req.device_id, - limit=SMS_LOGIN_MAX_PER_HOUR, + subject=subject_id, + limit=bind_limit, window_sec=3600, detail="登录尝试过于频繁,请稍后再试", + bucket_suffix=bind_policy.bucket_version, ) try: - ok = verify_code(req.phone, req.code) + ok = verify_code( + req.phone, + req.code, + max_failed_attempts=( + None + if verify_policy.override_id is None + and verify_policy.bucket_version == "default" + else ( + UNLIMITED_VERIFY_ATTEMPTS + if verify_policy.limit is None + else verify_policy.limit + ) + ), + ) except SmsError as e: # provider 校验降级(如阿里云接口异常)→ 原样透出其状态码(503),别误报「验证码错误」 raise HTTPException(status_code=e.status_code, detail=str(e)) from e if not ok: @@ -525,9 +641,23 @@ def wechat_conflict_continue( except TokenError as e: raise HTTPException(status_code=401, detail="操作超时,请重新用微信登录") from e + subject_id = _device_subject(req.device_id, request) + conflict_policy = limit_policy.resolve( + db, + "wechat.conflict.hourly", + phone=claims["phone"], + device=subject_id, + ) + conflict_limit = ( + SMS_LOGIN_MAX_PER_HOUR + if conflict_policy.override_id is None + and conflict_policy.bucket_version == "default" + else conflict_policy.limit + ) enforce_rate_limit( - request, scope="wechat-conflict-device", subject=req.device_id, - limit=SMS_LOGIN_MAX_PER_HOUR, window_sec=3600, detail="操作过于频繁,请稍后再试", + request, scope="wechat-conflict-device", subject=subject_id, + limit=conflict_limit, window_sec=3600, detail="操作过于频繁,请稍后再试", + bucket_suffix=conflict_policy.bucket_version, ) user = user_repo.get_user_by_phone(db, claims["phone"]) @@ -565,21 +695,50 @@ def wechat_conflict_continue( def wechat_conflict_rebind( req: WechatConflictRebindRequest, request: Request, db: DbSession ) -> WechatBindResultResponse: - from app.core.config import settings # 局部 import,避免循环 - try: claims = decode_conflict_ticket(req.conflict_ticket) except TokenError as e: raise HTTPException(status_code=401, detail="操作超时,请重新用微信登录") from e + subject_id = _device_subject(req.device_id, request) + conflict_policy = limit_policy.resolve( + db, + "wechat.conflict.hourly", + phone=claims["phone"], + device=subject_id, + ) + conflict_limit = ( + SMS_LOGIN_MAX_PER_HOUR + if conflict_policy.override_id is None + and conflict_policy.bucket_version == "default" + else conflict_policy.limit + ) enforce_rate_limit( - request, scope="wechat-conflict-device", subject=req.device_id, - limit=SMS_LOGIN_MAX_PER_HOUR, window_sec=3600, detail="操作过于频繁,请稍后再试", + request, scope="wechat-conflict-device", subject=subject_id, + limit=conflict_limit, window_sec=3600, detail="操作过于频繁,请稍后再试", + bucket_suffix=conflict_policy.bucket_version, ) phone = claims["phone"] - if rebind_repo.rebound_within_days(db, phone, settings.PHONE_REBIND_LIMIT_DAYS): - days = rebind_repo.remaining_block_days(db, phone, settings.PHONE_REBIND_LIMIT_DAYS) + rebind_policy = limit_policy.resolve( + db, + "phone.rebind.days", + phone=phone, + device=req.device_id, + ) + rebind_days = rebind_policy.limit or 0 + if rebind_repo.rebound_within_days( + db, + phone, + rebind_days, + reset_at=rebind_policy.reset_at, + ): + days = rebind_repo.remaining_block_days( + db, + phone, + rebind_days, + reset_at=rebind_policy.reset_at, + ) raise HTTPException(status_code=409, detail=f"该手机号 {days} 天内已换绑过,暂不能再次换绑") user = user_repo.rebind_account( diff --git a/app/api/v1/compare_record.py b/app/api/v1/compare_record.py index 3d2b484..174a087 100644 --- a/app/api/v1/compare_record.py +++ b/app/api/v1/compare_record.py @@ -16,6 +16,7 @@ import logging from fastapi import APIRouter, BackgroundTasks, HTTPException, Query, status from app.api.deps import CurrentUser, DbSession +from app.core import limit_policy from app.repositories import comparison as crud_compare from app.repositories import risk as risk_repo from app.schemas.compare_record import ( @@ -53,17 +54,29 @@ def reserve_compare_start( ): raise HTTPException(status_code=403, detail="账号存在异常,该功能暂不可用") try: + policy = limit_policy.resolve( + db, + "compare.start.daily", + phone=user.phone, + device=payload.device_id, + ) rec, used = crud_compare.reserve_daily_start( db, user_id=user.id, trace_id=payload.trace_id, business_type=payload.business_type, device_id=payload.device_id, + limit=policy.limit, + reset_at=policy.reset_at, ) except crud_compare.DailyCompareStartLimitExceeded: raise HTTPException( status_code=status.HTTP_429_TOO_MANY_REQUESTS, - detail="今日已比价超过100次,请明天再试", + detail=( + f"今日已比价超过{policy.limit}次,请明天再试" + if policy.limit is not None + else "今日比价次数已达上限,请明天再试" + ), ) from None except crud_compare.ComparisonTraceOwnershipError: raise HTTPException( @@ -71,11 +84,16 @@ def reserve_compare_start( detail="比价任务标识冲突,请重新发起", ) from None # 风控阈值由后台动态配置,不能再只在固定的 100 次业务上限处同步。 - risk_repo.sync_compare_incident(db, user_id=user.id, at=rec.created_at) + risk_repo.sync_compare_incident( + db, + user_id=user.id, + at=rec.created_at, + device_id=payload.device_id, + ) return CompareStartReserveOut( - limit=crud_compare.DAILY_COMPARE_START_LIMIT, + limit=policy.limit, used=used, - remaining=max(crud_compare.DAILY_COMPARE_START_LIMIT - used, 0), + remaining=max(policy.limit - used, 0) if policy.limit is not None else None, ) diff --git a/app/core/config_schema.py b/app/core/config_schema.py index 9a63720..b2289b3 100644 --- a/app/core/config_schema.py +++ b/app/core/config_schema.py @@ -14,12 +14,138 @@ from app.core import rewards as r RISK_SMS_HOURLY_THRESHOLD_KEY = "risk_sms_hourly_threshold" RISK_ONECLICK_DAILY_THRESHOLD_KEY = "risk_oneclick_daily_threshold" RISK_COMPARE_DAILY_THRESHOLD_KEY = "risk_compare_daily_threshold" +COMPARE_DAILY_LIMIT_KEY = "compare_daily_limit" +SMS_SEND_HOURLY_LIMIT_KEY = "sms_send_hourly_limit" +SMS_SEND_DAILY_LIMIT_KEY = "sms_send_daily_limit" +SMS_LOGIN_HOURLY_LIMIT_KEY = "sms_login_hourly_limit" +WECHAT_BIND_SMS_HOURLY_LIMIT_KEY = "wechat_bind_sms_hourly_limit" +WECHAT_CONFLICT_HOURLY_LIMIT_KEY = "wechat_conflict_hourly_limit" +SMS_PHONE_COOLDOWN_SECONDS_KEY = "sms_phone_cooldown_seconds" +SMS_CODE_MAX_FAILED_ATTEMPTS_KEY = "sms_code_max_failed_attempts" +AD_REWARD_VIDEO_DAILY_LIMIT_KEY = "ad_reward_video_daily_limit" +AD_FEED_DAILY_LIMIT_KEY = "ad_feed_daily_limit" +PHONE_REBIND_DAYS_KEY = "phone_rebind_days" +GUIDE_VIDEO_MAX_PLAYS_KEY = "guide_video_max_plays" +LIMIT_POLICY_GLOBAL_KEY = "limit_policy_global" # type 约定(给前端渲染编辑控件用):int / int_list / dict_str_int / bool / enum # hidden=True:仍是合法可配项(业务照常 get_value / admin 可经专用端点读写),但**不在通用 # 「系统配置」页渲染**(admin/routers/config.py:list_config 按此过滤)。用于把已下线/已改由 # 专用页管理的项从福利页 Tab 收起,同时保留后端默认值与写入能力。 CONFIG_DEFS: dict[str, dict[str, Any]] = { + COMPARE_DAILY_LIMIT_KEY: { + "default": 100, + "label": "账号每日比价次数上限", + "group": "限制策略", + "type": "int", + "min": 1, + "max": 100_000, + "hidden": True, + "help": "北京时间自然日内每个账号可发起的比价次数。", + }, + SMS_SEND_HOURLY_LIMIT_KEY: { + "default": 5, + "label": "短信每小时发送上限", + "group": "限制策略", + "type": "int", + "min": 1, + "max": 10_000, + "hidden": True, + "help": "同一设备和 IP 在固定 1 小时窗口内成功发送短信的次数。", + }, + SMS_SEND_DAILY_LIMIT_KEY: { + "default": 20, + "label": "短信 24 小时发送上限", + "group": "限制策略", + "type": "int", + "min": 1, + "max": 100_000, + "hidden": True, + "help": "同一设备和 IP 在固定 24 小时窗口内成功发送短信的次数。", + }, + SMS_LOGIN_HOURLY_LIMIT_KEY: { + "default": 5, + "label": "短信登录每小时尝试上限", + "group": "限制策略", + "type": "int", + "min": 1, + "max": 10_000, + "hidden": True, + "help": "成功与失败均计入。", + }, + WECHAT_BIND_SMS_HOURLY_LIMIT_KEY: { + "default": 5, + "label": "微信短信绑定每小时尝试上限", + "group": "限制策略", + "type": "int", + "min": 1, + "max": 10_000, + "hidden": True, + }, + WECHAT_CONFLICT_HOURLY_LIMIT_KEY: { + "default": 5, + "label": "微信冲突处理每小时尝试上限", + "group": "限制策略", + "type": "int", + "min": 1, + "max": 10_000, + "hidden": True, + "help": "继续登录与重新绑定共享此额度。", + }, + SMS_PHONE_COOLDOWN_SECONDS_KEY: { + "default": 60, + "label": "同手机号短信发送冷却秒数", + "group": "限制策略", + "type": "int", + "min": 0, + "max": 86_400, + "hidden": True, + }, + SMS_CODE_MAX_FAILED_ATTEMPTS_KEY: { + "default": 5, + "label": "单验证码最大失败次数", + "group": "限制策略", + "type": "int", + "min": 1, + "max": 100, + "hidden": True, + }, + AD_REWARD_VIDEO_DAILY_LIMIT_KEY: { + "default": r.DAILY_AD_REWARD_LIMIT, + "label": "激励视频每日发奖次数", + "group": "限制策略", + "type": "int", + "min": 1, + "max": 100_000, + "hidden": True, + }, + AD_FEED_DAILY_LIMIT_KEY: { + "default": r.DAILY_AD_REWARD_LIMIT, + "label": "Draw 信息流每日发奖次数", + "group": "限制策略", + "type": "int", + "min": 1, + "max": 100_000, + "hidden": True, + }, + PHONE_REBIND_DAYS_KEY: { + "default": 30, + "label": "手机/微信换绑冷却天数", + "group": "限制策略", + "type": "int", + "min": 0, + "max": 3650, + "hidden": True, + }, + GUIDE_VIDEO_MAX_PLAYS_KEY: { + "default": 3, + "label": "领券引导视频最大播放次数", + "group": "限制策略", + "type": "int", + "min": 0, + "max": 50, + "hidden": True, + }, "signin_rewards": { "default": list(r.SIGNIN_REWARDS), "label": "签到 7 天金币档位", "group": "签到", "type": "int_list", @@ -56,7 +182,8 @@ CONFIG_DEFS: dict[str, dict[str, Any]] = { }, "ad_daily_limit": { "default": r.DAILY_AD_REWARD_LIMIT, "label": "看广告每日上限(次)", - "group": "看广告", "type": "int", "help": "福利页激励视频每日可发奖次数上限,默认 500。", + "group": "看广告", "type": "int", "min": 1, "max": 100_000, "hidden": True, + "help": "历史共享上限;新配置由白名单页分别管理激励视频与 Draw 信息流。", }, "ad_max_coin": { "default": r.MAX_AD_REWARD_COIN, "label": "看广告单次金币上限", @@ -68,7 +195,8 @@ CONFIG_DEFS: dict[str, dict[str, Any]] = { }, "ad_cooldown_sec": { "default": r.VIDEO_ROUND_COOLDOWN_SECONDS, "label": "广告关闭后冷却(秒)", - "group": "看广告", "type": "int", "help": "点击退出广告后,下次点击观看前的冷却时间,默认 3 秒。", + "group": "看广告", "type": "int", "min": 0, "max": 86_400, + "help": "点击退出广告后,下次点击观看前的冷却时间,默认 3 秒。", }, "comparing_ad_enabled": { "default": True, "label": "比价/领券期信息流广告", @@ -118,9 +246,9 @@ CONFIG_DEFS: dict[str, dict[str, Any]] = { "group": "风控", "type": "int", "min": 1, - "max": 5, + "max": 100_000, "hidden": True, - "help": "同一设备在北京时间同一自然小时内成功下发短信达到该次数时告警;不得高于现有每小时 5 次的发送上限。", + "help": "同一设备在北京时间同一自然小时内成功下发短信达到该次数时告警。", }, RISK_ONECLICK_DAILY_THRESHOLD_KEY: { "default": 20, @@ -138,7 +266,7 @@ CONFIG_DEFS: dict[str, dict[str, Any]] = { "group": "风控", "type": "int", "min": 1, - "max": 100, + "max": 100_000, "hidden": True, "help": "同一账户在北京时间同一自然日内发起比价达到该次数时告警。", }, diff --git a/app/core/limit_policy.py b/app/core/limit_policy.py new file mode 100644 index 0000000..d91d571 --- /dev/null +++ b/app/core/limit_policy.py @@ -0,0 +1,623 @@ +"""Unified global limits and per-phone/device policy overrides. + +The registry is the single source of truth for the whitelist page. Existing +constants remain as backwards-compatible defaults, while business call sites +resolve an effective value here. +""" +from __future__ import annotations + +from collections.abc import Iterable +from dataclasses import dataclass +from datetime import UTC, datetime + +from sqlalchemy import delete, or_, select +from sqlalchemy.orm import Session + +from app.core.config_schema import ( + AD_FEED_DAILY_LIMIT_KEY, + AD_REWARD_VIDEO_DAILY_LIMIT_KEY, + COMPARE_DAILY_LIMIT_KEY, + GUIDE_VIDEO_MAX_PLAYS_KEY, + LIMIT_POLICY_GLOBAL_KEY, + PHONE_REBIND_DAYS_KEY, + RISK_COMPARE_DAILY_THRESHOLD_KEY, + RISK_ONECLICK_DAILY_THRESHOLD_KEY, + RISK_SMS_HOURLY_THRESHOLD_KEY, + SMS_CODE_MAX_FAILED_ATTEMPTS_KEY, + SMS_LOGIN_HOURLY_LIMIT_KEY, + SMS_PHONE_COOLDOWN_SECONDS_KEY, + SMS_SEND_DAILY_LIMIT_KEY, + SMS_SEND_HOURLY_LIMIT_KEY, + WECHAT_BIND_SMS_HOURLY_LIMIT_KEY, + WECHAT_CONFLICT_HOURLY_LIMIT_KEY, +) +from app.models.app_config import AppConfig +from app.models.limit_policy import LimitPolicyOverride +from app.models.user import User +from app.repositories import app_config + +MODE_INHERIT = "inherit" +MODE_OVERRIDE = "override" +MODE_UNLIMITED = "unlimited" +MODE_SUPPRESS_ALERT = "suppress_alert" + +SUBJECT_PHONE = "phone" +SUBJECT_DEVICE = "device" +SUBJECT_TYPES = (SUBJECT_PHONE, SUBJECT_DEVICE) +SUBJECT_PRECEDENCE = {SUBJECT_DEVICE: 0, SUBJECT_PHONE: 1} +LEGACY_IP_DEVICE_PREFIX = "legacy-ip:" + + +@dataclass(frozen=True) +class RuleDefinition: + code: str + label: str + group: str + config_key: str + default_limit: int + window_label: str + subject_types: tuple[str, ...] + min_value: int = 1 + max_value: int = 100_000 + allow_unlimited: bool = True + supports_reset: bool = True + alert_only: bool = False + legacy_config_key: str | None = None + legacy_json_field: str | None = None + + @property + def allowed_modes(self) -> tuple[str, ...]: + if self.alert_only: + return (MODE_SUPPRESS_ALERT,) + return (MODE_UNLIMITED,) if self.allow_unlimited else () + + +RULES: tuple[RuleDefinition, ...] = ( + RuleDefinition( + "compare.start.daily", + "每日发起比价次数", + "比价", + COMPARE_DAILY_LIMIT_KEY, + 100, + "北京时间自然日", + SUBJECT_TYPES, + ), + RuleDefinition( + "sms.send.hourly", + "短信每小时成功发送次数", + "短信与登录", + SMS_SEND_HOURLY_LIMIT_KEY, + 5, + "固定 1 小时窗口", + SUBJECT_TYPES, + max_value=10_000, + ), + RuleDefinition( + "sms.send.daily", + "短信 24 小时成功发送次数", + "短信与登录", + SMS_SEND_DAILY_LIMIT_KEY, + 20, + "固定 24 小时窗口", + SUBJECT_TYPES, + ), + RuleDefinition( + "sms.phone.cooldown", + "同手机号短信发送冷却", + "短信与登录", + SMS_PHONE_COOLDOWN_SECONDS_KEY, + 60, + "秒", + (SUBJECT_PHONE,), + min_value=0, + max_value=86_400, + ), + RuleDefinition( + "sms.code.failed_attempts", + "单验证码最大失败次数", + "短信与登录", + SMS_CODE_MAX_FAILED_ATTEMPTS_KEY, + 5, + "单个验证码", + (SUBJECT_PHONE,), + max_value=100, + ), + RuleDefinition( + "sms.login.hourly", + "短信登录每小时尝试次数", + "短信与登录", + SMS_LOGIN_HOURLY_LIMIT_KEY, + 5, + "固定 1 小时窗口", + SUBJECT_TYPES, + max_value=10_000, + ), + RuleDefinition( + "wechat.bind.hourly", + "微信短信绑定每小时尝试次数", + "短信与登录", + WECHAT_BIND_SMS_HOURLY_LIMIT_KEY, + 5, + "固定 1 小时窗口", + SUBJECT_TYPES, + max_value=10_000, + ), + RuleDefinition( + "wechat.conflict.hourly", + "微信冲突处理每小时尝试次数", + "短信与登录", + WECHAT_CONFLICT_HOURLY_LIMIT_KEY, + 5, + "固定 1 小时窗口", + SUBJECT_TYPES, + max_value=10_000, + ), + RuleDefinition( + "ad.reward_video.daily", + "激励视频每日发奖次数", + "广告", + AD_REWARD_VIDEO_DAILY_LIMIT_KEY, + 500, + "北京时间自然日", + (SUBJECT_PHONE,), + legacy_config_key="ad_daily_limit", + ), + RuleDefinition( + "ad.feed.daily", + "Draw 信息流每日发奖次数", + "广告", + AD_FEED_DAILY_LIMIT_KEY, + 500, + "北京时间自然日", + (SUBJECT_PHONE,), + legacy_config_key="ad_daily_limit", + ), + RuleDefinition( + "ad.reward_video.cooldown", + "激励视频发奖后冷却秒数", + "广告", + "ad_cooldown_sec", + 3, + "秒", + (SUBJECT_PHONE,), + min_value=0, + max_value=86_400, + ), + RuleDefinition( + "guide.video.lifetime", + "领券引导视频最大播放次数", + "引导与账号", + GUIDE_VIDEO_MAX_PLAYS_KEY, + 3, + "账号生命周期", + (SUBJECT_PHONE,), + min_value=0, + max_value=50, + legacy_config_key="coupon_guide_video", + legacy_json_field="max_plays", + ), + RuleDefinition( + "phone.rebind.days", + "手机/微信换绑冷却天数", + "引导与账号", + PHONE_REBIND_DAYS_KEY, + 30, + "自然日", + (SUBJECT_PHONE,), + min_value=0, + max_value=3650, + ), + RuleDefinition( + "risk.sms.hourly", + "短信设备每小时告警", + "风控免告警", + RISK_SMS_HOURLY_THRESHOLD_KEY, + 5, + "北京时间自然小时", + (SUBJECT_DEVICE,), + max_value=100_000, + allow_unlimited=False, + alert_only=True, + ), + RuleDefinition( + "risk.oneclick.daily", + "一键登录设备每日告警", + "风控免告警", + RISK_ONECLICK_DAILY_THRESHOLD_KEY, + 20, + "北京时间自然日", + (SUBJECT_DEVICE,), + max_value=100_000, + allow_unlimited=False, + alert_only=True, + ), + RuleDefinition( + "risk.compare.daily", + "比价账号每日告警", + "风控免告警", + RISK_COMPARE_DAILY_THRESHOLD_KEY, + 100, + "北京时间自然日", + (SUBJECT_PHONE,), + max_value=100_000, + allow_unlimited=False, + alert_only=True, + ), +) +RULE_MAP = {rule.code: rule for rule in RULES} +RULE_CODE_BY_CONFIG_KEY = {rule.config_key: rule.code for rule in RULES} +LIMIT_CONFIG_KEYS = tuple(RULE_CODE_BY_CONFIG_KEY) + + +@dataclass(frozen=True) +class EffectiveLimit: + rule_code: str + global_limit: int + limit: int | None + mode: str + suppressed: bool + override_id: int | None + matched_subject_type: str | None + matched_subject_value: str | None + reset_at: datetime | None + bucket_version: str + + @property + def unlimited(self) -> bool: + return self.limit is None + + +def normalize_subject(subject_type: str, value: str) -> str: + value = (value or "").strip() + if subject_type == SUBJECT_PHONE: + value = "".join(ch for ch in value if ch.isdigit()) + if subject_type not in SUBJECT_TYPES: + raise ValueError(f"unsupported subject type: {subject_type}") + if not value: + raise ValueError("subject value is empty") + return value[:128] + + +def validate_whitelist_subject(subject_type: str, value: str) -> str: + """Normalize a whitelist subject and reject unsafe pseudo-devices. + + Old clients without a device ID are grouped by public IP for rate limiting. + That fallback remains valid at runtime, but it is not a stable, unique + device identity and must never be persisted as a device whitelist target. + """ + + normalized = normalize_subject(subject_type, value) + if ( + subject_type == SUBJECT_DEVICE + and normalized.startswith(LEGACY_IP_DEVICE_PREFIX) + ): + raise ValueError( + "旧客户端未上报真实设备 ID,不能加入设备白名单,请升级客户端后重试" + ) + return normalized + + +def get_rule(rule_code: str) -> RuleDefinition: + try: + return RULE_MAP[rule_code] + except KeyError as exc: + raise ValueError(f"unknown rule: {rule_code}") from exc + + +def device_source_scope(rule_code: str) -> str: + """返回设备候选数据所属命名空间,防止一个设备 ID 跨来源误套规则。""" + rule = get_rule(rule_code) + if SUBJECT_DEVICE not in rule.subject_types: + raise ValueError("当前限制项不支持设备白名单") + return "comparison" if rule_code == "compare.start.daily" else "auth" + + +def default_global_limits() -> dict[str, int]: + """Return the complete 16-rule default snapshot keyed by rule code.""" + return {rule.code: rule.default_limit for rule in RULES} + + +def _normalise_global_limits(value: object) -> dict[str, int]: + """Merge a stored JSON object with safe code defaults. + + The migration and every admin write persist all rules. Defaults are still + merged here so a manually damaged/older partial JSON cannot take the + service down after deployment. + """ + + values = default_global_limits() + if not isinstance(value, dict): + return values + for rule_code, raw in value.items(): + rule = RULE_MAP.get(str(rule_code)) + if rule is None or isinstance(raw, bool): + continue + try: + parsed = int(raw) + except (TypeError, ValueError): + continue + if rule.min_value <= parsed <= rule.max_value: + values[rule.code] = parsed + return values + + +def _legacy_global_limit(db: Session, rule: RuleDefinition) -> tuple[int, str]: + """Read the pre-bundle representation while upgrading old/test databases.""" + + row = db.get(AppConfig, rule.config_key) + if row is not None: + return int(row.value), "legacy-key" + + if rule.legacy_config_key: + legacy = db.get(AppConfig, rule.legacy_config_key) + if legacy is not None: + value = legacy.value + if rule.legacy_json_field: + value = value.get(rule.legacy_json_field) if isinstance(value, dict) else None + if value is not None: + return int(value), "legacy" + + try: + return int(app_config.get_value(db, rule.config_key)), "default" + except KeyError: + return rule.default_limit, "default" + + +def get_global_limits(db: Session) -> dict[str, int]: + """Read the complete global-limit JSON, with a pre-migration fallback.""" + + row = db.get(AppConfig, LIMIT_POLICY_GLOBAL_KEY) + if row is not None: + return _normalise_global_limits(row.value) + return {rule.code: _legacy_global_limit(db, rule)[0] for rule in RULES} + + +def set_global_limits( + db: Session, + updates: dict[str, int], + *, + admin_id: int, + commit: bool = True, +) -> dict[str, int]: + """Atomically update selected rules inside the single complete JSON row.""" + + parsed_updates: dict[str, int] = {} + for rule_code, raw_value in updates.items(): + rule = get_rule(rule_code) + value = int(raw_value) + if not rule.min_value <= value <= rule.max_value: + raise ValueError( + f"limit for {rule_code} must be between " + f"{rule.min_value} and {rule.max_value}" + ) + parsed_updates[rule.code] = value + + row = db.scalar( + select(AppConfig) + .where(AppConfig.key == LIMIT_POLICY_GLOBAL_KEY) + .with_for_update() + ) + values = ( + _normalise_global_limits(row.value) + if row is not None + else {rule.code: _legacy_global_limit(db, rule)[0] for rule in RULES} + ) + values.update(parsed_updates) + + if row is None: + row = AppConfig( + key=LIMIT_POLICY_GLOBAL_KEY, + value=values, + updated_by_admin_id=admin_id, + ) + db.add(row) + else: + row.value = dict(values) + row.updated_by_admin_id = admin_id + + # Once the bundle exists, stale sparse rows must not become a second source + # of truth. The data migration performs the same cleanup for production. + db.execute(delete(AppConfig).where(AppConfig.key.in_(LIMIT_CONFIG_KEYS))) + if commit: + db.commit() + db.refresh(row) + else: + db.flush() + return values + + +def _global_limit(db: Session, rule: RuleDefinition) -> tuple[int, str]: + row = db.get(AppConfig, LIMIT_POLICY_GLOBAL_KEY) + if row is not None: + return _normalise_global_limits(row.value)[rule.code], "configured" + return _legacy_global_limit(db, rule) + + +def _aware(value: datetime | None) -> datetime | None: + if value is None: + return None + return value.replace(tzinfo=UTC) if value.tzinfo is None else value + + +def _matching_overrides( + db: Session, + rule: RuleDefinition, + subjects: dict[str, str | None], + now: datetime, +) -> list[LimitPolicyOverride]: + pairs: list[tuple[str, str]] = [] + for subject_type in rule.subject_types: + raw = subjects.get(subject_type) + if raw: + pairs.append((subject_type, normalize_subject(subject_type, raw))) + if not pairs: + return [] + clauses = [ + ( + (LimitPolicyOverride.subject_type == subject_type) + & (LimitPolicyOverride.subject_value == subject_value) + ) + for subject_type, subject_value in pairs + ] + rows = list( + db.execute( + select(LimitPolicyOverride).where( + LimitPolicyOverride.rule_code == rule.code, + LimitPolicyOverride.enabled.is_(True), + or_(*clauses), + ) + ).scalars() + ) + active = [ + row + for row in rows + if (_aware(row.starts_at) is None or _aware(row.starts_at) <= now) + and (_aware(row.expires_at) is None or _aware(row.expires_at) > now) + ] + return sorted(active, key=lambda row: SUBJECT_PRECEDENCE[row.subject_type]) + + +def resolve( + db: Session, + rule_code: str, + *, + phone: str | None = None, + device: str | None = None, + now: datetime | None = None, +) -> EffectiveLimit: + """Resolve global config plus the most specific active override. + + Device overrides win over phone overrides when both match. + """ + + rule = get_rule(rule_code) + now = _aware(now) or datetime.now(UTC) + global_limit, global_version = _global_limit(db, rule) + matches = _matching_overrides( + db, rule, {SUBJECT_PHONE: phone, SUBJECT_DEVICE: device}, now + ) + row = matches[0] if matches else None + if row is None: + return EffectiveLimit( + rule.code, + global_limit, + global_limit, + MODE_INHERIT, + False, + None, + None, + None, + None, + global_version, + ) + + limit: int | None = global_limit + suppressed = False + if row.mode == MODE_OVERRIDE: + limit = int(row.limit_value) if row.limit_value is not None else global_limit + elif row.mode == MODE_UNLIMITED: + limit = None + elif row.mode == MODE_SUPPRESS_ALERT: + suppressed = True + # 编辑限制值/备注不应隐式清空计数;只有显式“重置状态”才切换桶。 + reset_version = _aware(row.reset_at) + version = ( + f"{global_version}:o:{row.id}:" + f"{reset_version.isoformat() if reset_version else '0'}" + ) + return EffectiveLimit( + rule.code, + global_limit, + limit, + row.mode, + suppressed, + row.id, + row.subject_type, + row.subject_value, + _aware(row.reset_at), + version, + ) + + +def resolve_for_user( + db: Session, + rule_code: str, + user_id: int, + *, + device: str | None = None, +) -> EffectiveLimit: + user = db.get(User, user_id) + return resolve( + db, + rule_code, + phone=user.phone if user is not None else None, + device=device, + ) + + +def rule_catalog(db: Session) -> list[dict]: + out: list[dict] = [] + values = get_global_limits(db) + for rule in RULES: + out.append( + { + "code": rule.code, + "label": rule.label, + "group": rule.group, + "global_limit": values[rule.code], + "default_limit": rule.default_limit, + "window_label": rule.window_label, + "subject_types": list(rule.subject_types), + "allowed_modes": list(rule.allowed_modes), + "min_value": rule.min_value, + "max_value": rule.max_value, + "supports_reset": rule.supports_reset, + "alert_only": rule.alert_only, + } + ) + return out + + +def validate_override( + rule: RuleDefinition, + *, + subject_type: str, + mode: str, + limit_value: int | None, + starts_at: datetime | None, + expires_at: datetime | None, +) -> None: + if subject_type not in rule.subject_types: + raise ValueError("该规则不支持此主体类型") + if mode not in rule.allowed_modes: + raise ValueError("该规则不支持此策略模式") + if limit_value is not None: + raise ValueError("白名单不支持覆盖指定值") + if mode in {MODE_UNLIMITED, MODE_SUPPRESS_ALERT} and expires_at is None: + raise ValueError("临时白名单必须设置失效时间") + if starts_at and expires_at and _aware(expires_at) <= _aware(starts_at): + raise ValueError("失效时间必须晚于生效时间") + if ( + mode in {MODE_UNLIMITED, MODE_SUPPRESS_ALERT} + and expires_at is not None + and _aware(expires_at) <= datetime.now(UTC) + ): + raise ValueError("临时白名单的失效时间必须晚于当前时间") + if mode == MODE_SUPPRESS_ALERT and not rule.alert_only: + raise ValueError("免告警只支持风控监控的三项规则") + + +def active_overrides( + db: Session, + *, + subject_type: str | None = None, + keyword: str | None = None, +) -> Iterable[LimitPolicyOverride]: + stmt = select(LimitPolicyOverride) + if subject_type: + stmt = stmt.where(LimitPolicyOverride.subject_type == subject_type) + if keyword: + stmt = stmt.where(LimitPolicyOverride.subject_value.ilike(f"%{keyword.strip()}%")) + return db.execute( + stmt.order_by(LimitPolicyOverride.updated_at.desc(), LimitPolicyOverride.id.desc()) + ).scalars() diff --git a/app/core/ratelimit.py b/app/core/ratelimit.py index f516ae2..e2fc88d 100644 --- a/app/core/ratelimit.py +++ b/app/core/ratelimit.py @@ -75,10 +75,11 @@ def enforce_rate_limit( request: Request, scope: str, subject: str, - limit: int, + limit: int | None, window_sec: float, *, detail: str = "操作过于频繁,请稍后再试", + bucket_suffix: str = "", ) -> None: """在路由内部手动限流,按 (subject, 客户端 IP) 计数。 @@ -87,9 +88,9 @@ def enforce_rate_limit( key = `scope:subject:client_ip`;同一 (subject, IP) 在 window_sec 内超过 limit 次 → 抛 429。 受 [settings.RATE_LIMIT_ENABLED] 总开关控制(与 [rate_limit] 一致)。 """ - if not settings.RATE_LIMIT_ENABLED: + if not settings.RATE_LIMIT_ENABLED or limit is None: return - key = f"{scope}:{subject}:{_client_ip(request)}" + key = f"{scope}:{subject}:{_client_ip(request)}:{bucket_suffix}" if not _hit(key, limit, window_sec): raise HTTPException( status_code=status.HTTP_429_TOO_MANY_REQUESTS, @@ -112,9 +113,10 @@ class RateLimitRule(NamedTuple): """ scope: str - limit: int + limit: int | None window_sec: float detail: str = "操作过于频繁,请稍后再试" + bucket_suffix: str = "" def _peek(key: str, limit: int, window_sec: float) -> bool: @@ -150,7 +152,13 @@ def check_rate_limits(request: Request, subject: str, rules: list[RateLimitRule] return ip = _client_ip(request) for rule in rules: - if not _peek(f"{rule.scope}:{subject}:{ip}", rule.limit, rule.window_sec): + if rule.limit is None: + continue + if not _peek( + f"{rule.scope}:{subject}:{ip}:{rule.bucket_suffix}", + rule.limit, + rule.window_sec, + ): raise HTTPException( status_code=status.HTTP_429_TOO_MANY_REQUESTS, detail=rule.detail, @@ -167,4 +175,9 @@ def record_rate_limits(request: Request, subject: str, rules: list[RateLimitRule return ip = _client_ip(request) for rule in rules: - _commit(f"{rule.scope}:{subject}:{ip}", rule.window_sec) + if rule.limit is None: + continue + _commit( + f"{rule.scope}:{subject}:{ip}:{rule.bucket_suffix}", + rule.window_sec, + ) diff --git a/app/integrations/sms/__init__.py b/app/integrations/sms/__init__.py index 90191bb..98e9149 100644 --- a/app/integrations/sms/__init__.py +++ b/app/integrations/sms/__init__.py @@ -35,38 +35,62 @@ def _fallback(): return _ALL.get(name) -def send_code(phone: str) -> SendResult: +def _send_with_cooldown(provider, phone: str, cooldown_sec: int | None) -> int: + if cooldown_sec is None: + return provider.send_code(phone) + return provider.send_code(phone, cooldown_sec=cooldown_sec) + + +def send_code( + phone: str, + *, + cooldown_sec: int | None = None, +) -> SendResult: """发码:主成功即返回;仅主「供应商不可用(503)」且配置了备时转备补发。 429(本地冷却/超频)、400(手机号无效)不转——不绕过防刷、不为无效号白烧。 - 备也失败则抛备的 SmsError。返回 SendResult(cooldown + 实际渠道 + 是否 fallback)。 + 备也失败则抛备的 SmsError。调用方传入的动态冷却值在主、备渠道保持一致。 """ primary = _primary() fb = _fallback() try: - cooldown = primary.send_code(phone) + cooldown = _send_with_cooldown(primary, phone, cooldown_sec) return SendResult(cooldown_sec=cooldown, provider=_NAME[primary], fallback=False) except SmsError as e: if fb is not None and e.status_code == 503: logger.warning("[SMS] primary=%s 不可用(%s),fallback→%s", _NAME[primary], e, _NAME[fb]) - cooldown = fb.send_code(phone) # 备的冷却/错误码原样透出 + cooldown = _send_with_cooldown(fb, phone, cooldown_sec) return SendResult(cooldown_sec=cooldown, provider=_NAME[fb], fallback=True) raise -def verify_code(phone: str, code: str) -> bool: +def verify_code( + phone: str, + code: str, + *, + max_failed_attempts: int | None = None, +) -> bool: """校验:try-both,遍历「启用的 fallback 链」(主→备),任一命中即 True。 码只存在实际发码那家(fallback 前主已 pop 掉自己的码),另一家 rec is None 即 False、 - 不误判、不累加其防爆破计数。关闭 fallback 时链中只有主,备完全不参与。 + 不误判、不累加其防爆破计数。动态失败次数上限会一致传给主备渠道。 """ chain = [_primary()] fb = _fallback() if fb is not None: chain.append(fb) for prov in chain: - if prov.verify_code(phone, code): # Mode B:纯本地内存比对,不联网 + verified = ( + prov.verify_code(phone, code) + if max_failed_attempts is None + else prov.verify_code( + phone, + code, + max_failed_attempts=max_failed_attempts, + ) + ) + if verified: logger.info("[SMS] verify hit provider=%s", _NAME[prov]) return True return False diff --git a/app/integrations/sms/aliyun.py b/app/integrations/sms/aliyun.py index 9cd7be0..c7c66fc 100644 --- a/app/integrations/sms/aliyun.py +++ b/app/integrations/sms/aliyun.py @@ -47,19 +47,26 @@ _client = None # 惰性构建的 SDK client(模块级缓存) # ============================ 对外:发码 / 校验 ============================ -def send_code(phone: str) -> int: +def send_code(phone: str, *, cooldown_sec: int | None = None) -> int: """发送验证码(阿里云生成+下发)。 Returns: 距下次可发的秒数(= ALIYUN_SMS_INTERVAL_SEC,冷却由阿里云 Interval 侧执行)。 Raises: SmsError(手机号无效 400 / 过频·天级流控 429 / 未配置·未开通·其他 503)。 """ + effective_cooldown = ( + settings.ALIYUN_SMS_INTERVAL_SEC if cooldown_sec is None else cooldown_sec + ) if settings.SMS_MOCK: logger.info("[SMS-aliyun-MOCK] to %s**** (不真发)", phone[:3]) - return settings.ALIYUN_SMS_INTERVAL_SEC + return effective_cooldown if not settings.aliyun_sms_configured: raise SmsError("短信服务未配置(缺阿里云凭证)", status_code=503) - result = _call_send(phone) # 传输/SDK 异常在内部抛 SmsError(503) + result = ( + _call_send(phone) + if cooldown_sec is None + else _call_send(phone, cooldown_sec=effective_cooldown) + ) if result["success"] and result["code"] == "OK": now = time.time() @@ -68,7 +75,7 @@ def send_code(phone: str) -> int: _verify_attempts.pop(phone, None) # 新码 = 新失败预算 _verify_seen.pop(phone, None) logger.info("[SMS-aliyun] sent to %s****", phone[:3]) - return settings.ALIYUN_SMS_INTERVAL_SEC + return effective_cooldown code = result["code"] logger.error("[SMS-aliyun] send failed code=%s msg=%s", code, result["message"]) @@ -78,7 +85,12 @@ def send_code(phone: str) -> int: raise SmsError(msg, status_code=status) -def verify_code(phone: str, code: str) -> bool: +def verify_code( + phone: str, + code: str, + *, + max_failed_attempts: int | None = None, +) -> bool: """校验验证码(阿里云裁决)。 - **mock**:放行任意 N 位数字(provider 无关,同极光)。 @@ -92,8 +104,13 @@ def verify_code(phone: str, code: str) -> bool: # 失败计数是 best-effort:网络调用不持锁(不能锁跨 IO),故并发下同号可能多放行个位数次。 # 无碍——API 层登录频控(设备+IP 5/时)是硬上限,阿里云码有效期 + DuplicatePolicy 亦兜底。 + effective_max_attempts = ( + settings.SMS_MAX_VERIFY_ATTEMPTS + if max_failed_attempts is None + else max_failed_attempts + ) with _lock: - if _verify_attempts.get(phone, 0) >= settings.SMS_MAX_VERIFY_ATTEMPTS: + if _verify_attempts.get(phone, 0) >= effective_max_attempts: return False # 已作废:保持计数(直到 send_code 重置),与极光「达上限即作废」一致 result = _call_check(phone, code) # 传输/SDK 异常在内部抛 SmsError(503) @@ -146,7 +163,7 @@ def _get_client(): return _client -def _call_send(phone: str) -> dict: +def _call_send(phone: str, *, cooldown_sec: int | None = None) -> dict: """调 SendSmsVerifyCode。返回归一化 {success, code, message};import/建 client/调用 任一失败抛 SmsError(503)。""" valid_min = max(1, settings.ALIYUN_SMS_VALID_TIME_SEC // 60) template_param = json.dumps({"code": "##code##", "min": str(valid_min)}, ensure_ascii=False) @@ -160,7 +177,11 @@ def _call_send(phone: str) -> dict: template_param=template_param, code_length=settings.ALIYUN_SMS_CODE_LENGTH, valid_time=settings.ALIYUN_SMS_VALID_TIME_SEC, - interval=settings.ALIYUN_SMS_INTERVAL_SEC, + interval=( + settings.ALIYUN_SMS_INTERVAL_SEC + if cooldown_sec is None + else cooldown_sec + ), scheme_name=settings.ALIYUN_SMS_SCHEME_NAME or None, ) body = _get_client().send_sms_verify_code(req).body diff --git a/app/integrations/sms/chuanglan.py b/app/integrations/sms/chuanglan.py index b3afd28..dfc5092 100644 --- a/app/integrations/sms/chuanglan.py +++ b/app/integrations/sms/chuanglan.py @@ -84,20 +84,23 @@ def _gc(now: float) -> None: _last_sent.pop(p, None) -def send_code(phone: str) -> int: +def send_code(phone: str, *, cooldown_sec: int | None = None) -> int: """发送验证码。 Returns: 距下次可发的秒数(= SMS_SEND_INTERVAL_SEC) Raises: SmsError(过频 429 / 手机号无效 400 / 供应商失败 503) """ now = time.time() + effective_cooldown = ( + settings.SMS_SEND_INTERVAL_SEC if cooldown_sec is None else cooldown_sec + ) # --- lock 内:防刷检查 + 预占(防并发重复发烧钱)--- with _lock: _gc(now) # 顺手清过期内存(超阈值才扫) elapsed = now - _last_sent.get(phone, 0.0) - if elapsed < settings.SMS_SEND_INTERVAL_SEC: - remain = int(settings.SMS_SEND_INTERVAL_SEC - elapsed) + if elapsed < effective_cooldown: + remain = int(effective_cooldown - elapsed) raise SmsError(f"发送过于频繁,请 {remain}s 后再试") code = _gen_code() @@ -122,10 +125,15 @@ def send_code(phone: str) -> int: logger.exception("[SMS-chuanglan] send failed phone=%s****", phone[:3]) raise SmsError("验证码发送失败,请稍后重试", status_code=503) from e - return settings.SMS_SEND_INTERVAL_SEC + return effective_cooldown -def verify_code(phone: str, code: str) -> bool: +def verify_code( + phone: str, + code: str, + *, + max_failed_attempts: int | None = None, +) -> bool: """校验验证码。 - **mock 模式**:放行任意 N 位数字(测试/开发便利,不真校验)。 @@ -136,6 +144,11 @@ def verify_code(phone: str, code: str) -> bool: logger.info("[SMS-chuanglan-MOCK] verify %s for %s****", "ok" if ok else "fail", phone[:3]) return ok + effective_max_attempts = ( + settings.SMS_MAX_VERIFY_ATTEMPTS + if max_failed_attempts is None + else max_failed_attempts + ) with _lock: rec = _codes.get(phone) if rec is None: @@ -143,7 +156,7 @@ def verify_code(phone: str, code: str) -> bool: if time.time() > rec.expires_at: _codes.pop(phone, None) return False - if rec.attempts >= settings.SMS_MAX_VERIFY_ATTEMPTS: + if rec.attempts >= effective_max_attempts: _codes.pop(phone, None) # 试错过多,作废 return False if secrets.compare_digest(code.encode("utf-8"), rec.code.encode("utf-8")): diff --git a/app/integrations/sms/jiguang.py b/app/integrations/sms/jiguang.py index 7406963..c4d9826 100644 --- a/app/integrations/sms/jiguang.py +++ b/app/integrations/sms/jiguang.py @@ -70,20 +70,23 @@ def _gc(now: float) -> None: _last_sent.pop(p, None) -def send_code(phone: str) -> int: +def send_code(phone: str, *, cooldown_sec: int | None = None) -> int: """发送验证码。 Returns: 距下次可发的秒数(= SMS_SEND_INTERVAL_SEC) Raises: SmsError(过频 429 / 当日超限 429 / 供应商失败 503 / 手机号无效 400) """ now = time.time() + effective_cooldown = ( + settings.SMS_SEND_INTERVAL_SEC if cooldown_sec is None else cooldown_sec + ) # --- lock 内:防刷检查 + 预占(防并发重复发烧钱)--- with _lock: _gc(now) # 顺手清过期内存(超阈值才扫) elapsed = now - _last_sent.get(phone, 0.0) - if elapsed < settings.SMS_SEND_INTERVAL_SEC: - remain = int(settings.SMS_SEND_INTERVAL_SEC - elapsed) + if elapsed < effective_cooldown: + remain = int(effective_cooldown - elapsed) raise SmsError(f"发送过于频繁,请 {remain}s 后再试") code = _gen_code() @@ -108,10 +111,15 @@ def send_code(phone: str) -> int: logger.exception("[SMS] send failed phone=%s****", phone[:3]) raise SmsError("验证码发送失败,请稍后重试", status_code=503) from e - return settings.SMS_SEND_INTERVAL_SEC + return effective_cooldown -def verify_code(phone: str, code: str) -> bool: +def verify_code( + phone: str, + code: str, + *, + max_failed_attempts: int | None = None, +) -> bool: """校验验证码。 - **mock 模式**:放行任意 N 位数字(测试/开发便利,不真校验)。 @@ -122,6 +130,11 @@ def verify_code(phone: str, code: str) -> bool: logger.info("[SMS-MOCK] verify %s for %s****", "ok" if ok else "fail", phone[:3]) return ok + effective_max_attempts = ( + settings.SMS_MAX_VERIFY_ATTEMPTS + if max_failed_attempts is None + else max_failed_attempts + ) with _lock: rec = _codes.get(phone) if rec is None: @@ -129,7 +142,7 @@ def verify_code(phone: str, code: str) -> bool: if time.time() > rec.expires_at: _codes.pop(phone, None) return False - if rec.attempts >= settings.SMS_MAX_VERIFY_ATTEMPTS: + if rec.attempts >= effective_max_attempts: _codes.pop(phone, None) # 试错过多,作废 return False if secrets.compare_digest(code.encode("utf-8"), rec.code.encode("utf-8")): diff --git a/app/models/__init__.py b/app/models/__init__.py index 040ce07..9220931 100644 --- a/app/models/__init__.py +++ b/app/models/__init__.py @@ -36,6 +36,7 @@ from app.models.inactivity import ( # noqa: F401 from app.models.invite import InviteRelation # noqa: F401 from app.models.invite_fingerprint import InviteFingerprint # noqa: F401 from app.models.launch_confirm_sample import LaunchConfirmSample # noqa: F401 +from app.models.limit_policy import LimitPolicyOverride # noqa: F401 from app.models.meituan_coupon import MeituanCoupon # noqa: F401 from app.models.notification import Notification # noqa: F401 from app.models.onboarding import OnboardingCompletion # noqa: F401 diff --git a/app/models/limit_policy.py b/app/models/limit_policy.py new file mode 100644 index 0000000..ee1ae7b --- /dev/null +++ b/app/models/limit_policy.py @@ -0,0 +1,77 @@ +"""Per-subject limit policy overrides used by the admin whitelist page.""" +from __future__ import annotations + +from datetime import datetime + +from sqlalchemy import ( + Boolean, + DateTime, + Index, + Integer, + String, + UniqueConstraint, + func, + true, +) +from sqlalchemy.orm import Mapped, mapped_column + +from app.db.base import Base + + +class LimitPolicyOverride(Base): + """One rule override for one phone or device. + + ``reset_at`` is a non-destructive usage baseline. Business records and + security events remain intact; quota readers only count rows at or after + this timestamp. + """ + + __tablename__ = "limit_policy_override" + __table_args__ = ( + UniqueConstraint( + "subject_type", + "subject_value", + "rule_code", + name="uq_limit_policy_subject_rule", + ), + Index( + "ix_limit_policy_lookup", + "subject_type", + "subject_value", + "rule_code", + "enabled", + ), + Index("ix_limit_policy_expires", "expires_at"), + ) + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + subject_type: Mapped[str] = mapped_column(String(16), nullable=False) + subject_value: Mapped[str] = mapped_column(String(128), nullable=False) + rule_code: Mapped[str] = mapped_column(String(64), nullable=False) + # 产品只保留“临时不限/免告警”。即使有内部脚本绕过 API 直接建 ORM + # 对象,也不能再悄悄落成已经下线的 override 模式。 + mode: Mapped[str] = mapped_column(String(24), nullable=False, default="unlimited") + limit_value: Mapped[int | None] = mapped_column(Integer, nullable=True) + enabled: Mapped[bool] = mapped_column( + Boolean, nullable=False, default=True, server_default=true() + ) + starts_at: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True + ) + expires_at: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True + ) + reset_at: Mapped[datetime | None] = mapped_column( + DateTime(timezone=True), nullable=True + ) + reason: Mapped[str | None] = mapped_column(String(256), nullable=True) + created_by_admin_id: Mapped[int | None] = mapped_column(Integer, nullable=True) + created_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), server_default=func.now(), nullable=False + ) + updated_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), + server_default=func.now(), + onupdate=func.now(), + nullable=False, + ) diff --git a/app/repositories/ad_feed_reward.py b/app/repositories/ad_feed_reward.py index ea6ebc9..a90c0c7 100644 --- a/app/repositories/ad_feed_reward.py +++ b/app/repositories/ad_feed_reward.py @@ -5,11 +5,13 @@ """ from __future__ import annotations +from datetime import datetime + from sqlalchemy import func, select from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session -from app.core import rewards +from app.core import limit_policy, rewards from app.core.rewards import cn_today from app.models.ad_feed_reward import AdFeedRewardRecord from app.repositories import wallet as crud_wallet @@ -29,8 +31,14 @@ def _find_by_event(db: Session, client_event_id: str) -> AdFeedRewardRecord | No ).scalar_one_or_none() -def _granted_today(db: Session, user_id: int, reward_date: str) -> int: - return db.execute( +def _granted_today( + db: Session, + user_id: int, + reward_date: str, + *, + reset_at: datetime | None = None, +) -> int: + stmt = ( select(func.count()) .select_from(AdFeedRewardRecord) .where( @@ -38,7 +46,10 @@ def _granted_today(db: Session, user_id: int, reward_date: str) -> int: AdFeedRewardRecord.reward_date == reward_date, AdFeedRewardRecord.status == "granted", ) - ).scalar_one() + ) + if reset_at is not None: + stmt = stmt.where(AdFeedRewardRecord.created_at >= reset_at) + return db.execute(stmt).scalar_one() def granted_unit_total(db: Session, user_id: int) -> int: @@ -121,7 +132,27 @@ def grant_feed_reward( ) return _commit_record(db, rec, client_event_id) - if _granted_today(db, user_id, today) >= rewards.get_ad_daily_limit(db): + daily_policy = limit_policy.resolve_for_user( + db, + "ad.feed.daily", + user_id, + ) + daily_limit = ( + rewards.get_ad_daily_limit(db) + if daily_policy.override_id is None + and daily_policy.bucket_version == "default" + else daily_policy.limit + ) + if ( + daily_limit is not None + and _granted_today( + db, + user_id, + today, + reset_at=daily_policy.reset_at, + ) + >= daily_limit + ): rec = AdFeedRewardRecord( client_event_id=client_event_id, user_id=user_id, diff --git a/app/repositories/ad_reward.py b/app/repositories/ad_reward.py index 3854f12..cd6c8f1 100644 --- a/app/repositories/ad_reward.py +++ b/app/repositories/ad_reward.py @@ -10,13 +10,13 @@ """ from __future__ import annotations -from datetime import datetime, timezone +from datetime import UTC, datetime from sqlalchemy import func, select from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session -from app.core import rewards +from app.core import limit_policy, rewards from app.core.ad_cooldown import compute_cooldown from app.core.rewards import DAILY_AD_WATCH_SECONDS_LIMIT, cn_today from app.models.ad_reward import AdRewardRecord @@ -96,8 +96,14 @@ def round_coin_total(db: Session, user_id: int, boost_round_id: str) -> int: ) -def _granted_today(db: Session, user_id: int, reward_date: str) -> int: - return db.execute( +def _granted_today( + db: Session, + user_id: int, + reward_date: str, + *, + reset_at: datetime | None = None, +) -> int: + stmt = ( select(func.count()) .select_from(AdRewardRecord) .where( @@ -106,7 +112,10 @@ def _granted_today(db: Session, user_id: int, reward_date: str) -> int: AdRewardRecord.status == "granted", AdRewardRecord.reward_scene == "reward_video", ) - ).scalar_one() + ) + if reset_at is not None: + stmt = stmt.where(AdRewardRecord.created_at >= reset_at) + return db.execute(stmt).scalar_one() def _granted_cumulative(db: Session, user_id: int) -> int: @@ -165,7 +174,27 @@ def grant_ad_reward( DAILY_AD_WATCH_SECONDS_LIMIT > 0 and watched_seconds_today(db, user_id, today=today) >= DAILY_AD_WATCH_SECONDS_LIMIT ) - over_count = _granted_today(db, user_id, today) >= rewards.get_ad_daily_limit(db) + daily_policy = limit_policy.resolve_for_user( + db, + "ad.reward_video.daily", + user_id, + ) + daily_limit = ( + rewards.get_ad_daily_limit(db) + if daily_policy.override_id is None + and daily_policy.bucket_version == "default" + else daily_policy.limit + ) + over_count = ( + daily_limit is not None + and _granted_today( + db, + user_id, + today, + reset_at=daily_policy.reset_at, + ) + >= daily_limit + ) if over_time or over_count: rec = AdRewardRecord( trans_id=trans_id, user_id=user_id, coin=0, status="capped", @@ -321,19 +350,24 @@ def _commit_record(db: Session, rec: AdRewardRecord, trans_id: str) -> AdRewardR return rec -def _granted_times_today_desc(db: Session, user_id: int, reward_date: str) -> list[datetime]: +def _granted_times_today_desc( + db: Session, + user_id: int, + reward_date: str, + *, + reset_at: datetime | None = None, +) -> list[datetime]: """当日 status=granted 记录的 created_at,按时间倒序(最新在前)——冷却策略的输入数据。""" - return list( - db.execute( - select(AdRewardRecord.created_at) - .where( + stmt = select(AdRewardRecord.created_at).where( AdRewardRecord.user_id == user_id, AdRewardRecord.reward_date == reward_date, AdRewardRecord.status == "granted", AdRewardRecord.reward_scene == "reward_video", ) - .order_by(AdRewardRecord.created_at.desc()) - ).scalars() + if reset_at is not None: + stmt = stmt.where(AdRewardRecord.created_at >= reset_at) + return list( + db.execute(stmt.order_by(AdRewardRecord.created_at.desc())).scalars() ) @@ -349,16 +383,47 @@ def today_status( 旧客户端兼容,当前 DAILY_AD_WATCH_SECONDS_LIMIT=0 表示不启用时长闸。 """ today = cn_today().isoformat() - granted_desc = _granted_times_today_desc(db, user_id, today) + daily_policy = limit_policy.resolve_for_user( + db, + "ad.reward_video.daily", + user_id, + ) + cooldown_policy = limit_policy.resolve_for_user( + db, + "ad.reward_video.cooldown", + user_id, + ) + daily_limit = ( + rewards.get_ad_daily_limit(db) + if daily_policy.override_id is None + and daily_policy.bucket_version == "default" + else daily_policy.limit + ) + cooldown_seconds = ( + rewards.get_ad_cooldown_sec(db) + if cooldown_policy.override_id is None + and cooldown_policy.bucket_version == "default" + else (cooldown_policy.limit or 0) + ) + granted_desc = _granted_times_today_desc( + db, + user_id, + today, + reset_at=daily_policy.reset_at, + ) state = compute_cooldown( granted_desc, - datetime.now(timezone.utc), + datetime.now(UTC), round_size=rewards.get_ad_round_count(db), - cooldown_seconds=rewards.get_ad_cooldown_sec(db), + cooldown_seconds=cooldown_seconds, ) return ( len(granted_desc), - rewards.get_ad_daily_limit(db), + ( + daily_limit + if daily_limit is not None + else limit_policy.get_rule("ad.reward_video.daily").max_value + ), 0, state.round_count, state.cooldown_until, diff --git a/app/repositories/comparison.py b/app/repositories/comparison.py index 39ccf4b..1ad23bd 100644 --- a/app/repositories/comparison.py +++ b/app/repositories/comparison.py @@ -444,6 +444,8 @@ def reserve_daily_start( business_type: str = "food", device_id: str | None = None, now: datetime | None = None, + limit: int | None = DAILY_COMPARE_START_LIMIT, + reset_at: datetime | None = None, ) -> tuple[ComparisonRecord, int]: """Atomically reserve one of a user's 100 Beijing-day comparison starts. @@ -469,6 +471,11 @@ def reserve_daily_start( if existing_at.tzinfo is not None: existing_at = existing_at.astimezone(CN_TZ).replace(tzinfo=None) day_start = existing_at.replace(hour=0, minute=0, second=0, microsecond=0) + if reset_at is not None: + reset_start = reset_at + if reset_start.tzinfo is not None: + reset_start = reset_start.astimezone(CN_TZ).replace(tzinfo=None) + day_start = max(day_start, reset_start) day_end = day_start + timedelta(days=1) used = db.scalar( select(func.count(ComparisonRecord.id)).where( @@ -483,6 +490,11 @@ def reserve_daily_start( if current.tzinfo is not None: current = current.astimezone(CN_TZ).replace(tzinfo=None) day_start = current.replace(hour=0, minute=0, second=0, microsecond=0) + if reset_at is not None: + reset_start = reset_at + if reset_start.tzinfo is not None: + reset_start = reset_start.astimezone(CN_TZ).replace(tzinfo=None) + day_start = max(day_start, reset_start) day_end = day_start + timedelta(days=1) used = db.scalar( select(func.count(ComparisonRecord.id)).where( @@ -491,7 +503,7 @@ def reserve_daily_start( ComparisonRecord.created_at < day_end, ) ) or 0 - if used >= DAILY_COMPARE_START_LIMIT: + if limit is not None and used >= limit: raise DailyCompareStartLimitExceeded rec = ComparisonRecord( diff --git a/app/repositories/guide_video.py b/app/repositories/guide_video.py index 5051fcb..c4c1592 100644 --- a/app/repositories/guide_video.py +++ b/app/repositories/guide_video.py @@ -10,7 +10,7 @@ from sqlalchemy import func, select, update from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session -from app.core import media, rewards +from app.core import limit_policy, media, rewards from app.core.config import settings from app.models.app_config import AppConfig from app.models.guide_video import GuideVideoPlay @@ -25,7 +25,9 @@ _KEY_BY_SCENE = { BIZ_TYPE = "guide_video" CIRCLE_COUNT = 10 PLAN_TTL = timedelta(minutes=10) -MIN_PLAYS = 1 +# 0 表示全局暂停播放;白名单页与旧引导视频配置页必须接受同一口径, +# 否则在白名单页设为 0 后,旧页面连金币/开关等无关字段也无法保存。 +MIN_PLAYS = 0 MAX_PLAYS_LIMIT = 50 MIN_REWARD_COIN = 10 REWARD_COIN_LIMIT = 10_000 @@ -95,6 +97,10 @@ def _public_config(cfg: dict[str, Any]) -> dict[str, Any]: def get_config(db: Session, scene: str = "coupon") -> dict[str, Any]: row = db.get(AppConfig, _config_key(scene)) cfg = _public_config(_merge(row.value if row is not None else None)) + if scene == "coupon": + cfg["max_plays"] = limit_policy.resolve( + db, "guide.video.lifetime" + ).global_limit cfg["updated_at"] = row.updated_at.isoformat() if row is not None and row.updated_at else None return cfg @@ -153,7 +159,27 @@ def update_config( raw["max_plays"] = candidate_plays raw["reward_coin"] = candidate_reward raw["config_version"] = int(raw.get("config_version") or 0) + 1 - after = _write(db, raw, scene=scene, admin_id=admin_id, commit=commit) + after = _write( + db, + raw, + scene=scene, + admin_id=admin_id, + commit=False, + ) + if scene == "coupon" and max_plays is not None: + # 兼容仍调用旧专用接口的客户端/脚本,并把统一策略全局值一并更新。 + limit_policy.set_global_limits( + db, + {"guide.video.lifetime": candidate_plays}, + admin_id=admin_id, + commit=False, + ) + # set_global_limits 已 flush;即使由 admin 路由统一在外层提交,也要把 + # 同一事务内的最新统一策略值写进审计 after 快照。 + after = get_config(db, scene) + if commit: + db.commit() + after = get_config(db, scene) return before, after @@ -191,20 +217,52 @@ def set_video( return before, after -def used_plays(db: Session, user_id: int, scene: str = "coupon") -> int: - return int( - db.execute( - select(func.count()).select_from(GuideVideoPlay).where( - GuideVideoPlay.user_id == user_id, - GuideVideoPlay.scene == scene, - GuideVideoPlay.status != "prepared", - ) - ).scalar_one() +def used_plays( + db: Session, + user_id: int, + scene: str = "coupon", + *, + reset_at: datetime | None = None, +) -> int: + stmt = select(func.count()).select_from(GuideVideoPlay).where( + GuideVideoPlay.user_id == user_id, + GuideVideoPlay.scene == scene, + GuideVideoPlay.status != "prepared", ) + if reset_at is not None: + # guide_video_play 使用北京时间 naive 墙钟;白名单重置点是带时区时间。 + reset_value = reset_at.astimezone(rewards.CN_TZ).replace(tzinfo=None) + stmt = stmt.where(GuideVideoPlay.started_at >= reset_value) + return int(db.execute(stmt).scalar_one()) -def _prepare_miss(scene: str, reason: str, cfg: dict[str, Any], used: int) -> dict[str, Any]: - maximum = int(cfg.get("max_plays") or 0) +def _effective_quota( + db: Session, + user_id: int, + scene: str, + cfg: dict[str, Any], +) -> tuple[int | None, datetime | None]: + if scene != "coupon": + return int(cfg.get("max_plays") or 0), None + policy = limit_policy.resolve_for_user( + db, "guide.video.lifetime", user_id + ) + return policy.limit, policy.reset_at + + +def _remaining(maximum: int | None, used: int) -> int: + # API 字段保持整数兼容;不限时使用足够大的展示值,不参与服务端判定。 + exposed_maximum = maximum if maximum is not None else 1_000_000 + return max(0, exposed_maximum - used) + + +def _prepare_miss( + scene: str, + reason: str, + cfg: dict[str, Any], + used: int, + maximum: int | None, +) -> dict[str, Any]: return { "should_play": False, "reason": reason, @@ -218,23 +276,27 @@ def _prepare_miss(scene: str, reason: str, cfg: dict[str, Any], used: int) -> di "reward_coin": int(cfg.get("reward_coin") or 0), "reward_per_circle": int(cfg.get("reward_coin") or 0) // CIRCLE_COUNT, "seq": used, - "remaining": max(0, maximum - used), + "remaining": _remaining(maximum, used), "expires_at": None, } def prepare_play(db: Session, user_id: int, *, scene: str = "coupon") -> dict[str, Any]: cfg = get_config(db, scene) - used = used_plays(db, user_id, scene) + maximum, reset_at = _effective_quota(db, user_id, scene, cfg) + used = used_plays(db, user_id, scene, reset_at=reset_at) video_url = str(cfg.get("video_url") or "").strip() duration = int(cfg.get("duration_ms") or 0) - maximum = int(cfg.get("max_plays") or 0) if not cfg.get("enabled"): - return _prepare_miss(scene, "disabled", cfg, used) + return _prepare_miss(scene, "disabled", cfg, used, maximum) if not video_url or cfg.get("analysis_status") != "valid" or duration <= 0: - return _prepare_miss(scene, "video_unavailable", cfg, used) - if used >= maximum: - return _prepare_miss(scene, "play_limit_reached", cfg, used) + return _prepare_miss( + scene, "video_unavailable", cfg, used, maximum + ) + if maximum is not None and used >= maximum: + return _prepare_miss( + scene, "play_limit_reached", cfg, used, maximum + ) now = _now() play = GuideVideoPlay( @@ -268,7 +330,7 @@ def prepare_play(db: Session, user_id: int, *, scene: str = "coupon") -> dict[st "reward_coin": play.coin, "reward_per_circle": play.coin // CIRCLE_COUNT, "seq": used + 1, - "remaining": max(0, maximum - used - 1), + "remaining": _remaining(maximum, used + 1), "expires_at": play.expires_at.isoformat(), } @@ -282,7 +344,12 @@ def _find_play(db: Session, user_id: int, token: str) -> GuideVideoPlay | None: ).scalar_one_or_none() -def _start_out(play: GuideVideoPlay, maximum: int, status: str) -> dict[str, Any]: +def _start_out( + play: GuideVideoPlay, + maximum: int | None, + used: int, + status: str, +) -> dict[str, Any]: assert play.started_at is not None and play.seq is not None and play.video_url return { "started": True, @@ -297,7 +364,7 @@ def _start_out(play: GuideVideoPlay, maximum: int, status: str) -> dict[str, Any "reward_coin": play.coin, "reward_per_circle": play.coin // CIRCLE_COUNT, "seq": play.seq, - "remaining": max(0, maximum - play.seq), + "remaining": _remaining(maximum, used), "started_at": play.started_at.isoformat(), } @@ -307,9 +374,14 @@ def start_play(db: Session, user_id: int, *, play_token: str) -> dict[str, Any]: if play is None: raise PlayStateError("play_not_found", "播放计划不存在") cfg = get_config(db, play.scene) - maximum = int(cfg["max_plays"]) + maximum, reset_at = _effective_quota( + db, user_id, play.scene, cfg + ) + used = used_plays( + db, user_id, play.scene, reset_at=reset_at + ) if play.status in {"started", "completed"}: - return _start_out(play, maximum, "already_started") + return _start_out(play, maximum, used, "already_started") if play.status != "prepared": raise PlayStateError("play_not_found", "播放计划不可用") now = _now() @@ -323,10 +395,11 @@ def start_play(db: Session, user_id: int, *, play_token: str) -> dict[str, Any]: "config_changed", "视频配置已变化,请重新获取", reprepare_required=True, ) - used = used_plays(db, user_id, play.scene) - if used >= maximum: + if maximum is not None and used >= maximum: raise PlayStateError("play_limit_reached", "播放次数已用完") - play.seq = used + 1 + # seq 在数据库中是账号+场景生命周期唯一值;白名单重置只重置额度, + # 不能从 1 重新编号,否则会与历史记录冲突。 + play.seq = used_plays(db, user_id, play.scene) + 1 play.status = "started" play.started_at = now try: @@ -335,7 +408,7 @@ def start_play(db: Session, user_id: int, *, play_token: str) -> dict[str, Any]: except IntegrityError as exc: db.rollback() raise PlayStateError("play_limit_reached", "并发起播冲突,请重新获取") from exc - return _start_out(play, maximum, "started") + return _start_out(play, maximum, used + 1, "started") def _coin_balance(db: Session, user_id: int) -> int: diff --git a/app/repositories/phone_rebind.py b/app/repositories/phone_rebind.py index 34275d3..0a52e96 100644 --- a/app/repositories/phone_rebind.py +++ b/app/repositories/phone_rebind.py @@ -2,7 +2,7 @@ from __future__ import annotations import math -from datetime import datetime, timedelta, timezone +from datetime import UTC, datetime, timedelta from sqlalchemy import func, select from sqlalchemy.orm import Session @@ -10,9 +10,22 @@ from sqlalchemy.orm import Session from app.models.phone_rebind_log import PhoneRebindLog -def rebound_within_days(db: Session, phone: str, days: int) -> bool: +def rebound_within_days( + db: Session, + phone: str, + days: int, + *, + reset_at: datetime | None = None, +) -> bool: """该手机号在最近 days 天内是否换绑过(命中 → 禁止再次换绑)。""" - since = datetime.now(timezone.utc) - timedelta(days=days) + since = datetime.now(UTC) - timedelta(days=days) + if reset_at is not None: + reset_value = ( + reset_at.replace(tzinfo=UTC) + if reset_at.tzinfo is None + else reset_at.astimezone(UTC) + ) + since = max(since, reset_value) stmt = ( select(PhoneRebindLog.id) .where(PhoneRebindLog.phone == phone, PhoneRebindLog.rebound_at >= since) @@ -21,16 +34,25 @@ def rebound_within_days(db: Session, phone: str, days: int) -> bool: return db.execute(stmt).first() is not None -def remaining_block_days(db: Session, phone: str, days: int) -> int: +def remaining_block_days( + db: Session, + phone: str, + days: int, + *, + reset_at: datetime | None = None, +) -> int: """距离该手机号可再次换绑还剩几天(向上取整;无记录返回 0)。""" + conditions = [PhoneRebindLog.phone == phone] + if reset_at is not None: + conditions.append(PhoneRebindLog.rebound_at >= reset_at) last = db.execute( - select(func.max(PhoneRebindLog.rebound_at)).where(PhoneRebindLog.phone == phone) + select(func.max(PhoneRebindLog.rebound_at)).where(*conditions) ).scalar_one_or_none() if last is None: return 0 if last.tzinfo is None: # SQLite 取回 naive datetime,按 UTC 归一 - last = last.replace(tzinfo=timezone.utc) - remaining = (last + timedelta(days=days) - datetime.now(timezone.utc)).total_seconds() + last = last.replace(tzinfo=UTC) + remaining = (last + timedelta(days=days) - datetime.now(UTC)).total_seconds() return max(0, math.ceil(remaining / 86400)) diff --git a/app/repositories/risk.py b/app/repositories/risk.py index d3315f1..336f325 100644 --- a/app/repositories/risk.py +++ b/app/repositories/risk.py @@ -9,15 +9,10 @@ from sqlalchemy import func, select from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session -from app.core.config_schema import ( - RISK_COMPARE_DAILY_THRESHOLD_KEY, - RISK_ONECLICK_DAILY_THRESHOLD_KEY, - RISK_SMS_HOURLY_THRESHOLD_KEY, -) +from app.core import limit_policy from app.models.app_config import AppConfig from app.models.comparison import ComparisonRecord from app.models.risk import BehaviorEvent, RiskIncident, SubjectRestriction -from app.repositories import app_config CN_TZ = ZoneInfo("Asia/Shanghai") @@ -41,7 +36,6 @@ class RuleSpec: code: str event_type: str subject_type: str - threshold_key: str window: str count_outcomes: tuple[str, ...] @@ -51,7 +45,6 @@ RULES: dict[str, RuleSpec] = { code=RULE_SMS_HOURLY, event_type=EVENT_SMS_SEND, subject_type="device", - threshold_key=RISK_SMS_HOURLY_THRESHOLD_KEY, window="hour", count_outcomes=("success",), ), @@ -59,22 +52,23 @@ RULES: dict[str, RuleSpec] = { code=RULE_ONECLICK_DAILY, event_type=EVENT_ONECLICK_LOGIN, subject_type="device", - threshold_key=RISK_ONECLICK_DAILY_THRESHOLD_KEY, window="day", count_outcomes=("success", "failed"), ), } -RULE_THRESHOLD_KEYS: dict[str, str] = { - RULE_SMS_HOURLY: RISK_SMS_HOURLY_THRESHOLD_KEY, - RULE_ONECLICK_DAILY: RISK_ONECLICK_DAILY_THRESHOLD_KEY, - RULE_COMPARE_DAILY: RISK_COMPARE_DAILY_THRESHOLD_KEY, +RISK_LIMIT_RULE_CODES: dict[str, str] = { + RULE_SMS_HOURLY: "risk.sms.hourly", + RULE_ONECLICK_DAILY: "risk.oneclick.daily", + RULE_COMPARE_DAILY: "risk.compare.daily", } def get_rule_threshold(db: Session, rule_code: str) -> int: """读取规则当前阈值;配置表为空时回退上线前的 5/20/100 默认值。""" - return int(app_config.get_value(db, RULE_THRESHOLD_KEYS[rule_code])) + return limit_policy.resolve( + db, RISK_LIMIT_RULE_CODES[rule_code] + ).global_limit def utcnow() -> datetime: @@ -257,22 +251,34 @@ def _upsert_incident( def evaluate_behavior_rule( - db: Session, *, rule_code: str, subject_id: str, at: datetime + db: Session, + *, + rule_code: str, + subject_id: str, + at: datetime, + threshold: int | None = None, + subject_reset_at: datetime | None = None, ) -> RiskIncident | None: spec = RULES[rule_code] - threshold = get_rule_threshold(db, rule_code) + effective_threshold = threshold or get_rule_threshold(db, rule_code) window_key, window_start, end = _window_bounds(at, spec.window) reset_at = get_rule_reset_at(db, rule_code) - start = max(window_start, reset_at) if reset_at else window_start + baselines = [value for value in (reset_at, subject_reset_at) if value is not None] + start = max(window_start, *baselines) if baselines else window_start count, first_at, last_at, triggered_at = _event_stats( db, spec, subject_id=subject_id, start=start, end=end, - threshold=threshold, + threshold=effective_threshold, ) - if count < threshold or first_at is None or last_at is None or triggered_at is None: + if ( + count < effective_threshold + or first_at is None + or last_at is None + or triggered_at is None + ): return None return _upsert_incident( db, @@ -300,7 +306,6 @@ def reconcile_behavior_rule( """按当前阈值重算短信/一键登录当前窗口,并收起已不再命中的待处理告警。""" spec = RULES[rule_code] current = at or utcnow() - threshold = get_rule_threshold(db, rule_code) window_key, window_start, end = _window_bounds(current, spec.window) reset_at = get_rule_reset_at(db, rule_code) start = max(window_start, reset_at) if reset_at else window_start @@ -311,23 +316,39 @@ def reconcile_behavior_rule( BehaviorEvent.occurred_at >= start, BehaviorEvent.occurred_at < end, ) - qualifying_rows = db.execute( + subject_rows = db.execute( select( BehaviorEvent.subject_id, func.max(BehaviorEvent.occurred_at), + func.max(BehaviorEvent.phone), ) .where(*filters) .group_by(BehaviorEvent.subject_id) - .having(func.count(BehaviorEvent.id) >= threshold) ).all() - qualifying = {str(subject_id) for subject_id, _ in qualifying_rows} - for subject_id, last_at in qualifying_rows: - evaluate_behavior_rule( + qualifying: set[str] = set() + policy_code = { + RULE_SMS_HOURLY: "risk.sms.hourly", + RULE_ONECLICK_DAILY: "risk.oneclick.daily", + }[rule_code] + for subject_id, last_at, phone in subject_rows: + policy = limit_policy.resolve( + db, + policy_code, + phone=phone, + device=str(subject_id), + ) + if policy.suppressed or policy.limit is None: + continue + incident = evaluate_behavior_rule( db, rule_code=rule_code, subject_id=str(subject_id), at=last_at or current, + threshold=policy.limit, + subject_reset_at=policy.reset_at, ) + if incident is not None: + qualifying.add(str(subject_id)) open_incidents = db.scalars( select(RiskIncident).where( @@ -383,7 +404,29 @@ def record_behavior_event( db.add(event) db.flush() if evaluate_rule: - evaluate_behavior_rule(db, rule_code=evaluate_rule, subject_id=subject_id, at=at) + policy_code = { + RULE_SMS_HOURLY: "risk.sms.hourly", + RULE_ONECLICK_DAILY: "risk.oneclick.daily", + }.get(evaluate_rule) + policy = ( + limit_policy.resolve( + db, + policy_code, + phone=phone, + device=device_id or subject_id, + ) + if policy_code + else None + ) + if policy is None or not policy.suppressed: + evaluate_behavior_rule( + db, + rule_code=evaluate_rule, + subject_id=subject_id, + at=at, + threshold=policy.limit if policy else None, + subject_reset_at=policy.reset_at if policy else None, + ) if commit: db.commit() db.refresh(event) @@ -396,20 +439,33 @@ def sync_compare_incident( user_id: int, at: datetime, threshold: int | None = None, + device_id: str | None = None, commit: bool = True, ) -> RiskIncident | None: - effective_threshold = threshold or get_rule_threshold(db, RULE_COMPARE_DAILY) + policy = limit_policy.resolve_for_user( + db, + "risk.compare.daily", + user_id, + device=device_id, + ) + if policy.suppressed: + return None + effective_threshold = threshold or policy.limit or get_rule_threshold( + db, RULE_COMPARE_DAILY + ) # comparison_record 的既有写入口统一落“北京时间 naive”时间;这里必须沿用同一 # 口径,否则 SQLite/PG session timezone 不同时会把凌晨记录算到前一天。 local = at.astimezone(CN_TZ).replace(tzinfo=None) if at.tzinfo else at window_start = local.replace(hour=0, minute=0, second=0, microsecond=0) end = window_start + timedelta(days=1) window_key = window_start.strftime("%Y-%m-%d") - reset_at = get_rule_reset_at(db, RULE_COMPARE_DAILY) - reset_local = ( - reset_at.astimezone(CN_TZ).replace(tzinfo=None) if reset_at else None - ) - start = max(window_start, reset_local) if reset_local else window_start + global_reset_at = get_rule_reset_at(db, RULE_COMPARE_DAILY) + baselines = [ + value.astimezone(CN_TZ).replace(tzinfo=None) + for value in (global_reset_at, policy.reset_at) + if value is not None + ] + start = max(window_start, *baselines) if baselines else window_start filters = ( ComparisonRecord.user_id == user_id, ComparisonRecord.created_at >= start, @@ -464,7 +520,6 @@ def reconcile_compare_rule( reset_at.astimezone(CN_TZ).replace(tzinfo=None) if reset_at else None ) start = max(window_start, reset_local) if reset_local else window_start - threshold = get_rule_threshold(db, RULE_COMPARE_DAILY) rows = db.execute( select( ComparisonRecord.user_id, @@ -476,17 +531,17 @@ def reconcile_compare_rule( ComparisonRecord.created_at < end, ) .group_by(ComparisonRecord.user_id) - .having(func.count(ComparisonRecord.id) >= threshold) ).all() - qualifying = {str(user_id) for user_id, _ in rows} + qualifying: set[str] = set() for user_id, last_at in rows: - sync_compare_incident( + incident = sync_compare_incident( db, user_id=int(user_id), at=last_at or current, - threshold=threshold, commit=False, ) + if incident is not None: + qualifying.add(str(user_id)) open_incidents = db.scalars( select(RiskIncident).where( diff --git a/app/schemas/compare_record.py b/app/schemas/compare_record.py index c30e495..4a6b617 100644 --- a/app/schemas/compare_record.py +++ b/app/schemas/compare_record.py @@ -222,9 +222,9 @@ class CompareStartReserveIn(BaseModel): class CompareStartReserveOut(BaseModel): - limit: int + limit: int | None used: int - remaining: int + remaining: int | None class CompareStatsOut(BaseModel): diff --git a/tests/test_admin_config.py b/tests/test_admin_config.py index cfdf29a..76992d3 100644 --- a/tests/test_admin_config.py +++ b/tests/test_admin_config.py @@ -68,12 +68,12 @@ def test_list_config(admin_client: TestClient, token: str) -> None: r = admin_client.get("/admin/api/config", headers=_auth(token)) assert r.status_code == 200, r.text items = {i["key"]: i for i in r.json()} - # 非 hidden 项照常返回;看广告组保留可见的:每日上限 / 单次金币上限 / 关闭后冷却。 - assert "signin_rewards" in items and "ad_daily_limit" in items and "ad_cooldown_sec" in items - # hidden 项(任务/里程碑、首页轮播数据源、看广告组的单次金币/每轮次数/信息流广告开关)不在配置页返回。 + # 非 hidden 项照常返回;广告次数上限迁到「白名单」统一配置。 + assert "signin_rewards" in items and "ad_cooldown_sec" in items + # hidden 项(任务/里程碑、首页轮播数据源、广告次数/单次金币/每轮次数/信息流广告开关)不在配置页返回。 for hidden_key in ( "task_rewards", "record_milestones", "marquee_feed_mode", - "ad_reward_coin", "ad_round_count", "comparing_ad_enabled", + "ad_daily_limit", "ad_reward_coin", "ad_round_count", "comparing_ad_enabled", ): assert hidden_key not in items, f"{hidden_key} 应被 hidden 过滤" assert items["signin_rewards"]["value"] == [ @@ -124,6 +124,50 @@ def test_update_ad_limit_takes_effect(admin_client: TestClient, token: str) -> N db.close() +def test_list_config_reads_limit_values_from_global_bundle( + admin_client: TestClient, + token: str, +) -> None: + changed = admin_client.patch( + "/admin/api/config/ad_cooldown_sec", + json={"value": 17}, + headers=_auth(token), + ) + assert changed.status_code == 200, changed.text + + items = { + item["key"]: item + for item in admin_client.get( + "/admin/api/config", + headers=_auth(token), + ).json() + } + assert items["ad_cooldown_sec"]["value"] == 17 + assert items["ad_cooldown_sec"]["overridden"] is True + + +@pytest.mark.parametrize( + ("key", "value"), + [ + ("ad_daily_limit", 0), + ("ad_daily_limit", 100_001), + ("ad_cooldown_sec", 86_401), + ], +) +def test_update_limit_config_rejects_out_of_range_values( + admin_client: TestClient, + token: str, + key: str, + value: int, +) -> None: + response = admin_client.patch( + f"/admin/api/config/{key}", + json={"value": value}, + headers=_auth(token), + ) + assert response.status_code == 400, response.text + + def test_update_bool_config(admin_client: TestClient, token: str) -> None: # 提现自动对账开关默认 True items = { diff --git a/tests/test_admin_roles.py b/tests/test_admin_roles.py index 430ae3e..7a5eb82 100644 --- a/tests/test_admin_roles.py +++ b/tests/test_admin_roles.py @@ -73,6 +73,7 @@ def test_monitoring_audit_catalog_and_api_permissions( "analytics-health", "event-logs", "audit-logs", + "limit-whitelist", ] # 运营默认可查风控和设备存活,不能绕过导航直调技术/审计接口。 @@ -82,6 +83,9 @@ def test_monitoring_audit_catalog_and_api_permissions( assert admin_client.get( "/admin/api/device-liveness/stats", headers=_auth(operator_token) ).status_code == 200 + assert admin_client.get( + "/admin/api/limit-whitelist/rules", headers=_auth(operator_token) + ).status_code == 200 for path in ( "/admin/api/analytics-health/overview?date_from=2026-07-01T00:00:00Z&date_to=2026-07-02T00:00:00Z", "/admin/api/event-logs", @@ -89,13 +93,14 @@ def test_monitoring_audit_catalog_and_api_permissions( ): assert admin_client.get(path, headers=_auth(operator_token)).status_code == 403 - # 技术角色默认拥有监控审计组全部五项权限。 + # 技术角色默认拥有监控审计组全部六项权限。 for path in ( "/admin/api/risk-monitor/summary", "/admin/api/device-liveness/stats", "/admin/api/analytics-health/overview?date_from=2026-07-01T00:00:00Z&date_to=2026-07-02T00:00:00Z", "/admin/api/event-logs", "/admin/api/audit-logs", + "/admin/api/limit-whitelist/rules", ): assert admin_client.get(path, headers=_auth(tech_token)).status_code == 200 @@ -199,7 +204,7 @@ def test_builtin_roles_labels_and_pages(admin_client, super_token) -> None: } assert set(roles["tech"]["pages"]) == { "dashboard", "risk-monitor", "device-liveness", "analytics-health", "config", - "ad-revenue", "huawei-review", "event-logs", "audit-logs", + "ad-revenue", "huawei-review", "event-logs", "audit-logs", "limit-whitelist", } diff --git a/tests/test_limit_whitelist.py b/tests/test_limit_whitelist.py new file mode 100644 index 0000000..87a3655 --- /dev/null +++ b/tests/test_limit_whitelist.py @@ -0,0 +1,1547 @@ +from __future__ import annotations + +from copy import deepcopy +from datetime import UTC, datetime, timedelta +from uuid import uuid4 + +import pytest +from fastapi.testclient import TestClient +from sqlalchemy import select + +from app.admin.main import admin_app +from app.admin.repositories import admin_user as admin_repo +from app.admin.security import create_admin_token +from app.api.v1 import auth as auth_api +from app.core import limit_policy +from app.core.config_schema import LIMIT_POLICY_GLOBAL_KEY +from app.core.security import create_token +from app.db.session import SessionLocal +from app.main import app +from app.models.admin import AdminAuditLog +from app.models.app_config import AppConfig +from app.models.comparison import ComparisonRecord +from app.models.limit_policy import LimitPolicyOverride +from app.models.risk import RiskIncident +from app.repositories import guide_video as guide_video_repo +from app.repositories import risk as risk_repo +from app.repositories import user as user_repo + + +def _snapshot_configs(keys: list[str]) -> dict[str, tuple[bool, object, int | None]]: + with SessionLocal() as db: + snapshot = {} + for key in keys: + row = db.get(AppConfig, key) + snapshot[key] = ( + row is not None, + deepcopy(row.value) if row is not None else None, + row.updated_by_admin_id if row is not None else None, + ) + return snapshot + + +def _restore_configs(snapshot: dict[str, tuple[bool, object, int | None]]) -> None: + with SessionLocal() as db: + for key, (existed, value, admin_id) in snapshot.items(): + row = db.get(AppConfig, key) + if not existed: + if row is not None: + db.delete(row) + continue + if row is None: + db.add( + AppConfig( + key=key, + value=deepcopy(value), + updated_by_admin_id=admin_id, + ) + ) + else: + row.value = deepcopy(value) + row.updated_by_admin_id = admin_id + db.commit() + + +@pytest.fixture() +def admin_headers() -> dict[str, str]: + username = f"limit_admin_{uuid4().hex[:8]}" + with SessionLocal() as db: + admin = admin_repo.create_admin( + db, + username=username, + password="limit-admin-pass", + role="super_admin", + ) + admin_id = admin.id + token, _ = create_admin_token(admin_id=admin_id, role="super_admin") + return {"Authorization": f"Bearer {token}"} + + +@pytest.fixture() +def custom_admin_headers() -> dict[str, str]: + username = f"limit_custom_{uuid4().hex[:8]}" + with SessionLocal() as db: + admin = admin_repo.create_admin( + db, + username=username, + password="limit-admin-pass", + role="custom", + pages_override=["limit-whitelist"], + ) + admin_id = admin.id + token, _ = create_admin_token(admin_id=admin_id, role="custom") + return {"Authorization": f"Bearer {token}"} + + +def test_custom_admin_with_page_permission_can_manage_whitelist( + custom_admin_headers, +) -> None: + phone = f"139{int(uuid4().hex[:8], 16) % 100000000:08d}" + expires_at = (datetime.now(UTC) + timedelta(hours=2)).isoformat() + + with TestClient(admin_app) as client: + created = client.post( + "/admin/api/limit-whitelist/bulk", + headers=custom_admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_codes": ["compare.start.daily"], + "expires_at": expires_at, + "reason": "custom role page permission", + }, + ) + assert created.status_code == 201, created.text + override_id = created.json()[0]["id"] + deleted = client.delete( + f"/admin/api/limit-whitelist/{override_id}", + headers=custom_admin_headers, + ) + assert deleted.status_code == 204, deleted.text + + +def test_bulk_create_reactivates_existing_rows(admin_headers) -> None: + phone = f"139{int(uuid4().hex[:8], 16) % 100000000:08d}" + first_expiry = (datetime.now(UTC) + timedelta(hours=1)).isoformat() + second_expiry = (datetime.now(UTC) + timedelta(hours=3)).isoformat() + + with TestClient(admin_app) as client: + first = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_codes": ["compare.start.daily"], + "expires_at": first_expiry, + "reason": "first period", + }, + ) + assert first.status_code == 201, first.text + first_id = first.json()[0]["id"] + reset = client.post( + f"/admin/api/limit-whitelist/{first_id}/reset", + headers=admin_headers, + ) + assert reset.status_code == 200, reset.text + + recreated = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_codes": [ + "compare.start.daily", + "sms.phone.cooldown", + ], + "expires_at": second_expiry, + "reason": "renewed period", + }, + ) + assert recreated.status_code == 201, recreated.text + rows = {item["rule_code"]: item for item in recreated.json()} + assert rows["compare.start.daily"]["id"] == first_id + assert rows["compare.start.daily"]["enabled"] is True + assert rows["compare.start.daily"]["status"] == "active" + assert rows["compare.start.daily"]["reason"] == "renewed period" + assert rows["sms.phone.cooldown"]["enabled"] is True + + +def test_bulk_create_appends_rules_and_unifies_subject_period(admin_headers) -> None: + phone = f"139{int(uuid4().hex[:8], 16) % 100000000:08d}" + first_start = (datetime.now(UTC) + timedelta(minutes=5)).isoformat() + first_expiry = (datetime.now(UTC) + timedelta(hours=1)).isoformat() + latest_start = (datetime.now(UTC) + timedelta(minutes=10)).isoformat() + latest_expiry = (datetime.now(UTC) + timedelta(hours=3)).isoformat() + + with TestClient(admin_app) as client: + first = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_codes": ["compare.start.daily"], + "starts_at": first_start, + "expires_at": first_expiry, + "reason": "original rule", + }, + ) + assert first.status_code == 201, first.text + + appended = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_codes": ["sms.phone.cooldown"], + "starts_at": latest_start, + "expires_at": latest_expiry, + "reason": "appended rule", + }, + ) + assert appended.status_code == 201, appended.text + rows = {item["rule_code"]: item for item in appended.json()} + assert set(rows) == { + "compare.start.daily", + "sms.phone.cooldown", + } + assert all( + datetime.fromisoformat(item["starts_at"]) == datetime.fromisoformat(latest_start) + for item in rows.values() + ) + assert all( + datetime.fromisoformat(item["expires_at"]) == datetime.fromisoformat(latest_expiry) + for item in rows.values() + ) + assert all(item["enabled"] is True for item in rows.values()) + assert rows["compare.start.daily"]["reason"] == "original rule" + assert rows["sms.phone.cooldown"]["reason"] == "appended rule" + with SessionLocal() as db: + audit = db.scalar( + select(AdminAuditLog) + .where(AdminAuditLog.action == "limit.override.bulk_create") + .order_by(AdminAuditLog.id.desc()) + ) + assert audit is not None + assert len(audit.detail["before"]) == 1 + assert len(audit.detail["after"]) == 2 + + +def test_legacy_writes_preserve_subject_level_period_and_enabled_state( + admin_headers, +) -> None: + phone = f"138{int(uuid4().hex[:8], 16) % 100000000:08d}" + first_expiry = datetime.now(UTC) + timedelta(hours=2) + latest_start = datetime.now(UTC) + timedelta(minutes=5) + latest_expiry = datetime.now(UTC) + timedelta(hours=4) + + with TestClient(admin_app) as client: + created = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_codes": [ + "compare.start.daily", + "sms.phone.cooldown", + ], + "expires_at": first_expiry.isoformat(), + "reason": "initial subject", + }, + ) + assert created.status_code == 201, created.text + first_id = created.json()[0]["id"] + + patched = client.patch( + f"/admin/api/limit-whitelist/{first_id}", + headers=admin_headers, + json={ + "starts_at": latest_start.isoformat(), + "expires_at": latest_expiry.isoformat(), + }, + ) + assert patched.status_code == 200, patched.text + + legacy_added = client.post( + "/admin/api/limit-whitelist", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_code": "ad.feed.daily", + "mode": "unlimited", + "enabled": True, + "starts_at": latest_start.isoformat(), + "expires_at": latest_expiry.isoformat(), + "reason": "legacy append", + }, + ) + assert legacy_added.status_code == 201, legacy_added.text + + listed = client.get( + "/admin/api/limit-whitelist/subjects", + headers=admin_headers, + params={"keyword": phone}, + ) + assert listed.status_code == 200, listed.text + rows = listed.json()["items"][0]["items"] + assert len(rows) == 3 + assert { + datetime.fromisoformat(item["starts_at"]).replace(tzinfo=None) for item in rows + } == {latest_start.replace(tzinfo=None)} + assert { + datetime.fromisoformat(item["expires_at"]).replace(tzinfo=None) for item in rows + } == {latest_expiry.replace(tzinfo=None)} + assert all(item["enabled"] is True for item in rows) + + reset = client.post( + f"/admin/api/limit-whitelist/{first_id}/reset", + headers=admin_headers, + ) + assert reset.status_code == 200, reset.text + listed = client.get( + "/admin/api/limit-whitelist/subjects", + headers=admin_headers, + params={"keyword": phone}, + ) + assert all(item["enabled"] is False for item in listed.json()["items"][0]["items"]) + + +def test_subject_list_groups_rules_and_paginates_by_subject(admin_headers) -> None: + suffix = f"{int(uuid4().hex[:7], 16) % 1000000:06d}" + phones = [f"13870{suffix}", f"13871{suffix}"] + expires_at = (datetime.now(UTC) + timedelta(hours=2)).isoformat() + + with TestClient(admin_app) as client: + first = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phones[0], + "rule_codes": [ + "compare.start.daily", + "sms.phone.cooldown", + ], + "expires_at": expires_at, + "reason": "grouped list first", + }, + ) + assert first.status_code == 201, first.text + second = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phones[1], + "rule_codes": ["compare.start.daily"], + "expires_at": expires_at, + "reason": "grouped list second", + }, + ) + assert second.status_code == 201, second.text + + page_one = client.get( + "/admin/api/limit-whitelist/subjects", + headers=admin_headers, + params={"keyword": suffix, "limit": 1, "offset": 0}, + ) + page_two = client.get( + "/admin/api/limit-whitelist/subjects", + headers=admin_headers, + params={"keyword": suffix, "limit": 1, "offset": 1}, + ) + assert page_one.status_code == 200, page_one.text + assert page_two.status_code == 200, page_two.text + assert page_one.json()["total"] == 2 + assert len(page_one.json()["items"]) == 1 + assert len(page_two.json()["items"]) == 1 + listed = page_one.json()["items"] + page_two.json()["items"] + assert {item["subject_value"] for item in listed} == set(phones) + first_subject = next(item for item in listed if item["subject_value"] == phones[0]) + assert first_subject["total_rules"] == 2 + assert sum(first_subject["group_counts"].values()) == 2 + assert {item["rule_code"] for item in first_subject["items"]} == { + "compare.start.daily", + "sms.phone.cooldown", + } + + +def test_subject_replace_updates_rule_selection_atomically(admin_headers) -> None: + phone = f"137{int(uuid4().hex[:8], 16) % 100000000:08d}" + expires_at = (datetime.now(UTC) + timedelta(hours=2)).isoformat() + updated_expiry = (datetime.now(UTC) + timedelta(hours=4)).isoformat() + + with TestClient(admin_app) as client: + created = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_codes": [ + "compare.start.daily", + "sms.phone.cooldown", + ], + "expires_at": expires_at, + "reason": "before subject edit", + }, + ) + assert created.status_code == 201, created.text + original_created_at = min(item["created_at"] for item in created.json()) + + replaced = client.put( + "/admin/api/limit-whitelist/subjects", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_codes": [ + "ad.reward_video.daily", + ], + "expires_at": updated_expiry, + "reason": "after subject edit", + }, + ) + assert replaced.status_code == 200, replaced.text + payload = replaced.json() + assert payload["subject_value"] == phone + assert payload["total_rules"] == 1 + assert {item["rule_code"] for item in payload["items"]} == { + "ad.reward_video.daily", + } + assert payload["created_at"] == original_created_at + assert all(item["reason"] == "after subject edit" for item in payload["items"]) + + listed = client.get( + "/admin/api/limit-whitelist/subjects", + headers=admin_headers, + params={"keyword": phone}, + ) + assert listed.status_code == 200, listed.text + assert listed.json()["total"] == 1 + assert {item["rule_code"] for item in listed.json()["items"][0]["items"]} == { + "ad.reward_video.daily", + } + assert listed.json()["items"][0]["created_at"] == original_created_at + + +def test_subject_enabled_patch_updates_all_rules_together(admin_headers) -> None: + phone = f"136{int(uuid4().hex[:8], 16) % 100000000:08d}" + expires_at = (datetime.now(UTC) + timedelta(hours=2)).isoformat() + + with TestClient(admin_app) as client: + created = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_codes": [ + "compare.start.daily", + "sms.phone.cooldown", + "ad.reward_video.daily", + ], + "expires_at": expires_at, + "reason": "subject switch", + }, + ) + assert created.status_code == 201, created.text + + disabled = client.patch( + "/admin/api/limit-whitelist/subjects/enabled", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "enabled": False, + }, + ) + assert disabled.status_code == 200, disabled.text + assert disabled.json()["total_rules"] == 3 + assert all(item["enabled"] is False for item in disabled.json()["items"]) + assert all(item["status"] == "disabled" for item in disabled.json()["items"]) + + enabled = client.patch( + "/admin/api/limit-whitelist/subjects/enabled", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "enabled": True, + }, + ) + assert enabled.status_code == 200, enabled.text + assert all(item["enabled"] is True for item in enabled.json()["items"]) + assert all(item["status"] == "active" for item in enabled.json()["items"]) + + disabled_again = client.patch( + "/admin/api/limit-whitelist/subjects/enabled", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "enabled": False, + }, + ) + assert disabled_again.status_code == 200, disabled_again.text + oldest_id = min(item["id"] for item in created.json()) + expired = client.patch( + f"/admin/api/limit-whitelist/{oldest_id}", + headers=admin_headers, + json={ + "expires_at": (datetime.now(UTC) - timedelta(minutes=1)).isoformat(), + }, + ) + assert expired.status_code == 200, expired.text + + rejected = client.patch( + "/admin/api/limit-whitelist/subjects/enabled", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "enabled": True, + }, + ) + assert rejected.status_code == 400, rejected.text + listed = client.get( + "/admin/api/limit-whitelist/subjects", + headers=admin_headers, + params={"keyword": phone}, + ) + assert listed.status_code == 200, listed.text + assert all(item["enabled"] is False for item in listed.json()["items"][0]["items"]) + + +def test_admin_whitelist_crud_and_policy_precedence(admin_headers) -> None: + suffix = uuid4().hex[:10] + phone = f"139{int(suffix[:8], 16) % 100000000:08d}" + device = f"limit-device-{suffix}" + expires_at = (datetime.now(UTC) + timedelta(hours=2)).isoformat() + + with TestClient(admin_app) as client: + rules = client.get( + "/admin/api/limit-whitelist/rules", + headers=admin_headers, + ) + assert rules.status_code == 200, rules.text + assert {item["code"] for item in rules.json()} >= { + "compare.start.daily", + "sms.send.hourly", + "ad.reward_video.daily", + "risk.compare.daily", + } + rules_by_code = {item["code"]: item for item in rules.json()} + assert all("inherit" not in item["allowed_modes"] for item in rules.json()) + assert all("override" not in item["allowed_modes"] for item in rules.json()) + assert "unlimited" in rules_by_code["sms.phone.cooldown"]["allowed_modes"] + assert "unlimited" in rules_by_code["sms.code.failed_attempts"]["allowed_modes"] + assert "suppress_alert" in rules_by_code["risk.compare.daily"]["allowed_modes"] + + removed_inherit = client.post( + "/admin/api/limit-whitelist", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_code": "compare.start.daily", + "mode": "inherit", + "enabled": True, + }, + ) + assert removed_inherit.status_code == 422 + + invalid_alert = client.post( + "/admin/api/limit-whitelist", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_code": "compare.start.daily", + "mode": "suppress_alert", + "enabled": True, + "reason": "invalid", + }, + ) + assert invalid_alert.status_code == 400 + + missing_expiry = client.post( + "/admin/api/limit-whitelist", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_code": "compare.start.daily", + "mode": "unlimited", + "enabled": True, + "reason": "temporary QA", + }, + ) + assert missing_expiry.status_code == 400 + + expired_unlimited = client.post( + "/admin/api/limit-whitelist", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_code": "compare.start.daily", + "mode": "unlimited", + "enabled": True, + "expires_at": "2020-01-01T00:00:00Z", + "reason": "expired temporary policy", + }, + ) + assert expired_unlimited.status_code == 400 + assert "晚于当前时间" in expired_unlimited.json()["detail"] + + removed_override = client.post( + "/admin/api/limit-whitelist", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_code": "compare.start.daily", + "mode": "override", + "limit_value": 2, + "enabled": True, + "expires_at": expires_at, + }, + ) + assert removed_override.status_code == 422 + + phone_created = client.post( + "/admin/api/limit-whitelist", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_code": "compare.start.daily", + "mode": "unlimited", + "enabled": True, + "expires_at": expires_at, + "reason": "temporary QA", + }, + ) + assert phone_created.status_code == 201, phone_created.text + phone_id = phone_created.json()["id"] + assert phone_created.json()["effective_limit"] is None + + phone_updated = client.patch( + f"/admin/api/limit-whitelist/{phone_id}", + headers=admin_headers, + json={ + "enabled": True, + "reason": "updated from admin", + }, + ) + assert phone_updated.status_code == 200, phone_updated.text + assert phone_updated.json()["effective_limit"] is None + assert phone_updated.json()["reason"] == "updated from admin" + with SessionLocal() as db: + persisted = db.get(LimitPolicyOverride, phone_id) + assert persisted is not None + assert persisted.mode == "unlimited" + assert persisted.limit_value is None + assert persisted.reason == "updated from admin" + + removed_patch_fields = client.patch( + f"/admin/api/limit-whitelist/{phone_id}", + headers=admin_headers, + json={"mode": "override", "limit_value": 4}, + ) + assert removed_patch_fields.status_code == 422 + + inverted_time = client.patch( + f"/admin/api/limit-whitelist/{phone_id}", + headers=admin_headers, + json={ + "starts_at": "2030-01-02T00:00:00Z", + "expires_at": "2030-01-01T00:00:00Z", + }, + ) + assert inverted_time.status_code == 400 + assert "晚于生效时间" in inverted_time.json()["detail"] + + duplicate = client.post( + "/admin/api/limit-whitelist", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_code": "compare.start.daily", + "mode": "unlimited", + "enabled": True, + "expires_at": expires_at, + }, + ) + assert duplicate.status_code == 409 + + device_created = client.post( + "/admin/api/limit-whitelist", + headers=admin_headers, + json={ + "subject_type": "device", + "subject_value": device, + "rule_code": "compare.start.daily", + "mode": "unlimited", + "enabled": True, + "expires_at": expires_at, + "reason": "device QA", + }, + ) + assert device_created.status_code == 201, device_created.text + device_id = device_created.json()["id"] + assert device_created.json()["effective_limit"] is None + + with SessionLocal() as db: + effective = limit_policy.resolve( + db, + "compare.start.daily", + phone=phone, + device=device, + ) + assert effective.unlimited is True + assert effective.matched_subject_type == "device" + + restored = client.post( + f"/admin/api/limit-whitelist/{phone_id}/reset", + headers=admin_headers, + ) + assert restored.status_code == 200, restored.text + assert restored.json()["enabled"] is False + assert restored.json()["status"] == "disabled" + assert restored.json()["effective_limit"] == restored.json()["global_limit"] + assert restored.json()["reset_at"] is None + with SessionLocal() as db: + effective = limit_policy.resolve( + db, + "compare.start.daily", + phone=phone, + device=None, + ) + assert effective.override_id is None + assert effective.limit == effective.global_limit + + ordered = client.get( + "/admin/api/limit-whitelist", + headers=admin_headers, + params={"rule_code": "compare.start.daily", "limit": 500}, + ) + assert ordered.status_code == 200, ordered.text + ordered_ids = [item["id"] for item in ordered.json()["items"]] + assert ordered_ids.index(device_id) < ordered_ids.index(phone_id) + + listed = client.get( + "/admin/api/limit-whitelist", + headers=admin_headers, + params={"keyword": suffix[:5]}, + ) + assert listed.status_code == 200 + assert listed.json()["total"] >= 1 + + deleted = client.delete( + f"/admin/api/limit-whitelist/{device_id}", + headers=admin_headers, + ) + assert deleted.status_code == 204 + client.delete( + f"/admin/api/limit-whitelist/{phone_id}", + headers=admin_headers, + ) + + +def test_device_candidates_are_rule_aware_and_searchable(admin_headers) -> None: + suffix = uuid4().hex[:8] + phone = f"135{int(suffix, 16) % 100000000:08d}" + auth_device = f"android-id-{suffix}" + compare_device = f"device-PJZ110-{suffix}" + now = datetime.now(UTC).replace(microsecond=0) + with SessionLocal() as db: + user = user_repo.upsert_user_for_login( + db, + phone=phone, + register_channel="sms", + ) + user.nickname = f"候选设备用户{suffix}" + risk_repo.record_behavior_event( + db, + event_type=risk_repo.EVENT_ONECLICK_LOGIN, + subject_type="device", + subject_id=auth_device, + user_id=user.id, + device_id=auth_device, + device_model="OPPO Find X8", + phone=phone, + outcome="success", + occurred_at=now, + ) + db.add( + ComparisonRecord( + user_id=user.id, + device_id=compare_device, + trace_id=f"device-candidate-{suffix}", + device_model="PJZ110", + status="success", + items=[], + comparison_results=[], + skipped_dish_names=[], + created_at=now.replace(tzinfo=None), + ) + ) + db.commit() + user_id = user.id + username = user.username + + with TestClient(admin_app) as client: + for keyword in (phone, str(user_id), username, "Find X8"): + response = client.get( + "/admin/api/limit-whitelist/device-candidates", + headers=admin_headers, + params={ + "rule_code": "risk.oneclick.daily", + "keyword": keyword, + }, + ) + assert response.status_code == 200, response.text + item = next(row for row in response.json() if row["device_id"] == auth_device) + assert item["source_label"] == "一键登录" + assert item["user_id"] == user_id + assert item["phone"] == phone + assert item["device_model"] == "OPPO Find X8" + + comparison = client.get( + "/admin/api/limit-whitelist/device-candidates", + headers=admin_headers, + params={ + "rule_code": "compare.start.daily", + "keyword": "PJZ110", + }, + ) + assert comparison.status_code == 200, comparison.text + assert len(comparison.json()) == 1 + comparison_item = comparison.json()[0] + assert comparison_item["device_id"] == compare_device + assert comparison_item["source"] == "comparison_record" + assert comparison_item["source_label"] == "比价记录" + assert comparison_item["user_id"] == user_id + assert comparison_item["username"] == username + assert comparison_item["phone"] == phone + assert comparison_item["nickname"] == f"候选设备用户{suffix}" + assert comparison_item["device_model"] == "PJZ110" + assert comparison_item["last_active_at"].startswith(now.replace(tzinfo=None).isoformat()) + + phone_only = client.get( + "/admin/api/limit-whitelist/device-candidates", + headers=admin_headers, + params={"rule_code": "risk.compare.daily"}, + ) + assert phone_only.status_code == 400 + assert "不支持设备白名单" in phone_only.json()["detail"] + + +def test_legacy_ip_is_not_a_device_candidate_or_valid_whitelist_subject( + admin_headers, +) -> None: + suffix = uuid4().hex[:8] + legacy_device = f"legacy-ip:203.0.113.{int(suffix[:2], 16) % 200 + 1}" + expires_at = (datetime.now(UTC) + timedelta(hours=2)).isoformat() + with SessionLocal() as db: + risk_repo.record_behavior_event( + db, + event_type=risk_repo.EVENT_SMS_SEND, + subject_type="device", + subject_id=legacy_device, + device_id=None, + phone=f"136{int(suffix, 16) % 100000000:08d}", + outcome="success", + ) + + with TestClient(admin_app) as client: + candidates = client.get( + "/admin/api/limit-whitelist/device-candidates", + headers=admin_headers, + params={ + "rule_code": "sms.send.hourly", + "keyword": legacy_device, + }, + ) + assert candidates.status_code == 200, candidates.text + assert all(row["device_id"] != legacy_device for row in candidates.json()) + + single = client.post( + "/admin/api/limit-whitelist", + headers=admin_headers, + json={ + "subject_type": "device", + "subject_value": legacy_device, + "rule_code": "sms.send.hourly", + "mode": "unlimited", + "enabled": True, + "expires_at": expires_at, + }, + ) + assert single.status_code == 400, single.text + assert "未上报真实设备 ID" in single.json()["detail"] + + bulk = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "device", + "subject_value": legacy_device, + "rule_codes": ["sms.send.hourly", "sms.send.daily"], + "enabled": True, + "expires_at": expires_at, + }, + ) + assert bulk.status_code == 400, bulk.text + assert "未上报真实设备 ID" in bulk.json()["detail"] + + +def test_bulk_create_automatically_selects_unlimited_and_suppress_modes( + admin_headers, +) -> None: + suffix = uuid4().hex[:8] + phone = f"137{int(suffix, 16) % 100000000:08d}" + expires_at = (datetime.now(UTC) + timedelta(hours=2)).isoformat() + + with TestClient(admin_app) as client: + created = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_codes": [ + "compare.start.daily", + "risk.compare.daily", + ], + "expires_at": expires_at, + "reason": "批量白名单测试", + }, + ) + assert created.status_code == 201, created.text + rows = {item["rule_code"]: item for item in created.json()} + assert rows["compare.start.daily"]["mode"] == "unlimited" + assert rows["compare.start.daily"]["effective_limit"] is None + assert rows["risk.compare.daily"]["mode"] == "suppress_alert" + assert ( + rows["risk.compare.daily"]["effective_limit"] + == rows["risk.compare.daily"]["global_limit"] + ) + + appended_batch = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "phone", + "subject_value": phone, + "rule_codes": [ + "compare.start.daily", + "sms.send.daily", + ], + "expires_at": expires_at, + }, + ) + assert appended_batch.status_code == 201, appended_batch.text + appended_rows = {item["rule_code"]: item for item in appended_batch.json()} + assert set(appended_rows) == { + "compare.start.daily", + "risk.compare.daily", + "sms.send.daily", + } + + for item in appended_batch.json(): + deleted = client.delete( + f"/admin/api/limit-whitelist/{item['id']}", + headers=admin_headers, + ) + assert deleted.status_code == 204 + + +def test_device_bulk_accepts_rules_from_different_business_categories( + admin_headers, +) -> None: + device_id = f"manually-entered-device-{uuid4().hex[:8]}" + expires_at = (datetime.now(UTC) + timedelta(hours=2)).isoformat() + with TestClient(admin_app) as client: + response = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "device", + "subject_value": device_id, + "rule_codes": [ + "compare.start.daily", + "sms.send.hourly", + "risk.oneclick.daily", + ], + "expires_at": expires_at, + }, + ) + assert response.status_code == 201, response.text + rows = {item["rule_code"]: item for item in response.json()} + assert set(rows) == { + "compare.start.daily", + "sms.send.hourly", + "risk.oneclick.daily", + } + assert all(item["subject_value"] == device_id for item in rows.values()) + + with SessionLocal() as db: + compare = limit_policy.resolve( + db, + "compare.start.daily", + device=device_id, + ) + sms = limit_policy.resolve( + db, + "sms.send.hourly", + device=device_id, + ) + alert = limit_policy.resolve( + db, + "risk.oneclick.daily", + device=device_id, + ) + assert compare.unlimited is True + assert sms.unlimited is True + assert alert.suppressed is True + + +def test_device_subject_can_append_and_replace_cross_category_rules( + admin_headers, +) -> None: + device_id = f"opaque-device-value-{uuid4().hex[:8]}" + expires_at = (datetime.now(UTC) + timedelta(hours=2)).isoformat() + with TestClient(admin_app) as client: + initial = client.post( + "/admin/api/limit-whitelist/bulk", + headers=admin_headers, + json={ + "subject_type": "device", + "subject_value": device_id, + "rule_codes": ["compare.start.daily"], + "expires_at": expires_at, + }, + ) + assert initial.status_code == 201, initial.text + + appended = client.post( + "/admin/api/limit-whitelist", + headers=admin_headers, + json={ + "subject_type": "device", + "subject_value": device_id, + "rule_code": "risk.oneclick.daily", + "mode": "suppress_alert", + "enabled": True, + "expires_at": expires_at, + }, + ) + assert appended.status_code == 201, appended.text + + replaced = client.put( + "/admin/api/limit-whitelist/subjects", + headers=admin_headers, + json={ + "subject_type": "device", + "subject_value": device_id, + "rule_codes": [ + "compare.start.daily", + "sms.send.daily", + "risk.oneclick.daily", + ], + "enabled": True, + "expires_at": expires_at, + }, + ) + assert replaced.status_code == 200, replaced.text + assert { + item["rule_code"] for item in replaced.json()["items"] + } == { + "compare.start.daily", + "sms.send.daily", + "risk.oneclick.daily", + } + + +def test_legacy_ip_device_whitelist_matches_sms_policy(monkeypatch) -> None: + suffix = uuid4().hex[:8] + client_ip = f"203.0.113.{int(suffix[:2], 16) % 200 + 1}" + legacy_device = f"legacy-ip:{client_ip}" + phone = f"134{int(suffix, 16) % 100000000:08d}" + expires_at = datetime.now(UTC) + timedelta(hours=2) + with SessionLocal() as db: + row = LimitPolicyOverride( + subject_type="device", + subject_value=legacy_device, + rule_code="sms.send.hourly", + mode="unlimited", + enabled=True, + expires_at=expires_at, + ) + db.add(row) + db.commit() + override_id = row.id + + original_resolve = auth_api.limit_policy.resolve + matched_override_ids: list[int | None] = [] + + def spy_resolve(db, rule_code, **subjects): + result = original_resolve(db, rule_code, **subjects) + if rule_code == "sms.send.hourly": + matched_override_ids.append(result.override_id) + return result + + monkeypatch.setattr(auth_api.limit_policy, "resolve", spy_resolve) + try: + with TestClient(app) as client: + response = client.post( + "/api/v1/auth/sms/send", + headers={"x-forwarded-for": client_ip}, + json={"phone": phone}, + ) + assert response.status_code == 200, response.text + assert matched_override_ids == [override_id] + finally: + with SessionLocal() as db: + row = db.get(LimitPolicyOverride, override_id) + if row is not None: + db.delete(row) + db.commit() + + +def test_global_guide_limit_stays_in_sync_with_legacy_config(admin_headers) -> None: + snapshot = _snapshot_configs( + [ + LIMIT_POLICY_GLOBAL_KEY, + "coupon_guide_video", + "comparison_guide_video", + ] + ) + try: + with TestClient(admin_app) as client: + changed = client.patch( + "/admin/api/limit-whitelist/rules/guide.video.lifetime", + headers=admin_headers, + json={"value": 7}, + ) + assert changed.status_code == 200, changed.text + assert changed.json()["global_limit"] == 7 + + legacy_page = client.get( + "/admin/api/guide-video", + headers=admin_headers, + ) + assert legacy_page.status_code == 200, legacy_page.text + assert legacy_page.json()["max_plays"] == 7 + assert legacy_page.json()["scene"] == "coupon" + + comparison_changed = client.patch( + "/admin/api/guide-video", + headers=admin_headers, + params={"scene": "comparison"}, + json={"max_plays": 9}, + ) + assert comparison_changed.status_code == 200, comparison_changed.text + assert comparison_changed.json()["scene"] == "comparison" + assert comparison_changed.json()["max_plays"] == 9 + + coupon_again = client.get( + "/admin/api/guide-video", + headers=admin_headers, + params={"scene": "coupon"}, + ) + assert coupon_again.status_code == 200, coupon_again.text + assert coupon_again.json()["max_plays"] == 7 + finally: + _restore_configs(snapshot) + + +def test_legacy_guide_limit_audit_uses_synced_global_value(admin_headers) -> None: + snapshot = _snapshot_configs([LIMIT_POLICY_GLOBAL_KEY, "coupon_guide_video"]) + try: + with TestClient(admin_app) as client: + changed = client.patch( + "/admin/api/guide-video", + headers=admin_headers, + params={"scene": "coupon"}, + json={"max_plays": 11}, + ) + assert changed.status_code == 200, changed.text + assert changed.json()["max_plays"] == 11 + + with SessionLocal() as db: + audit = db.scalar( + select(AdminAuditLog) + .where(AdminAuditLog.action == "guide_video.update") + .order_by(AdminAuditLog.id.desc()) + ) + assert audit is not None + assert audit.detail["after"]["max_plays"] == 11 + finally: + _restore_configs(snapshot) + + +def test_zero_global_guide_limit_keeps_legacy_config_editable(admin_headers) -> None: + snapshot = _snapshot_configs([LIMIT_POLICY_GLOBAL_KEY, "coupon_guide_video"]) + try: + with TestClient(admin_app) as client: + zeroed = client.patch( + "/admin/api/limit-whitelist/rules/guide.video.lifetime", + headers=admin_headers, + json={"value": 0}, + ) + assert zeroed.status_code == 200, zeroed.text + assert zeroed.json()["global_limit"] == 0 + + legacy_changed = client.patch( + "/admin/api/guide-video", + headers=admin_headers, + params={"scene": "coupon"}, + json={"reward_coin": 200}, + ) + assert legacy_changed.status_code == 200, legacy_changed.text + assert legacy_changed.json()["max_plays"] == 0 + assert legacy_changed.json()["reward_coin"] == 200 + finally: + _restore_configs(snapshot) + + +def test_guide_video_v2_honours_unlimited_override_without_reusing_seq() -> None: + snapshot = _snapshot_configs([LIMIT_POLICY_GLOBAL_KEY, "coupon_guide_video"]) + suffix = uuid4().hex[:8] + phone = f"139{int(suffix, 16) % 100_000_000:08d}" + override_id: int | None = None + try: + with SessionLocal() as db: + user = user_repo.upsert_user_for_login( + db, + phone=phone, + register_channel="sms", + ) + user_id = user.id + guide_video_repo.set_video( + db, + "/media/guide_video/whitelist-v2-test.mp4", + analysis={ + "duration_ms": 10_000, + "video_codec": "h264", + "audio_codec": "aac", + "analysis_status": "valid", + "analysis_error": None, + }, + scene="coupon", + admin_id=1, + ) + guide_video_repo.update_config( + db, + scene="coupon", + enabled=True, + max_plays=1, + reward_coin=100, + admin_id=1, + ) + + first_plan = guide_video_repo.prepare_play(db, user_id, scene="coupon") + first_start = guide_video_repo.start_play( + db, user_id, play_token=first_plan["play_token"] + ) + assert first_start["seq"] == 1 + assert ( + guide_video_repo.prepare_play(db, user_id, scene="coupon")["reason"] + == "play_limit_reached" + ) + + override = LimitPolicyOverride( + rule_code="guide.video.lifetime", + subject_type="phone", + subject_value=phone, + mode="unlimited", + limit_value=None, + starts_at=datetime.now(UTC) - timedelta(seconds=1), + expires_at=datetime.now(UTC) + timedelta(hours=1), + reset_at=datetime.now(UTC), + enabled=True, + reason="新版十圈视频白名单回归测试", + ) + db.add(override) + db.commit() + db.refresh(override) + override_id = override.id + + second_plan = guide_video_repo.prepare_play(db, user_id, scene="coupon") + assert second_plan["should_play"] is True + second_start = guide_video_repo.start_play( + db, user_id, play_token=second_plan["play_token"] + ) + assert second_start["seq"] == 2 + assert second_start["remaining"] > 0 + finally: + with SessionLocal() as db: + if override_id is not None: + override = db.get(LimitPolicyOverride, override_id) + if override is not None: + db.delete(override) + db.commit() + _restore_configs(snapshot) + + +def test_legacy_ad_limit_update_syncs_split_global_rules(admin_headers) -> None: + snapshot = _snapshot_configs(["ad_daily_limit", LIMIT_POLICY_GLOBAL_KEY]) + try: + with TestClient(admin_app) as client: + changed = client.patch( + "/admin/api/config/ad_daily_limit", + headers=admin_headers, + json={"value": 321}, + ) + assert changed.status_code == 200, changed.text + assert changed.json()["value"] == 321 + + rules = client.get( + "/admin/api/limit-whitelist/rules", + headers=admin_headers, + ) + by_code = {item["code"]: item for item in rules.json()} + assert by_code["ad.reward_video.daily"]["global_limit"] == 321 + assert by_code["ad.feed.daily"]["global_limit"] == 321 + finally: + _restore_configs(snapshot) + + +def test_zero_disables_cooldown_style_global_rules(admin_headers) -> None: + snapshot = _snapshot_configs([LIMIT_POLICY_GLOBAL_KEY]) + try: + with TestClient(admin_app) as client: + for rule_code in ("sms.phone.cooldown", "phone.rebind.days"): + changed = client.patch( + f"/admin/api/limit-whitelist/rules/{rule_code}", + headers=admin_headers, + json={"value": 0}, + ) + assert changed.status_code == 200, changed.text + assert changed.json()["global_limit"] == 0 + finally: + _restore_configs(snapshot) + + +def test_global_limits_are_stored_in_one_complete_json(admin_headers) -> None: + keys = [LIMIT_POLICY_GLOBAL_KEY, *limit_policy.LIMIT_CONFIG_KEYS] + snapshot = _snapshot_configs(keys) + try: + with TestClient(admin_app) as client: + first = client.patch( + "/admin/api/limit-whitelist/rules/compare.start.daily", + headers=admin_headers, + json={"value": 73}, + ) + second = client.patch( + "/admin/api/limit-whitelist/rules/sms.send.hourly", + headers=admin_headers, + json={"value": 9}, + ) + assert first.status_code == 200, first.text + assert second.status_code == 200, second.text + direct_bundle_write = client.patch( + f"/admin/api/config/{LIMIT_POLICY_GLOBAL_KEY}", + headers=admin_headers, + json={"value": {"compare.start.daily": 1}}, + ) + assert direct_bundle_write.status_code == 404 + + with SessionLocal() as db: + bundle = db.get(AppConfig, LIMIT_POLICY_GLOBAL_KEY) + assert bundle is not None + assert isinstance(bundle.value, dict) + assert set(bundle.value) == set(limit_policy.RULE_MAP) + assert len(bundle.value) == 16 + assert bundle.value["compare.start.daily"] == 73 + assert bundle.value["sms.send.hourly"] == 9 + sparse_rows = ( + db.query(AppConfig) + .filter(AppConfig.key.in_(limit_policy.LIMIT_CONFIG_KEYS)) + .count() + ) + assert sparse_rows == 0 + finally: + _restore_configs(snapshot) + + +def test_compare_start_uses_phone_override_and_reset_baseline() -> None: + suffix = uuid4().hex[:8] + phone = f"136{int(suffix, 16) % 100000000:08d}" + device = f"compare-policy-{suffix}" + with SessionLocal() as db: + user = user_repo.upsert_user_for_login( + db, + phone=phone, + register_channel="sms", + ) + db.add( + LimitPolicyOverride( + subject_type="phone", + subject_value=phone, + rule_code="compare.start.daily", + mode="override", + limit_value=1, + enabled=True, + ) + ) + db.commit() + user_id = user.id + + token, _ = create_token(user_id=user_id, token_type="access") + headers = {"Authorization": f"Bearer {token}"} + with TestClient(app) as client: + first = client.post( + "/api/v1/compare/start", + headers=headers, + json={ + "trace_id": f"limit-policy-first-{suffix}", + "business_type": "food", + "device_id": device, + }, + ) + assert first.status_code == 200, first.text + assert first.json() == {"limit": 1, "used": 1, "remaining": 0} + + blocked = client.post( + "/api/v1/compare/start", + headers=headers, + json={ + "trace_id": f"limit-policy-blocked-{suffix}", + "business_type": "food", + "device_id": device, + }, + ) + assert blocked.status_code == 429 + + with SessionLocal() as db: + override = ( + db.query(LimitPolicyOverride) + .filter( + LimitPolicyOverride.subject_type == "phone", + LimitPolicyOverride.subject_value == phone, + LimitPolicyOverride.rule_code == "compare.start.daily", + ) + .one() + ) + override.reset_at = datetime.now(UTC) + db.commit() + + after_reset = client.post( + "/api/v1/compare/start", + headers=headers, + json={ + "trace_id": f"limit-policy-reset-{suffix}", + "business_type": "food", + "device_id": device, + }, + ) + assert after_reset.status_code == 200, after_reset.text + assert after_reset.json() == {"limit": 1, "used": 1, "remaining": 0} + + +def test_suppress_alert_keeps_events_but_creates_no_incident() -> None: + suffix = uuid4().hex[:10] + device = f"suppress-device-{suffix}" + phone = f"138{int(suffix[:8], 16) % 100000000:08d}" + now = risk_repo.utcnow().replace(microsecond=0) + with SessionLocal() as db: + db.add( + LimitPolicyOverride( + subject_type="device", + subject_value=device, + rule_code="risk.sms.hourly", + mode="suppress_alert", + enabled=True, + ) + ) + db.commit() + for index in range(10): + risk_repo.record_behavior_event( + db, + event_type=risk_repo.EVENT_SMS_SEND, + subject_type="device", + subject_id=device, + device_id=device, + phone=phone, + outcome="success", + occurred_at=now + timedelta(seconds=index), + evaluate_rule=risk_repo.RULE_SMS_HOURLY, + ) + + with SessionLocal() as db: + incident = ( + db.query(RiskIncident) + .filter( + RiskIncident.rule_code == risk_repo.RULE_SMS_HOURLY, + RiskIncident.subject_id == device, + ) + .one_or_none() + ) + assert incident is None + + +def test_adding_suppress_alert_resolves_existing_incident(admin_headers) -> None: + suffix = uuid4().hex[:10] + device = f"suppress-existing-{suffix}" + now = risk_repo.utcnow().replace(microsecond=0) + with SessionLocal() as db: + for index in range(5): + risk_repo.record_behavior_event( + db, + event_type=risk_repo.EVENT_SMS_SEND, + subject_type="device", + subject_id=device, + device_id=device, + phone=f"13700001{index:03d}", + outcome="success", + occurred_at=now + timedelta(seconds=index), + evaluate_rule=risk_repo.RULE_SMS_HOURLY, + ) + incident = ( + db.query(RiskIncident) + .filter( + RiskIncident.rule_code == risk_repo.RULE_SMS_HOURLY, + RiskIncident.subject_id == device, + ) + .one() + ) + assert incident.status == "open" + + with TestClient(admin_app) as client: + created = client.post( + "/admin/api/limit-whitelist", + headers=admin_headers, + json={ + "subject_type": "device", + "subject_value": device, + "rule_code": "risk.sms.hourly", + "mode": "suppress_alert", + "enabled": True, + "expires_at": (datetime.now(UTC) + timedelta(hours=2)).isoformat(), + "reason": "QA device", + }, + ) + assert created.status_code == 201, created.text + + with SessionLocal() as db: + incident = ( + db.query(RiskIncident) + .filter( + RiskIncident.rule_code == risk_repo.RULE_SMS_HOURLY, + RiskIncident.subject_id == device, + ) + .one() + ) + assert incident.status == "resolved" + assert incident.action_reason == risk_repo.AUTO_RESOLVED_REASON diff --git a/tests/test_risk_monitor.py b/tests/test_risk_monitor.py index b358cd3..4e41244 100644 --- a/tests/test_risk_monitor.py +++ b/tests/test_risk_monitor.py @@ -10,11 +10,7 @@ from fastapi.testclient import TestClient from app.admin.main import admin_app from app.admin.repositories import admin_user as admin_repo from app.api.v1 import compare as compare_api -from app.core.config_schema import ( - RISK_COMPARE_DAILY_THRESHOLD_KEY, - RISK_ONECLICK_DAILY_THRESHOLD_KEY, - RISK_SMS_HOURLY_THRESHOLD_KEY, -) +from app.core.config_schema import LIMIT_POLICY_GLOBAL_KEY from app.core.security import issue_token_pair from app.db.session import SessionLocal from app.models.app_config import AppConfig @@ -28,9 +24,7 @@ from app.repositories import user as user_repo @pytest.fixture(autouse=True) def _reset_risk_rule_config(): keys = ( - RISK_SMS_HOURLY_THRESHOLD_KEY, - RISK_ONECLICK_DAILY_THRESHOLD_KEY, - RISK_COMPARE_DAILY_THRESHOLD_KEY, + LIMIT_POLICY_GLOBAL_KEY, risk_repo.RISK_RESET_BASELINES_KEY, ) with SessionLocal() as db: @@ -533,7 +527,7 @@ def test_admin_can_edit_rules_and_current_window_is_reconciled() -> None: "/admin/api/risk-monitor/rules", headers=headers, json={ - "sms_hourly_threshold": 21, + "sms_hourly_threshold": 100_001, "oneclick_daily_threshold": 20, "compare_daily_threshold": 100, }, diff --git a/tests/test_sms_fallback.py b/tests/test_sms_fallback.py index cb511b7..91c56e1 100644 --- a/tests/test_sms_fallback.py +++ b/tests/test_sms_fallback.py @@ -94,6 +94,30 @@ def test_send_code_backup_also_fails_raises_backup_error(monkeypatch): assert "创蓝" in str(ei.value) +def test_dynamic_cooldown_is_forwarded_to_primary_and_backup(monkeypatch): + calls = [] + + def primary(phone, *, cooldown_sec=None): + calls.append(("jiguang", cooldown_sec)) + raise SmsError("极光不可用", status_code=503) + + def backup(phone, *, cooldown_sec=None): + calls.append(("chuanglan", cooldown_sec)) + return cooldown_sec + + monkeypatch.setattr(settings, "SMS_PROVIDER", "jiguang") + monkeypatch.setattr(settings, "SMS_FALLBACK_PROVIDER", "chuanglan") + monkeypatch.setattr(jiguang, "send_code", primary) + monkeypatch.setattr(chuanglan, "send_code", backup) + + result = sms.send_code(PHONE, cooldown_sec=15) + + assert result.cooldown_sec == 15 + assert result.provider == "chuanglan" + assert result.fallback is True + assert calls == [("jiguang", 15), ("chuanglan", 15)] + + def test_verify_hits_primary_without_touching_backup(monkeypatch): calls = [] monkeypatch.setattr(settings, "SMS_PROVIDER", "jiguang") @@ -140,3 +164,22 @@ def test_verify_backup_not_touched_when_fallback_off(monkeypatch): assert sms.verify_code(PHONE, "123456") is False assert calls == ["jiguang"] + + +def test_dynamic_verify_attempt_limit_is_forwarded_to_fallback_chain(monkeypatch): + calls = [] + + def verify(provider, result): + def _verify(phone, code, *, max_failed_attempts=None): + calls.append((provider, max_failed_attempts)) + return result + + return _verify + + monkeypatch.setattr(settings, "SMS_PROVIDER", "jiguang") + monkeypatch.setattr(settings, "SMS_FALLBACK_PROVIDER", "chuanglan") + monkeypatch.setattr(jiguang, "verify_code", verify("jiguang", False)) + monkeypatch.setattr(chuanglan, "verify_code", verify("chuanglan", True)) + + assert sms.verify_code(PHONE, "123456", max_failed_attempts=2) is True + assert calls == [("jiguang", 2), ("chuanglan", 2)]