"""邀请归因服务:邀请码生成、分享/扫码归因、小程序码生成、邀请统计。""" 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=&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}