Files
mindfulness/server/app/api/v1/push.py
2026-03-12 18:02:25 +08:00

419 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from __future__ import annotations
from datetime import datetime, timezone
from typing import Any, Literal, Optional
import httpx
import redis
from fastapi import APIRouter, Depends, Header, HTTPException, Query
from pydantic import BaseModel, Field
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.api.limits import rate_limit_push_by_ip
from app.core.config import get_settings
from app.db.models.push_preference import PushPreference
from app.db.models.push_token import PushToken
from app.db.models.push_send_log import PushSendLog
from app.db.session import get_db
from app.features.push_payload import build_home_push_data
from app.features.user_profile_scoring.types import UserProfileV1_2
from app.worker import celery_app
router = APIRouter(
prefix="/v1/push",
tags=["push"],
dependencies=[Depends(rate_limit_push_by_ip)],
)
PushEnv = Literal["dev", "prod"]
PushPlatform = Literal["ios", "android"]
class PushDeviceMeta(BaseModel):
model: Optional[str] = None
os_version: Optional[str] = None
app_version: Optional[str] = None
locale: Optional[str] = None
timezone: Optional[str] = None
class PushRegisterRequest(BaseModel):
client_user_id: str = Field(min_length=8, max_length=64)
platform: PushPlatform
push_token: str = Field(min_length=8, max_length=255)
app_id: str = Field(min_length=1, max_length=255)
env: PushEnv
device_meta: Optional[PushDeviceMeta] = None
class PushPreferencesRequest(BaseModel):
client_user_id: str = Field(min_length=8, max_length=64)
enabled: bool
times_per_day: int = Field(ge=0, le=5)
timezone: Optional[str] = None
locale: Optional[str] = None
# 可选:用户画像(用于 Push 场景推荐文案生成)
user_profile: Optional[UserProfileV1_2] = None
class PushPreferencesResponse(BaseModel):
client_user_id: str
enabled: bool
times_per_day: int
timezone: Optional[str] = None
locale: Optional[str] = None
updated_at: Optional[str] = None
class PushTestRequest(BaseModel):
client_user_id: str = Field(min_length=8, max_length=64)
title: Optional[str] = None
body: Optional[str] = None
def _ensure_utc(dt: datetime) -> datetime:
if dt.tzinfo is None:
return dt.replace(tzinfo=timezone.utc)
return dt.astimezone(timezone.utc)
async def _pick_active_token(db: AsyncSession, *, client_user_id: str) -> Optional[PushToken]:
# 取最近一次上报的 active token
q = (
select(PushToken)
.where(PushToken.client_user_id == client_user_id, PushToken.is_active == True) # noqa: E712
.order_by(PushToken.last_seen_at.desc())
.limit(1)
)
row = await db.execute(q)
return row.scalar_one_or_none()
async def _send_expo_push(*, to: str, title: str, body: str, data: Optional[dict[str, Any]] = None) -> dict[str, Any]:
"""
发送 Expo Push。
说明:
- V1最小可用实现满足 test 与后续定时任务调用
- 失败处理与 token 停用在后续任务逻辑中完善
"""
settings = get_settings()
url = "https://exp.host/--/api/v2/push/send"
headers: dict[str, str] = {
"Content-Type": "application/json",
}
if settings.expo_access_token:
headers["Authorization"] = f"Bearer {settings.expo_access_token}"
payload: dict[str, Any] = {"to": to, "title": title, "body": body}
if data:
payload["data"] = data
async with httpx.AsyncClient(timeout=10.0) as client:
res = await client.post(url, headers=headers, json=payload)
if res.status_code >= 400:
raise HTTPException(status_code=502, detail=f"expo_push_failed:{res.status_code}")
return res.json()
@router.post("/register")
async def register(req: PushRegisterRequest, db: AsyncSession = Depends(get_db)) -> dict[str, str]:
"""
注册/更新 Push Token幂等
"""
now = datetime.now(timezone.utc)
# 以 env+app_id+push_token 唯一:存在则更新归属与 last_seen
q = select(PushToken).where(
PushToken.env == req.env,
PushToken.app_id == req.app_id,
PushToken.push_token == req.push_token,
)
row = await db.execute(q)
token = row.scalar_one_or_none()
if token is None:
token = PushToken(
client_user_id=req.client_user_id,
platform=req.platform,
push_token=req.push_token,
app_id=req.app_id,
env=req.env,
is_active=True,
last_seen_at=_ensure_utc(now),
)
db.add(token)
else:
token.client_user_id = req.client_user_id
token.platform = req.platform
token.is_active = True
token.last_seen_at = _ensure_utc(now)
# 额外:尽早写入/补齐时区与语言(用于按用户时区生成排程)
# 说明:
# - 用户首次授权后会立即调用 /register但不一定马上进入“每日提醒”确认页
# - 若 push_preferences 里 timezone 为空,会导致排程回退到 UTC体验不符合预期
if req.device_meta:
tz = (req.device_meta.timezone or "").strip() or None
loc = (req.device_meta.locale or "").strip() or None
if tz or loc:
qpref = select(PushPreference).where(PushPreference.client_user_id == req.client_user_id)
rpref = await db.execute(qpref)
pref = rpref.scalar_one_or_none()
if pref is None:
pref = PushPreference(
client_user_id=req.client_user_id,
enabled=False,
times_per_day=0,
timezone=tz,
locale=loc,
)
db.add(pref)
else:
if tz and not (pref.timezone or "").strip():
pref.timezone = tz
if loc and not (pref.locale or "").strip():
pref.locale = loc
await db.commit()
return {"status": "ok"}
@router.put("/preferences", response_model=PushPreferencesResponse)
async def put_preferences(req: PushPreferencesRequest, db: AsyncSession = Depends(get_db)) -> PushPreferencesResponse:
"""
设置每日提醒偏好(幂等)。
"""
# 规范化enabled=false 或 times=0 均视为关闭
times = int(req.times_per_day)
enabled = bool(req.enabled) and times > 0
q = select(PushPreference).where(PushPreference.client_user_id == req.client_user_id)
row = await db.execute(q)
pref = row.scalar_one_or_none()
if pref is None:
pref = PushPreference(
client_user_id=req.client_user_id,
enabled=enabled,
times_per_day=times,
timezone=req.timezone,
locale=req.locale,
# 注意Pydantic 会把 ISO8601 字符串解析成 datetime
# SQLAlchemy JSON 列默认使用 json.dumps无法序列化 datetime。
# 这里用 mode="json" 保证写库内容都是可 JSON 序列化的基础类型datetime → ISO 字符串)。
user_profile_json=req.user_profile.model_dump(mode="json") if req.user_profile else None,
)
db.add(pref)
else:
pref.enabled = enabled
pref.times_per_day = times
# 注意:只在客户端显式传入时覆盖,避免把已保存的 timezone/locale 清空导致排程回退到 UTC
if req.timezone is not None:
pref.timezone = req.timezone
if req.locale is not None:
pref.locale = req.locale
if req.user_profile is not None:
pref.user_profile_json = req.user_profile.model_dump(mode="json")
await db.commit()
# 返回更新时间:
# - 某些运行环境/驱动组合下commit 后访问 ORM 字段可能触发隐式 IO导致 async 下报 MissingGreenlet。
# - 这里直接用当前时间兜底(字段本身为可选,仅用于前端展示)。
updated_at_iso = datetime.now(timezone.utc).isoformat()
return PushPreferencesResponse(
client_user_id=req.client_user_id,
enabled=enabled,
times_per_day=times,
timezone=req.timezone,
locale=req.locale,
updated_at=updated_at_iso,
)
@router.get("/preferences", response_model=PushPreferencesResponse)
async def get_preferences(
client_user_id: str = Query(min_length=8, max_length=64),
db: AsyncSession = Depends(get_db),
) -> PushPreferencesResponse:
q = select(PushPreference).where(PushPreference.client_user_id == client_user_id)
row = await db.execute(q)
pref = row.scalar_one_or_none()
if pref is None:
return PushPreferencesResponse(client_user_id=client_user_id, enabled=False, times_per_day=0)
updated_at_iso = pref.updated_at.isoformat() if pref.updated_at else None
return PushPreferencesResponse(
client_user_id=client_user_id,
enabled=bool(pref.enabled) and int(pref.times_per_day) > 0,
times_per_day=int(pref.times_per_day),
timezone=pref.timezone,
locale=pref.locale,
updated_at=updated_at_iso,
)
@router.post("/test")
async def test_push(
req: PushTestRequest,
db: AsyncSession = Depends(get_db),
accept_language: Optional[str] = Header(default=None, alias="Accept-Language"),
) -> dict[str, Any]:
"""
立即测试推送(仅用于 dev 联调)。
"""
settings = get_settings()
if settings.app_env != "dev":
raise HTTPException(status_code=403, detail="test_only_in_dev")
token = await _pick_active_token(db, client_user_id=req.client_user_id)
if token is None:
raise HTTPException(status_code=404, detail="no_active_token")
# 如果没有用户画像,也不影响 test 推送;文案按请求/默认文案发送。
q = select(PushPreference).where(PushPreference.client_user_id == req.client_user_id)
row = await db.execute(q)
pref = row.scalar_one_or_none()
if pref and pref.user_profile_json:
_ = UserProfileV1_2.model_validate(pref.user_profile_json)
# V1先发固定测试文案后续在定时任务中替换为推荐模块的 push 场景模板
title = req.title or "Dear Mama"
body = req.body or "这是一条测试推送dev"
expo_res = await _send_expo_push(
to=token.push_token,
title=title,
body=body,
data=build_home_push_data(
client_user_id=req.client_user_id,
body=body,
scene="push",
),
)
_ = accept_language
return {"status": "ok", "expo": expo_res}
def _env_prefix(app_env: str) -> str:
"""
根据环境生成前缀:
- dev -> dev
- prod -> pro
"""
return "dev" if str(app_env) == "dev" else "pro"
@router.get("/scheduler/health")
async def scheduler_health(db: AsyncSession = Depends(get_db)) -> dict[str, Any]:
"""
推送“定时服务”健康检查(用于容器内验证)。
返回内容(尽量不暴露敏感信息):
- Redis是否可连通
- Worker是否至少有一个 worker 在线inspect ping
- Beat是否在跑beat 心跳 key 是否在持续刷新)
- DB是否可查询到 push_send_log 的最新时间(辅助定位排程是否生成)
"""
settings = get_settings()
prefix = _env_prefix(settings.app_env)
beat_key = f"{prefix}:beat:heartbeat"
out: dict[str, Any] = {
"env": settings.app_env,
"redis": {"ok": False},
"worker": {"ok": False, "worker_count": 0},
"beat": {"ok": False, "last_heartbeat_at": None, "age_seconds": None},
"db": {"ok": False, "push_send_log_latest_created_at": None},
"db_push_tokens": {"ok": False, "count": None, "latest": None},
"now_utc": datetime.now(timezone.utc).isoformat(),
}
# 1) Redis 连通性 + 读取 beat 心跳
try:
r = redis.Redis.from_url(settings.celery_broker_url, decode_responses=True)
r.ping()
out["redis"]["ok"] = True
hb = r.get(beat_key)
if hb:
out["beat"]["last_heartbeat_at"] = hb
try:
# Python 3.11+ 支持解析 ISO8601含 +00:00
hb_dt = datetime.fromisoformat(hb.replace("Z", "+00:00"))
now = datetime.now(timezone.utc)
age = int((now - hb_dt.astimezone(timezone.utc)).total_seconds())
out["beat"]["age_seconds"] = age
# 2 分钟内认为健康beat 每分钟刷新一次)
out["beat"]["ok"] = age <= 120
except Exception:
# 解析失败:至少说明 key 存在,但时间格式异常
out["beat"]["ok"] = False
except Exception as e:
out["redis"]["error"] = f"{type(e).__name__}: {e}"
# 2) Worker 在线性inspect ping
try:
insp = celery_app.control.inspect(timeout=1.0)
pings = insp.ping() or {}
if isinstance(pings, dict):
out["worker"]["worker_count"] = len(pings)
out["worker"]["ok"] = len(pings) > 0
except Exception as e:
out["worker"]["error"] = f"{type(e).__name__}: {e}"
# 3) DB查询 push_send_log 最新创建时间(用于判断排程是否有生成)
try:
q = select(PushSendLog.created_at).order_by(PushSendLog.created_at.desc()).limit(1)
row = await db.execute(q)
latest = row.scalar_one_or_none()
out["db"]["ok"] = True
out["db"]["push_send_log_latest_created_at"] = latest.isoformat() if latest else None
except Exception as e:
out["db"]["error"] = f"{type(e).__name__}: {e}"
# 4) DB查询 push_tokens 计数与最近一条(用于确认 /v1/push/register 是否真正落库)
try:
qcount = select(func.count()).select_from(PushToken)
rcount = await db.execute(qcount)
cnt = int(rcount.scalar_one() or 0)
qlatest = select(PushToken).order_by(PushToken.last_seen_at.desc()).limit(1)
rlatest = await db.execute(qlatest)
t = rlatest.scalar_one_or_none()
latest_obj = None
if t is not None:
tok = str(t.push_token or "")
masked = tok[:10] + "***" + tok[-6:] if len(tok) > 20 else (tok[:6] + "***" if tok else "")
latest_obj = {
"id": int(getattr(t, "id", 0) or 0),
"client_user_id": str(t.client_user_id),
"env": str(t.env),
"app_id": str(t.app_id),
"platform": str(t.platform),
"is_active": bool(t.is_active),
"last_seen_at": t.last_seen_at.isoformat() if t.last_seen_at else None,
"push_token_masked": masked,
}
out["db_push_tokens"]["ok"] = True
out["db_push_tokens"]["count"] = cnt
out["db_push_tokens"]["latest"] = latest_obj
except Exception as e:
out["db_push_tokens"]["error"] = f"{type(e).__name__}: {e}"
return out