181 lines
5.8 KiB
Python
181 lines
5.8 KiB
Python
"""邀请归因服务:邀请码生成、分享/扫码归因、小程序码生成、邀请统计。"""
|
|
from __future__ import annotations
|
|
|
|
import secrets
|
|
import string
|
|
import time
|
|
from datetime import datetime
|
|
|
|
from .. import config
|
|
from ..infrastructure.repositories import Database
|
|
|
|
|
|
def _now() -> str:
|
|
return datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
|
|
|
|
|
def gen_invite_code(length: int = 8) -> str:
|
|
"""生成易读邀请码(排除易混淆字符)。"""
|
|
alphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
|
|
return "".join(secrets.choice(alphabet) for _ in range(length))
|
|
|
|
|
|
async def ensure_invite_code(db: Database, user_id: str) -> str:
|
|
"""确保用户有邀请码,没有则生成。返回邀请码。"""
|
|
user = await db.users.get(user_id)
|
|
if not user:
|
|
return ""
|
|
code = user.get("invite_code", "")
|
|
if code:
|
|
return code
|
|
# 生成唯一邀请码
|
|
for _ in range(10):
|
|
code = gen_invite_code()
|
|
existing = await db.users.get_by_invite_code(code)
|
|
if not existing:
|
|
break
|
|
else:
|
|
code = gen_invite_code(10)
|
|
await db.users.update(user_id, invite_code=code)
|
|
return code
|
|
|
|
|
|
async def resolve_inviter(db: Database, inviter: str) -> dict | None:
|
|
"""解析邀请人:inviter 可以是用户ID或邀请码。"""
|
|
if not inviter:
|
|
return None
|
|
# 先按邀请码查
|
|
user = await db.users.get_by_invite_code(inviter)
|
|
if user:
|
|
return user
|
|
# 再按用户ID查
|
|
user = await db.users.get(inviter)
|
|
if user:
|
|
return user
|
|
return None
|
|
|
|
|
|
async def bind_invitee(db: Database, invitee_id: str, inviter: str, openid: str = "",
|
|
channel: str = "", scene: str = "", env: str = "trial", ip: str = "") -> bool:
|
|
"""新用户绑定邀请人(仅当用户尚未有邀请人时)。返回是否成功绑定。"""
|
|
if not inviter or not invitee_id:
|
|
return False
|
|
invitee = await db.users.get(invitee_id)
|
|
if not invitee:
|
|
return False
|
|
# 已有邀请人则不覆盖
|
|
if invitee.get("inviter_id"):
|
|
return False
|
|
inviter_user = await resolve_inviter(db, inviter)
|
|
if not inviter_user:
|
|
return False
|
|
# 不能邀请自己
|
|
if inviter_user["id"] == invitee_id:
|
|
return False
|
|
await db.users.update(invitee_id, inviter_id=inviter_user["id"], invited_at=_now())
|
|
# 记录归因
|
|
await db.invite_records.add(
|
|
inviter_id=inviter_user["id"],
|
|
invite_code=inviter_user.get("invite_code", ""),
|
|
invitee_id=invitee_id,
|
|
invitee_openid=openid,
|
|
action="register",
|
|
channel=channel,
|
|
scene=scene,
|
|
env=env,
|
|
ip=ip,
|
|
)
|
|
return True
|
|
|
|
|
|
async def record_share(db: Database, user_id: str, channel: str = "share_app",
|
|
scene: str = "", env: str = "trial", ip: str = "") -> None:
|
|
"""记录分享行为。"""
|
|
if not user_id:
|
|
return
|
|
code = await ensure_invite_code(db, user_id)
|
|
await db.invite_records.add(
|
|
inviter_id=user_id,
|
|
invite_code=code,
|
|
action="share",
|
|
channel=channel,
|
|
scene=scene,
|
|
env=env,
|
|
ip=ip,
|
|
)
|
|
# 递增分享次数
|
|
user = await db.users.get(user_id)
|
|
if user:
|
|
await db.users.update(user_id, invite_share_count=int(user.get("invite_share_count", 0)) + 1)
|
|
|
|
|
|
async def record_scan(db: Database, inviter: str, openid: str = "",
|
|
channel: str = "qrcode", scene: str = "", env: str = "trial", ip: str = "") -> None:
|
|
"""记录扫码行为(被邀请人扫码时)。"""
|
|
inviter_user = await resolve_inviter(db, inviter)
|
|
if not inviter_user:
|
|
return
|
|
await db.invite_records.add(
|
|
inviter_id=inviter_user["id"],
|
|
invite_code=inviter_user.get("invite_code", ""),
|
|
invitee_openid=openid,
|
|
action="scan",
|
|
channel=channel,
|
|
scene=scene,
|
|
env=env,
|
|
ip=ip,
|
|
)
|
|
|
|
|
|
async def get_invite_stats(db: Database, user_id: str) -> dict:
|
|
"""获取用户邀请统计。"""
|
|
code = await ensure_invite_code(db, user_id)
|
|
# 邀请的用户列表
|
|
invitees = await db.users.get_invitees(user_id)
|
|
# 分享次数
|
|
user = await db.users.get(user_id)
|
|
share_count = int(user.get("invite_share_count", 0)) if user else 0
|
|
# 扫码次数
|
|
scan_count = await db.invite_records.count_by(inviter_id=user_id, action="scan")
|
|
# 注册成功数
|
|
register_count = len(invitees)
|
|
from ..infrastructure.oss import resolve_url
|
|
return {
|
|
"invite_code": code,
|
|
"share_count": share_count,
|
|
"scan_count": scan_count,
|
|
"register_count": register_count,
|
|
"invitees": [
|
|
{
|
|
"id": u["id"],
|
|
"nickname": u.get("nickname", ""),
|
|
"avatar": resolve_url(u.get("avatar", "")),
|
|
"invited_at": u.get("invited_at", ""),
|
|
}
|
|
for u in invitees
|
|
],
|
|
}
|
|
|
|
|
|
async def get_wxacode(db: Database, user_id: str, scene: str = "home",
|
|
page: str = "pages/index/index", env: str = "trial",
|
|
width: int = 430) -> dict:
|
|
"""生成带邀请参数的小程序码。
|
|
|
|
scene 格式:inviter=<user_id>&scene=<场景>
|
|
体验版和正式版使用不同的 access_token(由 wechat 模块处理)。
|
|
"""
|
|
code = await ensure_invite_code(db, user_id)
|
|
# scene 最长32字符,用邀请码代替用户ID
|
|
scene_str = f"i={code[:8]}&s={scene[:10]}"
|
|
try:
|
|
from ..services import wechat
|
|
img_buffer = await wechat.get_wxacode(
|
|
scene=scene_str,
|
|
page=page,
|
|
env_version=env,
|
|
)
|
|
return {"success": True, "image": img_buffer, "scene": scene_str, "invite_code": code}
|
|
except Exception as exc:
|
|
return {"success": False, "error": str(exc), "invite_code": code}
|