diff --git a/achievement/tasks.py b/achievement/tasks.py index 969708c..c271e18 100644 --- a/achievement/tasks.py +++ b/achievement/tasks.py @@ -24,3 +24,39 @@ def check_achievements(user_id, submission_id): notify_achievements(user_id, records) except Exception as e: logger.exception(f"check_achievements failed: user_id={user_id}, submission_id={submission_id}, error={e}") + + +@dramatiq.actor(**DRAMATIQ_WORKER_ARGS()) +def rescan_achievement(achievement_id): + """新建成就或调低阈值后,补发给已达标的存量用户。 + + 判定只在判题时发生,因此后台改了阈值不会自动补发,必须显式扫一遍。 + """ + from django.db.models import IntegerField + from django.db.models.fields.json import KeyTextTransform + from django.db.models.functions import Cast + + from achievement.models import Achievement, Operator, UserAchievement, UserStat + + try: + achievement = Achievement.objects.get(id=achievement_id, visible=True) + except Achievement.DoesNotExist: + return + + # JSONField 的默认比较是 JSON 值比较,数字会按字符串序比("9" > "50"), + # 必须显式 cast 成整数,否则筛出来的用户是错的 + qs = UserStat.objects.filter(metrics__has_key=achievement.metric).annotate(v=Cast(KeyTextTransform(achievement.metric, "metrics"), IntegerField())) + if achievement.operator == Operator.GTE: + qs = qs.filter(v__gte=achievement.threshold) + else: + qs = qs.filter(v__lte=achievement.threshold) + + already = set(UserAchievement.objects.filter(achievement=achievement).values_list("user_id", flat=True)) + for stat in qs.select_related("user").iterator(): + if stat.user_id in already: + continue + try: + records = checker.unlock(stat.user, [achievement]) + notify_achievements(stat.user_id, records) + except Exception as e: + logger.error(f"rescan_achievement failed for user {stat.user_id}: {e}") diff --git a/achievement/urls/admin.py b/achievement/urls/admin.py new file mode 100644 index 0000000..abe83ae --- /dev/null +++ b/achievement/urls/admin.py @@ -0,0 +1,8 @@ +from django.urls import path + +from achievement.views.admin import AchievementAdminAPI, AchievementMetricAdminAPI + +urlpatterns = [ + path("achievement", AchievementAdminAPI.as_view(), name="achievement_admin_api"), + path("achievement/metrics", AchievementMetricAdminAPI.as_view(), name="achievement_metric_admin_api"), +] diff --git a/achievement/views/admin.py b/achievement/views/admin.py new file mode 100644 index 0000000..88646ff --- /dev/null +++ b/achievement/views/admin.py @@ -0,0 +1,120 @@ +from account.decorators import super_admin_required +from achievement.metrics import METRIC_REGISTRY +from achievement.models import Achievement +from achievement.tasks import rescan_achievement +from utils.api import APIView +from utils.shortcuts import check_is_id + + +class AchievementAdminAPI(APIView): + @super_admin_required + def get(self, request): + achievement_id = request.GET.get("id") + if achievement_id: + try: + achievement = Achievement.objects.get(id=achievement_id) + except Achievement.DoesNotExist: + return self.error("成就不存在") + return self.success(_serialize(achievement)) + return self.success([_serialize(a) for a in Achievement.objects.all()]) + + @super_admin_required + def post(self, request): + data = request.data + error = _validate(data) + if error: + return self.error(error) + achievement = Achievement.objects.create( + name=data["name"], + description=data["description"], + icon=data["icon"], + rarity=data["rarity"], + hidden=data.get("hidden", False), + metric=data["metric"], + operator=data["operator"], + threshold=data["threshold"], + visible=data.get("visible", True), + order=data.get("order", 0), + ) + # 新建成就需要补发给已达标的存量用户 + rescan_achievement.send(achievement.id) + return self.success(_serialize(achievement)) + + @super_admin_required + def put(self, request): + data = request.data + if not check_is_id(data.get("id")): + return self.error("参数错误") + try: + achievement = Achievement.objects.get(id=data["id"]) + except Achievement.DoesNotExist: + return self.error("成就不存在") + error = _validate(data) + if error: + return self.error(error) + + old_threshold = achievement.threshold + old_operator = achievement.operator + for field in ("name", "description", "icon", "rarity", "hidden", "metric", "operator", "threshold", "visible", "order"): + if field in data: + setattr(achievement, field, data[field]) + achievement.save() + + # 条件放宽(gte 调低阈值 / lte 调高阈值 / 换了比较符)时补发 + loosened = ( + achievement.operator != old_operator + or (achievement.operator == "gte" and achievement.threshold < old_threshold) + or (achievement.operator == "lte" and achievement.threshold > old_threshold) + ) + if loosened and achievement.visible: + rescan_achievement.send(achievement.id) + return self.success(_serialize(achievement)) + + @super_admin_required + def delete(self, request): + achievement_id = request.GET.get("id") + if not check_is_id(achievement_id): + return self.error("参数错误") + Achievement.objects.filter(id=achievement_id).delete() + return self.success("删除成功") + + +class AchievementMetricAdminAPI(APIView): + @super_admin_required + def get(self, request): + """供后台指标下拉框使用。这里的列表就是代码里注册了什么。""" + return self.success([{"key": key, "name": m.name, "help_text": m.help_text} for key, m in METRIC_REGISTRY.items()]) + + +def _validate(data): + for field in ("name", "description", "icon", "rarity", "metric", "operator"): + if not data.get(field): + return f"{field} 不能为空" + if data["metric"] not in METRIC_REGISTRY: + return "指标不存在" + if data["operator"] not in ("gte", "lte"): + return "比较符不合法" + if not isinstance(data.get("threshold"), int): + return "阈值必须是整数" + return None + + +def _serialize(a): + return { + "id": a.id, + "name": a.name, + "description": a.description, + "icon": a.icon, + "rarity": a.rarity, + "hidden": a.hidden, + "metric": a.metric, + "metric_name": METRIC_REGISTRY[a.metric].name if a.metric in METRIC_REGISTRY else a.metric, + "operator": a.operator, + "threshold": a.threshold, + "visible": a.visible, + # 后台列表必须显示这个:阈值配错时学生永远拿不到也永远不会来问, + # 这个计数器是唯一的仪表盘 + "unlock_count": a.unlock_count, + "order": a.order, + "create_time": a.create_time, + } diff --git a/oj/urls.py b/oj/urls.py index 1a89038..fa77037 100644 --- a/oj/urls.py +++ b/oj/urls.py @@ -27,4 +27,5 @@ urlpatterns = [ path("api/admin/", include("problemset.urls.admin")), path("api/", include("class_pk.urls.oj")), path("api/", include("achievement.urls.oj")), + path("api/admin/", include("achievement.urls.admin")), ]