Compare commits

..

6 Commits

Author SHA1 Message Date
unknown 3f2ec19ec2 功能:统一限制策略与白名单管理
新增统一限制规则、临时不限和风控免告警白名单,接入比价、短信登录、广告、引导视频与账号冷却等业务链路。补充设备选择、权限、审计、迁移及回归测试。
2026-07-29 19:07:51 +08:00
zuochenyong 53c3b7f60f feat(guide-video): 支持领券和比价独立视频奖励配置 (#196)
Co-authored-by: exinglang <exinglang@qq.com>
Reviewed-on: #196
Co-authored-by: zuochenyong <zuochenyong@wonderable.ai>
Co-committed-by: zuochenyong <zuochenyong@wonderable.ai>
2026-07-29 16:12:05 +08:00
marco 4bd4e66678 重构比价结果页取数据逻辑 (#195)
Reviewed-on: #195
2026-07-29 01:59:31 +08:00
linkeyu 90c6fe599a 修复:统一用户Draw信息流eCPM统计口径 (#190)
## 问题

业务收益详情的平均 Draw eCPM 仅平均成功发奖记录,会排除未发奖的真实展示,导致数值系统性偏高,且与广告收益页口径不一致。

## 修复

- `feed_avg_ecpm` 改为从 `ad_ecpm_record` 的全部 `draw/feed` 实际展示计算
- 成功发奖、未发奖展示均纳入,每次展示等权
- 日期、正式/测试环境、业务代码位、领券/比价场景支持与广告收益页对齐
- 奖励份数仍基于成功发奖表,不混用展示数据源
- 复用广告收益报表的业务代码位集合

## 线上数据复算

2026-07-25、正式业务、用户 #33:

- 旧口径(只看成功发奖):`29.9117 元/千次`
- 新口径(333 次真实展示):`19.9926 元/千次`
- 新值与广告收益报表一致

## 验证

- 新增成功/未发奖、场景、环境、业务代码位回归用例
- `tests/test_admin_read.py` + `tests/test_admin_ad_revenue_scope.py`:29 项全通过
- Ruff 改动文件检查通过

## 上线顺序

本 PR 需先于管理后台配套 PR 上线。

---------

Co-authored-by: guke <guke@wonderable.ai>
Co-authored-by: unknown <798648091@qq.com>
Reviewed-on: #190
Co-authored-by: linkeyu <linkeyu@wonderable.ai>
Co-committed-by: linkeyu <linkeyu@wonderable.ai>
2026-07-28 17:58:23 +08:00
linkeyu e529112a90 修复中途退出比价的 LLM 成本回填 (#191)
## 问题

比价记录进入中途退出后未触发 LLM 成本回填,周期补偿也未扫描 cancelled,导致实际已有 LLM 调用的记录长期显示成本、LLM、TOKEN 为空。

## 修改

- finalize 落库后立即追加 LLM 成本回填
- 周期补偿范围加入 cancelled
- 保持无有效调用和全调用失败记录不伪造成本
- 增加即时回填和周期补偿回归测试

## 验证

- ruff 检查通过
- 相关测试 28 项通过
- 全仓 626 项通过;主干既有失败已在未修改的 origin/main 复现

---------

Co-authored-by: unknown <798648091@qq.com>
Reviewed-on: #191
Co-authored-by: linkeyu <linkeyu@wonderable.ai>
Co-committed-by: linkeyu <linkeyu@wonderable.ai>
2026-07-28 17:57:52 +08:00
guke 50da718e35 比价记录失败卡展示具体原因(新增 fail_reason) (#189)
失败记录不再一律「网络开小差」:新增记录级 fail_reason 派生列——information
具体则直出,笼统则从 platform_results 救出业务原因(找不到店/菜、未起送、打烊、
单点不配送等),纯系统失败为 None → 端侧品牌兜底。store_closed/no_delivery 被
pricebot 漏成 status=failed 的按 reason 补判,打烊脏店名统一简短模板。接入
harvest_done 与灰度期 upsert_record 两条写路径。

- models: comparison_record.fail_reason 列
- repositories: _derive_fail_display + 补判/清洗 helper,两条写路径接入
- schemas: ComparisonRecordOut 暴露 fail_reason
- alembic: 加列 + 回填老 specific 失败记录
- tests: _derive_fail_display 单测(8 例)+ harvest 失败落库集成测试

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>

---------

Co-authored-by: guke <guke@autohome.com.cn>
Reviewed-on: #189
2026-07-28 14:04:36 +08:00
66 changed files with 4382 additions and 1929 deletions
+2 -11
View File
@@ -6,17 +6,8 @@ APP_NAME=shaguabijia-app-server
APP_DEBUG=true
# ===== 数据库 =====
# 本地开发/测试统一用 Docker PostgreSQL:run.bat/run.sh 会自动拉起容器
# (docker-compose.yml + scripts/ensure_pg.py)。详见 docs/database/postgres-migration.md。
# 生产用原生 PG,由 scripts/init_postgres.py 写入强随机密码的连接串。
# ⚠️ scheme 必须是 postgresql+psycopg://(psycopg3);不要写成 postgresql://(会去找未装的 psycopg2)。
DATABASE_URL=postgresql+psycopg://shaguabijia_app:shaguabijia_dev_pw@localhost:5432/shaguabijia
# Docker 自动定位:run.bat/run.sh 会自动找 Docker Desktop 并启动——优先从 PATH 上的 docker CLI 反推
# 安装目录(装在 D 盘等非默认盘符也能找到),再退到注册表 / 常见目录。仅当你的安装位置极特殊、自动
# 探测失败时,才需下面这行显式指到 exe(值可含空格,直接写到行尾即可,无需引号):
# DOCKER_DESKTOP_EXE=D:\Program Files\Docker\Docker\Docker Desktop.exe
# 实在不想装/启 Docker → 把上面 DATABASE_URL 改成 sqlite 可降级跑(仅救急,PG 专有 SQL/严格性不被验证):
# DATABASE_URL=sqlite:///./data/app.db
# SQLite 本地文件路径。生产环境用 /opt/shaguabijia-app-server/data.db
DATABASE_URL=sqlite:///./data/app.db
# ===== JWT =====
# 生产部署务必改成随机长字符串,可用:python -c "import secrets; print(secrets.token_urlsafe(64))"
-3
View File
@@ -60,6 +60,3 @@ tests/meituan_coupon_bj.tsv
tests/meituan_coupon_data.tsv
tests/meituan_coupon_fz.tsv
tests/meituan_coupon_xm.tsv
# git worktrees (superpowers 隔离工作区)
.worktrees/
+3 -3
View File
@@ -76,8 +76,8 @@ Endpoints under `app/api/internal/` are for server-to-server communication (pric
## Database
- **Dev/Test**: Docker PostgreSQL 16 — `run.sh`/`run.bat` auto-start it via `scripts/ensure_pg.py` + `docker-compose.yml`; `.env.example` ships the PG URL by default; pytest uses the same container's `shaguabijia_test` DB. **Local no longer uses SQLite** (the SQLite branch in `db/session.py` is retained as a fallback only).
- **Prod**: native PostgreSQL — bootstrap with `scripts/init_postgres.py` (no Docker). Pool size 10 + max overflow 20, pool_recycle 3600.
- **Dev**: SQLite (`sqlite:///./data/app.db`), `check_same_thread=False`, no connection pool.
- **Prod**: PostgreSQL — just change `DATABASE_URL` in `.env`. Pool size 10 + max overflow 20, pool_recycle 3600.
- **Migrations**: Alembic with `render_as_batch` for SQLite compatibility. ~60+ migration files in `alembic/versions/` (filenames are descriptive, not hex prefixes). Migration chain uses `down_revision` within each file.
- **New models**: Define in `app/models/`, import in `app/models/__init__.py`, then run `alembic revision --autogenerate`.
@@ -89,7 +89,7 @@ All config via `pydantic-settings` in `app/core/config.py`. Single `Settings` cl
## Testing
- `tests/conftest.py`: Sets env vars BEFORE imports, ensures the Docker PG `shaguabijia_test` DB via `scripts/ensure_pg.py`, builds all tables with `Base.metadata.create_all()` (drop+create for a clean start), tears down with `drop_all()`.
- `tests/conftest.py`: Sets env vars BEFORE imports, creates temp SQLite file, builds all tables with `Base.metadata.create_all()`, tears down with `drop_all()` + unlink.
- External integrations are monkeypatched in tests (e.g., WeChat Pay, Jiguang, Pangle callbacks) — tests never make real HTTP calls.
- `TestClient` from FastAPI is used for all tests. Rate limiting is disabled globally in tests.
@@ -0,0 +1,26 @@
"""merge comparison platforms + fail_reason heads
Revision ID: 6d2309208549
Revises: comparison_platforms_col, comparison_record_fail_reason
Create Date: 2026-07-29 01:48:41.868083
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = '6d2309208549'
down_revision: Union[str, Sequence[str], None] = ('comparison_platforms_col', 'comparison_record_fail_reason')
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
pass
def downgrade() -> None:
pass
@@ -0,0 +1,47 @@
"""add platforms unified array column to comparison_record
展示模型统一数组(pricebot done.params.platforms 原样存): 每平台一行、自带
status/is_best/display, 记录页据此直接渲染, 不再靠 comparison_results + 客户端合并 + 前端派生。
纯新增列, 老记录为空 → 前端回退老 comparison_results。
Revision ID: comparison_platforms_col
Revises: user_manual_risk_fields
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy.dialects import postgresql
from alembic import op
revision: str = "comparison_platforms_col"
down_revision: str | None = "user_manual_risk_fields"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
_JSON = sa.JSON().with_variant(postgresql.JSONB(), "postgresql")
def upgrade() -> None:
# 幂等: 线上为了提前给历史数据补 platforms(2026-07-29), 已手动
# `ALTER TABLE comparison_record ADD COLUMN IF NOT EXISTS platforms jsonb
# NOT NULL DEFAULT '[]'::jsonb`(与本 migration 定义一致)。列已存在时跳过,
# 否则上线 alembic upgrade head 会撞 DuplicateColumn 直接部署失败。
bind = op.get_bind()
cols = {c["name"] for c in sa.inspect(bind).get_columns("comparison_record")}
if "platforms" in cols:
return
with op.batch_alter_table("comparison_record") as batch_op:
batch_op.add_column(
sa.Column(
"platforms", _JSON, nullable=False,
server_default=sa.text("'[]'"),
)
)
def downgrade() -> None:
with op.batch_alter_table("comparison_record") as batch_op:
batch_op.drop_column("platforms")
@@ -0,0 +1,49 @@
"""comparison_record.fail_reason (失败卡展示原因)
Revision ID: comparison_record_fail_reason
Revises: user_manual_risk_fields
Create Date: 2026-07-28 12:00:00.000000
失败记录的展示原因:information 具体则=它;笼统则由写路径从 platform_results 捞出的
业务原因;纯系统失败为 None(端侧品牌兜底)。见 repositories.comparison._derive_fail_display。
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = 'comparison_record_fail_reason'
down_revision: Union[str, Sequence[str], None] = 'user_manual_risk_fields'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
with op.batch_alter_table('comparison_record', schema=None) as batch_op:
batch_op.add_column(sa.Column('fail_reason', sa.String(length=256), nullable=True))
# 回填老失败记录:information 具体的直接搬过来(笼统/系统失败留 None → 端侧品牌兜底)。
# 新记录由写路径 _derive_fail_display 落库(含 platform_results 救援/补判),不走这条。
# platform_results 只在 raw_payload 里,SQL 里不易解析,故老记录不做救援/补判(可接受:
# 老 mixed/打烊记录回退品牌兜底);具体 information 的老记录本次即可显示真实原因。
op.execute(
"""
UPDATE comparison_record
SET fail_reason = information
WHERE status = 'failed'
AND information IS NOT NULL
AND information <> ''
AND information NOT IN (
'比价过程出错,请稍后重试',
'比价出错',
'比价未完成',
'done 参数缺少可验证的目标平台结果'
)
"""
)
def downgrade() -> None:
with op.batch_alter_table('comparison_record', schema=None) as batch_op:
batch_op.drop_column('fail_reason')
@@ -0,0 +1,31 @@
"""guide video play count is independent for coupon and comparison
Revision ID: guide_video_scene_unique
Revises: 6d2309208549
"""
from alembic import op
revision = "guide_video_scene_unique"
down_revision = "6d2309208549"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.drop_index("uq_guide_video_play_user_seq", table_name="guide_video_play")
op.create_index(
"uq_guide_video_play_user_scene_seq",
"guide_video_play",
["user_id", "scene", "seq"],
unique=True,
)
def downgrade() -> None:
op.drop_index("uq_guide_video_play_user_scene_seq", table_name="guide_video_play")
op.create_index(
"uq_guide_video_play_user_seq",
"guide_video_play",
["user_id", "seq"],
unique=True,
)
+128
View File
@@ -0,0 +1,128 @@
"""add per-subject limit policy whitelist
Revision ID: limit_policy_whitelist
Revises: guide_video_scene_unique
"""
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 = "guide_video_scene_unique"
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")
+2
View File
@@ -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)
+3 -1
View File
@@ -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",
]},
]
+5 -5
View File
@@ -43,7 +43,7 @@ _KNOWN_PROD_BUSINESS_CODE_IDS = frozenset({"104098712", "104099389"})
_TEST_BUSINESS_CODE_IDS = frozenset({"104127529", "104127626", "104137445"})
def _business_code_ids(db: Session, app_env: str | None) -> set[str]:
def business_code_ids(db: Session, app_env: str | None) -> set[str]:
"""返回指定应用环境下可用于业务收益对账的 GroMore 聚合代码位。"""
prod_config = app_config.get_ad_config(db)
prod_ids = set(_KNOWN_PROD_BUSINESS_CODE_IDS) | {
@@ -320,10 +320,10 @@ def ad_revenue_report(
# 业务口径仅保留正式配置/测试业务链路实际使用的代码位。穿山甲“全量”还包含广告测试
# demo、插屏等没有客户端收益上报的曝光,两边直接比较会天然产生假差额。
business_code_ids: set[str] | None = None
business_ids: set[str] | None = None
if revenue_scope == "business":
business_code_ids = _business_code_ids(db, app_env)
events = [e for e in events if e.get("our_code_id") in business_code_ids]
business_ids = business_code_ids(db, app_env)
events = [e for e in events if e.get("our_code_id") in business_ids]
# 排序:time=按时间倒序(新→旧);ecpm=按 eCPM 数值倒序(eCPM 原值是字符串「分」,转数值排;
# 纯发奖行用其发奖采用的 eCPM,缺失/非法计 0 排末尾)。
@@ -381,7 +381,7 @@ def ad_revenue_report(
date_from=date_from,
date_to=date_to,
app_env=app_env,
our_code_ids=business_code_ids,
our_code_ids=business_ids,
)
if pangle_aggs:
by_date = {a["date"]: a for a in pangle_aggs}
+455
View File
@@ -0,0 +1,455 @@
"""CRUD and presentation helpers for limit policy overrides."""
from __future__ import annotations
from datetime import UTC, datetime
from sqlalchemy import String, cast, func, or_, select
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 app_config
from app.repositories import risk as risk_repo
class DuplicateOverrideError(Exception):
pass
_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 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
app_config.set_value(db, rule.config_key, 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
+57 -14
View File
@@ -14,6 +14,7 @@ from sqlalchemy.orm import Session
from app.core import rewards
from app.core.config import settings
from app.models.ad_ecpm import AdEcpmRecord
from app.models.ad_feed_reward import AdFeedRewardRecord
from app.models.ad_reward import AdRewardRecord
from app.models.admin import AdminAuditLog
@@ -1257,11 +1258,15 @@ def user_reward_stats(
date_from: datetime | None = None,
date_to: datetime | None = None,
withdraw_source: str | None = None,
app_env: str | None = None,
revenue_scope: str = "all",
feed_scene: str | None = None,
) -> dict:
"""提现详情「用户统计区」10 项。窗口作用于除「现金余额」外的所有项(余额是当前快照)。
口径:激励视频/信息流只统计 granted;数量——视频按条数、信息流按份数(unit_count 累加);
平均 eCPM 用原始分值(分/千次)按记录取算术平均;各「提现」= 该来源累计金币折现。
口径:激励视频/信息流奖励数量只统计 granted;数量——视频按条数、信息流按份数(unit_count 累加)
平均 Draw eCPM 与广告收益报表一致:基于 ad_ecpm_record 的全部 draw/feed 展示记录计算,
不以是否发奖为筛选条件。各「提现」= 该来源累计金币折现。
传统任务 = 窗口内正向金币中,排除广告(reward_video/feed_ad_reward)与人工调整后的折现。
"""
withdraw_source_conds = (
@@ -1296,30 +1301,68 @@ def user_reward_stats(
# 只投影本统计实际使用的列。避免滚动发布或旧本地库尚未补齐无关新列时,
# SQLAlchemy 因 select(ORM) 自动展开整表字段而让提现详情整体 500。
business_ids: set[str] | None = None
if revenue_scope == "business":
# 与广告收益报表共用正式/测试业务代码位集合,避免两个页面随配置切换后再次漂移。
from app.admin.repositories.ad_revenue import business_code_ids
business_ids = business_code_ids(db, app_env)
rv_conds = [
AdRewardRecord.user_id == user_id,
AdRewardRecord.reward_scene == "reward_video",
AdRewardRecord.status == "granted",
*_window_conds(AdRewardRecord.created_at, date_from, date_to),
]
if app_env is not None:
rv_conds.append(AdRewardRecord.app_env == app_env)
if business_ids is not None:
rv_conds.append(AdRewardRecord.our_code_id.in_(business_ids))
rv = db.execute(
select(AdRewardRecord.ecpm_raw, AdRewardRecord.coin).where(
AdRewardRecord.user_id == user_id,
AdRewardRecord.reward_scene == "reward_video",
AdRewardRecord.status == "granted",
*_window_conds(AdRewardRecord.created_at, date_from, date_to),
*rv_conds,
)
).all()
rv_ecpms = [rewards.parse_ecpm_fen(r.ecpm_raw) for r in rv if r.ecpm_raw]
rv_coins = sum(r.coin for r in rv)
feed = db.execute(
feed_reward_conds = [
AdFeedRewardRecord.user_id == user_id,
AdFeedRewardRecord.status == "granted",
*_window_conds(AdFeedRewardRecord.created_at, date_from, date_to),
]
if app_env is not None:
feed_reward_conds.append(AdFeedRewardRecord.app_env == app_env)
if feed_scene is not None:
feed_reward_conds.append(AdFeedRewardRecord.feed_scene == feed_scene)
if business_ids is not None:
feed_reward_conds.append(AdFeedRewardRecord.our_code_id.in_(business_ids))
feed_rewards = db.execute(
select(
AdFeedRewardRecord.unit_count,
AdFeedRewardRecord.ecpm_raw,
AdFeedRewardRecord.coin,
).where(
AdFeedRewardRecord.user_id == user_id,
AdFeedRewardRecord.status == "granted",
*_window_conds(AdFeedRewardRecord.created_at, date_from, date_to),
*feed_reward_conds,
)
).all()
feed_ecpms = [rewards.parse_ecpm_fen(f.ecpm_raw) for f in feed if f.ecpm_raw]
feed_coins = sum(f.coin for f in feed)
feed_coins = sum(f.coin for f in feed_rewards)
feed_impression_conds = [
AdEcpmRecord.user_id == user_id,
AdEcpmRecord.ad_type.in_(("draw", "feed")),
*_window_conds(AdEcpmRecord.created_at, date_from, date_to),
]
if app_env is not None:
feed_impression_conds.append(AdEcpmRecord.app_env == app_env)
if feed_scene is not None:
feed_impression_conds.append(AdEcpmRecord.feed_scene == feed_scene)
if business_ids is not None:
feed_impression_conds.append(AdEcpmRecord.our_code_id.in_(business_ids))
feed_impressions = db.execute(
select(AdEcpmRecord.ecpm_raw).where(*feed_impression_conds)
).all()
# 与 ad_revenue.category_stats 相同:每次展示权重相同,非法原值按 parse_ecpm_fen 记 0。
feed_ecpms = [rewards.parse_ecpm_fen(row.ecpm_raw) for row in feed_impressions]
trad_coins = db.execute(
select(func.coalesce(func.sum(CoinTransaction.amount), 0)).where(
@@ -1338,7 +1381,7 @@ def user_reward_stats(
"reward_video_count": len(rv),
"reward_video_avg_ecpm": round(sum(rv_ecpms) / len(rv_ecpms), 2) if rv_ecpms else 0.0,
"reward_video_cash_cents": _coins_to_cents(rv_coins),
"feed_count": int(sum(f.unit_count for f in feed)),
"feed_count": int(sum(f.unit_count for f in feed_rewards)),
"feed_avg_ecpm": round(sum(feed_ecpms) / len(feed_ecpms), 2) if feed_ecpms else 0.0,
"feed_cash_cents": _coins_to_cents(feed_coins),
}
+22 -1
View File
@@ -12,7 +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.config_schema import CONFIG_DEFS
from app.core.config_schema import (
AD_FEED_DAILY_LIMIT_KEY,
AD_REWARD_VIDEO_DAILY_LIMIT_KEY,
CONFIG_DEFS,
)
from app.core.rewards import SIGNIN_CYCLE_LEN
from app.models.admin import AdminUser
from app.repositories import app_config
@@ -88,6 +92,23 @@ def update_config(
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 两项,避免旧入口写入后业务实际值不变。
app_config.set_value(
db,
AD_REWARD_VIDEO_DAILY_LIMIT_KEY,
body.value,
admin_id=admin.id,
commit=False,
)
app_config.set_value(
db,
AD_FEED_DAILY_LIMIT_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,
+35 -14
View File
@@ -9,7 +9,7 @@ client_max_body_size,见 shaguabijia-admin-web/deploy/nginx/admin.shaguabijia.co
"""
from __future__ import annotations
from typing import Annotated
from typing import Annotated, Literal
from fastapi import APIRouter, Depends, File, HTTPException, Request, UploadFile
@@ -26,15 +26,21 @@ router = APIRouter(
dependencies=[Depends(get_current_admin)],
)
GuideScene = Literal["coupon", "comparison"]
def _out(db: AdminDb) -> GuideVideoConfigOut:
def _out(db: AdminDb, scene: GuideScene) -> GuideVideoConfigOut:
"""配置 + 播放统计合成响应(四个写接口都以最新状态返回,前端一次同步到位)。"""
return GuideVideoConfigOut(**guide_video.get_config(db), **guide_video.play_stats(db))
return GuideVideoConfigOut(
scene=scene,
**guide_video.get_config(db, scene),
**guide_video.play_stats(db, scene),
)
@router.get("", response_model=GuideVideoConfigOut, summary="新手引导视频配置(领券浮层)")
def get_config(db: AdminDb) -> GuideVideoConfigOut:
return _out(db)
def get_config(db: AdminDb, scene: GuideScene = "coupon") -> GuideVideoConfigOut:
return _out(db, scene)
@router.patch("", response_model=GuideVideoConfigOut, summary="改开关/次数/金币(带审计)")
@@ -43,21 +49,24 @@ def update_config(
request: Request,
admin: Annotated[AdminUser, Depends(require_role("operator"))],
db: AdminDb,
scene: GuideScene = "coupon",
) -> GuideVideoConfigOut:
before, after = guide_video.update_config(
db,
enabled=body.enabled,
max_plays=body.max_plays,
reward_coin=body.reward_coin,
scene=scene,
admin_id=admin.id,
commit=False,
)
write_audit(
db, admin, action="guide_video.update", target_type="guide_video", target_id=None,
detail={"before": before, "after": after}, ip=get_client_ip(request), commit=False,
detail={"scene": scene, "before": before, "after": after},
ip=get_client_ip(request), commit=False,
)
db.commit()
return _out(db)
return _out(db, scene)
@router.post("/video", response_model=GuideVideoConfigOut, summary="上传新手引导视频(MP4,带审计)")
@@ -65,23 +74,31 @@ async def upload_video(
request: Request,
admin: Annotated[AdminUser, Depends(require_role("operator"))],
db: AdminDb,
file: UploadFile = File(...),
file: Annotated[UploadFile, File()],
scene: GuideScene = "coupon",
) -> GuideVideoConfigOut:
data = await file.read()
try:
url = media.save_guide_video(data)
except media.MediaError as e:
raise HTTPException(status_code=400, detail=str(e)) from e
before, after = guide_video.set_video(db, url, admin_id=admin.id, commit=False)
before, after = guide_video.set_video(
db, url, scene=scene, admin_id=admin.id, commit=False
)
write_audit(
db, admin, action="guide_video.set_video", target_type="guide_video", target_id=None,
detail={"before": before.get("video_url"), "after": url, "bytes": len(data)},
detail={
"scene": scene,
"before": before.get("video_url"),
"after": url,
"bytes": len(data),
},
ip=get_client_ip(request), commit=False,
)
db.commit()
# 提交成功后再删旧片,避免新片没落库就把旧片丢了
media.delete_guide_video(before.get("video_url"))
return _out(db)
return _out(db, scene)
@router.delete("/video", response_model=GuideVideoConfigOut, summary="移除新手引导视频(带审计)")
@@ -89,13 +106,17 @@ def delete_video(
request: Request,
admin: Annotated[AdminUser, Depends(require_role("operator"))],
db: AdminDb,
scene: GuideScene = "coupon",
) -> GuideVideoConfigOut:
"""移除后 /guide-video/start 一律返回 should_play=false,领券浮层回到「只放广告」。"""
before, after = guide_video.set_video(db, None, admin_id=admin.id, commit=False)
before, after = guide_video.set_video(
db, None, scene=scene, admin_id=admin.id, commit=False
)
write_audit(
db, admin, action="guide_video.delete_video", target_type="guide_video", target_id=None,
detail={"before": before.get("video_url")}, ip=get_client_ip(request), commit=False,
detail={"scene": scene, "before": before.get("video_url")},
ip=get_client_ip(request), commit=False,
)
db.commit()
media.delete_guide_video(before.get("video_url"))
return _out(db)
return _out(db, scene)
+361
View File
@@ -0,0 +1,361 @@
"""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,
)
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"))],
)
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)
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")
@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.post(
"",
response_model=LimitOverrideOut,
status_code=status.HTTP_201_CREATED,
summary="新增白名单配置",
)
def create_override(
body: LimitOverrideWrite,
request: Request,
admin: CurrentAdmin,
db: AdminDb,
) -> LimitOverrideOut:
try:
row = repo.create(
db,
**body.model_dump(),
limit_value=None,
admin_id=admin.id,
)
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
_reconcile_risk_rule(db, row.rule_code)
result = _out(db, row)
write_audit(
db,
admin,
action="limit.override.create",
target_type="limit_override",
target_id=str(row.id),
detail={"after": _audit_payload(result)},
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] = []
current_rule_code = ""
try:
if body.subject_type == limit_policy.SUBJECT_DEVICE:
limit_policy.validate_device_rule_scope(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
)
rows.append(
repo.create(
db,
subject_type=body.subject_type,
subject_value=body.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,
reactivate_inactive=True,
)
)
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={"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)
before = _audit_payload(_out(db, row))
try:
repo.update(
db,
row,
**body.model_dump(),
fields_set=set(body.model_fields_set),
)
except (KeyError, ValueError) as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
_reconcile_risk_rule(db, row.rule_code)
after = _out(db, row)
write_audit(
db,
admin,
action="limit.override.update",
target_type="limit_override",
target_id=str(row.id),
detail={"before": before, "after": _audit_payload(after)},
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)
before = _audit_payload(_out(db, row))
repo.restore_global(row)
_reconcile_risk_rule(db, row.rule_code)
after = _out(db, row)
write_audit(
db,
admin,
action="limit.override.restore_global",
target_type="limit_override",
target_id=str(row.id),
detail={"before": before, "after": _audit_payload(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)
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()
+8
View File
@@ -88,6 +88,11 @@ def get_user_reward_stats(
withdraw_source: Annotated[
str | None, Query(pattern="^(coin_cash|invite_cash)$")
] = None,
app_env: Annotated[str | None, Query(pattern="^(prod|test)$")] = None,
revenue_scope: Annotated[str, Query(pattern="^(business|all)$")] = "all",
feed_scene: Annotated[
str | None, Query(pattern="^(comparison|coupon|welfare)$")
] = None,
) -> UserRewardStats:
"""提现详情抽屉「用户统计区」。date_from/date_to 都不传 = 注册至今(全量)。"""
if not user_repo.user_exists(db, user_id):
@@ -99,6 +104,9 @@ def get_user_reward_stats(
date_from=date_from,
date_to=date_to,
withdraw_source=withdraw_source,
app_env=app_env,
revenue_scope=revenue_scope,
feed_scene=feed_scene,
)
)
+1
View File
@@ -7,6 +7,7 @@ from app.repositories.guide_video import MAX_PLAYS_LIMIT, REWARD_COIN_LIMIT
class GuideVideoConfigOut(BaseModel):
scene: str
enabled: bool
video_url: str | None = None # 相对地址 /media/guide_video/xxx.mp4;未配片 = None
max_plays: int
+127
View File
@@ -0,0 +1,127 @@
"""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 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
+2 -2
View File
@@ -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):
+1 -1
View File
@@ -60,7 +60,7 @@ class UserRewardStats(BaseModel):
reward_video_avg_ecpm: float # 平均激励视频 eCPM(分/千次)
reward_video_cash_cents: int # 激励视频提现(金币折现)
feed_count: int # 累计信息流广告数(granted 份数,unit_count 累加)
feed_avg_ecpm: float # 平均信息流广告 eCPM(分/千次)
feed_avg_ecpm: float # 全部 Draw/feed 实际展示的平均 eCPM(分/千次,含未发奖展示)
feed_cash_cents: int # 信息流广告提现(金币折现)
+13 -2
View File
@@ -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
),
)
+185 -26
View File
@@ -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:
cooldown = send_code(req.phone)
cooldown = 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)
),
)
except SmsError as e:
risk_repo.record_behavior_event(
db,
@@ -224,7 +262,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 避免循环
@@ -283,17 +321,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:
@@ -394,15 +463,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),
@@ -418,7 +497,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
),
)
@@ -451,18 +535,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:
@@ -522,9 +638,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"])
@@ -562,21 +692,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(
+11 -3
View File
@@ -108,7 +108,7 @@ def _harvest_done_blocking(
def _harvest_abort_blocking(
trace_id: str, status_hint: str, reason: str | None, trace_url: str | None,
) -> None:
) -> int | None:
with SessionLocal() as db:
rec = crud_compare.harvest_abort(
db, trace_id=trace_id, status=status_hint, reason=reason, trace_url=trace_url,
@@ -118,6 +118,7 @@ def _harvest_abort_blocking(
extra={"phase": "harvest_abort",
"status": (rec.status if rec else None), "reason": reason},
)
return rec.id if rec is not None else None
async def _forward(
@@ -291,7 +292,10 @@ async def trace_epilogue(
@router.post("/trace/finalize", summary="比价 trace 收尾上云 (透传 + 夭折落库)")
async def trace_finalize(
request: Request, user: OptionalUser, db: DbSession
request: Request,
background_tasks: BackgroundTasks,
user: OptionalUser,
db: DbSession,
) -> dict[str, Any]:
_ensure_compare_allowed(user, db)
# 用户终止 / Phase1 未识别没到 done 帧: pricebot 打包半截上云返回 {trace_url};
@@ -302,12 +306,16 @@ async def trace_finalize(
request, "/api/trace/finalize", user, harvest_first_frame=False,
)
try:
await run_in_threadpool(
record_id = await run_in_threadpool(
_harvest_abort_blocking, trace_id,
(meta.get("status") or "cancelled"),
(meta.get("reason") or meta.get("information")),
(resp.get("trace_url") if isinstance(resp, dict) else None),
)
if record_id is not None:
background_tasks.add_task(
backfill_comparison_llm_cost, record_id, trace_id
)
except Exception as e: # noqa: BLE001
logger.warning("harvest_abort failed trace=%s: %s", trace_id, e)
return resp
+22 -4
View File
@@ -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,
)
+130 -4
View File
@@ -14,12 +14,137 @@ 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"
# 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 +181,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", "hidden": True,
"help": "历史共享上限;新配置由白名单页分别管理激励视频与 Draw 信息流。",
},
"ad_max_coin": {
"default": r.MAX_AD_REWARD_COIN, "label": "看广告单次金币上限",
@@ -118,9 +244,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 +264,7 @@ CONFIG_DEFS: dict[str, dict[str, Any]] = {
"group": "风控",
"type": "int",
"min": 1,
"max": 100,
"max": 100_000,
"hidden": True,
"help": "同一账户在北京时间同一自然日内发起比价达到该次数时告警。",
},
+525
View File
@@ -0,0 +1,525 @@
"""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 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,
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}
@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 validate_device_rule_scope(rule_codes: Iterable[str]) -> None:
scopes = {device_source_scope(rule_code) for rule_code in rule_codes}
if len(scopes) > 1:
raise ValueError("比价设备限制与短信/登录设备限制不能在同一白名单中混选")
def _global_limit(db: Session, rule: RuleDefinition) -> tuple[int, str]:
row = db.get(AppConfig, rule.config_key)
if row is not None:
return int(row.value), "configured"
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 _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] = []
for rule in RULES:
value, _ = _global_limit(db, rule)
out.append(
{
"code": rule.code,
"label": rule.label,
"group": rule.group,
"global_limit": value,
"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()
+19 -6
View File
@@ -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,
)
+15 -4
View File
@@ -27,11 +27,22 @@ def _provider():
return _PROVIDERS.get(settings.SMS_PROVIDER, jiguang)
def send_code(phone: str) -> int:
def send_code(phone: str, *, cooldown_sec: int | None = None) -> int:
"""发送验证码,返回距下次可发的冷却秒数;失败抛 SmsError。委托给当前 provider。"""
return _provider().send_code(phone)
provider = _provider()
if cooldown_sec is None:
return provider.send_code(phone)
return provider.send_code(phone, cooldown_sec=cooldown_sec)
def verify_code(phone: str, code: str) -> bool:
def verify_code(
phone: str,
code: str,
*,
max_failed_attempts: int | None = None,
) -> bool:
"""校验验证码,返回是否通过;provider 异常降级抛 SmsError。委托给当前 provider。"""
return _provider().verify_code(phone, code)
provider = _provider()
if max_failed_attempts is None:
return provider.verify_code(phone, code)
return provider.verify_code(phone, code, max_failed_attempts=max_failed_attempts)
+29 -8
View File
@@ -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
+19 -6
View File
@@ -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")):
+19 -6
View File
@@ -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")):
+1
View File
@@ -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
+9
View File
@@ -102,12 +102,21 @@ class ComparisonRecord(Base):
# done 帧 information 文案。成功:"在美团找到同店,到手价 ¥X…";
# 失败:具体原因(如"美团、京东外卖均未找到该商品")。前端在比价失败时当原因展示。
information: Mapped[str | None] = mapped_column(String(256), nullable=True)
# 失败卡「原因」行的展示文案(仅 status=failed 时非空):information 具体则=它;笼统则从
# platform_results 捞出的业务原因(打烊/未起送/找不到店或菜/单点不配送);纯系统失败为 None
# → 端侧显示品牌兜底「网络开小差…」。写路径(harvest_done / upsert_record)落库时派生。
# 见 repositories.comparison._derive_fail_display。
fail_reason: Mapped[str | None] = mapped_column(String(256), nullable=True)
# ===== 明细(JSON,越详细越好)=====
# 下单菜品 [{name, qty, specs?}]
items: Mapped[list] = mapped_column(_JSON, nullable=False, default=list)
# 逐平台对比 [{platform_id, platform_name, package, price, is_source, rank, coupon_saved, coupon_name, applied_coupons}](price/coupon_saved 单位:元,原样存;coupon_name=优惠来源名;applied_coupons=[{name,amount}] 多券明细)
comparison_results: Mapped[list] = mapped_column(_JSON, nullable=False, default=list)
# 展示模型统一数组(pricebot done.params.platforms 原样存): 每平台一行、自带
# status/is_best/display/display_order, 记录页据此直接渲染, 不再靠 comparison_results
# + 客户端合并 + 前端派生。老记录/旧客户端为空 → 前端回退老 comparison_results 渲染。
platforms: Mapped[list] = mapped_column(_JSON, nullable=False, default=list)
# 目标平台未找到、跳过的菜名
skipped_dish_names: Mapped[list] = mapped_column(_JSON, nullable=False, default=list)
# 客户端上报的原始 payload(calibration + done.params 全量),未来取数兜底
+8 -2
View File
@@ -1,7 +1,7 @@
"""新手引导视频播放记录(领券浮层前 N 次用它替代广告)。
产品规则(2026-07 拍板):新用户点「一键自动领取」后的等候浮层,**前 3 次**不放广告,
改放运营后台上传的引导视频;每次固定 120 金币,中途关闭也算看完照发。
改放运营后台上传的引导视频;默认每次固定 100 金币,中途关闭也算看完照发。
口径:
- **计次按账号**(user_id),与设备无关 —— 换设备不重新送 3 次。
@@ -38,7 +38,13 @@ class GuideVideoPlay(Base):
# start_play 捕获 IntegrityError 降级成"这次不放视频"。
# 用 unique Index 而非 UniqueConstraint:与迁移里的 create_index 对齐(SQLite 加约束
# 要整表重建),autogenerate 才不会每次报一条假 diff。
Index("uq_guide_video_play_user_seq", "user_id", "seq", unique=True),
Index(
"uq_guide_video_play_user_scene_seq",
"user_id",
"scene",
"seq",
unique=True,
),
)
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
+77
View File
@@ -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,
)
+36 -5
View File
@@ -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,
+82 -17
View File
@@ -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,
+202 -13
View File
@@ -52,6 +52,81 @@ def _product_names_from_items(items: list | None) -> str | None:
return joined[:500] or None
# ---- 失败记录的展示文案(记录页失败卡「原因」行)------------------------------
# information 具体就直出;笼统(_GENERIC_INFO)则从 platform_results 捞一条用户可读的业务
# 原因;捞不到 → None(端侧显示品牌兜底「网络开小差…」)。pricebot 把 store_closed /
# no_delivery 漏成了 status=failed,这里按 reason 关键字补判;打烊类 reason 常带一坨脏店名
# (店名+月售+起送+配送…),统一成简短模板。自动化黑话(搜索失败/读价失败/购物车残留/裸
# FAILED…)不给用户看 → 归入品牌兜底。
# pricebot 组不出具体原因时的笼统 information(线上统计的大头),一律走品牌兜底。
_GENERIC_INFO = {
"比价过程出错,请稍后重试",
"比价出错",
"比价未完成",
"done 参数缺少可验证的目标平台结果",
}
# 干净业务结局 status(直接可信),按展示优先级(越靠前越先选)。
_BIZ_STATUS_PRIORITY = (
"below_minimum",
"no_delivery",
"store_closed",
"items_not_found",
"store_not_found",
)
def _store_closed_text(reason: str | None) -> str:
"""打烊/暂停营业/休息类 reason 常带脏店名元数据 → 只留结论,套简短模板。"""
r = reason or ""
if "暂停营业" in r:
state = "暂停营业"
elif "休息" in r:
state = "休息中"
else:
state = "已打烊"
return f"门店{state},无法比价"
def _target_display_reason(platform_results: dict | None) -> str | None:
"""从逐平台结果里挑一条"可展示给用户"的失败原因;挑不到返回 None。
status 命中干净业务结局集 直接采信(打烊套模板,其余用 reason);
补判 pricebot 漏成 status=failed 的两类:打烊(套模板)单点不配送(reason 本身干净);
自动化黑话(搜索失败/读价失败/购物车残留/ FAILED)一律不展示 None"""
pr = platform_results or {}
targets = [
v for v in pr.values() if isinstance(v, dict) and not v.get("is_source")
]
for want in _BIZ_STATUS_PRIORITY: # ① 干净 status 优先
for v in targets:
if v.get("status") == want:
if want == "store_closed":
return _store_closed_text(v.get("reason"))
if v.get("reason"):
return v["reason"]
for v in targets: # ② 漏成 failed 的业务结局补判
if v.get("status") != "failed":
continue
reason = (v.get("reason") or "").strip()
if any(k in reason for k in ("打烊", "暂停营业", "休息")):
return _store_closed_text(reason)
if "单点不配送" in reason:
return reason
return None
def _derive_fail_display(
information: str | None, platform_results: dict | None
) -> str | None:
"""失败记录展示文案:information 具体则直出;笼统则从 platform_results 捞/补判;
都拿不到 None(端侧品牌兜底)仅在 status=failed 时调用"""
info = (information or "").strip()
text = info if (info and info not in _GENERIC_INFO) else _target_display_reason(
platform_results
)
return text[:256] if text else None
def _derive(payload: ComparisonRecordIn) -> dict:
"""从上报 payload 派生结构化列(best/saved/is_source_best/status)。"""
results = payload.comparison_results
@@ -89,8 +164,9 @@ def _derive(payload: ComparisonRecordIn) -> dict:
is_source_best = best.is_source if best is not None else None
# status:客户端显式给了就用;否则有"非源且有价"的结果=success,否则 failed
status = payload.status
# status:优先 pricebot record_status(区分 below_minimum/store_closed) → 客户端显式 status
# → 兜底"非源且有价"=success/否则 failed。record_status 让"未满起送"不再塌缩成 failed。
status = payload.record_status or payload.status
if status is None:
has_valid_target = any(
(not r.is_source) and r.price is not None for r in results
@@ -105,6 +181,11 @@ def _derive(payload: ComparisonRecordIn) -> dict:
"saved_amount_cents": saved_amount_cents,
"is_source_best": is_source_best,
"status": status,
"fail_reason": (
_derive_fail_display(payload.information, _pr)
if status == "failed"
else None
),
}
@@ -117,16 +198,35 @@ def upsert_record(
灰度期老客户端 POST /compare/record 走这条,与后端 harvest trace_id reconcile;
新客户端不再 POST(改由 compare.py 透传壳 harvest 落库)
"""
derived = _derive(payload)
# 单源派生: 与 harvest_done 一致, payload 带 platforms 时从它派生(唯一真相源
# _derive_from_platforms), 老客户端不带 platforms 时回退 _derive(从 comparison_results)。
if payload.platforms:
derived = _derive_from_platforms(payload.platforms, payload.record_status)
# 对齐 _derive 返回键(#189 fail_reason): 两路径 fields 键集一致, 覆盖已有行时不残留旧值
derived["fail_reason"] = (
_derive_fail_display(payload.information, payload.platform_results or {})
if derived["status"] == "failed"
else None
)
# 单源派生取自 platforms 源行(常无源平台元数据/店名)→ 空则用 payload 兜底不丢字段。
# 下面 fields 不再显式写这四个键, 统一由 derived 提供(否则 dict(store_name=..., **derived)
# 与 _derive_from_platforms 同名键撞键 TypeError)。
for _k in ("store_name", "source_platform_id", "source_platform_name", "source_package"):
if not derived.get(_k):
derived[_k] = getattr(payload, _k)
else:
derived = _derive(payload)
# _derive 只从 comparison_results 派生, 不含源平台四件套 / store_name → 从 payload 补,
# 与上面 platforms 分支键集对齐(fields 统一靠 **derived 提供这些列)。
for _k in ("store_name", "source_platform_id", "source_platform_name", "source_package"):
derived[_k] = getattr(payload, _k)
items = [it.model_dump(exclude_none=True) for it in payload.items]
fields = dict(
device_id=payload.device_id,
business_type=payload.business_type,
store_name=payload.store_name,
product_names=_product_names_from_items(items),
source_platform_id=payload.source_platform_id,
source_platform_name=payload.source_platform_name,
source_package=payload.source_package,
# store_name / source_platform_id / source_platform_name / source_package 统一由
# derived 提供(见上方两分支补齐), 不在此显式写 —— 否则与 _derive_from_platforms 撞键。
information=payload.information,
best_deeplink=payload.best_deeplink,
trace_url=payload.trace_url,
@@ -134,6 +234,7 @@ def upsert_record(
skipped_dish_count=payload.skipped_dish_count,
items=items,
comparison_results=[r.model_dump() for r in payload.comparison_results],
platforms=list(payload.platforms or []),
skipped_dish_names=list(payload.skipped_dish_names),
# 客户端环境 / 性能(debug,客户端上报;旧客户端为 None)
device_model=payload.device_model,
@@ -206,7 +307,8 @@ def upsert_record(
def _derive_from_results(
results: list[dict], platform_results: dict | None = None
results: list[dict], platform_results: dict | None = None,
record_status: str | None = None,
) -> dict:
"""从 done 帧 comparison_results(pricebot 原始 dict 列表)派生结构化列。
等价 _derive,但吃原始字段(is_source/price/rank/platform_id/store_name...)而非 pydantic 对象
@@ -254,7 +356,53 @@ def _derive_from_results(
"saved_amount_cents": saved_amount_cents,
"is_source_best": best.get("is_source") if best else None,
"store_name": (src_row or {}).get("store_name") or None,
"status": "success" if has_valid_target else "failed",
# 记录级结局: 优先用 pricebot 下发的 record_status(区分 below_minimum/store_closed,
# 不再把"未满起送"塌缩成 failed → 记录页不再误报"网络开小差"); 旧 pricebot 未下发时
# 回退老的 success/failed 二态派生, 向后兼容。
"status": record_status or ("success" if has_valid_target else "failed"),
}
def _derive_from_platforms(
platforms: list, record_status: str | None = None,
) -> dict:
"""从 done 帧 platforms(每平台一行、渲染就绪)派生结构化列——**单一真相源**。
best_* 直接取 platforms is_best 的那一行source_* role=source ,与前端读的
platforms 天然一致(不再像 _derive_from_results 那样从 comparison_results 二次评最优,
消除"标量列 vs platforms"双源不一致)platforms 非空时优先走这里; pricebot
platforms 时调用方回退 _derive_from_results(向后兼容)"""
rows = [p for p in (platforms or []) if isinstance(p, dict)]
src = next((p for p in rows if p.get("role") == "source"), None)
best = next((p for p in rows if p.get("is_best")), None)
source_price_cents = _yuan_to_cents(src.get("price")) if src else None
best_price_cents = _yuan_to_cents(best.get("price")) if best else None
saved_amount_cents = None
if source_price_cents is not None and best_price_cents is not None:
saved_amount_cents = source_price_cents - best_price_cents
has_valid_target = any(
p.get("role") != "source" and p.get("price") is not None for p in rows
)
# store_name: 优先源行; recompare 场景源平台自己当目标、源行被目标覆盖(pricebot
# _build_platform_rows 有意去重, platforms 无 role=source 行)→ 回退 best 行 → 首个有店名
# 的行(显示现场实际比到的店), 免得记录页店名空掉兜底显示成"比价"。正常比价有源行不走回退。
store_name = (
(src or {}).get("store_name")
or (best or {}).get("store_name")
or next((p.get("store_name") for p in rows if p.get("store_name")), None)
)
return {
"source_platform_id": (src or {}).get("platform_id"),
"source_platform_name": (src or {}).get("platform_name"),
"source_package": (src or {}).get("package"),
"source_price_cents": source_price_cents,
"best_platform_id": (best or {}).get("platform_id"),
"best_platform_name": (best or {}).get("platform_name"),
"best_price_cents": best_price_cents,
"saved_amount_cents": saved_amount_cents,
"is_source_best": (best.get("role") == "source") if best else None,
"store_name": store_name or None,
"status": record_status or ("success" if has_valid_target else "failed"),
}
@@ -286,6 +434,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.
@@ -311,6 +461,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(
@@ -325,6 +480,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(
@@ -333,7 +493,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(
@@ -413,12 +573,40 @@ def harvest_done(
返回 (记录, 是否本次****落成 success)供调用方据此幂等发一次邀请奖
行不存在(理论上帧0已建;防御)则新建"""
results = done_params.get("comparison_results") or []
derived = _derive_from_results(results, done_params.get("platform_results"))
# 菜品:pricebot 已把源单菜品塞进 comparison_results[源行].items
items = next((r.get("items") or [] for r in results if r.get("is_source")), [])
# 展示模型统一数组(pricebot 新增, 每平台一行自带 status/is_best): 原样存, 记录页据此直渲染。
# record_status: 记录级结局(success/below_minimum/store_closed/failed), 覆盖老二态派生。
platforms = done_params.get("platforms") or []
record_status = done_params.get("record_status")
# 单源派生: platforms(含 pricebot 权威 is_best)是唯一真相源, best_*/source_*/saved/status
# 全从它取 → 与前端读的 platforms 天然一致; 菜品也取 platforms 源行。老 pricebot 无
# platforms 时回退从 comparison_results 派生(向后兼容)。
if platforms:
derived = _derive_from_platforms(platforms, record_status)
# 菜品优先源行; recompare 无源行 → 回退 best 行 → 首个有菜品的行(同 store_name 回退)
_item_row = (
next((p for p in platforms if isinstance(p, dict) and p.get("role") == "source"), None)
or next((p for p in platforms if isinstance(p, dict) and p.get("is_best")), None)
or next((p for p in platforms if isinstance(p, dict) and p.get("items")), None)
)
items = (_item_row or {}).get("items") or []
else:
derived = _derive_from_results(
results, done_params.get("platform_results"), record_status
)
# pricebot 已把源单菜品塞进 comparison_results[源行].items
items = next((r.get("items") or [] for r in results if r.get("is_source")), [])
# 失败展示原因(#189): platforms / results 两个派生分支的 status 都可能 failed, 统一在此算
fail_reason = (
_derive_fail_display(
done_params.get("information"), done_params.get("platform_results")
)
if derived["status"] == "failed"
else None
)
fields = dict(
business_type=business_type or "food",
information=done_params.get("information") or None,
fail_reason=fail_reason,
# best_deeplink 来自客户端剪贴板采集,harvest 拿不到 → 留空(灰度期 fromComparison 会补;
# 纯 harvest 行「再次比价」退化为按 package 拉起 App。要精确深链需客户端另传,后续)。
trace_url=trace_url or done_params.get("trace_url"),
@@ -426,6 +614,7 @@ def harvest_done(
skipped_dish_count=done_params.get("skipped_dish_count"),
skipped_dish_names=list(done_params.get("skipped_dish_names") or []),
comparison_results=results,
platforms=platforms,
items=items,
product_names=_product_names_from_items(items),
raw_payload=done_params,
+133 -37
View File
@@ -1,8 +1,8 @@
"""新手引导视频:运营配置读写 + 播放计次 + 发币。
**配置**(开关 / 视频地址 / 前几次 / 每次金币)整体作为一个 JSON 存进通用 app_config
(key=coupon_guide_video),写法完全对齐 feedback_qr 不进 CONFIG_DEFS,所以不会污染
系统配置页的通用列表,由本模块独占维护
**配置**(开关 / 视频地址 / 前几次 / 每次金币)按场景分别作为一个 JSON 存进通用
app_config 领券使用 coupon_guide_video比价使用 comparison_guide_video
写法完全对齐 feedback_qr不进入系统配置页的通用列表
**计次**按账号(user_id)**开播即计数**:客户端每次要展示领券等候浮层时调
`/api/v1/guide-video/start`,命中则当场写一行 guide_video_play(status='playing')
@@ -25,12 +25,16 @@ from sqlalchemy import func, select, update
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from app.core import rewards
from app.core import limit_policy, rewards
from app.models.app_config import AppConfig
from app.models.guide_video import GuideVideoPlay
from app.repositories import wallet as crud_wallet
_KEY = "coupon_guide_video"
SCENES = ("coupon", "comparison")
_KEY_BY_SCENE = {
"coupon": "coupon_guide_video",
"comparison": "comparison_guide_video",
}
#: 金币流水 biz_type。客户端收益明细按它显示「新手引导视频奖励」。
BIZ_TYPE = "guide_video"
@@ -41,7 +45,7 @@ _DEFAULTS: dict[str, Any] = {
"enabled": True,
"video_url": None, # None/空 = 未配片 → 不下发,浮层照旧放广告
"max_plays": 3, # 每个账号前 N 次浮层放引导视频
"reward_coin": 120, # 每次固定金币
"reward_coin": 100, # 每次固定金币
}
_FIELDS = tuple(_DEFAULTS.keys())
@@ -65,19 +69,38 @@ def _merge(raw: Any) -> dict[str, Any]:
return out
def get_config(db: Session) -> dict[str, Any]:
def _config_key(scene: str) -> str:
if scene not in _KEY_BY_SCENE:
raise ValueError(f"unsupported guide video scene: {scene}")
return _KEY_BY_SCENE[scene]
def get_config(db: Session, scene: str = "coupon") -> dict[str, Any]:
"""完整配置 + updated_at(admin 读 / 业务读共用)。"""
row = db.get(AppConfig, _KEY)
row = db.get(AppConfig, _config_key(scene))
cfg = _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
def _write(db: Session, value: dict[str, Any], *, admin_id: int, commit: bool) -> dict[str, Any]:
def _write(
db: Session,
value: dict[str, Any],
*,
scene: str,
admin_id: int,
commit: bool,
) -> dict[str, Any]:
"""整体覆写该行(value 须为完整字段 dict),返回合并后的完整配置(含 updated_at)。"""
row = db.get(AppConfig, _KEY)
key = _config_key(scene)
row = db.get(AppConfig, key)
if row is None:
row = AppConfig(key=_KEY, value=value, updated_by_admin_id=admin_id)
row = AppConfig(key=key, value=value, updated_by_admin_id=admin_id)
db.add(row)
else:
row.value = value # 整体重新赋值,SQLAlchemy 才侦测得到变更
@@ -98,11 +121,12 @@ def update_config(
enabled: bool | None = None,
max_plays: int | None = None,
reward_coin: int | None = None,
scene: str = "coupon",
admin_id: int,
commit: bool = True,
) -> tuple[dict[str, Any], dict[str, Any]]:
"""改开关 / 次数 / 金币(只改传了的字段;视频走 set_video)。返回 (before, after) 供审计。"""
row = db.get(AppConfig, _KEY)
row = db.get(AppConfig, _config_key(scene))
before = _merge(row.value if row is not None else None)
new_value = {k: before[k] for k in _FIELDS}
if enabled is not None:
@@ -111,45 +135,91 @@ def update_config(
new_value["max_plays"] = max(0, min(int(max_plays), MAX_PLAYS_LIMIT))
if reward_coin is not None:
new_value["reward_coin"] = max(0, min(int(reward_coin), REWARD_COIN_LIMIT))
after = _write(db, new_value, admin_id=admin_id, commit=commit)
after = _write(
db,
new_value,
scene=scene,
admin_id=admin_id,
commit=False,
)
if scene == "coupon" and max_plays is not None:
# 兼容仍调用旧专用接口的客户端/脚本,并把统一策略全局值一并更新。
from app.core.config_schema import GUIDE_VIDEO_MAX_PLAYS_KEY
from app.repositories import app_config
app_config.set_value(
db,
GUIDE_VIDEO_MAX_PLAYS_KEY,
new_value["max_plays"],
admin_id=admin_id,
commit=False,
)
if commit:
db.commit()
after = get_config(db, scene)
return before, after
def set_video(
db: Session, video_url: str | None, *, admin_id: int, commit: bool = True
db: Session,
video_url: str | None,
*,
scene: str = "coupon",
admin_id: int,
commit: bool = True,
) -> tuple[dict[str, Any], dict[str, Any]]:
"""设置/清空引导视频地址。返回 (before, after);before['video_url'] 供调用方删旧文件。"""
row = db.get(AppConfig, _KEY)
row = db.get(AppConfig, _config_key(scene))
before = _merge(row.value if row is not None else None)
new_value = {k: before[k] for k in _FIELDS}
new_value["video_url"] = video_url
after = _write(db, new_value, admin_id=admin_id, commit=commit)
after = _write(
db,
new_value,
scene=scene,
admin_id=admin_id,
commit=commit,
)
return before, after
# ===== 播放计次 =====
def used_plays(db: Session, user_id: int) -> int:
def used_plays(
db: Session,
user_id: int,
*,
scene: str = "coupon",
reset_at: datetime | None = None,
) -> int:
"""该账号已用掉的引导视频次数(开播即算,含未发币的)。"""
return int(
db.execute(
select(func.count()).select_from(GuideVideoPlay).where(
GuideVideoPlay.user_id == user_id
)
).scalar_one()
stmt = select(func.count()).select_from(GuideVideoPlay).where(
GuideVideoPlay.user_id == user_id,
GuideVideoPlay.scene == scene,
)
if reset_at is not None:
# 该表的既有写入口使用北京时间 naive 墙钟,重置基线必须转换成
# 同一存储口径再比较;单独改成 UTC 会让存量/增量记录偏移 8 小时。
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 play_stats(db: Session) -> dict[str, int]:
def play_stats(db: Session, scene: str = "coupon") -> dict[str, int]:
"""全站播放统计(admin 页展示):总播放次数 / 其中已发币次数。"""
total = int(
db.execute(select(func.count()).select_from(GuideVideoPlay)).scalar_one()
db.execute(
select(func.count()).select_from(GuideVideoPlay).where(
GuideVideoPlay.scene == scene
)
).scalar_one()
)
granted = int(
db.execute(
select(func.count()).select_from(GuideVideoPlay).where(
GuideVideoPlay.status == "granted"
GuideVideoPlay.status == "granted",
GuideVideoPlay.scene == scene,
)
).scalar_one()
)
@@ -168,11 +238,22 @@ def start_play(
reward_coin 播完/中途关闭都发的固定金币
seq / remaining 第几次 / 发完这次还剩几次(仅展示与排查用)
"""
cfg = get_config(db)
_config_key(scene)
cfg = get_config(db, scene)
video_url = (cfg.get("video_url") or "").strip()
max_plays = int(cfg.get("max_plays") or 0)
policy = (
limit_policy.resolve_for_user(db, "guide.video.lifetime", user_id)
if scene == "coupon"
else None
)
max_plays = policy.limit if policy is not None else int(cfg.get("max_plays") or 0)
exposed_max_plays = max_plays if max_plays is not None else 1_000_000
reward_coin = int(cfg.get("reward_coin") or 0)
used = used_plays(db, user_id)
used = (
used_plays(db, user_id, scene=scene, reset_at=policy.reset_at)
if policy is not None and policy.reset_at is not None
else used_plays(db, user_id, scene=scene)
)
def _miss(used_now: int) -> dict[str, Any]:
return {
@@ -181,20 +262,29 @@ def start_play(
"play_token": "",
"reward_coin": reward_coin,
"seq": used_now,
"remaining": max(0, max_plays - used_now),
"remaining": max(0, exposed_max_plays - used_now),
}
if not cfg.get("enabled") or not video_url or max_plays <= 0 or used >= max_plays:
if (
not cfg.get("enabled")
or not video_url
or (max_plays is not None and (max_plays <= 0 or used >= max_plays))
):
return _miss(used)
seq = used + 1
# ``seq`` remains scene 内 lifetime-monotonic because
# (user_id, scene, seq) is unique.
# A whitelist reset only changes the quota baseline; reusing seq=1 would
# collide with historical rows and make every post-reset start fail.
seq = used_plays(db, user_id, scene=scene) + 1
play = GuideVideoPlay(
user_id=user_id,
play_token=uuid.uuid4().hex,
scene=scene,
seq=seq,
video_url=video_url,
coin=0,
# 固化本次承诺发放的金币,避免运营改价后已开播记录按新价结算。
coin=reward_coin,
status="playing",
completed=0,
started_at=datetime.now(rewards.CN_TZ).replace(tzinfo=None),
@@ -210,14 +300,19 @@ def start_play(
db.flush()
except IntegrityError:
db.rollback()
return _miss(used_plays(db, user_id))
used_now = (
used_plays(db, user_id, scene=scene, reset_at=policy.reset_at)
if policy is not None and policy.reset_at is not None
else used_plays(db, user_id, scene=scene)
)
return _miss(used_now)
return {
"should_play": True,
"video_url": video_url,
"play_token": play.play_token,
"reward_coin": reward_coin,
"seq": seq,
"remaining": max(0, max_plays - seq),
"remaining": max(0, exposed_max_plays - (used + 1)),
}
@@ -240,8 +335,9 @@ def grant_play(
重复上报返回 granted=False + 已发金币(客户端据此不重复累加 toast 金额)
"""
token = (play_token or "").strip()
# 金币额度以**服务端配置**为准,不信客户端(客户端只上报"播完/关闭")
coin = int(get_config(db).get("reward_coin") or 0)
# 金币额度取开播时由服务端固化的值,不信客户端,也不受后续配置变更影响
play = _find_play(db, user_id, token)
coin = int(play.coin if play is not None else 0)
# 幂等核心:把 status 放进 WHERE 做条件更新(compare-and-set),而不是"先读再判再写"。
# 「播完」与「✕ 关闭」抢跑、或客户端超时重试时,两个请求会都读到 status='playing',
+29 -7
View File
@@ -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))
+85 -23
View File
@@ -9,6 +9,7 @@ from sqlalchemy import func, select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from app.core import limit_policy
from app.core.config_schema import (
RISK_COMPARE_DAILY_THRESHOLD_KEY,
RISK_ONECLICK_DAILY_THRESHOLD_KEY,
@@ -257,22 +258,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 +313,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 +323,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 +411,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 +446,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 +527,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 +538,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(
+14 -2
View File
@@ -107,6 +107,13 @@ class ComparisonRecordIn(BaseModel):
# 明细
items: list[ComparisonItemIn] = Field(default_factory=list)
comparison_results: list[ComparisonResultIn] = Field(default_factory=list)
# 展示模型统一数组(pricebot done.params.platforms 原样透传): 每平台一行、自带
# status/is_best/display/display_order,记录页据此直渲染。宽松 list[dict] 存(结构由
# pricebot 定,server 只原样落库),前端读它、老记录空时回退 comparison_results。
platforms: list[dict] = Field(default_factory=list)
# 记录级结局(pricebot 下发): success/below_minimum/store_closed/failed。让"未满起送"不再
# 被塌缩成 failed。_derive 优先用它、其次客户端 status、再兜底二态派生。
record_status: str | None = None
# 逐平台结局摘要(含失败平台的细分原因 status: store_not_found/items_not_found/below_minimum/
# unsupported/...)。来自 done.params.platform_results,客户端透传;落 raw_payload(不单列),
# admin「卡在哪一步」从这里读。dict{platform_id: {...}} 宽松存(结构由 pricebot 定——是
@@ -173,8 +180,13 @@ class ComparisonRecordOut(BaseModel):
skipped_dish_count: int | None = None
status: str
information: str | None = None
# 失败卡「原因」文案:具体失败给具体原因,纯系统失败为 None(端侧品牌兜底)。见模型 fail_reason。
fail_reason: str | None = None
items: list = []
comparison_results: list = []
# 展示模型统一数组(每平台一行、自带 status/is_best/display/display_order): 记录页据此
# 直渲染, 不再靠 comparison_results + 前端派生。老记录为空 → 前端回退 comparison_results。
platforms: list = []
skipped_dish_names: list = []
total_ms: int | None = None
# 「已下单」(店级):该店名在该用户真实下单(source='compare')里出现过即 True。
@@ -210,9 +222,9 @@ class CompareStartReserveIn(BaseModel):
class CompareStartReserveOut(BaseModel):
limit: int
limit: int | None
used: int
remaining: int
remaining: int | None
class CompareStatsOut(BaseModel):
+4 -2
View File
@@ -1,13 +1,15 @@
"""新手引导视频(领券等候浮层前 N 次替代广告)的客户端请求/响应契约。"""
from __future__ import annotations
from typing import Literal
from pydantic import BaseModel, Field
class GuideVideoStartIn(BaseModel):
"""开播询问。scene 目前只有 coupon(领券浮层);预留给日后比价等场景"""
"""开播询问。领券与比价分别使用独立配置和独立次数"""
scene: str = Field(default="coupon", max_length=16)
scene: Literal["coupon", "comparison"] = "coupon"
class GuideVideoStartOut(BaseModel):
+1 -1
View File
@@ -135,7 +135,7 @@ def repair_missing_comparison_llm_costs(
select(ComparisonRecord.id, ComparisonRecord.trace_id)
.where(
*date_conditions,
ComparisonRecord.status.in_(("success", "failed")),
ComparisonRecord.status.in_(("success", "failed", "cancelled")),
ComparisonRecord.llm_cost_yuan.is_(None),
)
.order_by(ComparisonRecord.created_at.desc(), ComparisonRecord.id.desc())
-22
View File
@@ -1,22 +0,0 @@
# 本地开发/测试用 PostgreSQL。生产用原生 PG(scripts/init_postgres.py),不使用本文件。
services:
postgres:
image: postgres:16-alpine
container_name: shaguabijia-pg
environment:
POSTGRES_USER: shaguabijia_app
POSTGRES_PASSWORD: shaguabijia_dev_pw
POSTGRES_DB: shaguabijia
ports:
- "5432:5432"
volumes:
- pgdata:/var/lib/postgresql/data
- ./docker/initdb:/docker-entrypoint-initdb.d:ro
healthcheck:
test: ["CMD-SHELL", "pg_isready -U shaguabijia_app -d shaguabijia"]
interval: 3s
timeout: 3s
retries: 20
volumes:
pgdata:
-3
View File
@@ -1,3 +0,0 @@
-- 仅在 pgdata 卷首次初始化时执行一次(以 shaguabijia_app 连 shaguabijia 库运行)。
-- 幂等兜底见 scripts/ensure_pg.py 的 _ensure_test_db()。
CREATE DATABASE shaguabijia_test OWNER shaguabijia_app;
+1 -15
View File
@@ -27,21 +27,7 @@ PG 默认上 16 版(工具链最齐),驱动用 **psycopg3**(SQLAlchemy 2.0 时
## 1. 本地起 PG + 跑通空库(半天)
### 1.0 推荐:Docker 一键起(本地开发/测试)
本地开发不必手动装 PG。已提供 `docker-compose.yml` + `scripts/ensure_pg.py`:
```bash
cp .env.example .env # DATABASE_URL 默认已是 Docker PG 连接串
./run.sh # 或 run.bat;会自动:探测 PG → 没起则启 Docker → 起 PG 容器 → 建库 → alembic → uvicorn
pytest # conftest 自动引导同一容器的 shaguabijia_test 库
```
容器:`postgres:16-alpine`(名 `shaguabijia-pg`,端口 5432,命名卷 `pgdata` 持久化),
首启即建业务库 `shaguabijia` 与测试库 `shaguabijia_test`。下面 1.1-1.5 的手动装 PG 步骤仅在
不用 Docker 时才需要;生产仍走 §4 的原生 PG。
### 1.1 装 PG(不用 Docker 时的手动方式)
### 1.1 装 PG
macOS:
```bash
@@ -1,740 +0,0 @@
# 本地开发切 Docker PostgreSQL 实现计划
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
**Goal:** 让 app-server 本地开发运行与 pytest 都跑在 Docker 化的 PostgreSQL 16 上,退掉 SQLite,`run.bat`/`run.sh` 启动时自动检测并拉起 PG(必要时先启 Docker Desktop、缺镜像先拉)。
**Architecture:** 新增 `docker-compose.yml`(声明 PG 服务/卷/健康检查)+ `scripts/ensure_pg.py`(跨平台引导:探测→启 Docker→compose up→等就绪→幂等建测试库),由 `run.sh`/`run.bat`/`tests/conftest.py` 三处共用。`.env.example` 默认切 PG。`app/db/session.py``alembic/env.py` 已天然支持 PG,无需改。
**Tech Stack:** Docker Compose、`postgres:16-alpine`、Python 3.10+ 标准库(`socket`/`subprocess`/`urllib.parse`)、psycopg3(已装)、SQLAlchemy 2.0 + Alembic、pytest。
**工作目录:** 本计划在 worktree `.worktrees/local-dev-postgres-docker`(分支 `chore/local-dev-postgres-docker`,基于 `main`)内执行。下面所有路径相对该 worktree 根(= 仓库根)。
**设计依据:** [docs/superpowers/specs/2026-07-08-local-dev-postgres-docker-design.md](../specs/2026-07-08-local-dev-postgres-docker-design.md)
**连接参数(全程固定值):**
- 镜像 `postgres:16-alpine`,容器名 `shaguabijia-pg`,宿主端口 `5432`
- 用户 `shaguabijia_app`,dev 密码 `shaguabijia_dev_pw`(本地非机密)
- 业务库 `shaguabijia`,测试库 `shaguabijia_test`,命名卷 `pgdata`
- dev URL:`postgresql+psycopg://shaguabijia_app:shaguabijia_dev_pw@localhost:5432/shaguabijia`
- test URL:`...@localhost:5432/shaguabijia_test`
**前置:** 执行机已安装 Docker Desktop。
---
## Task 1: Docker Compose + 测试库 initdb 脚本
**Files:**
- Create: `docker-compose.yml`
- Create: `docker/initdb/01-create-test-db.sql`
- [ ] **Step 1: 写 `docker-compose.yml`**
```yaml
# 本地开发/测试用 PostgreSQL。生产用原生 PG(scripts/init_postgres.py),不使用本文件。
services:
postgres:
image: postgres:16-alpine
container_name: shaguabijia-pg
environment:
POSTGRES_USER: shaguabijia_app
POSTGRES_PASSWORD: shaguabijia_dev_pw
POSTGRES_DB: shaguabijia
ports:
- "5432:5432"
volumes:
- pgdata:/var/lib/postgresql/data
- ./docker/initdb:/docker-entrypoint-initdb.d:ro
healthcheck:
test: ["CMD-SHELL", "pg_isready -U shaguabijia_app -d shaguabijia"]
interval: 3s
timeout: 3s
retries: 20
volumes:
pgdata:
```
- [ ] **Step 2: 写 `docker/initdb/01-create-test-db.sql`**
```sql
-- 仅在 pgdata 卷首次初始化时执行一次(以 shaguabijia_app 连 shaguabijia 库运行)。
-- 幂等兜底见 scripts/ensure_pg.py 的 _ensure_test_db()。
CREATE DATABASE shaguabijia_test OWNER shaguabijia_app;
```
- [ ] **Step 3: 起容器验证**
Run: `docker compose up -d`
Expected: 拉取 `postgres:16-alpine`(首次)后 `Container shaguabijia-pg Started`
- [ ] **Step 4: 验证两个库都在 + 健康**
Run: `docker compose exec -T postgres psql -U shaguabijia_app -d shaguabijia -tAc "SELECT datname FROM pg_database WHERE datname IN ('shaguabijia','shaguabijia_test') ORDER BY 1"`
Expected 输出:
```
shaguabijia
shaguabijia_test
```
- [ ] **Step 5: 验证 `.worktrees/``data/` 忽略不受影响、compose 无落盘到项目目录**
Run: `git status --short`
Expected: 只列出本任务新增的 `docker-compose.yml``docker/initdb/01-create-test-db.sql`(数据在命名卷 `pgdata`,不在项目目录;`.worktrees/` 已忽略)。
- [ ] **Step 6: Commit**
```bash
git add docker-compose.yml docker/initdb/01-create-test-db.sql
git commit -m "feat(dev): docker-compose 起本地 PostgreSQL(含测试库 initdb)"
```
---
## Task 2: `scripts/ensure_pg.py` 引导脚本(TDD)
**Files:**
- Create: `scripts/ensure_pg.py`
- Test: `tests/test_ensure_pg.py`
> 说明:纯函数(URL 解析 / sqlite 判定 / 端口探测 / 平台命令映射 / sqlite 守卫 / 端口通时短路)走 TDD 单测;真正拉 Docker 的编排 `ensure()` 全链路靠 Task 3/5 的运行来验证(需真 Docker,不做单测)。此时 `conftest.py` 仍是 SQLite,不依赖 PG,单测可独立跑。
- [ ] **Step 1: 写失败测试 `tests/test_ensure_pg.py`**
```python
"""scripts/ensure_pg.py 纯函数单测(不需要 Docker/PG)。"""
from __future__ import annotations
import socket
from scripts.ensure_pg import (
_docker_desktop_cmd,
_is_sqlite,
_parse_host_port,
_port_open,
ensure,
)
def test_is_sqlite():
assert _is_sqlite("sqlite:///./data/app.db")
assert _is_sqlite(" SQLite:///x ")
assert not _is_sqlite("postgresql+psycopg://u:p@localhost:5432/db")
def test_parse_host_port_full():
assert _parse_host_port(
"postgresql+psycopg://u:p@localhost:5432/shaguabijia"
) == ("localhost", 5432)
def test_parse_host_port_defaults():
# 缺端口 → 5432
assert _parse_host_port("postgresql+psycopg://u:p@db.example/x")[1] == 5432
# 缺 host → localhost
assert _parse_host_port("postgresql+psycopg:///x") == ("localhost", 5432)
def test_parse_host_port_testdb():
assert _parse_host_port(
"postgresql+psycopg://u:p@localhost:5432/shaguabijia_test"
) == ("localhost", 5432)
def test_port_open_true():
srv = socket.socket()
srv.bind(("127.0.0.1", 0))
srv.listen(1)
port = srv.getsockname()[1]
try:
assert _port_open("127.0.0.1", port, timeout=1.0)
finally:
srv.close()
def test_port_open_false():
s = socket.socket()
s.bind(("127.0.0.1", 0))
port = s.getsockname()[1]
s.close() # 释放端口,无人监听 → 连接应失败
assert not _port_open("127.0.0.1", port, timeout=0.3)
def test_docker_desktop_cmd_windows():
cmd = _docker_desktop_cmd("win32", r"C:\Program Files")
assert cmd is not None
assert cmd[0].endswith("Docker Desktop.exe")
assert "Docker" in cmd[0]
def test_docker_desktop_cmd_darwin():
assert _docker_desktop_cmd("darwin", "") == ["open", "-a", "Docker"]
def test_docker_desktop_cmd_linux():
assert _docker_desktop_cmd("linux", "") is None
def test_ensure_rejects_sqlite():
# dev 守卫:sqlite 直接 False(不碰 Docker)
assert ensure("sqlite:///./data/app.db") is False
def test_ensure_shortcircuits_when_pg_up(monkeypatch):
# 端口通 → 直接 True,绝不触碰 docker
monkeypatch.setattr("scripts.ensure_pg._port_open", lambda *a, **k: True)
def _boom():
raise AssertionError("端口通时不应调用 docker")
monkeypatch.setattr("scripts.ensure_pg._docker_cli_ok", _boom)
assert ensure("postgresql+psycopg://u:p@localhost:5432/shaguabijia") is True
```
- [ ] **Step 2: 跑测试确认失败**
Run: `pytest tests/test_ensure_pg.py -v`
Expected: FAIL —— `ModuleNotFoundError: No module named 'scripts.ensure_pg'`(还没建)。
- [ ] **Step 3: 写实现 `scripts/ensure_pg.py`**
```python
"""确保本地 PostgreSQL 就绪(开发/测试统一用 Docker PG)。
被三处复用:
- run.sh / run.bat:`python -m scripts.ensure_pg`(CLI,失败退非 0)
- tests/conftest.py:`from scripts.ensure_pg import ensure; ensure(test_url)`
流程:读 DATABASE_URL → TCP 探测 → 没起就(必要时启 Docker Desktop)→
`docker compose up -d` → 等 PG ready → 幂等确保测试库存在。全程无 SQLite 兜底。
生产用原生 PG(scripts/init_postgres.py),不走本模块。
"""
from __future__ import annotations
import os
import socket
import subprocess
import sys
import time
from pathlib import Path
from urllib.parse import urlsplit
ROOT = Path(__file__).resolve().parent.parent
# 日志里可能含 emoji(如 ✅);Windows GBK 控制台(cmd.exe)无法编码会抛 UnicodeEncodeError → 脚本崩、
# run.bat 误判 ensure_pg 失败。用 backslashreplace 保底:中文仍正常,仅不可编码字符被转义,不崩。
for _stream in (sys.stdout, sys.stderr):
try:
_stream.reconfigure(errors="backslashreplace")
except (AttributeError, ValueError):
pass
APP_DB = "shaguabijia"
TEST_DB = "shaguabijia_test"
DB_USER = "shaguabijia_app"
COMPOSE_SERVICE = "postgres"
DOCKER_START_TIMEOUT = int(os.environ.get("ENSURE_PG_DOCKER_TIMEOUT", "120"))
PG_READY_TIMEOUT = int(os.environ.get("ENSURE_PG_READY_TIMEOUT", "60"))
POLL_INTERVAL = 3.0
SQLITE_FIX_HINT = (
"postgresql+psycopg://shaguabijia_app:shaguabijia_dev_pw@localhost:5432/shaguabijia"
)
def _log(msg: str) -> None:
print(f"[ensure_pg] {msg}", flush=True)
def _is_sqlite(url: str) -> bool:
return url.strip().lower().startswith("sqlite")
def _parse_host_port(url: str) -> tuple[str, int]:
"""从 SQLAlchemy URL 取 host/port,缺省 localhost:5432。"""
parts = urlsplit(url)
return (parts.hostname or "localhost"), (parts.port or 5432)
def _port_open(host: str, port: int, timeout: float = 1.0) -> bool:
try:
with socket.create_connection((host, port), timeout=timeout):
return True
except OSError:
return False
def _docker_desktop_cmd(platform: str, program_files: str) -> list[str] | None:
"""按平台给出启动 Docker Desktop 的命令;Linux 返回 None(daemon 需 sudo,让用户手动)。"""
if platform.startswith("win"):
return [str(Path(program_files) / "Docker" / "Docker" / "Docker Desktop.exe")]
if platform == "darwin":
return ["open", "-a", "Docker"]
return None
def _docker_ok(subcmd: str) -> bool:
"""`docker version`(CLI 在不在)/`docker info`(daemon 起没起)成功与否。"""
try:
subprocess.run(
["docker", subcmd],
cwd=ROOT,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
check=True,
)
return True
except (OSError, subprocess.CalledProcessError):
return False
def _docker_cli_ok() -> bool:
return _docker_ok("version")
def _docker_daemon_ok() -> bool:
return _docker_ok("info")
def _start_docker_daemon() -> bool:
"""守护进程没起时按平台拉起,轮询到就绪。返回是否成功。"""
if _docker_daemon_ok():
return True
cmd = _docker_desktop_cmd(
sys.platform, os.environ.get("ProgramFiles", r"C:\Program Files")
)
if cmd is None:
_log("Docker 守护进程未运行。Linux 请手动:sudo systemctl start docker,然后重试。")
return False
if sys.platform.startswith("win") and not Path(cmd[0]).exists():
_log(f"找不到 Docker Desktop:{cmd[0]}。请手动启动 Docker Desktop 后重试。")
return False
_log(f"启动 Docker Desktop(首次冷启可能 30-60s)…")
try:
subprocess.Popen(cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
except OSError as e:
_log(f"启动 Docker Desktop 失败:{e}")
return False
deadline = time.monotonic() + DOCKER_START_TIMEOUT
while time.monotonic() < deadline:
if _docker_daemon_ok():
_log("Docker 守护进程已就绪。")
return True
_log("等待 Docker 守护进程…")
time.sleep(POLL_INTERVAL)
_log(f"等待 Docker 守护进程超时({DOCKER_START_TIMEOUT}s)。")
return False
def _compose_up() -> bool:
_log("docker compose up -d(镜像缺失会自动拉取,首用约几十秒)…")
try:
subprocess.run(["docker", "compose", "up", "-d"], cwd=ROOT, check=True)
return True
except (OSError, subprocess.CalledProcessError) as e:
_log(f"docker compose up 失败:{e}")
return False
def _pg_isready() -> bool:
r = subprocess.run(
["docker", "compose", "exec", "-T", COMPOSE_SERVICE,
"pg_isready", "-U", DB_USER, "-d", APP_DB],
cwd=ROOT, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
)
return r.returncode == 0
def _wait_pg_ready(host: str, port: int) -> bool:
deadline = time.monotonic() + PG_READY_TIMEOUT
while time.monotonic() < deadline:
if _port_open(host, port) and _pg_isready():
_log("PostgreSQL 已就绪。")
return True
_log("等待 PostgreSQL 就绪…")
time.sleep(POLL_INTERVAL)
_log(f"等待 PostgreSQL 就绪超时({PG_READY_TIMEOUT}s)。")
return False
def _ensure_test_db() -> None:
"""幂等建测试库(兼容老 pgdata 卷首启没跑 initdb 的情况)。"""
check = subprocess.run(
["docker", "compose", "exec", "-T", COMPOSE_SERVICE,
"psql", "-U", DB_USER, "-d", APP_DB, "-tAc",
f"SELECT 1 FROM pg_database WHERE datname='{TEST_DB}'"],
cwd=ROOT, capture_output=True, text=True,
)
if check.returncode == 0 and check.stdout.strip() == "1":
return
_log(f"建测试库 {TEST_DB}…")
subprocess.run(
["docker", "compose", "exec", "-T", COMPOSE_SERVICE,
"psql", "-U", DB_USER, "-d", APP_DB, "-c",
f"CREATE DATABASE {TEST_DB} OWNER {DB_USER}"],
cwd=ROOT, check=False,
)
def ensure(database_url: str | None = None) -> bool:
"""确保 PG 就绪,返回 True/False。database_url 缺省从 settings 读(尊重 .env)。"""
if database_url is None:
from app.core.config import settings # 延迟导入,避免过早固化 settings
database_url = settings.DATABASE_URL
if _is_sqlite(database_url):
_log("检测到 DATABASE_URL 仍是 SQLite。本地开发/测试已切 PostgreSQL,请改成:")
_log(f" DATABASE_URL={SQLITE_FIX_HINT}")
return False
host, port = _parse_host_port(database_url)
if _port_open(host, port):
_log(f"✅ PostgreSQL 已在 {host}:{port} 运行,跳过 Docker。")
return True
_log(f"{host}:{port} 无 PostgreSQL,准备用 Docker 拉起…")
if not _docker_cli_ok():
_log("未检测到 docker 命令。请先安装 Docker Desktop:")
_log(" https://www.docker.com/products/docker-desktop/")
return False
if not _start_docker_daemon():
return False
if not _compose_up():
return False
if not _wait_pg_ready(host, port):
return False
_ensure_test_db()
return True
if __name__ == "__main__":
sys.exit(0 if ensure() else 1)
```
- [ ] **Step 4: 跑测试确认通过**
Run: `pytest tests/test_ensure_pg.py -v`
Expected: 11 passed。
- [ ] **Step 5: 手动冒烟(PG 已在跑时应秒过短路)**
Run: `python -m scripts.ensure_pg`
Expected: 打印 `[ensure_pg] ✅ PostgreSQL 已在 localhost:5432 运行,跳过 Docker。`,退出码 0。
- [ ] **Step 6: Commit**
```bash
git add scripts/ensure_pg.py tests/test_ensure_pg.py
git commit -m "feat(dev): scripts/ensure_pg.py 探测/拉起本地 Docker PostgreSQL"
```
---
## Task 3: `run.sh` / `run.bat` 接入 ensure_pg
**Files:**
- Modify: `run.sh`(在 `alembic upgrade head` 前插一步)
- Modify: `run.bat`(同上)
- [ ] **Step 1: 改 `run.sh`**
`mkdir -p data` 之后、`"$PY" -m alembic upgrade head` 之前插入:
```bash
"$PY" -m scripts.ensure_pg # 确保本地 Docker PostgreSQL 就绪(没起会自动拉起;失败即退出)
```
(`set -e` 已在文件顶部,ensure_pg 失败会自动终止脚本。)
- [ ] **Step 2: 改 `run.bat`**
`if not exist data mkdir data` 之后、`call "%PY%" -m alembic upgrade head` 之前插入:
```bat
REM 确保本地 Docker PostgreSQL 就绪(没起会自动拉起 Docker + PG 容器)
call "%PY%" -m scripts.ensure_pg
if errorlevel 1 (
echo [X] ensure_pg failed ^(PostgreSQL 未就绪^)
exit /b %errorlevel%
)
```
- [ ] **Step 3: 验证 `run.sh`(PG 已在跑,应短路后继续 alembic + uvicorn)**
Run(Git Bash):`bash run.sh 8770`
Expected: 依次出现 `[ensure_pg] ✅ PostgreSQL 已在 localhost:5432 运行` → alembic 无报错 → uvicorn `Application startup complete``Ctrl-C` 停。
- [ ] **Step 4: 验证 `run.bat`(同上,Windows 原生)**
Run(cmd/PowerShell):`.\run.bat 8770`
Expected: 同 Step 3。`Ctrl-C` 停。
- [ ] **Step 5: Commit**
```bash
git add run.sh run.bat
git commit -m "feat(dev): run.sh/run.bat 启动前确保 Docker PostgreSQL 就绪"
```
---
## Task 4: `.env.example` 默认切 PostgreSQL
**Files:**
- Modify: `.env.example`(第 8-10 行「数据库」段)
- [ ] **Step 1: 改 `.env.example` 的 DATABASE_URL**
把:
```ini
# ===== 数据库 =====
# SQLite 本地文件路径。生产环境用 /opt/shaguabijia-app-server/data.db
DATABASE_URL=sqlite:///./data/app.db
```
改成:
```ini
# ===== 数据库 =====
# 本地开发/测试统一用 Docker PostgreSQL:run.bat/run.sh 会自动拉起容器
# (docker-compose.yml + scripts/ensure_pg.py)。详见 docs/database/postgres-migration.md。
# 生产用原生 PG,由 scripts/init_postgres.py 写入强随机密码的连接串。
# ⚠️ scheme 必须是 postgresql+psycopg://(psycopg3);不要写成 postgresql://(会去找未装的 psycopg2)。
DATABASE_URL=postgresql+psycopg://shaguabijia_app:shaguabijia_dev_pw@localhost:5432/shaguabijia
```
- [ ] **Step 2: 验证(新 .env 从模板复制后能起服务)**
Run: `cp .env.example /tmp/env.check && grep '^DATABASE_URL=' /tmp/env.check`
Expected: `DATABASE_URL=postgresql+psycopg://shaguabijia_app:shaguabijia_dev_pw@localhost:5432/shaguabijia`
- [ ] **Step 3: Commit**
```bash
git add .env.example
git commit -m "feat(dev): .env.example 默认 DATABASE_URL 切 Docker PostgreSQL"
```
---
## Task 5: `tests/conftest.py` 切 PostgreSQL 测试库
**Files:**
- Modify: `tests/conftest.py`(整体替换:去掉临时 SQLite,改指 `shaguabijia_test` + 调 ensure_pg + fixture 改 drop/create)
- [ ] **Step 1: 整体替换 `tests/conftest.py`**
```python
"""测试用 fixtures。
测试库用 Docker PG 的 shaguabijia_test(与 dev 业务库 shaguabijia 隔离)。
顺序(必须):设 test DATABASE_URL(在 import app.* 之前)→ ensure PG 就绪 →
import app → 建表。持久卷可能残留上次的表 → session 开头先 drop 再 create。
"""
from __future__ import annotations
import os
from collections.abc import Iterator
# 1) 测试库连接串——必须在 import app.* 之前设好(app.db.session 在 import 期建 engine)
_TEST_DB_URL = (
"postgresql+psycopg://shaguabijia_app:shaguabijia_dev_pw@localhost:5432/shaguabijia_test"
)
os.environ["DATABASE_URL"] = _TEST_DB_URL
os.environ.setdefault("JWT_SECRET_KEY", "test-secret-please-ignore-this-is-only-for-pytest-not-real")
os.environ.setdefault("ADMIN_JWT_SECRET", "test-admin-secret-please-ignore-only-for-pytest-not-real")
os.environ.setdefault("JG_APP_KEY", "test-key")
os.environ.setdefault("JG_MASTER_SECRET", "test-secret")
os.environ.setdefault("SMS_MOCK", "true")
os.environ.setdefault("APP_ENV", "dev")
os.environ.setdefault("APP_DEBUG", "false") # 测试不打 SQL 日志
os.environ.setdefault("WECHAT_APP_ID", "wxtest0000000000")
os.environ.setdefault("WECHAT_APP_SECRET", "test-secret")
os.environ.setdefault("WXPAY_MCH_ID", "test-mch")
os.environ.setdefault("WXPAY_MCH_SERIAL_NO", "test-serial")
os.environ.setdefault("WXPAY_PUBLIC_KEY_ID", "test-pubkey-id")
os.environ.setdefault("RATE_LIMIT_ENABLED", "false")
os.environ.setdefault("PANGLE_CALLBACK_ENABLED", "true")
os.environ.setdefault("PANGLE_REWARD_SECRET", "test-pangle-secret-only-for-pytest")
# 2) 保证 Docker PG 就绪 + 测试库存在(必须在 import app.db.session 建 engine 之前)
from scripts.ensure_pg import ensure
if not ensure(_TEST_DB_URL):
raise RuntimeError(
"测试需要 Docker PostgreSQL 就绪。请确认已装 Docker Desktop;"
"或先跑一次 run.bat/run.sh 把 PG 拉起,再重试 pytest。"
)
import pytest
from fastapi.testclient import TestClient
from app.db.base import Base
from app.db.session import engine
from app.main import app
@pytest.fixture(scope="session", autouse=True)
def _setup_db() -> Iterator[None]:
# 持久卷可能残留上次跑崩后的表/数据 → 先 drop 再 create,保证干净起点
Base.metadata.drop_all(engine)
Base.metadata.create_all(engine)
yield
Base.metadata.drop_all(engine)
@pytest.fixture()
def client() -> TestClient:
return TestClient(app)
```
- [ ] **Step 2: 验证 conftest 能引导 PG 并收集用例(选一个不涉 DB 的测试文件)**
Run: `pytest tests/test_ensure_pg.py -v`
Expected: conftest 先打印 `[ensure_pg] ✅ PostgreSQL 已在 localhost:5432 运行`(或拉起过程),随后 12 passed。说明「测试走 PG 引导」链路通、且纯函数测试不受影响。
- [ ] **Step 3: 验证建表落到 PG 测试库(跑一个 DB 相关用例)**
Run: `pytest tests/test_invite.py -v`
Expected: 用例在 `shaguabijia_test` 上建表并执行(可能有个别红,留待 Task 6);关键是不再出现 SQLite 临时文件、engine 连的是 PG。
- [ ] **Step 4: Commit**
```bash
git add tests/conftest.py
git commit -m "test(dev): conftest 切 shaguabijia_test(Docker PG),引导+drop/create"
```
---
## Task 6: 全量跑 pytest on PG,逐个修红用例
> SQLite 宽松、PG 严格,切库会暴露一批真 bug(迁移指南 §2.2 已列)。本任务是**发现驱动**:先跑全量、按类别归因、按下述配方修,直到全绿。修改范围限被测业务/模型代码,不改测试来掩盖真 bug(除非测试本身依赖 SQLite 特性,如秒级时间精度)。
**Files:**
- Modify: 视失败而定(常见:`app/models/*.py``app/**/repositories/*.py`、少量 `tests/*.py`)
- [ ] **Step 1: 全量跑,拿到失败清单**
Run: `pytest -q`
Expected: 大部分通过;记录所有 FAIL 的用例名与报错文本,按下面类别归因。
- [ ] **Step 2: 修「naive datetime / 时区」类**
定位:`git grep -n "utcnow()" app/`。把 `datetime.utcnow()` 改成 `datetime.now(timezone.utc)`(并 `from datetime import timezone`)。
症状:PG `TIMESTAMPTZ` 与 naive datetime 比较/写入报错或结果错位;`tests/test_cps_admin.py` 已注释过 SQLite 忽略 tzinfo 的行为。
例:
```python
# 改前
from datetime import datetime
ts = datetime.utcnow()
# 改后
from datetime import datetime, timezone
ts = datetime.now(timezone.utc)
```
- [ ] **Step 3: 修「字符串/整数隐式比较」类**
症状:SQLite 允许 `WHERE phone = 13800138000`(自动转型),PG 直接报类型错。定位报错用例引用的查询,确保比较两侧类型一致(手机号等一律按字符串传参 `:phone`,不要传裸 int)。
- [ ] **Step 3.5: 修「事务已中止」类**
症状:某用例后续报 `current transaction is aborted, commands ignored until end of transaction block`,根因是前一句 SQL 出错后业务代码缺 `db.rollback()`/`db.commit()` 边界。补上正确的 commit/rollback。
- [ ] **Step 4: 修「测试依赖 SQLite 特性」类(仅此类可改测试)**
症状:测试断言依赖 SQLite 秒级时间精度或 FK 不强制(见 `test_invite.py:328``test_compare_harvest.py:153` 的注释)。PG 下时间精度更高/FK 更严——调整测试数据(如手动拉开时间间隔、用合法 FK)使断言在 PG 下成立,不改业务逻辑。
- [ ] **Step 5: 反复跑到全绿**
Run: `pytest -q`
Expected: `N passed`(0 failed)。若仍有红,回到 Step 2-4 继续归因。
- [ ] **Step 6: Commit**
```bash
git add -A
git commit -m "fix(db): 测试套件切 PostgreSQL 后修复严格性暴露的用例"
```
---
## Task 7: 文档更新
**Files:**
- Modify: `docs/database/postgres-migration.md`(§1 增「本地 Docker 一键起」小节)
- Modify: `CLAUDE.md`(DB 段注明 dev/test = Docker PG)
- Modify: `scripts/init_postgres.py`(顶部注释区分生产/本地)
- [ ] **Step 1: `postgres-migration.md` 在「## 1. 本地起 PG」开头插入推荐做法**
`## 1. 本地起 PG + 跑通空库(半天)` 标题下、`### 1.1 装 PG` 之前插入:
```markdown
### 1.0 推荐:Docker 一键起(本地开发/测试)
本地开发不必手动装 PG。已提供 `docker-compose.yml` + `scripts/ensure_pg.py`:
```bash
cp .env.example .env # DATABASE_URL 默认已是 Docker PG 连接串
./run.sh # 或 run.bat;会自动:探测 PG → 没起则启 Docker → 起 PG 容器 → 建库 → alembic → uvicorn
pytest # conftest 自动引导同一容器的 shaguabijia_test 库
```
容器:`postgres:16-alpine`(名 `shaguabijia-pg`,端口 5432,命名卷 `pgdata` 持久化),
首启即建业务库 `shaguabijia` 与测试库 `shaguabijia_test`。下面 1.1-1.5 的手动装 PG 步骤仅在
不用 Docker 时才需要;生产仍走 §4 的原生 PG。
```
- [ ] **Step 2: `CLAUDE.md` DB 段补充**
找到 DB 相关行(`**Prod**: PostgreSQL — just change DATABASE_URL...`)所在段,在其上方加一行:
```markdown
- **Dev/Test**: Docker PostgreSQL 16 — `run.sh`/`run.bat``scripts/ensure_pg.py` + `docker-compose.yml` 自动拉起;`.env.example` 默认即 PG 连接串;pytest 用同容器的 `shaguabijia_test` 库。**本地不再用 SQLite**。
```
- [ ] **Step 3: `scripts/init_postgres.py` 顶部注释区分场景**
把模块 docstring 第一行下方(`新机器初始化用。前置:...` 那行)改为:
```python
新机器初始化用(面向【生产原生 PG】:apt/systemd 装好的 PostgreSQL)。
本地开发/测试请改用 docker-compose.yml + scripts/ensure_pg.py(run.sh/run.bat 自动拉起),不必跑本脚本。
前置:已装 PostgreSQL 16 + 知道 postgres 超级用户密码。
```
- [ ] **Step 4: 验证无坏链接/格式**
Run: `git diff --stat`
Expected: 三个文档文件有改动,无其他文件被误改。
- [ ] **Step 5: Commit**
```bash
git add docs/database/postgres-migration.md CLAUDE.md scripts/init_postgres.py
git commit -m "docs(dev): 记录本地 Docker PostgreSQL 用法,区分生产原生 PG 路径"
```
---
## 完成标准(对齐 spec §8 验收)
- [ ] 全新机器(装了 Docker Desktop、`.env``.env.example` 复制)跑 `run.bat`/`run.sh` 全自动拉起 PG 并起服务,无手动装 PG。
- [ ] `docker ps``shaguabijia-pg` healthy;`shaguabijia``shaguabijia_test` 两库都在。
- [ ] PG 已在跑时再跑 `run`,ensure_pg 秒过短路。
- [ ] `pytest -q` 全绿(连 `shaguabijia_test`)。
- [ ] `.env` 改回 sqlite 时,`python -m scripts.ensure_pg` 硬失败并打印正确 PG 串。
- [ ] 能在 `admin/repositories` 写一段 PG 专有聚合(如 `count(*) FILTER (WHERE ...)`),`run` 手动跑通且相应 pytest 通过。
## Self-Review 记录(计划作者已核)
- **Spec 覆盖:** spec §9 待实现清单 8 项 → Task 1(compose+initdb)、Task 2(ensure_pg)、Task 3(run 接线)、Task 4(.env.example)、Task 5(conftest)、Task 6(修红用例)、Task 7(文档);「session.py 无需改」在 Header 与 spec §4.7 说明;「data/ 忽略」在 Task 1 Step 5 验证。无遗漏。
- **占位符:** 全部步骤含真实代码/命令/期望输出。Task 6 是发现驱动,已用「类别+具体转换配方+定位命令」代替不可预知的逐条 diff——非占位。
- **类型/命名一致:** `ensure(database_url=None)` 签名在 Task 2 定义,Task 5 以 `ensure(_TEST_DB_URL)` 调用一致;库名/用户/密码/端口全程为 Header 固定值;compose 服务名 `postgres` 与 ensure_pg `COMPOSE_SERVICE` 一致;`shaguabijia_test` 在 initdb SQL、`_ensure_test_db()`、conftest 三处一致。
@@ -1,281 +0,0 @@
# 本地开发切 Docker PostgreSQL —— 设计文档
> 让本地开发与测试统一跑在 Docker 化的 PostgreSQL 上,彻底退掉 SQLite。
> `run.bat` / `run.sh` 启动时自动检测本机 PG,没起就拉起 Docker → 起 PG 容器(镜像缺失先拉),
> 目的是让开发/大模型能放心用 PG 专有的高效聚合函数,不再为兼容 SQLite 而退化成"取基础数据后内存聚合"。
>
> 状态:已定稿(待用户复核)。作者对话日期:2026-07-08。
> 关联:[postgres-migration.md](../../database/postgres-migration.md)(切引擎完整步骤)、`scripts/init_postgres.py`(生产原生 PG 初始化)。
> ⚠️ **2026-07-27 增补(见 §10)**:D4 已从「sqlite 硬失败」松为「显式 SQLite 逃生舱」。§1-9 描述的是初版「彻底退掉 SQLite」设计;凡涉及「无 Docker / DATABASE_URL 是 sqlite 时如何处理」,**以 §10 为准**(测试仍只跑 PG 不变)。
---
## 1. 背景与目标
### 问题
当前开发环境默认用 SQLite(`DATABASE_URL=sqlite:///./data/app.db`),生产用 PostgreSQL 16。两套引擎并存,导致写数据访问代码时(尤其 `app/admin/repositories/` 的报表聚合)为了"两边都能跑",放弃 PG 专有能力(窗口函数、`FILTER``JSONB` 操作符、`GROUPING SETS` 等),改成"先查基础数据、再在 Python 内存里聚合"——既慢又啰嗦。
### 目标
本地开发与测试都跑在 PG 上,SQLite 退出本地开发闭环。之后写 PG 专有 SQL 时:
- 开发运行时(`run.bat`/`run.sh`)直接连 PG,手动验证可行;
- `pytest` 也连 PG,PG 专有 SQL 在被测代码路径里也安全,不会因 SQLite 而挂——**这是"双库兼容代码彻底消失"的必要条件**。
### 非目标(本期不做)
- 不动**生产**部署(生产仍是原生 PG16 + systemd,无 Docker;`init_postgres.py` 保持不变)。
- 不做数据搬迁(MVP 阶段无真实用户数据,详见迁移指南背景假设)。
- 不接 CI(仓库当前无 `.github/workflows`;若将来加 CI,再单独让 CI 起 PG service)。
- 不引入 Redis / testcontainers / 连接池中间件。
---
## 2. 决策记录(本次对话已拍板)
| # | 决策点 | 结论 | 理由 |
|---|---|---|---|
| D1 | PG 覆盖范围 | **dev 运行 + 测试都切 PG** | 只切运行时的话,被 SQLite 测试覆盖的代码路径(如 `test_cps_admin.py` 覆盖的 `admin/repositories/cps.py`)仍不能用 PG 专有 SQL,双库代码不会真正消失 |
| D2 | 打包方式 | **方案 A:Compose + `scripts/ensure_pg.py`** | 唯一真正需要定制的部分(启动 Docker 守护进程、等 PG 就绪)集中到一个跨平台模块,`run.bat`/`run.sh`/`conftest.py` 共用;声明式的容器/卷/健康检查交给 Compose |
| D3 | 宿主端口 | **5432**(与生产/文档一致) | 边界:若本机已有原生 PG 占 5432,`ensure_pg` 会探测到"PG 已在"直接复用它(可能连到不带业务库的实例)——见 §7 风险,文档提示 |
| D4 | dev 下 `DATABASE_URL` 仍是 sqlite | **硬失败**(打印一行 fix 后非 0 退出) | 彻底断掉 SQLite 退路,符合"让大家都用 PG"的目标 |
| D5 | dev 数据库密码 | 固定 `shaguabijia_dev_pw`,写进 compose + `.env.example` | 本地容器仅绑 `localhost`,非机密;保证 `.env.example` 复制即可用。生产密码另由 `init_postgres.py` 强随机生成,不复用 |
| D6 | 改哪些启动脚本 | `run.bat``run.sh` **都改** | 仓库一贯保持两者同步 |
| D7 | 镜像 | `postgres:16-alpine` | 对齐生产 PG16;alpine 体积小 |
---
## 3. 现状(改动前)
- **配置**:`app/core/config.py` `DATABASE_URL` 默认 `sqlite:///./data/app.db`,pydantic-settings 从 `.env` 读(环境变量优先级高于 `.env` 文件)。
- **引擎**:`app/db/session.py``_is_sqlite = DATABASE_URL.startswith("sqlite")` 分流——SQLite 加 `check_same_thread=False`、不建池;非 SQLite 加 `pool_size=10/max_overflow=20/pool_recycle=3600`。**已天然支持 PG,无需改。**
- **迁移**:`alembic/env.py``settings.DATABASE_URL` 读连接串,`render_as_batch` 仅对 sqlite 开;PG 下自动关。**无 psycopg2 硬编码,切 PG 无需改。**
- **驱动**:`pyproject.toml` 已装 `psycopg[binary]>=3.1`(psycopg3)。URL scheme 必须 `postgresql+psycopg://`(裸 `postgresql://` 会被 SQLAlchemy 路由到未安装的 psycopg2 → ModuleNotFoundError)。
- **测试**:`tests/conftest.py` 在 import app 前把 `DATABASE_URL` 设成临时文件 SQLite;session 级 autouse fixture 做 `Base.metadata.create_all(engine)` / 结束 `drop_all`(schema 来自 model 而非 alembic,无逐用例 rollback,全会话共享一个库)。
- **启动脚本**:`run.bat` / `run.sh` 均为:校验 `.env` 存在 → `mkdir data``alembic upgrade head` → uvicorn 监听 `0.0.0.0:8770`(`.sh``--reload --reload-dir app`)。
- **现有 PG 资产**:`scripts/init_postgres.py`(交互式:建用户/建库/授权/写 .env/跑迁移,面向**已装好的原生 PG**)、`docs/database/postgres-migration.md`(切引擎完整步骤,含 §2 测试切 PG、§2.2 会暴露的真 bug 清单)。
- **CI**:无(`.github/workflows` 不存在),故测试切 PG 无 CI 联动负担。
---
## 4. 方案详解
### 4.1 新增 `docker-compose.yml`(app-server 根目录)
```yaml
services:
postgres:
image: postgres:16-alpine
container_name: shaguabijia-pg
environment:
POSTGRES_USER: shaguabijia_app
POSTGRES_PASSWORD: shaguabijia_dev_pw
POSTGRES_DB: shaguabijia
ports:
- "5432:5432"
volumes:
- pgdata:/var/lib/postgresql/data
- ./docker/initdb:/docker-entrypoint-initdb.d:ro
healthcheck:
test: ["CMD-SHELL", "pg_isready -U shaguabijia_app -d shaguabijia"]
interval: 3s
timeout: 3s
retries: 20
volumes:
pgdata:
```
- `POSTGRES_USER` 设定后,该用户以超级用户身份创建并拥有 `POSTGRES_DB`,故能再建测试库。
- 首启 initdb 脚本建测试库(见 4.2)。命名卷 `pgdata` 让数据跨重启留存。
### 4.2 新增 `docker/initdb/01-create-test-db.sql`
```sql
-- 仅在 pgdata 卷首次初始化时执行一次。以 shaguabijia_app(超级用户)连 shaguabijia 库运行。
CREATE DATABASE shaguabijia_test OWNER shaguabijia_app;
```
### 4.3 新增 `scripts/ensure_pg.py`(纯标准库 + docker CLI,跨平台)
对外同时暴露**可导入函数** `ensure()`(供 `conftest.py` 直接调)和 **CLI 入口** `if __name__ == "__main__": sys.exit(0 if ensure() else 1)`(供 `run``python -m scripts.ensure_pg` 跑)。`ensure()``app.core.config.settings``DATABASE_URL`,解析 host/port,主流程:
1. **sqlite 守卫**:若 `DATABASE_URL``sqlite` 开头 → 打印"dev 已切 PG,请把 .env 的 DATABASE_URL 改成 `postgresql+psycopg://shaguabijia_app:shaguabijia_dev_pw@localhost:5432/shaguabijia`"→ 非 0 退出(D4)。
2. **TCP 探测** `host:port`(stdlib `socket`,超时 1s)。通 → 打印"✅ PG 已就绪"直接返回(幂等:PG 已在跑时开销≈一次握手)。
3. 不通 → `docker version` 探 CLI;缺失 → 中文报错"请先安装 Docker Desktop:https://www.docker.com/products/docker-desktop/" → 非 0 退出。
4. `docker info` 探守护进程;不通 → 按平台启动:
- Windows:`start "" "%ProgramFiles%\Docker\Docker\Docker Desktop.exe"`(找不到则报错让用户手动开)
- macOS:`open -a Docker`
- Linux:不自动 sudo,打印 `sudo systemctl start docker` 让用户执行后重试
然后轮询 `docker info` 直到就绪或超时(默认 120s,每 3s 一次,打印进度)。
5. `docker compose up -d`(Compose 在镜像缺失时**自动拉取**,首用拉 alpine ~90MB;有进度输出)。
6. 轮询 healthcheck(`docker inspect` 的 health 状态)/ TCP 直到 PG 接受连接(默认 60s 超时)。
7. **幂等确保测试库存在**(兼容"老 pgdata 卷没跑过 initdb"的情况):
`docker compose exec -T postgres psql -U shaguabijia_app -tc "SELECT 1 FROM pg_database WHERE datname='shaguabijia_test'"`,不存在则 `CREATE DATABASE shaguabijia_test OWNER shaguabijia_app`
失败即清晰中文报错 + 非 0 退出,**全程不回退 SQLite**。所有超时可用环境变量覆盖(如 `ENSURE_PG_DOCKER_TIMEOUT`)。
### 4.4 `run.bat` / `run.sh` 接线
`alembic upgrade head` **之前**插一行调用,失败即退出:
- `run.sh`:`"$PY" -m scripts.ensure_pg`(`set -e` 已在,失败自动退出)
- `run.bat`:`call "%PY%" -m scripts.ensure_pg` + `if errorlevel 1 exit /b 1`
其余逻辑不动(`mkdir data` 保留给 media 等落盘目录)。
### 4.5 `.env.example` 默认切 PG
```ini
DATABASE_URL=sqlite:///./data/app.db
```
改为
```ini
# 本地开发/测试统一用 Docker PG(run.bat/run.sh 会自动拉起容器;详见 docs/database/postgres-migration.md §本地 Docker 一键起)。
# 生产用原生 PG,由 scripts/init_postgres.py 写入强随机密码的连接串。
DATABASE_URL=postgresql+psycopg://shaguabijia_app:shaguabijia_dev_pw@localhost:5432/shaguabijia
```
### 4.6 `tests/conftest.py` 切 PG
调整顶部顺序(仍必须在 `import app.*` 之前完成 env 设定):
1. 设 `os.environ["DATABASE_URL"] = "postgresql+psycopg://shaguabijia_app:shaguabijia_dev_pw@localhost:5432/shaguabijia_test"`(测试库,永不碰 dev 业务库)。
2. 调 `scripts.ensure_pg.ensure()`(保证容器在 + 测试库在;PG 已在时几乎零开销)。
3. `import app...`
session 级 autouse fixture:改为 **`Base.metadata.drop_all(engine)``create_all(engine)`(开头先清干净,防持久卷里上一次跑残留的表/数据)→ yield → 结束 `drop_all`**;删掉临时 SQLite 文件相关代码。
> 预期:部分用例会因 PG 的严格性变红(SQLite 宽松、PG 严格),按迁移指南 §2.2 逐个修——常见为:字符串/整数隐式比较、`datetime.utcnow()` naive vs `TIMESTAMPTZ`、事务边界(`current transaction is aborted`)。这既是工作量也是本次改造的**直接收益**(暴露真 bug)。实现阶段需为"跑 pytest 并修红用例"单列步骤。
### 4.7 `app/db/session.py`
**无需改动**——`_is_sqlite` 为假时自动走 PG 池化分支。
### 4.8 文档
- `docs/database/postgres-migration.md` 增一节「本地 Docker 一键起 PG(推荐)」,指向 compose + `ensure_pg`,并说明它替代了 §1.1 的手动 brew/apt 装 PG。
- `CLAUDE.md` 的 DB 段注明:dev/test = Docker PG(`run` 自动拉起);prod = 原生 PG(`init_postgres.py`)。
- `scripts/init_postgres.py` 顶部注释补一句"本脚本面向生产原生 PG;本地开发用 docker-compose + scripts/ensure_pg"。
---
## 5. 连接参数汇总
| 项 | 值 |
|---|---|
| 镜像 | `postgres:16-alpine` |
| 容器名 | `shaguabijia-pg` |
| 宿主端口 | `5432` |
| 超级/业务用户 | `shaguabijia_app` |
| dev 密码 | `shaguabijia_dev_pw`(本地非机密) |
| 业务库(dev 运行) | `shaguabijia` |
| 测试库(pytest) | `shaguabijia_test` |
| dev `DATABASE_URL` | `postgresql+psycopg://shaguabijia_app:shaguabijia_dev_pw@localhost:5432/shaguabijia` |
| test `DATABASE_URL` | `postgresql+psycopg://shaguabijia_app:shaguabijia_dev_pw@localhost:5432/shaguabijia_test` |
| 数据持久化 | 命名卷 `pgdata` |
---
## 6. 失败处理矩阵(无 SQLite 兜底)
| 情形 | ensure_pg 行为 |
|---|---|
| `DATABASE_URL` 是 sqlite | 打印应改成的 PG 串 → 非 0 退出 |
| PG 已在跑(TCP 通) | 打印"已就绪" → 返回 0(跳过 docker) |
| 无 docker CLI | 提示装 Docker Desktop + 官网链接 → 非 0 退出 |
| docker 守护进程未起 | 尝试按平台启动 Docker Desktop,轮询到就绪;超时则报错 → 非 0 退出 |
| 镜像缺失 | `docker compose up -d` 自动拉取(不额外处理) |
| 容器起了但 PG 未 ready | 轮询 healthcheck 到超时;超时报错 → 非 0 退出 |
| 老 pgdata 卷缺测试库 | 幂等 `CREATE DATABASE shaguabijia_test` |
---
## 7. 风险与边界
- **端口占用(原生 PG 撞 5432)**:D3 选了 5432。若开发机已有原生 PG 监听 5432,step 2 的 TCP 探测会判"PG 已在"并复用它——但那个实例可能没有 `shaguabijia`/`shaguabijia_test` 库或用户,后续 `alembic upgrade head` / 测试会报连不上库或认证失败。**缓解**:文档提示"本机别再单独跑原生 PG";报错信息里提示检查是不是撞了原生 PG。
- **首次启动慢**:首用需 Docker Desktop 冷启(~3060s)+ 拉镜像(~数十秒~数分钟,视网络)。`ensure_pg` 全程打印进度,超时可配。
- **`DATABASE_URL` 环境变量优先级**:pydantic-settings 里 shell 环境变量优先于 `.env`。若开发者 shell 残留旧的 `DATABASE_URL`(如指向 sqlite),会盖过 `.env`。sqlite 守卫(D4)能挡住 sqlite 残留;但若残留的是另一个 PG 串,则以它为准——文档提示。
- **持久卷脏状态**:测试用 drop_all→create_all 开头清库,避免上次崩溃残留污染;dev 业务库随卷留存(符合预期)。
- **Docker 未安装/公司网络拉镜像受限**:硬失败并给出明确指引;不提供 SQLite 退路是刻意选择(D4/目标)。
---
## 8. 验收标准
1. 全新机器(装了 Docker Desktop、`.env``.env.example` 复制)执行 `run.bat`(或 `run.sh`):自动拉起 Docker→起 PG 容器→建库→`alembic upgrade head`→uvicorn 起在 8770,无手动装 PG 步骤。
2. `docker ps``shaguabijia-pg` 健康;`psql`/客户端能连 `shaguabijia``shaguabijia_test` 两个库。
3. PG 已在跑时再次 `run`,`ensure_pg` 秒过(不重复拉容器)。
4. `pytest``shaguabijia_test` 跑;红用例全部修绿(PG 严格性暴露的问题)。
5. `.env``DATABASE_URL` 改回 sqlite 时,`run`/`pytest` 硬失败并打印正确的 PG 串。
6. 能在 `admin/repositories/` 里写一段 PG 专有聚合 SQL(如带 `FILTER (WHERE ...)` 的聚合),`run` 下手动跑通、相应 pytest 也通过——即"双库兼容负担消失"的实证。
---
## 9. 待实现清单(供 writing-plans 拆解)
- [ ] 新增 `docker-compose.yml`
- [ ] 新增 `docker/initdb/01-create-test-db.sql`
- [ ] 新增 `scripts/ensure_pg.py`(TCP 探测 / 启 Docker Desktop 轮询 / compose up / 等 healthy / 幂等建测试库 / sqlite 守卫)
- [ ] `run.sh``run.bat` 接入 `ensure_pg`
- [ ] `.env.example``DATABASE_URL` 切 PG
- [ ] `tests/conftest.py``shaguabijia_test` + 调 `ensure_pg` + fixture 改 drop/create
- [ ] 跑 `pytest`,按迁移指南 §2.2 修红用例
- [ ] 文档:`postgres-migration.md` 增「本地 Docker 一键起」节;`CLAUDE.md` DB 段;`init_postgres.py` 注释
- [ ] `.gitignore` 确认 `data/` 已忽略(compose 用命名卷,不落项目目录,无需额外忽略)
---
## 10. 增补(2026-07-27):D4 反转 —— 显式 SQLite 逃生舱
> 背景:§2 的 D4 定为「dev 下 `DATABASE_URL` 仍是 sqlite → 硬失败」,目的是彻底断掉 SQLite 退路。实践中这对「本机装不了 Docker」的开发者过于刚性——直接被卡死、连跑都跑不起来。本次(2026-07-27 对话)把 D4 从「硬失败」松成「**显式逃生舱**」:工具**从不替你静默切库**,但会在没 Docker 时告诉你怎么手动降级,且降级时每次启动都醒目告警。
### 10.1 决策更新
| # | 原决策 | 新决策 | 理由 |
|---|---|---|---|
| D4 | dev sqlite URL → 硬失败退出 | **放行 + 每次打印醒目降级横幅**(仍非静默) | 已手动改 `.env`=sqlite = 开发者的显式选择,尊重它;但吼一嗓子防止忘了自己在降级、把 PG 专有 SQL 提交上去 |
| D8(新) | (无) | 无 docker CLI 时,报错里**追加逃生舱指路**(改 `.env`=sqlite),但仍非 0 退出 | 「显式」的关键:工具不替你切库,只指路;开发者改完 `.env` 再跑一次才真正降级 |
**未变**:D1(测试仍只跑 PG)、D2-D3、D5-D7 全部保留。逃生舱**只作用于 `run.sh`/`run.bat` 运行时**;`pytest` 仍写死连 PG 测试库(`conftest.py` 传 PG URL,sqlite 分支根本不触发),没 Docker 就 `raise`、跑不了完整套件——这正是 D1「测试上 PG 才能暴露真 bug」的初衷,刻意不给逃生舱。
### 10.2 代码改动(仅 `scripts/ensure_pg.py``ensure()`)
1. **sqlite 分支**(原 `return False`)→ 打印多行降级横幅后 `return True`。横幅点明:PG 专有 SQL/严格类型在此模式**不被验证**、提交前须在有 Docker 的机器上用 PG 复跑、装好 Docker 后把 `DATABASE_URL` 改回 PG 串。
2. **无 docker CLI 分支**(原仅提示装 Docker + `return False`)→ 追加一句「装不了 Docker?把 `.env``DATABASE_URL` 改成 `sqlite:///./data/app.db` 可降级运行」;**仍 `return False`**(run 脚本照常退出,开发者需显式改 .env 再跑)。
3. **常量**:新增 `SQLITE_URL = "sqlite:///./data/app.db"`(逃生舱指路用);`SQLITE_FIX_HINT` 重命名 `PG_URL`(降级横幅"改回 PG"引用)。
4. 更新模块 docstring 中「全程无 SQLite 兜底」一句,改述为「无 Docker/sqlite URL 时【显式】降级 SQLite(带醒目告警),测试侧不降级」。
**其余全不动**:`run.sh`/`run.bat`(sqlite 下 `ensure` 返 True → 照常 `alembic upgrade head` + uvicorn)、`docker-compose.yml``app/db/session.py`(SQLite 引擎分支本就保留为 fallback)、`tests/conftest.py``.env.example`(默认仍 PG)。
### 10.3 改完后行为矩阵(覆盖用户列的 5 场景)
| 场景 | `DATABASE_URL` | ensure_pg 行为 |
|---|---|---|
| ① 无 Docker | PG(默认) | 报错 + 指逃生舱 → 退出;开发者改 `.env`=sqlite → 再跑 → **放行 + 降级横幅**,alembic/uvicorn 跑 SQLite |
| ② 有 Docker 未启动 | PG | 启 Docker Desktop → `compose up` → 等 ready → 建测试库(**不变**) |
| ③ 有 Docker 已启动 | PG | `compose up` → 等 ready(**不变**) |
| ④ PG 已在跑 | PG | TCP 通 → 秒过跳过 Docker(**不变**) |
| ⑤ PG 起来后 | 任意 | run 脚本 `alembic upgrade head`(**不变**;SQLite 走 `render_as_batch`) |
### 10.4 风险
- **降级被忽视**:横幅仅在 `run` 启动时打印一次;若开发者用 IDE 直接起 uvicorn(绕过 run 脚本)则看不到。缓解:横幅足够醒目 + 文档强调;**不**引入 app 启动期重复告警(YAGNI)。
- **测试无 Docker 跑不了**:刻意保留(D1)。文档提示无 Docker 者:要么装 Docker 跑全量测试,要么只在 CI/有 Docker 的机器上验证 PG 相关改动。
### 10.5 Redis 前瞻(不在本次)
§2 未涉及 Redis。②③ 场景未来若加 Redis 实例:在 `docker-compose.yml``redis` 服务即可,`docker compose up -d` 天然带起;仅当启动期有组件依赖 Redis 才需给 `ensure_pg` 加 redis readiness 探测。本次不做,方案对它友好。
### 10.6 附带修复:`_docker_cli_ok` 守护进程误判(2026-07-27)
诊断「装了 Docker Desktop 却报未检测到 docker」时发现的真 bug:`_docker_cli_ok()` 原用 `docker version`
判断 CLI 是否存在,但该命令**要连 daemon**,守护进程没起时退非零 → 把「Docker 装了但没启动」
误判成「没装 CLI」,`ensure()` 直接打印"请安装 Docker Desktop"并 `return False`,**绕过了专为需求②
写的 `_start_docker_daemon()` 自动拉起逻辑**——需求②(有 Docker 未启动 → 自动启动)因此从未真正生效。
修复:改用 `docker --version`(纯客户端、不连 daemon、退 0)。`_docker_daemon_ok()` 仍用 `docker info`
(正确,该检查本就依赖 daemon)。实测机器:Docker Desktop 20.10.12 已装但引擎未起,修复前 `_docker_cli_ok()`
误报 False,修复后 True。
### 10.7 附带修复:固定 compose 项目名 + 清理残留同名容器(2026-07-27)
诊断「`docker compose up``container name "/shaguabijia-pg" already in use`」时发现的又一 bug:compose
项目名默认取运行目录 basename,在不同目录/worktree(如 `local-dev-postgres-docker` vs `shaguabijia-app-server`)
之间切换会各自成一个项目;而 `docker-compose.yml` 写死了 `container_name: shaguabijia-pg`(全局唯一名),
于是新项目 `up` 时要创建同名容器 → 撞上旧项目留下的那个 → 冲突。副作用:`pgdata` 卷也按项目名分裂
`local-dev-postgres-docker_pgdata` / `shaguabijia-app-server_pgdata`,数据被切成两半。
修复(均在 `scripts/ensure_pg.py`,`docker-compose.yml` 不动、容器名仍是 `shaguabijia-pg`):
1. 模块级 `os.environ.setdefault("COMPOSE_PROJECT_NAME", "shaguabijia")` —— 钉死项目名,无论从哪个
目录/worktree 跑都是同一个项目、同一个卷 `shaguabijia_pgdata`,所有 `docker compose up/exec` 一致。
2. `_compose_up()` 前置 `_remove_stale_container()`:若存在「同名但不属于本项目」的残留容器,先 `docker rm -f`
再 up(靠 `docker ps --filter name/label` 判归属;数据在命名卷里,删容器不丢)。旧目录/worktree 留下的
残留容器就此自动清掉,不需手动干预。
影响:本次修复后首跑,旧的 `shaguabijia-pg`(属项目 `local-dev-postgres-docker`)会被自动删除、在项目
`shaguabijia` 下重建,挂载全新的 `shaguabijia_pgdata`(空库,`alembic upgrade head` 重建表)。旧数据仍留在
`local-dev-postgres-docker_pgdata` 卷里(未删,可恢复);确认不需要后可 `docker volume rm` 清理两个旧卷。
+1 -8
View File
@@ -30,14 +30,7 @@ if not exist .env (
if not exist data mkdir data
REM Ensure local Docker PostgreSQL is up (auto-starts Docker + PG container if needed)
call "%PY%" -m scripts.ensure_pg
if errorlevel 1 (
echo [X] ensure_pg failed ^(PostgreSQL not ready^)
exit /b %errorlevel%
)
REM Build/upgrade schema (idempotent; no-op if already at head)
REM Build/upgrade SQLite schema (idempotent; no-op if already at head)
call "%PY%" -m alembic upgrade head
if errorlevel 1 (
echo [X] alembic upgrade head failed
+1 -2
View File
@@ -18,8 +18,7 @@ if [ ! -f .env ]; then
exit 1
fi
mkdir -p data # 运行期落盘目录(媒体上传等)
"$PY" -m scripts.ensure_pg # 确保本地 Docker PostgreSQL 就绪(没起会自动拉起;失败即退出)
mkdir -p data # sqlite 文件所在目录
"$PY" -m alembic upgrade head # 确保表已建(幂等,已是最新则 no-op)
# --reload 只盯源码目录 app/:别去监视 logs/(日志写入触发"检测→再写日志"回环)和
-50
View File
@@ -1,50 +0,0 @@
@echo off
REM Admin backend startup (Windows) - the :8771 peer of run.bat.
REM
REM Usage:
REM cd shaguabijia-app-server
REM run8771.bat
REM
REM Runs the ADMIN FastAPI app (app.admin.main:admin_app) on 127.0.0.1:8771 —
REM a SEPARATE process from run.bat (which runs app.main:app on 8770). The admin
REM web frontend (Next.js :3001) points at http://localhost:8771. Auto-reload on
REM code change.
REM
REM Prerequisite (first time):
REM conda activate pricebot ^&^& pip install -e .
REM copy .env.example .env ^&^& fill JWT_SECRET_KEY
REM
REM Tip: shaguabijia-admin-web\start.bat starts user-api(8770) + admin-api(8771)
REM + frontend(3001) in one go, if you prefer a single command.
cd /d "%~dp0"
REM Prefer the project virtualenv (.venv) so we never inherit a wrong
REM global/conda interpreter. FastAPI<0.115 on Pydantic 2.12 crashes at import
REM with "'FieldInfo' object has no attribute 'in_'". Falls back to PATH python.
set "PY=python"
if exist "%~dp0.venv\Scripts\python.exe" set "PY=%~dp0.venv\Scripts\python.exe"
if not exist .env (
echo [X] Missing .env. Run: copy .env.example .env and fill JWT_SECRET_KEY ^(plus MT_CPS_* if you test Meituan^)
exit /b 1
)
if not exist data mkdir data
REM Ensure local Docker PostgreSQL is up (auto-starts Docker + PG container if needed)
call "%PY%" -m scripts.ensure_pg
if errorlevel 1 (
echo [X] ensure_pg failed ^(PostgreSQL not ready^)
exit /b %errorlevel%
)
REM Build/upgrade schema (idempotent; no-op if already at head)
call "%PY%" -m alembic upgrade head
if errorlevel 1 (
echo [X] alembic upgrade head failed
exit /b %errorlevel%
)
REM Long-running foreground process. Ctrl+C to stop.
"%PY%" -m uvicorn app.admin.main:admin_app --host 127.0.0.1 --port 8771 --reload
-394
View File
@@ -1,394 +0,0 @@
"""确保本地 PostgreSQL 就绪(开发/测试统一用 Docker PG)。
被三处复用:
- run.sh / run.bat:`python -m scripts.ensure_pg`(CLI,失败退非 0)
- tests/conftest.py:`from scripts.ensure_pg import ensure; ensure(test_url)`
流程: DATABASE_URL TCP 探测 没起就(必要时启 Docker Desktop)
`docker compose up -d` PG ready 幂等确保测试库存在
运行时(run.sh/run.bat)支持显式SQLite 逃生舱:DATABASE_URL 设为 sqlite 放行并打印
醒目降级横幅(绝不静默替你切库); docker CLI 时报错里也指路该逃生舱测试侧
(conftest PG URL)不降级sqlite 分支不触发, PG 直接 raise详见设计文档 §10
生产用原生 PG(scripts/init_postgres.py),不走本模块
"""
from __future__ import annotations
import os
import shutil
import socket
import subprocess
import sys
import time
from collections.abc import Mapping
from pathlib import Path, PureWindowsPath
from urllib.parse import urlsplit
ROOT = Path(__file__).resolve().parent.parent
# 日志里可能含 emoji(如 ✅);Windows GBK 控制台(cmd.exe)无法编码会抛 UnicodeEncodeError → 脚本崩、
# run.bat 误判 ensure_pg 失败。用 backslashreplace 保底:中文仍正常,仅不可编码字符被转义,不崩。
for _stream in (sys.stdout, sys.stderr):
try:
_stream.reconfigure(errors="backslashreplace")
except (AttributeError, ValueError):
pass
APP_DB = "shaguabijia"
TEST_DB = "shaguabijia_test"
DB_USER = "shaguabijia_app"
COMPOSE_SERVICE = "postgres"
CONTAINER_NAME = "shaguabijia-pg" # 必须与 docker-compose.yml 的 container_name 一致
# 钉死 compose 项目名:否则它默认取运行目录 basename,在不同目录/worktree 之间切会各自
# 成一个项目 → 同一个固定 container_name 撞名报错、pgdata 卷还会按项目名分裂成多份。
# 钉成 app 名后,无论从哪个目录/worktree 跑都是同一个项目、同一个卷。setdefault:尊重外部覆盖。
os.environ.setdefault("COMPOSE_PROJECT_NAME", "shaguabijia")
DOCKER_START_TIMEOUT = int(os.environ.get("ENSURE_PG_DOCKER_TIMEOUT", "120"))
PG_READY_TIMEOUT = int(os.environ.get("ENSURE_PG_READY_TIMEOUT", "60"))
# 单条 docker 探测/exec 命令的超时:防 Docker 守护进程半死(尤其 Windows 冷启)时
# docker info / exec 无限挂起、绕过上面的总超时。
DOCKER_CMD_TIMEOUT = int(os.environ.get("ENSURE_PG_CMD_TIMEOUT", "15"))
POLL_INTERVAL = 3.0
PG_URL = (
"postgresql+psycopg://shaguabijia_app:shaguabijia_dev_pw@localhost:5432/shaguabijia"
)
SQLITE_URL = "sqlite:///./data/app.db" # 无 Docker 时的显式降级逃生舱(仅 run 运行时)
def _log(msg: str) -> None:
print(f"[ensure_pg] {msg}", flush=True)
def _is_sqlite(url: str) -> bool:
return url.strip().lower().startswith("sqlite")
def _parse_host_port(url: str) -> tuple[str, int]:
"""从 SQLAlchemy URL 取 host/port,缺省 localhost:5432。"""
parts = urlsplit(url)
return (parts.hostname or "localhost"), (parts.port or 5432)
def _port_open(host: str, port: int, timeout: float = 1.0) -> bool:
try:
with socket.create_connection((host, port), timeout=timeout):
return True
except OSError:
return False
def _win_docker_desktop_candidates(
env: Mapping[str, str], docker_cli: str | None
) -> list[PureWindowsPath]:
"""Windows 上 Docker Desktop.exe 的候选路径(按优先级)。纯函数:不碰文件系统、不读注册表。
PureWindowsPath 解析,故在任意 OS 上跑单测都按 Windows 语义(反斜杠分隔),行为确定
优先级:
环境变量 DOCKER_DESKTOP_EXE 显式指定(终极逃生舱,盘符随你);
PATH 上的 docker CLI 反推Docker Desktop CLI
<安装目录>\\resources\\bin\\docker.exe,往上几级即安装目录,天然跟随实际盘符
(装在 D 盘就反推出 D ,不再写死 C );多取几级容忍未来目录布局微调;
Program Files 变体下的标准安装路径兜底(覆盖常规 C 盘装)
调用方按序取第一个真实存在的
"""
out: list[PureWindowsPath] = []
override = (env.get("DOCKER_DESKTOP_EXE") or "").strip().strip('"')
if override:
out.append(PureWindowsPath(override))
if docker_cli:
for parent in list(PureWindowsPath(docker_cli).parents)[:4]:
out.append(parent / "Docker Desktop.exe")
for var in ("ProgramFiles", "ProgramW6432", "ProgramFiles(x86)"):
root = env.get(var)
if root:
out.append(PureWindowsPath(root) / "Docker" / "Docker" / "Docker Desktop.exe")
return out
def _docker_desktop_from_registry() -> Path | None:
"""从注册表尽力取 Docker Desktop.exe 位置(best-effort;非 Windows / 任何异常都当没找到)。
比路径猜测更权威且完全跟随实际盘符探两处:
- App Paths\\Docker Desktop.exe 的默认值(通常就是 exe 全路径);
- Uninstall\\Docker Desktop InstallLocation(安装目录,需再拼 exe )
"""
try:
import winreg
except ImportError: # 非 Windows
return None
probes = (
(winreg.HKEY_LOCAL_MACHINE,
r"SOFTWARE\Microsoft\Windows\CurrentVersion\App Paths\Docker Desktop.exe", "", False),
(winreg.HKEY_LOCAL_MACHINE,
r"SOFTWARE\Microsoft\Windows\CurrentVersion\Uninstall\Docker Desktop",
"InstallLocation", True),
)
for hive, subkey, value_name, join_exe in probes:
try:
with winreg.OpenKey(hive, subkey) as key:
val, _ = winreg.QueryValueEx(key, value_name)
except OSError:
continue # 键不存在/无权限 → 下一个
if not val:
continue
exe = Path(val) / "Docker Desktop.exe" if join_exe else Path(val)
if exe.exists():
return exe
return None
def _find_docker_desktop_exe() -> Path | None:
"""Windows 上尽力定位【真实存在】的 Docker Desktop.exe;遍历候选 + 注册表兜底,找不到返回 None。"""
for cand in _win_docker_desktop_candidates(os.environ, shutil.which("docker")):
if Path(cand).exists():
return Path(cand)
return _docker_desktop_from_registry()
def _docker_desktop_cmd(platform: str) -> list[str] | None:
"""按平台给出启动 Docker Desktop 的命令。
Windows:智能定位 exe( _find_docker_desktop_exe),找不到 None
macOS:交给 `open -a Docker`Linux:None(daemon sudo,让用户手动)
"""
if platform.startswith("win"):
exe = _find_docker_desktop_exe()
return [str(exe)] if exe else None
if platform == "darwin":
return ["open", "-a", "Docker"]
return None
def _docker_ok(subcmd: str) -> bool:
"""`docker --version`(CLI 在不在,纯客户端)/`docker info`(daemon 起没起)成功与否。"""
try:
subprocess.run(
["docker", subcmd],
cwd=ROOT,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
check=True,
timeout=DOCKER_CMD_TIMEOUT,
)
return True
except (OSError, subprocess.CalledProcessError, subprocess.TimeoutExpired):
return False
def _docker_cli_ok() -> bool:
# 必须用 `docker --version`(纯客户端,不连 daemon)而非 `docker version`
# (后者要连 daemon,守护进程没起时退非零)——否则「装了 Docker 但没启动」
# 会被误判成「没装 CLI」,直接绕过下面 _start_docker_daemon() 的自动拉起(需求②)。
return _docker_ok("--version")
def _docker_daemon_ok() -> bool:
return _docker_ok("info")
def _start_docker_daemon() -> bool:
"""守护进程没起时按平台拉起,轮询到就绪。返回是否成功。"""
if _docker_daemon_ok():
return True
cmd = _docker_desktop_cmd(sys.platform)
if cmd is None:
if sys.platform.startswith("win"):
# docker CLI 在 PATH 上(否则走不到这)、却定位不到 Docker Desktop.exe:多为非标准安装位置
_log("找不到 Docker Desktop.exe(已试:PATH 上 docker CLI 反推、注册表、常见安装目录)。")
_log(" 确已安装 → 设环境变量 DOCKER_DESKTOP_EXE=<Docker Desktop.exe 全路径> 再重试,"
"或先手动启动 Docker Desktop。")
_log(f" 不想折腾 → 把 .env 的 DATABASE_URL 改成 {SQLITE_URL} 可降级用 SQLite 跑(仅救急)。")
else:
_log("Docker 守护进程未运行。Linux 请手动:sudo systemctl start docker,然后重试。")
return False
_log(f"启动 Docker Desktop(首次冷启可能 30-60s):{cmd[0]}")
try:
subprocess.Popen(cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
except OSError as e:
_log(f"启动 Docker Desktop 失败:{e}")
return False
deadline = time.monotonic() + DOCKER_START_TIMEOUT
while time.monotonic() < deadline:
if _docker_daemon_ok():
_log("Docker 守护进程已就绪。")
return True
_log("等待 Docker 守护进程…")
time.sleep(POLL_INTERVAL)
_log(f"等待 Docker 守护进程超时({DOCKER_START_TIMEOUT}s)。")
return False
def _ps_names(*filters: str) -> str:
"""docker ps -a 按 filter 查容器名(每行一个);失败返回空串。"""
args = ["docker", "ps", "-a", "--format", "{{.Names}}"]
for f in filters:
args += ["--filter", f]
try:
r = subprocess.run(
args, cwd=ROOT, capture_output=True, text=True, timeout=DOCKER_CMD_TIMEOUT,
)
except (OSError, subprocess.TimeoutExpired):
return ""
return r.stdout if r.returncode == 0 else ""
def _remove_stale_container() -> None:
"""删掉「同名但不属于本 compose 项目」的残留容器(旧目录/worktree 建的)。
固定的 container_name 是全局唯一名:若旧项目留下一个同名容器,`docker compose up`
会因撞名报 "container name already in use" 而失败这里在 up 之前主动清掉它
数据在命名卷(<project>_pgdata),删容器不删卷不丢数据
"""
project = os.environ.get("COMPOSE_PROJECT_NAME", "")
name_filter = f"name=^{CONTAINER_NAME}$"
if CONTAINER_NAME not in _ps_names(name_filter).split():
return # 没有同名容器
ours = _ps_names(name_filter, f"label=com.docker.compose.project={project}")
if CONTAINER_NAME in ours.split():
return # 就是本项目的容器,compose 会自己 start/复用,别删
_log(f"发现残留同名容器 {CONTAINER_NAME}(非本项目 '{project}'),删除以避免撞名"
f"(数据在卷里,不丢)…")
try:
subprocess.run(
["docker", "rm", "-f", CONTAINER_NAME], cwd=ROOT,
stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=DOCKER_CMD_TIMEOUT,
)
except (OSError, subprocess.TimeoutExpired):
_log(f"⚠️ 删除残留容器失败,可手动: docker rm -f {CONTAINER_NAME}")
def _compose_up() -> bool:
_remove_stale_container()
_log("docker compose up -d(镜像缺失会自动拉取,首用约几十秒)…")
try:
subprocess.run(["docker", "compose", "up", "-d"], cwd=ROOT, check=True)
return True
except (OSError, subprocess.CalledProcessError) as e:
_log(f"docker compose up 失败:{e}")
return False
def _pg_isready() -> bool:
try:
r = subprocess.run(
["docker", "compose", "exec", "-T", COMPOSE_SERVICE,
"pg_isready", "-U", DB_USER, "-d", APP_DB],
cwd=ROOT, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
timeout=DOCKER_CMD_TIMEOUT,
)
except (OSError, subprocess.TimeoutExpired):
return False
return r.returncode == 0
def _wait_pg_ready(host: str, port: int) -> bool:
deadline = time.monotonic() + PG_READY_TIMEOUT
while time.monotonic() < deadline:
if _port_open(host, port) and _pg_isready():
_log("PostgreSQL 已就绪。")
return True
_log("等待 PostgreSQL 就绪…")
time.sleep(POLL_INTERVAL)
_log(f"等待 PostgreSQL 就绪超时({PG_READY_TIMEOUT}s)。")
return False
def _test_db_exists() -> bool:
"""测试库是否已存在(连业务库 shaguabijia 查 pg_database)。"""
try:
r = subprocess.run(
["docker", "compose", "exec", "-T", COMPOSE_SERVICE,
"psql", "-U", DB_USER, "-d", APP_DB, "-tAc",
f"SELECT 1 FROM pg_database WHERE datname='{TEST_DB}'"],
cwd=ROOT, capture_output=True, text=True, timeout=DOCKER_CMD_TIMEOUT,
)
except (OSError, subprocess.TimeoutExpired):
return False
return r.returncode == 0 and r.stdout.strip() == "1"
def ensure_test_db() -> bool:
"""幂等建测试库(兼容老 pgdata 卷首启没跑 initdb 的情况)。返回测试库是否就绪。
公开给 conftest 单独调用:ensure() 端口已通时会短路返回不建测试库,
所以测试侧需在 ensure() 之后再显式补一刀(best-effort)
"""
if _test_db_exists():
return True
_log(f"建测试库 {TEST_DB}")
try:
create = subprocess.run(
["docker", "compose", "exec", "-T", COMPOSE_SERVICE,
"psql", "-U", DB_USER, "-d", APP_DB, "-c",
f"CREATE DATABASE {TEST_DB} OWNER {DB_USER}"],
cwd=ROOT, capture_output=True, text=True, timeout=DOCKER_CMD_TIMEOUT,
)
except (OSError, subprocess.TimeoutExpired) as e:
_log(f"⚠️ 建测试库 {TEST_DB} 失败:{e}")
return False
# returncode==0=建成功;非 0 但库已存在=与并发创建者竞争失败(42P04),仍算就绪
if create.returncode == 0 or _test_db_exists():
return True
_log(f"⚠️ 建测试库 {TEST_DB} 失败:{(create.stderr or '').strip()}")
return False
def _warn_sqlite_degraded() -> None:
"""DATABASE_URL 是 SQLite 时打印醒目降级横幅(显式逃生舱,非静默切库)。"""
for line in (
"⚠️ ================= 降级模式(SQLite) =================",
"⚠️ DATABASE_URL 是 SQLite,不是 PostgreSQL。",
"⚠️ PG 专有 SQL(窗口函数/FILTER/JSONB)与严格类型在此模式【不被验证】。",
"⚠️ 提交前请在装了 Docker 的机器上用 PG 复跑;装好后把 DATABASE_URL 改回:",
f"⚠️ {PG_URL}",
"⚠️ ===================================================",
):
_log(line)
def ensure(database_url: str | None = None) -> bool:
"""确保 PG 就绪,返回 True/False。database_url 缺省从 settings 读(尊重 .env)。
运行时若 DATABASE_URL SQLite 打印降级横幅并返回 True(显式逃生舱);
conftest 传的是 PG URL,故测试侧永不走此分支
"""
if database_url is None:
from app.core.config import settings # 延迟导入,避免过早固化 settings
database_url = settings.DATABASE_URL
if _is_sqlite(database_url):
_warn_sqlite_degraded()
return True
host, port = _parse_host_port(database_url)
if _port_open(host, port):
_log(f"✅ PostgreSQL 已在 {host}:{port} 运行,跳过 Docker。")
return True
_log(f"{host}:{port} 无 PostgreSQL,准备用 Docker 拉起…")
if not _docker_cli_ok():
_log("未检测到 docker 命令。请先安装 Docker Desktop:")
_log(" https://www.docker.com/products/docker-desktop/")
_log(f"装不了 Docker?把 .env 的 DATABASE_URL 改成 {SQLITE_URL} 可降级用 SQLite 跑")
_log(" (PG 专有 SQL/严格性不被验证,仅救急);改完重跑 run.sh/run.bat。")
return False
if not _start_docker_daemon():
return False
if not _compose_up():
return False
if not _wait_pg_ready(host, port):
return False
if not ensure_test_db():
return False
return True
if __name__ == "__main__":
sys.exit(0 if ensure() else 1)
+1 -3
View File
@@ -1,8 +1,6 @@
"""Bootstrap PostgreSQL: 建用户 + 建库 + 授权 + 写 .env + 跑迁移。
新机器初始化用(面向生产原生 PG:apt/systemd 装好的 PostgreSQL)
本地开发/测试请改用 docker-compose.yml + scripts/ensure_pg.py(run.sh/run.bat 自动拉起),不必跑本脚本
前置:已装 PostgreSQL 16 + 知道 postgres 超级用户密码
新机器初始化用前置:已装 PostgreSQL 16 + 知道 postgres 超级用户密码
用法:
python scripts/init_postgres.py
+15 -22
View File
@@ -1,19 +1,22 @@
"""测试用 fixtures。
测试库用 Docker PostgreSQL shaguabijia_test( dev 业务库 shaguabijia 隔离)
顺序(必须): test DATABASE_URL( import app.* 之前) ensure PG 就绪 + 测试库存在
import app 建表持久卷可能残留上次的表 session 开头先 drop create
测试 DB 用临时文件 SQLite
- 不用 in-memory:in-memory 默认 per-connection,跨连接看不到表
- 用临时文件保证 SessionLocal 每次新连都看到同一份 schema
顺序:set env(必须在 import app.* 之前) import app 建表 TestClient
"""
from __future__ import annotations
import os
import tempfile
from collections.abc import Iterator
# 1) 测试库连接串——必须在 import app.* 之前设好(app.db.session 在 import 期就建 engine)
_TEST_DB_URL = (
"postgresql+psycopg://shaguabijia_app:shaguabijia_dev_pw@localhost:5432/shaguabijia_test"
)
os.environ["DATABASE_URL"] = _TEST_DB_URL
# 临时 db 文件路径,进程退出后清理
_tmp_db = tempfile.NamedTemporaryFile(suffix=".db", delete=False)
_tmp_db.close()
os.environ["DATABASE_URL"] = f"sqlite:///{_tmp_db.name}"
os.environ.setdefault("JWT_SECRET_KEY", "test-secret-please-ignore-this-is-only-for-pytest-not-real")
os.environ.setdefault("ADMIN_JWT_SECRET", "test-admin-secret-please-ignore-only-for-pytest-not-real")
os.environ.setdefault("JG_APP_KEY", "test-key")
@@ -32,18 +35,6 @@ os.environ.setdefault("RATE_LIMIT_ENABLED", "false") # 限流内存计数会跨
os.environ.setdefault("PANGLE_CALLBACK_ENABLED", "true")
os.environ.setdefault("PANGLE_REWARD_SECRET", "test-pangle-secret-only-for-pytest")
# 2) 保证 Docker PG 就绪(必须在 import app.db.session 建 engine 之前)。
# ensure() 在「端口已通」时会短路、不建测试库,故随后再显式补一刀 ensure_test_db()。
from scripts.ensure_pg import ensure, ensure_test_db
if not ensure(_TEST_DB_URL):
raise RuntimeError(
"测试需要 Docker PostgreSQL 就绪。请确认已装 Docker Desktop;"
"或先跑一次 run.bat/run.sh 把 PG 拉起,再重试 pytest。"
)
# best-effort 兜底建测试库(PG 已在跑但测试库缺失=老卷)。真缺库时下面 create_all 会明确报错。
ensure_test_db()
import pytest
from fastapi.testclient import TestClient
@@ -54,11 +45,13 @@ from app.main import app
@pytest.fixture(scope="session", autouse=True)
def _setup_db() -> Iterator[None]:
# 持久卷可能残留上次跑崩后的表/数据 → 先 drop 再 create,保证干净起点
Base.metadata.drop_all(engine)
Base.metadata.create_all(engine)
yield
Base.metadata.drop_all(engine)
try:
os.unlink(_tmp_db.name)
except OSError:
pass
@pytest.fixture()
+4 -4
View File
@@ -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"] == [
+109
View File
@@ -219,6 +219,115 @@ def test_user_reward_stats_can_scope_withdrawals_by_account(
assert invite.json()["cash_balance_cents"] == 456
def test_user_reward_stats_draw_ecpm_uses_all_filtered_impressions(
admin_client: TestClient, admin_token: str
) -> None:
"""Draw 平均 eCPM 应与收益报表一致,不能只平均成功发奖记录。"""
from app.models.ad_ecpm import AdEcpmRecord
from app.models.ad_feed_reward import AdFeedRewardRecord
uid = _seed_user_with_data("13800000024")
created_at = datetime(2038, 1, 15, 4, tzinfo=UTC)
db = SessionLocal()
try:
db.add_all(
[
AdFeedRewardRecord(
client_event_id="reward-stats-granted-high",
ad_session_id="reward-stats-granted-high",
user_id=uid,
reward_date="2038-01-15",
duration_seconds=10,
unit_count=1,
ecpm_raw="9000",
ad_type="draw",
feed_scene="coupon",
app_env="prod",
our_code_id="104098712",
coin=9,
status="granted",
created_at=created_at,
),
AdEcpmRecord(
user_id=uid,
ad_type="draw",
feed_scene="coupon",
ad_session_id="reward-stats-impression-low",
app_env="prod",
our_code_id="104098712",
ecpm_raw="1000",
report_date="2038-01-15",
created_at=created_at,
),
AdEcpmRecord(
user_id=uid,
ad_type="feed",
feed_scene="coupon",
ad_session_id="reward-stats-impression-mid",
app_env="prod",
our_code_id="104098712",
ecpm_raw="3000",
report_date="2038-01-15",
created_at=created_at,
),
# 同用户但不同场景/环境/非业务代码位,均不应进入本次详情筛选。
AdEcpmRecord(
user_id=uid,
ad_type="draw",
feed_scene="comparison",
ad_session_id="reward-stats-other-scene",
app_env="prod",
our_code_id="104098712",
ecpm_raw="7000",
report_date="2038-01-15",
created_at=created_at,
),
AdEcpmRecord(
user_id=uid,
ad_type="draw",
feed_scene="coupon",
ad_session_id="reward-stats-test-env",
app_env="test",
our_code_id="104127529",
ecpm_raw="8000",
report_date="2038-01-15",
created_at=created_at,
),
AdEcpmRecord(
user_id=uid,
ad_type="draw",
feed_scene="coupon",
ad_session_id="reward-stats-non-business",
app_env="prod",
our_code_id="demo-slot",
ecpm_raw="9000",
report_date="2038-01-15",
created_at=created_at,
),
]
)
db.commit()
finally:
db.close()
response = admin_client.get(
f"/admin/api/users/{uid}/reward-stats",
params={
"date_from": "2038-01-15T00:00:00Z",
"date_to": "2038-01-15T23:59:59Z",
"app_env": "prod",
"revenue_scope": "business",
"feed_scene": "coupon",
},
headers=_auth(admin_token),
)
assert response.status_code == 200, response.text
data = response.json()
assert data["feed_count"] == 1
# 全部真实展示 (1000 + 3000) / 2;不能返回成功发奖记录的 9000。
assert data["feed_avg_ecpm"] == 2000.0
def test_user_coin_record_sort_accepts_mixed_timezone_datetimes() -> None:
"""线上 PostgreSQL 返回 awareSQLite/历史转换可能返回 naive,二者必须可混排。"""
naive = datetime(2038, 1, 1, 8, 0)
+7 -2
View File
@@ -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",
}
+38 -18
View File
@@ -18,7 +18,6 @@ from sqlalchemy import select
from app.db.session import SessionLocal
from app.models.comparison import ComparisonRecord
from app.repositories import comparison as crud
from app.repositories.user import get_user_by_phone
from app.schemas.compare_record import ComparisonRecordIn
@@ -47,17 +46,6 @@ def _get(db, trace_id: str) -> ComparisonRecord | None:
).scalar_one_or_none()
def _make_user(client, phone: str) -> int:
"""登录建号并返回其真实 user_id。
PG 强制 comparison_record.user_id user.id 外键,须引用真实存在的用户;
用本用例自己登录出的用户,不会与别的用例撞
"""
client.post("/api/v1/auth/sms/login", json={"phone": phone, "code": "123456"})
with SessionLocal() as db:
return get_user_by_phone(db, phone).id
# ============================================================
# repo 层
# ============================================================
@@ -114,6 +102,7 @@ def test_harvest_done_derives_and_newly_success_once(client) -> None:
assert rec.is_source_best is False
assert rec.store_name == "测试店"
assert rec.information == "美团更便宜"
assert rec.fail_reason is None # 成功记录不派生失败原因
assert rec.items == [{"name": "肥牛饭", "qty": 1}]
assert rec.trace_url.endswith("/done/")
# 再来一次(重试 done)→ 已 success,newly_success=False(发奖不重复触发)
@@ -122,6 +111,35 @@ def test_harvest_done_derives_and_newly_success_once(client) -> None:
assert newly2 is False
def test_harvest_done_failed_derives_fail_reason(client) -> None:
"""failed 记录:记录级 information 笼统,但 fail_reason 从 platform_results 救出具体原因
(id 3030 :美团系统失败 + 京东 items_not_found 展示京东那条)"""
tid = _tid()
done_failed = {
"comparison_results": [
{"platform_id": "taobao_flash", "platform_name": "淘宝闪购",
"package": "com.taobao.taobao", "price": 23.04, "is_source": True, "rank": 1,
"items": [{"name": "肥牛饭", "qty": 1}]},
],
"platform_results": {
"taobao_flash": {"is_source": True, "status": "source", "price": 23.04},
"meituan_waimai": {"is_source": False, "status": "failed",
"reason": "搜索店铺失败, 无法跳转到搜索页"},
"jd_waimai_standalone": {"is_source": False, "status": "items_not_found",
"reason": "京东外卖此店内未找到这些菜品"},
},
"information": "比价过程出错,请稍后重试",
}
with SessionLocal() as db:
crud.harvest_running(db, trace_id=tid, user_id=None)
rec, newly = crud.harvest_done(db, trace_id=tid, user_id=None,
done_params=done_failed)
assert newly is False # 没落成 success
assert rec.status == "failed"
assert rec.fail_reason == "京东外卖此店内未找到这些菜品"
assert rec.information == "比价过程出错,请稍后重试" # 原文案仍留存
def test_harvest_abort_cancels_running(client) -> None:
tid = _tid()
with SessionLocal() as db:
@@ -155,18 +173,17 @@ def test_harvest_abort_missing_row_returns_none(client) -> None:
def test_upsert_record_no_downgrade_after_harvest_success(client) -> None:
"""harvest 落 success 后,老客户端 fromFailure 的 cancelled 上报不许把它盖回去。"""
tid = _tid()
# PG 强制 comparison_record.user_id → user.id 外键(SQLite 不强制,老写法用合成 id
# 987654)。用本用例自己登录出的真实用户,既满足外键、又不与别的用例撞。
uid = _make_user(client, "13800007701")
with SessionLocal() as db:
crud.harvest_done(db, trace_id=tid, user_id=None, done_params=_done_params())
payload = ComparisonRecordIn(
trace_id=tid, business_type="food", status="cancelled",
information="用户终止", comparison_results=[],
)
rec = crud.upsert_record(db, user_id=uid, payload=payload)
# 用一个不会与顺序自增用户撞的合成 id(SQLite 测试库 FK 不强制;别用小整数,
# 否则会撞上别的测试 login 出来的真实 user_id → 记录混进那个用户的列表)。
rec = crud.upsert_record(db, user_id=987654, payload=payload)
assert rec.status == "success" # 不降级
assert rec.user_id == uid # 但补上了 user_id(原为 None)
assert rec.user_id == 987654 # 但补上了 user_id(原为 None)
# ============================================================
@@ -246,7 +263,9 @@ def test_trace_finalize_harvests_abort(client) -> None:
with SessionLocal() as db: # 先有 running 行(帧0建的)
crud.harvest_running(db, trace_id=tid, user_id=None)
p, _cap = _mock_pricebot({"trace_url": "https://price.shaguabijia.com/traces/fin/"})
with p:
with p, patch(
"app.api.v1.compare.backfill_comparison_llm_cost"
) as backfill:
r = client.post("/api/v1/trace/finalize",
json={"trace_id": tid, "status": "cancelled", "reason": "用户终止"})
assert r.status_code == 200
@@ -254,6 +273,7 @@ def test_trace_finalize_harvests_abort(client) -> None:
rec = _get(db, tid)
assert rec is not None and rec.status == "cancelled"
assert rec.trace_url.endswith("/fin/")
backfill.assert_called_once_with(rec.id, tid)
def test_price_step_binds_user_when_authed(client) -> None:
+4
View File
@@ -72,6 +72,7 @@ def test_backfill_retries_then_persists_cost(monkeypatch):
def test_repair_batch_only_targets_terminal_missing_rows(monkeypatch):
missing_id = _record("llm-repair-missing")
cancelled_id = _record("llm-repair-cancelled", status="cancelled")
running_id = _record("llm-repair-running", status="running")
calls = [
{
@@ -96,12 +97,15 @@ def test_repair_batch_only_targets_terminal_missing_rows(monkeypatch):
)
assert result["repaired"] >= 1
assert "llm-repair-missing" in seen
assert "llm-repair-cancelled" in seen
assert "llm-repair-running" not in seen
with SessionLocal() as db:
assert db.get(ComparisonRecord, missing_id).llm_cost_yuan is not None
assert db.get(ComparisonRecord, cancelled_id).llm_cost_yuan is not None
assert db.get(ComparisonRecord, running_id).llm_cost_yuan is None
finally:
_delete(missing_id)
_delete(cancelled_id)
_delete(running_id)
-129
View File
@@ -1,129 +0,0 @@
"""scripts/ensure_pg.py 纯函数单测(不需要 Docker/PG)。"""
from __future__ import annotations
import socket
from scripts.ensure_pg import (
_docker_desktop_cmd,
_docker_desktop_from_registry,
_is_sqlite,
_parse_host_port,
_port_open,
_win_docker_desktop_candidates,
ensure,
)
def test_is_sqlite():
assert _is_sqlite("sqlite:///./data/app.db")
assert _is_sqlite(" SQLite:///x ")
assert not _is_sqlite("postgresql+psycopg://u:p@localhost:5432/db")
def test_parse_host_port_full():
assert _parse_host_port(
"postgresql+psycopg://u:p@localhost:5432/shaguabijia"
) == ("localhost", 5432)
def test_parse_host_port_defaults():
# 缺端口 → 5432
assert _parse_host_port("postgresql+psycopg://u:p@db.example/x")[1] == 5432
# 缺 host → localhost
assert _parse_host_port("postgresql+psycopg:///x") == ("localhost", 5432)
def test_parse_host_port_testdb():
assert _parse_host_port(
"postgresql+psycopg://u:p@localhost:5432/shaguabijia_test"
) == ("localhost", 5432)
def test_port_open_true():
srv = socket.socket()
srv.bind(("127.0.0.1", 0))
srv.listen(1)
port = srv.getsockname()[1]
try:
assert _port_open("127.0.0.1", port, timeout=1.0)
finally:
srv.close()
def test_port_open_false():
s = socket.socket()
s.bind(("127.0.0.1", 0))
port = s.getsockname()[1]
s.close() # 释放端口,无人监听 → 连接应失败
assert not _port_open("127.0.0.1", port, timeout=0.3)
def test_win_candidates_override_wins():
env = {
"DOCKER_DESKTOP_EXE": r"X:\custom\Docker Desktop.exe",
"ProgramFiles": r"C:\Program Files",
}
cands = [str(p) for p in _win_docker_desktop_candidates(env, None)]
assert cands[0] == r"X:\custom\Docker Desktop.exe"
def test_win_candidates_follow_cli_drive():
# 核心:docker CLI 在 D 盘 → 反推出 D 盘的 Docker Desktop.exe(不再写死 C 盘)
env = {"ProgramFiles": r"C:\Program Files"}
cli = r"D:\Docker\Docker\resources\bin\docker.exe"
cands = [str(p) for p in _win_docker_desktop_candidates(env, cli)]
assert r"D:\Docker\Docker\Docker Desktop.exe" in cands
# 兜底的 C 盘常见路径也仍在
assert r"C:\Program Files\Docker\Docker\Docker Desktop.exe" in cands
def test_win_candidates_no_cli_uses_program_files():
env = {"ProgramFiles": r"C:\Program Files"}
cands = [str(p) for p in _win_docker_desktop_candidates(env, None)]
assert cands == [r"C:\Program Files\Docker\Docker\Docker Desktop.exe"]
def test_docker_desktop_cmd_windows_found(monkeypatch, tmp_path):
exe = tmp_path / "Docker Desktop.exe"
exe.write_text("") # 真实存在
monkeypatch.setattr("scripts.ensure_pg._find_docker_desktop_exe", lambda: exe)
assert _docker_desktop_cmd("win32") == [str(exe)]
def test_docker_desktop_cmd_windows_not_found(monkeypatch):
monkeypatch.setattr("scripts.ensure_pg._find_docker_desktop_exe", lambda: None)
assert _docker_desktop_cmd("win32") is None
def test_docker_desktop_cmd_darwin():
assert _docker_desktop_cmd("darwin") == ["open", "-a", "Docker"]
def test_docker_desktop_cmd_linux():
assert _docker_desktop_cmd("linux") is None
def test_registry_probe_never_raises():
# best-effort:无论平台/有无键,只返回 Path 或 None,绝不抛
r = _docker_desktop_from_registry()
assert r is None or hasattr(r, "exists")
def test_ensure_sqlite_escape_hatch(monkeypatch):
# sqlite 是【显式降级逃生舱】:打印横幅、返回 True,且绝不触碰 docker
def _boom():
raise AssertionError("sqlite 分支不应调用 docker")
monkeypatch.setattr("scripts.ensure_pg._docker_cli_ok", _boom)
assert ensure("sqlite:///./data/app.db") is True
def test_ensure_shortcircuits_when_pg_up(monkeypatch):
# 端口通 → 直接 True,绝不触碰 docker
monkeypatch.setattr("scripts.ensure_pg._port_open", lambda *a, **k: True)
def _boom():
raise AssertionError("端口通时不应调用 docker")
monkeypatch.setattr("scripts.ensure_pg._docker_cli_ok", _boom)
assert ensure("postgresql+psycopg://u:p@localhost:5432/shaguabijia") is True
+122
View File
@@ -0,0 +1,122 @@
"""失败卡展示原因派生(repositories.comparison._derive_fail_display)单元测试。
用例取自线上真实 failed 记录(platform_results 形态),覆盖:
- information 具体 直出
- information 笼统 + platform_results 有干净业务结局 救援出该原因(id 3030/2964 )
- information 笼统 + 仅系统失败(搜索失败等黑话) None(端侧品牌兜底,id 3027 )
- store_closed / no_delivery pricebot 漏成 status=failed reason 关键字补判
- 打烊类脏店名 blob 统一简短模板
- platform_results 为空 / 非对象 None
"""
from __future__ import annotations
from app.repositories import comparison as crud
def test_specific_information_passthrough() -> None:
# information 本身具体(未达起送/找不到菜等)→ 直出,不看 platform_results
assert (
crud._derive_fail_display("淘宝闪购未达起送门槛,可加菜凑单后下单", {})
== "淘宝闪购未达起送门槛,可加菜凑单后下单"
)
assert crud._derive_fail_display("未识别到商品", {}) == "未识别到商品"
def test_generic_info_rescued_from_items_not_found() -> None:
# id 3030 型:美团系统失败 + 京东 items_not_found,记录级 information 笼统 → 救出京东那条
pr = {
"eleme": {"is_source": True, "status": "source", "price": 23.04},
"meituan_waimai": {
"is_source": False, "status": "failed",
"reason": "搜索店铺失败, 无法跳转到搜索页",
},
"jd_waimai_standalone": {
"is_source": False, "status": "items_not_found",
"reason": "京东外卖此店内未找到这些菜品",
},
}
assert (
crud._derive_fail_display("比价过程出错,请稍后重试", pr)
== "京东外卖此店内未找到这些菜品"
)
def test_generic_info_rescued_from_store_not_found() -> None:
# id 2964 型:美团系统失败 + 京东 store_not_found → 救出京东相似店铺文案
pr = {
"taobao_flash": {"is_source": True, "status": "source", "price": 127.98},
"meituan": {
"is_source": False, "status": "failed",
"reason": "比价过程出错,请稍后重试",
},
"jd_waimai": {
"is_source": False, "status": "store_not_found",
"reason": "未在京东找到「黔珍味·贵州牛肉蘸水健康菜 (望京店)」相似店铺",
},
}
assert (
crud._derive_fail_display("比价过程出错,请稍后重试", pr)
== "未在京东找到「黔珍味·贵州牛肉蘸水健康菜 (望京店)」相似店铺"
)
def test_generic_info_pure_system_failure_returns_none() -> None:
# id 3027 型:唯一目标平台是自动化黑话失败 → 不给用户看 → None(端侧品牌兜底)
pr = {
"eleme": {"is_source": True, "status": "source", "price": 18.83},
"meituan_waimai": {
"is_source": False, "status": "failed",
"reason": "搜索店铺失败, 无法跳转到搜索页",
},
}
assert crud._derive_fail_display("比价过程出错,请稍后重试", pr) is None
def test_store_closed_leaked_to_failed_is_rescued_and_cleaned() -> None:
# 打烊被漏成 status=failed;reason 常带脏店名 blob → 统一简短模板
pr = {
"jd_waimai": {"is_source": True, "status": "source", "price": 25},
"taobao_flash": {
"is_source": False, "status": "failed",
"reason": "淘宝闪购「沙胆彪炭炉牛杂煲(...),蜂鸟准时达,月售300+,起送¥20」本店已休息,无法比价",
},
}
assert crud._derive_fail_display("比价过程出错,请稍后重试", pr) == "门店休息中,无法比价"
pr2 = {
"taobao_flash": {"is_source": True, "status": "source", "price": 31.83},
"meituan": {
"is_source": False, "status": "failed",
"reason": "美团「奈雪的茶(北京王府井奥莱·香江」门店已打烊,无法比价",
},
}
assert crud._derive_fail_display("比价出错", pr2) == "门店已打烊,无法比价"
def test_no_delivery_leaked_to_failed_is_rescued() -> None:
# 单点不配送被漏成 status=failed;reason 本身干净 → 直接用
reason = "京东外卖该商家所选商品单点不配送,无法进入结算比价"
pr = {
"taobao_flash": {"is_source": True, "status": "source", "price": 20.1},
"jd_waimai_standalone": {
"is_source": False, "status": "failed", "reason": reason,
},
}
assert crud._derive_fail_display("比价过程出错,请稍后重试", pr) == reason
def test_empty_or_missing_platform_results_returns_none() -> None:
# 「比价出错」+ 空 {} / None / 非对象:引擎早夭,无可展示原因 → None
assert crud._derive_fail_display("比价出错", {}) is None
assert crud._derive_fail_display("比价过程出错,请稍后重试", None) is None
assert crud._derive_fail_display("比价出错", []) is None # 老 array 形态,防御
def test_clean_status_wins_over_priority_order() -> None:
# 多个业务结局同现时按 _BIZ_STATUS_PRIORITY 选(below_minimum 优先于 store_not_found)
pr = {
"src": {"is_source": True, "status": "source", "price": 30},
"a": {"is_source": False, "status": "store_not_found", "reason": "未找到店铺A"},
"b": {"is_source": False, "status": "below_minimum", "reason": "B未达起送门槛"},
}
assert crud._derive_fail_display("比价过程出错,请稍后重试", pr) == "B未达起送门槛"
+102 -2
View File
@@ -75,7 +75,7 @@ def _assert_ledger_balanced(user_id: int) -> None:
@pytest.fixture()
def guide_configured():
"""给全局配置塞一支片子(默认 max_plays=3 / reward_coin=120),用完还原成未配片。"""
"""给全局配置塞一支片子(默认 max_plays=3 / reward_coin=100),用完还原成未配片。"""
with SessionLocal() as db:
crud_guide.set_video(db, VIDEO_URL, admin_id=1)
cfg = crud_guide.get_config(db)
@@ -201,7 +201,11 @@ def test_start_loses_seq_race_degrades_to_ad(client, guide_configured, monkeypat
assert first["should_play"] is True and first["seq"] == 1
# 本次请求读到的是过期计数 → 仍会算出 seq=1
monkeypatch.setattr(crud_guide, "used_plays", lambda db, user_id: 0)
monkeypatch.setattr(
crud_guide,
"used_plays",
lambda db, user_id, *, scene="coupon", reset_at=None: 0,
)
with SessionLocal() as db:
result = crud_guide.start_play(db, uid)
@@ -261,3 +265,99 @@ def test_reward_loses_race_does_not_double_mint(client, guide_configured) -> Non
assert _coin_balance(uid) == coin, "一次播放只能发一次币"
assert len(_guide_txns(uid)) == 1, "一个 play_token 只能有一条金币流水"
_assert_ledger_balanced(uid)
def test_coupon_and_comparison_configs_and_counts_are_independent(client) -> None:
phone = "13920000008"
token = _login(client, phone)
uid = _user_id(phone)
comparison_url = "/media/guide_video/pytest_comparison.mp4"
with SessionLocal() as db:
coupon_before = crud_guide.get_config(db, "coupon")
comparison_before = crud_guide.get_config(db, "comparison")
crud_guide.set_video(db, VIDEO_URL, scene="coupon", admin_id=1)
crud_guide.update_config(
db,
scene="comparison",
enabled=True,
max_plays=1,
reward_coin=321,
admin_id=1,
)
crud_guide.set_video(
db,
comparison_url,
scene="comparison",
admin_id=1,
)
try:
coupon = client.post(
"/api/v1/guide-video/start",
json={"scene": "coupon"},
headers=_auth(token),
).json()
comparison = client.post(
"/api/v1/guide-video/start",
json={"scene": "comparison"},
headers=_auth(token),
).json()
comparison_blocked = client.post(
"/api/v1/guide-video/start",
json={"scene": "comparison"},
headers=_auth(token),
).json()
assert coupon["should_play"] is True
assert coupon["video_url"] == VIDEO_URL
assert coupon["seq"] == 1
assert comparison["should_play"] is True
assert comparison["video_url"] == comparison_url
assert comparison["reward_coin"] == 321
assert comparison["seq"] == 1
assert comparison_blocked["should_play"] is False
with SessionLocal() as db:
assert crud_guide.play_stats(db, "coupon")["total_plays"] >= 1
assert crud_guide.play_stats(db, "comparison")["total_plays"] == 1
rows = list(
db.scalars(
select(GuideVideoPlay)
.where(GuideVideoPlay.user_id == uid)
.order_by(GuideVideoPlay.scene)
)
)
assert [(row.scene, row.seq) for row in rows] == [
("comparison", 1),
("coupon", 1),
]
finally:
with SessionLocal() as db:
crud_guide.update_config(
db,
scene="coupon",
enabled=coupon_before["enabled"],
reward_coin=coupon_before["reward_coin"],
admin_id=1,
)
crud_guide.set_video(
db,
coupon_before["video_url"],
scene="coupon",
admin_id=1,
)
crud_guide.update_config(
db,
scene="comparison",
enabled=comparison_before["enabled"],
max_plays=comparison_before["max_plays"],
reward_coin=comparison_before["reward_coin"],
admin_id=1,
)
crud_guide.set_video(
db,
comparison_before["video_url"],
scene="comparison",
admin_id=1,
)
+943
View File
@@ -0,0 +1,943 @@
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 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 (
AD_FEED_DAILY_LIMIT_KEY,
AD_REWARD_VIDEO_DAILY_LIMIT_KEY,
GUIDE_VIDEO_MAX_PLAYS_KEY,
PHONE_REBIND_DAYS_KEY,
SMS_PHONE_COOLDOWN_SECONDS_KEY,
)
from app.core.security import create_token
from app.db.session import SessionLocal
from app.main import app
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 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_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"]
duplicate_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 duplicate_batch.status_code == 409
with SessionLocal() as db:
rolled_back = (
db.query(LimitPolicyOverride)
.filter(
LimitPolicyOverride.subject_type == "phone",
LimitPolicyOverride.subject_value == phone,
LimitPolicyOverride.rule_code == "sms.send.daily",
)
.one_or_none()
)
assert rolled_back is None
for item in created.json():
deleted = client.delete(
f"/admin/api/limit-whitelist/{item['id']}",
headers=admin_headers,
)
assert deleted.status_code == 204
def test_device_bulk_rejects_mixed_candidate_namespaces(admin_headers) -> None:
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": f"mixed-device-{uuid4().hex[:8]}",
"rule_codes": [
"compare.start.daily",
"sms.send.hourly",
],
"expires_at": expires_at,
},
)
assert response.status_code == 400
assert "不能在同一白名单中混选" in response.json()["detail"]
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(
[
GUIDE_VIDEO_MAX_PLAYS_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_ad_limit_update_syncs_split_global_rules(admin_headers) -> None:
snapshot = _snapshot_configs(
["ad_daily_limit", AD_REWARD_VIDEO_DAILY_LIMIT_KEY, AD_FEED_DAILY_LIMIT_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
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(
[SMS_PHONE_COOLDOWN_SECONDS_KEY, PHONE_REBIND_DAYS_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_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
+1 -1
View File
@@ -533,7 +533,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,
},