Compare commits

..

10 Commits

Author SHA1 Message Date
a1b51ebb9e update 2025-07-14 21:41:27 +08:00
a9a6b87fef test for asgi 2025-07-14 21:33:03 +08:00
2d3588c755 revert 2025-06-15 20:26:43 +08:00
a2bfc28ac7 test 2025-06-15 20:21:37 +08:00
6aac767641 test 2025-06-15 20:18:24 +08:00
73af9d96b2 test 2025-06-15 20:15:49 +08:00
8a2fa11afc test 2025-06-15 20:12:48 +08:00
3f1c7250bd test 2025-06-15 20:06:50 +08:00
bd0a7f30f8 test 2025-06-15 19:35:11 +08:00
8a043d2ffa test 2025-06-15 19:26:45 +08:00
302 changed files with 5313 additions and 25840 deletions

View File

@@ -1,9 +1,4 @@
venv
.venv
.idea
.git
.DS_Store
__pycache__
*.pyc
.ruff_cache
.pytest_cache

10
.flake8 Normal file
View File

@@ -0,0 +1,10 @@
[flake8]
exclude =
xss_filter.py,
*/migrations/,
*settings.py
*/apps.py
venv/
max-line-length = 180
inline-quotes = "
no-accept-encodings = True

12
.github/issue_template.md vendored Normal file
View File

@@ -0,0 +1,12 @@
在提交issue之前请
- 认真阅读文档 http://docs.onlinejudge.me/#/
- 搜索和查看历史issues
- 安全类问题请不要在 GitHub 上公布,请发送邮件到 `admin@qduoj.com`,根据漏洞危害程度发送红包感谢。
然后提交issue请写清楚下列事项
 - 进行什么操作的时候遇到了什么问题,最好能有复现步骤
 - 错误提示是什么如果看不到错误提示请去data文件夹查看相应log文件。大段的错误提示请包在代码块标记里面。
- 你尝试修复问题的操作
- 页面问题请写清浏览器版本,尽量有截图

View File

@@ -1,32 +0,0 @@
name: Deploy
on:
push:
branches:
- yuetsh
permissions:
contents: read
jobs:
deploy:
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
include:
- name: debian
remote_port: 22
script: /root/OJDeploy/backend.sh
- name: school
remote_port: 8822
script: /root/OJ/backend.sh
steps:
- name: Deploy to ${{ matrix.name }}
uses: appleboy/ssh-action@v1
with:
host: ${{ secrets.HOST }}
port: ${{ matrix.remote_port }}
username: root
key: ${{ secrets.KEY }}
script: sh ${{ matrix.script }}

54
.github/workflows/release.yml vendored Normal file
View File

@@ -0,0 +1,54 @@
name: Release build
on:
push:
tags:
- v**
workflow_dispatch:
jobs:
build:
runs-on: ubuntu-latest
environment: release
permissions:
contents: read
packages: write
steps:
- name: Docker metadata
id: metadata
uses: docker/metadata-action@v5
with:
images: |
registry.cn-hongkong.aliyuncs.com/oj-image/backend
tags: |
type=semver,pattern={{version}}
type=semver,pattern={{major}}.{{minor}}
type=semver,pattern={{major}},enable=${{ !startsWith(github.ref, 'refs/tags/v0.') }}
- name: Set up QEMU
uses: docker/setup-qemu-action@v3
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
- name: Login to GitHub Container Registry
uses: docker/login-action@v3
with:
registry: ghcr.io
username: ${{ github.actor }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Login to Aliyun Container Registry
uses: docker/login-action@v3
with:
registry: registry.cn-hongkong.aliyuncs.com
username: ${{ secrets.ALIYUN_ACR_USERNAME }}
password: ${{ secrets.ALIYUN_ACR_PASSWORD }}
- name: Build and push
uses: docker/build-push-action@v5
with:
push: true
platforms: linux/amd64,linux/arm64
tags: ${{ steps.metadata.outputs.tags }}
annotations: ${{ steps.metadata.outputs.annotations }}

144
CLAUDE.md
View File

@@ -1,144 +0,0 @@
# CLAUDE.md
This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository.
## Project Overview
**OnlineJudge** is the backend for an Online Judge platform. Built with Django 6 + Django REST Framework, PostgreSQL, Redis, Django Channels (WebSocket), and Dramatiq (async task queue). Python 3.12+, managed with `uv`.
## Commands
```bash
# Development
python dev.py # Start dev server: Django on :8000 + Daphne WebSocket on :8001
python manage.py runserver # HTTP only (no WebSocket support)
python manage.py migrate # Apply database migrations
python manage.py makemigrations # Create new migrations
# Dependencies
uv sync # Install dependencies from uv.lock
uv add <package> # Add a dependency
# Testing
python manage.py test # Run all tests
python manage.py test account # Run tests for a single app
python manage.py test account.tests.TestClassName # Run a single test class
# Linting
ruff check . # Lint (E, F, I rules, 180-char line length)
ruff format . # Format (double quotes)
## Testing Policy
Do not write tests
# Initial setup
python manage.py inituser --username admin --password <pw> --action create_super_admin
python manage.py inituser --username admin --password <pw> --action reset
```
## Architecture
### App Modules
Each Django app follows the same structure:
```
<app>/
├── models.py # Django models
├── serializers.py # DRF serializers
├── views/
│ ├── oj.py # User-facing API views
│ └── admin.py # Admin API views
└── urls/
├── oj.py # User-facing URL patterns
└── admin.py # Admin URL patterns
```
Apps: `account`, `problem`, `submission`, `contest`, `ai`, `flowchart`, `problemset`, `class_pk`, `announcement`, `tutorial`, `message`, `reaction`, `conf`, `options`, `judge`
`utils/` is itself a Django app (listed in `INSTALLED_APPS`) — not just a helpers package. It provides `RichTextField` (XSS-sanitized `TextField`), `APIError`, the base `APIView`, caching, WebSocket helpers, and the `inituser` management command. Import shared utilities from `utils.*`.
### URL Routing
All routes are registered in `oj/urls.py`:
- `api/` — user-facing endpoints
- `api/admin/` — admin-only endpoints
WebSocket routing is in `oj/routing.py`.
### Settings Structure
- `oj/settings.py` — base configuration (imports dev or production settings based on `OJ_ENV`)
- `oj/dev_settings.py` — development overrides (imported when `OJ_ENV != "production"`)
- `oj/production_settings.py` — production overrides
### Base APIView & View Patterns
`utils/api/api.py` provides the custom base classes and decorators used by **all** views:
- **`APIView`** — base class for all views (not DRF's `APIView`). Key methods:
- `self.success(data)` — returns `{"error": null, "data": data}`
- `self.error(msg)` — returns `{"error": "error", "data": msg}`
- `self.paginate_data(request, query_set, serializer)` — offset/limit pagination
- `self.invalid_serializer(serializer)` — standard validation error response
- **`CSRFExemptAPIView`** — same as `APIView` but CSRF-exempt
- **`@validate_serializer(SerializerClass)`** — decorator for view methods that validates `request.data` against a serializer before the method runs. On success, `request.data` is replaced with validated data.
Typical view method pattern:
```python
@validate_serializer(CreateProblemSerializer)
@super_admin_required
def post(self, request):
# request.data is already validated
return self.success(...)
```
### Authentication & Permissions
`account/decorators.py` provides decorators used on view methods:
- `@login_required` / `@admin_role_required` / `@super_admin_required`
- `@problem_permission_required`
- `@check_contest_permission(check_type)` — validates contest access, sets `self.contest`
- `ensure_created_by(obj, user)` — helper that raises `APIError` if user doesn't own the object
### Judge System
- `judge/dispatcher.py` — dispatches submissions to the judge sandbox (JudgeServer)
- `judge/tasks.py` — Dramatiq async tasks for judging
- `judge/languages.py` — language configurations (compile/run commands, limits)
Judge status codes are defined in `submission/models.py` (`JudgeStatus` class, codes -2 to 8) and must match the frontend's `utils/constants.ts`.
### Site Configuration (SysOptions)
`options/options.py` provides `SysOptions` — a metaclass-based system for site-wide configuration stored in the database with thread-local caching. Access settings like `SysOptions.smtp_config`, `SysOptions.languages`, etc.
### WebSocket (Channels)
`submission/consumers.py` — WebSocket consumer for real-time submission status updates. Uses `channels-redis` as the channel layer backend. Push updates via `utils/websocket.py:push_submission_update()`.
### Caching
Redis-backed via `django-redis`. Cache keys use MD5 hashing for consistency. See `utils/cache.py`.
### AI Integration
`utils/openai.py` — OpenAI client wrapper configured to work with OpenAI-compatible APIs (e.g., DeepSeek). Used by `ai/` app for submission analysis.
### Data Directory
Test cases and submission outputs are stored in a separate data directory (configured in settings, not in the repo). The `data/` directory in the repo contains configuration templates and `secret.key`.
## Key Domain Concepts
| Concept | Details |
|---|---|
| Problem types | ACM (binary accept/reject) vs OI (partial scoring) |
| Judge statuses | COMPILE_ERROR(-2), WRONG_ANSWER(-1), ACCEPTED(0), CPU_TLE(1), REAL_TLE(2), MLE(3), RE(4), SE(5), PENDING(6), JUDGING(7), PARTIALLY_ACCEPTED(8) |
| User roles | Regular / Admin / Super Admin |
| Contest types | Public vs Password Protected |
| Supported languages | C, C++, Python3, Java, JavaScript, Golang |
## Related Repository
The frontend is at `../ojnext` — a Vue 3 + Rsbuild project. See its CLAUDE.md for frontend details.

View File

@@ -1,42 +1,29 @@
FROM python:3.13-slim
FROM python:3.12.2-alpine
ARG TARGETARCH
ARG TARGETVARIANT
ENV OJ_ENV=production
RUN sed -i 's/dl-cdn.alpinelinux.org/mirrors.ustc.edu.cn/g' /etc/apk/repositories
ENV OJ_ENV production
WORKDIR /app
COPY ./deploy/requirements.txt /app/deploy/
RUN --mount=type=cache,target=/var/cache/apt,sharing=locked,id=apt-cache-$TARGETARCH$TARGETVARIANT-final \
--mount=type=cache,target=/root/.cache/pip,id=pip-cache-$TARGETARCH$TARGETVARIANT-final \
# psycopg2: libpg-dev
# pillow: libjpeg-turbo-dev zlib-dev freetype-dev
RUN --mount=type=cache,target=/etc/apk/cache,id=apk-cahce-$TARGETARCH$TARGETVARIANT-final \
--mount=type=cache,target=/root/.cache/pip,id=pip-cahce-$TARGETARCH$TARGETVARIANT-final \
<<EOS
set -ex
pip config set global.index-url https://pypi.tuna.tsinghua.edu.cn/simple
if [ -f /etc/apt/sources.list.d/debian.sources ]; then
sed -i 's|deb.debian.org|mirrors.tuna.tsinghua.edu.cn|g; s|security.debian.org|mirrors.tuna.tsinghua.edu.cn|g' /etc/apt/sources.list.d/debian.sources
fi
if [ -f /etc/apt/sources.list ]; then
sed -i 's|deb.debian.org|mirrors.tuna.tsinghua.edu.cn|g; s|security.debian.org|mirrors.tuna.tsinghua.edu.cn|g' /etc/apt/sources.list
fi
apt-get update
# libpq / libjpeg 都在 wheel 里自带psycopg_binary.libs、pillow.libs不装系统版。
# zlib 由 base 镜像的 CPython 带着Pillow 用的是它。
apt-get install -y --no-install-recommends \
ca-certificates \
clang-format \
passwd \
supervisor
pip config set global.index-url https://mirrors.ustc.edu.cn/pypi/web/simple
apk add gcc libc-dev python3-dev libpq libpq-dev libjpeg-turbo libjpeg-turbo-dev zlib zlib-dev freetype freetype-dev supervisor openssl nginx curl unzip
pip install -r /app/deploy/requirements.txt
rm -rf /var/lib/apt/lists/*
apk del gcc libc-dev python3-dev libpq-dev libjpeg-turbo-dev zlib-dev freetype-dev
EOS
# Caddy 官方镜像里是静态链接的 Go 二进制,直接拷进 slim 就能跑,不需要额外依赖。
COPY --from=caddy:2-alpine /usr/bin/caddy /usr/bin/caddy
COPY --chmod=755 ./ /app/
COPY ./ /app/
RUN mkdir -p /app/dist/
RUN chmod -R u=rwX,go=rX ./ && chmod +x ./deploy/entrypoint.sh
HEALTHCHECK --interval=30s --timeout=3s --start-period=5s --retries=3 \
CMD python3 /app/deploy/health_check.py
HEALTHCHECK --interval=5s CMD [ "/usr/local/bin/python3", "/app/deploy/health_check.py" ]
EXPOSE 8000
ENTRYPOINT [ "/app/deploy/entrypoint.sh" ]

View File

@@ -1,85 +1,60 @@
import functools
import hashlib
import inspect
import time
from contest.models import Contest, ContestStatus, ContestType
from problem.models import Problem
from utils.api import APIError, JSONResponse
from contest.models import Contest, ContestType, ContestStatus, ContestRuleType
from utils.api import JSONResponse, APIError
from utils.constants import CONTEST_PASSWORD_SESSION_KEY
from .models import ProblemPermission
class BasePermissionDecorator(object):
def __init__(self, func):
self.func = func
functools.update_wrapper(self, func)
def __get__(self, obj, obj_type):
if inspect.iscoroutinefunction(self.func):
return functools.partial(self._async_call, obj)
return functools.partial(self.__call__, obj)
def error(self, data, err="permission-denied"):
return JSONResponse.response({"error": err, "data": data})
def _permission_error(self, request):
if not request.user.is_authenticated:
return self.error("请先登录", err="login-required")
return self.error("权限不足", err="permission-denied")
def error(self, data):
return JSONResponse.response({"error": "permission-denied", "data": data})
def __call__(self, *args, **kwargs):
request = args[1]
self.request = args[1]
if self.check_permission(request):
if request.user.is_disabled:
return self.error("账号已禁用")
if self.check_permission():
if self.request.user.is_disabled:
return self.error("Your account is disabled")
return self.func(*args, **kwargs)
else:
return self._permission_error(request)
return self.error("Please login first")
async def _async_call(self, *args, **kwargs):
request = args[1]
if self.check_permission(request):
if request.user.is_disabled:
return self.error("账号已禁用")
return await self.func(*args, **kwargs)
return self._permission_error(request)
def check_permission(self, request):
def check_permission(self):
raise NotImplementedError()
class login_required(BasePermissionDecorator):
def check_permission(self, request):
return request.user.is_authenticated
def check_permission(self):
return self.request.user.is_authenticated
class super_admin_required(BasePermissionDecorator):
def check_permission(self, request):
user = request.user
def check_permission(self):
user = self.request.user
return user.is_authenticated and user.is_super_admin()
class teacher_admin_required(BasePermissionDecorator):
def check_permission(self, request):
user = request.user
return user.is_authenticated and user.is_teacher_or_above()
class admin_role_required(BasePermissionDecorator):
def check_permission(self, request):
user = request.user
def check_permission(self):
user = self.request.user
return user.is_authenticated and user.is_admin_role()
class problem_permission_required(admin_role_required):
def check_permission(self, request):
if not super().check_permission(request):
def check_permission(self):
if not super(problem_permission_required, self).check_permission():
return False
if request.user.problem_permission == ProblemPermission.NONE:
if self.request.user.problem_permission == ProblemPermission.NONE:
return False
return True
@@ -116,44 +91,48 @@ def check_contest_permission(check_type="details"):
若通过验证在view中可通过self.contest获得该contest
"""
def _get_contest_id(request):
return request.data.get("contest_id") or request.GET.get("contest_id")
def _check_access(self, request, user):
if not user.is_authenticated:
return self.error("请先登录", err="login-required")
if user.is_contest_admin(self.contest):
return None
if self.contest.contest_type == ContestType.PASSWORD_PROTECTED_CONTEST:
if not check_contest_password(request.session.get(CONTEST_PASSWORD_SESSION_KEY, {}).get(str(self.contest.id)), self.contest.password):
return self.error("Wrong password or password expired")
if self.contest.status == ContestStatus.CONTEST_NOT_START and check_type != "details":
return self.error("Contest has not started yet.")
return None
def decorator(func):
@functools.wraps(func)
async def _wrapper(*args, **kwargs):
def _check_permission(*args, **kwargs):
self = args[0]
request = args[1]
contest_id = _get_contest_id(request)
user = request.user
if request.data.get("contest_id"):
contest_id = request.data["contest_id"]
else:
contest_id = request.GET.get("contest_id")
if not contest_id:
return self.error("Parameter error, contest_id is required")
try:
self.contest = await Contest.objects.select_related("created_by").aget(id=contest_id, visible=True)
# use self.contest to avoid query contest again in view.
self.contest = Contest.objects.select_related("created_by").get(id=contest_id, visible=True)
except Contest.DoesNotExist:
return self.error("Contest %s doesn't exist" % contest_id)
error = _check_access(self, request, request.user)
if error:
return error
return await func(*args, **kwargs)
return _wrapper
# Anonymous
if not user.is_authenticated:
return self.error("Please login first.")
# creator or owner
if user.is_contest_admin(self.contest):
return func(*args, **kwargs)
if self.contest.contest_type == ContestType.PASSWORD_PROTECTED_CONTEST:
# password error
if not check_contest_password(request.session.get(CONTEST_PASSWORD_SESSION_KEY, {}).get(self.contest.id), self.contest.password):
return self.error("Wrong password or password expired")
# regular user get contest problems, ranks etc. before contest started
if self.contest.status == ContestStatus.CONTEST_NOT_START and check_type != "details":
return self.error("Contest has not started yet.")
# check does user have permission to get ranks, submissions in OI Contest
if self.contest.status == ContestStatus.CONTEST_UNDERWAY and self.contest.rule_type == ContestRuleType.OI:
if not self.contest.real_time_rank and (check_type == "ranks" or check_type == "submissions"):
return self.error(f"No permission to get {check_type}")
return func(*args, **kwargs)
return _check_permission
return decorator

View File

@@ -1,60 +0,0 @@
from django.core.management.base import BaseCommand
from account.models import UserProfile
from problem.models import Problem
from submission.models import JudgeStatus
ACCEPTED_STATUSES = {JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED}
class Command(BaseCommand):
help = "从用户 Profile 中移除已被删除的题目记录,并同步修正 accepted_number"
def add_arguments(self, parser):
parser.add_argument("--dry-run", action="store_true", help="只检查,不写入数据库")
def handle(self, *args, **options):
dry_run = options["dry_run"]
# 所有现存非比赛题目的 PK 集合
existing_ids = set(Problem.objects.filter(contest__isnull=True).values_list("id", flat=True))
self.stdout.write(f"现存题库题目数: {len(existing_ids)}")
profiles = UserProfile.objects.select_related("user").exclude(acm_problems_status={})
total = profiles.count()
self.stdout.write(f"检查用户数: {total}{'dry-run 模式)' if dry_run else ''}")
fixed_count = 0
for profile in profiles:
problems = profile.acm_problems_status.get("problems", {})
if not problems:
continue
stale_keys = [k for k in problems if int(k) not in existing_ids]
if not stale_keys:
continue
removed_accepted = sum(1 for k in stale_keys if problems[k].get("status") in ACCEPTED_STATUSES)
stale_display = [problems[k].get("_id", k) for k in stale_keys]
self.stdout.write(
f" 用户 {profile.user.username} | 删除 {len(stale_keys)} 题: {', '.join(stale_display)}{f' | 其中已AC {removed_accepted}' if removed_accepted else ''}"
)
if dry_run:
continue
for k in stale_keys:
del profile.acm_problems_status["problems"][k]
if removed_accepted:
# 防止 accepted_number 变为负数
profile.accepted_number = max(0, profile.accepted_number - removed_accepted)
profile.save(update_fields=["acm_problems_status", "accepted_number"])
fixed_count += 1
if dry_run:
self.stdout.write(self.style.WARNING("dry-run 完成,未写入任何数据"))
else:
self.stdout.write(self.style.SUCCESS(f"完成,共修复 {fixed_count} 个用户 Profile"))

View File

@@ -1,10 +1,10 @@
from django.conf import settings
from django.db import connection
from django.utils.deprecation import MiddlewareMixin
from django.utils.timezone import now
from django.utils.deprecation import MiddlewareMixin
from account.models import User
from utils.api import JSONResponse
from account.models import User
class APITokenAuthMiddleware(MiddlewareMixin):
@@ -37,10 +37,8 @@ class AdminRoleRequiredMiddleware(MiddlewareMixin):
def process_request(self, request):
path = request.path_info
if path.startswith("/admin/") or path.startswith("/api/admin/"):
if not request.user.is_authenticated:
return JSONResponse.response({"error": "login-required", "data": "请先登录"})
if not request.user.is_admin_role():
return JSONResponse.response({"error": "permission-denied", "data": "权限不足"})
if not (request.user.is_authenticated and request.user.is_admin_role()):
return JSONResponse.response({"error": "login-required", "data": "Please login in first"})
class LogSqlMiddleware(MiddlewareMixin):

View File

@@ -1,18 +0,0 @@
# Generated by Django 5.2.3 on 2025-09-19 06:11
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('account', '0001_initial'),
]
operations = [
migrations.AddField(
model_name='userprofile',
name='class_name',
field=models.TextField(null=True),
),
]

View File

@@ -1,22 +0,0 @@
# Generated by Django 5.2.3 on 2025-09-19 06:14
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('account', '0002_userprofile_class_name'),
]
operations = [
migrations.RemoveField(
model_name='userprofile',
name='class_name',
),
migrations.AddField(
model_name='user',
name='class_name',
field=models.TextField(null=True),
),
]

View File

@@ -1,22 +0,0 @@
# Generated by Django 6.0.4 on 2026-05-09 08:18
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
("account", "0003_remove_userprofile_class_name_user_class_name"),
]
operations = [
migrations.AlterField(
model_name="user",
name="admin_type",
field=models.TextField(choices=[("Regular User", "Regular User"), ("Admin", "Admin"), ("Super Admin", "Super Admin")], default="Regular User"),
),
migrations.AlterField(
model_name="user",
name="problem_permission",
field=models.TextField(choices=[("None", "None"), ("Own", "Own"), ("All", "All")], default="None"),
),
]

View File

@@ -1,58 +0,0 @@
# Generated by Django 6.0.4 on 2026-05-09 11:53
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('account', '0004_alter_user_admin_type_alter_user_problem_permission'),
]
operations = [
migrations.AlterField(
model_name='user',
name='is_disabled',
field=models.BooleanField(db_default=False, default=False),
),
migrations.AlterField(
model_name='user',
name='open_api',
field=models.BooleanField(db_default=False, default=False),
),
migrations.AlterField(
model_name='user',
name='session_keys',
field=models.JSONField(db_default=models.Value([], output_field=models.JSONField()), default=list),
),
migrations.AlterField(
model_name='user',
name='two_factor_auth',
field=models.BooleanField(db_default=False, default=False),
),
migrations.AlterField(
model_name='userprofile',
name='accepted_number',
field=models.IntegerField(db_default=0, default=0),
),
migrations.AlterField(
model_name='userprofile',
name='acm_problems_status',
field=models.JSONField(db_default=models.Value({}, output_field=models.JSONField()), default=dict),
),
migrations.AlterField(
model_name='userprofile',
name='oi_problems_status',
field=models.JSONField(db_default=models.Value({}, output_field=models.JSONField()), default=dict),
),
migrations.AlterField(
model_name='userprofile',
name='submission_number',
field=models.IntegerField(db_default=0, default=0),
),
migrations.AlterField(
model_name='userprofile',
name='total_score',
field=models.BigIntegerField(db_default=0, default=0),
),
]

View File

@@ -1,18 +0,0 @@
# Generated by Django 6.0.4 on 2026-06-03 00:08
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('account', '0005_alter_user_is_disabled_alter_user_open_api_and_more'),
]
operations = [
migrations.AlterField(
model_name='user',
name='admin_type',
field=models.TextField(choices=[('Regular User', 'Regular User'), ('Student Admin', 'Student Admin'), ('Teacher Admin', 'Teacher Admin'), ('Super Admin', 'Super Admin')], default='Regular User'),
),
]

View File

@@ -1,27 +0,0 @@
# Generated by Django 6.0.4 on 2026-06-03 00:08
from django.db import migrations
def rename_admin_to_student_admin(apps, schema_editor):
User = apps.get_model("account", "User")
User.objects.filter(admin_type="Admin").update(admin_type="Student Admin")
def rename_student_admin_to_admin(apps, schema_editor):
User = apps.get_model("account", "User")
User.objects.filter(admin_type="Student Admin").update(admin_type="Admin")
class Migration(migrations.Migration):
dependencies = [
("account", "0006_alter_user_admin_type"),
]
operations = [
migrations.RunPython(
rename_admin_to_student_admin,
rename_student_admin_to_admin,
),
]

View File

@@ -1,21 +0,0 @@
# Generated by Django 6.0.4 on 2026-07-04 17:32
from django.db import migrations
class Migration(migrations.Migration):
dependencies = [
('account', '0007_rename_admin_to_student_admin'),
]
operations = [
migrations.RemoveField(
model_name='userprofile',
name='oi_problems_status',
),
migrations.RemoveField(
model_name='userprofile',
name='total_score',
),
]

View File

@@ -1,21 +0,0 @@
# Generated by Django 6.0.4 on 2026-08-06 04:22
from django.db import migrations
class Migration(migrations.Migration):
dependencies = [
('account', '0008_remove_userprofile_oi_problems_status_and_more'),
]
operations = [
migrations.RemoveField(
model_name='user',
name='tfa_token',
),
migrations.RemoveField(
model_name='user',
name='two_factor_auth',
),
]

View File

@@ -1,21 +0,0 @@
# Generated by Django 6.0.4 on 2026-08-06 05:15
from django.db import migrations
class Migration(migrations.Migration):
dependencies = [
('account', '0009_remove_user_tfa_token_remove_user_two_factor_auth'),
]
operations = [
migrations.RemoveField(
model_name='user',
name='reset_password_token',
),
migrations.RemoveField(
model_name='user',
name='reset_password_token_expire_time',
),
]

View File

@@ -1,21 +1,19 @@
from django.conf import settings
from django.contrib.auth.models import AbstractBaseUser
from django.conf import settings
from django.db import models
from utils.models import JSONField
class AdminType(models.TextChoices):
REGULAR_USER = "Regular User", "Regular User"
STUDENT_ADMIN = "Student Admin", "Student Admin"
TEACHER_ADMIN = "Teacher Admin", "Teacher Admin"
SUPER_ADMIN = "Super Admin", "Super Admin"
class AdminType(object):
REGULAR_USER = "Regular User"
ADMIN = "Admin"
SUPER_ADMIN = "Super Admin"
class ProblemPermission(models.TextChoices):
NONE = "None", "None"
OWN = "Own", "Own"
ALL = "All", "All"
class ProblemPermission(object):
NONE = "None"
OWN = "Own"
ALL = "All"
class UserManager(models.Manager):
@@ -24,59 +22,50 @@ class UserManager(models.Manager):
def get_by_natural_key(self, username):
return self.get(**{f"{self.model.USERNAME_FIELD}__iexact": username})
async def aget_by_natural_key(self, username):
return await self.aget(**{f"{self.model.USERNAME_FIELD}__iexact": username})
class User(AbstractBaseUser):
username = models.TextField(unique=True)
class_name = models.TextField(null=True)
email = models.TextField(null=True)
create_time = models.DateTimeField(auto_now_add=True, null=True)
# One of UserType
admin_type = models.TextField(default=AdminType.REGULAR_USER, choices=AdminType.choices)
problem_permission = models.TextField(default=ProblemPermission.NONE, choices=ProblemPermission.choices)
admin_type = models.TextField(default=AdminType.REGULAR_USER)
problem_permission = models.TextField(default=ProblemPermission.NONE)
reset_password_token = models.TextField(null=True)
reset_password_token_expire_time = models.DateTimeField(null=True)
# SSO auth token
auth_token = models.TextField(null=True)
session_keys = JSONField(default=list, db_default=models.Value([], output_field=models.JSONField()))
two_factor_auth = models.BooleanField(default=False)
tfa_token = models.TextField(null=True)
session_keys = JSONField(default=list)
# open api key
open_api = models.BooleanField(default=False, db_default=False)
open_api = models.BooleanField(default=False)
open_api_appkey = models.TextField(null=True)
is_disabled = models.BooleanField(default=False, db_default=False)
raw_password = models.CharField(max_length=20, null=True, blank=True, verbose_name="明文密码")
is_disabled = models.BooleanField(default=False)
raw_password = models.CharField(
max_length=20, null=True, blank=True, verbose_name="明文密码"
)
USERNAME_FIELD = "username"
REQUIRED_FIELDS = []
objects = UserManager()
def is_regular_user(self):
return self.admin_type == AdminType.REGULAR_USER
def is_student_admin(self):
return self.admin_type == AdminType.STUDENT_ADMIN
def is_teacher_admin(self):
return self.admin_type == AdminType.TEACHER_ADMIN
def is_admin(self):
return self.admin_type == AdminType.ADMIN
def is_super_admin(self):
return self.admin_type == AdminType.SUPER_ADMIN
def is_admin_role(self):
return self.admin_type in [
AdminType.STUDENT_ADMIN,
AdminType.TEACHER_ADMIN,
AdminType.SUPER_ADMIN,
]
def is_teacher_or_above(self):
return self.admin_type in [AdminType.TEACHER_ADMIN, AdminType.SUPER_ADMIN]
return self.admin_type in [AdminType.ADMIN, AdminType.SUPER_ADMIN]
def can_mgmt_all_problem(self):
return self.problem_permission == ProblemPermission.ALL
def is_contest_admin(self, contest):
return self.is_authenticated and (contest.created_by == self or self.admin_type == AdminType.SUPER_ADMIN)
return self.is_authenticated and (
contest.created_by == self or self.admin_type == AdminType.SUPER_ADMIN
)
def set_password(self, raw_password):
super().set_password(raw_password)
@@ -103,7 +92,9 @@ class UserProfile(models.Model):
# }
# }
# }
acm_problems_status = JSONField(default=dict, db_default=models.Value({}, output_field=models.JSONField()))
acm_problems_status = JSONField(default=dict)
# like acm_problems_status, merely add "score" field
oi_problems_status = JSONField(default=dict)
real_name = models.TextField(null=True)
avatar = models.TextField(default=f"{settings.AVATAR_URI_PREFIX}/default.png")
@@ -113,16 +104,25 @@ class UserProfile(models.Model):
school = models.TextField(null=True)
major = models.TextField(null=True)
language = models.TextField(null=True)
accepted_number = models.IntegerField(default=0, db_default=0)
submission_number = models.IntegerField(default=0, db_default=0)
# for ACM
accepted_number = models.IntegerField(default=0)
# for OI
total_score = models.BigIntegerField(default=0)
submission_number = models.IntegerField(default=0)
def add_accepted_problem_number(self):
self.accepted_number = models.F("accepted_number") + 1
self.save(update_fields=["accepted_number"])
self.save()
def add_submission_number(self):
self.submission_number = models.F("submission_number") + 1
self.save(update_fields=["submission_number"])
self.save()
# 计算总分时, 应先减掉上次该题所得分数, 然后再加上本次所得分数
def add_score(self, this_time_score, last_time_score=None):
last_time_score = last_time_score or 0
self.total_score = models.F("total_score") - last_time_score + this_time_score
self.save()
class Meta:
db_table = "user_profile"

View File

@@ -1,6 +1,6 @@
from django import forms
from utils.api import UsernameSerializer, serializers
from utils.api import serializers, UsernameSerializer
from .models import AdminType, ProblemPermission, User, UserProfile
@@ -8,16 +8,45 @@ from .models import AdminType, ProblemPermission, User, UserProfile
class UserLoginSerializer(serializers.Serializer):
username = serializers.CharField()
password = serializers.CharField()
tfa_code = serializers.CharField(required=False, allow_blank=True)
class UsernameOrEmailCheckSerializer(serializers.Serializer):
username = serializers.CharField(required=False)
email = serializers.EmailField(required=False)
class UserRegisterSerializer(serializers.Serializer):
username = serializers.CharField(max_length=32)
password = serializers.CharField(min_length=6)
email = serializers.EmailField(max_length=64)
captcha = serializers.CharField()
class UserChangePasswordSerializer(serializers.Serializer):
old_password = serializers.CharField()
new_password = serializers.CharField(min_length=6)
tfa_code = serializers.CharField(required=False, allow_blank=True)
class UserChangeEmailSerializer(serializers.Serializer):
password = serializers.CharField()
new_email = serializers.EmailField(max_length=64)
tfa_code = serializers.CharField(required=False, allow_blank=True)
class GenerateUserSerializer(serializers.Serializer):
prefix = serializers.CharField(max_length=16, allow_blank=True)
suffix = serializers.CharField(max_length=16, allow_blank=True)
number_from = serializers.IntegerField()
number_to = serializers.IntegerField()
password_length = serializers.IntegerField(max_value=16, default=8)
class ImportUserSerializer(serializers.Serializer):
users = serializers.ListField(child=serializers.ListField(child=serializers.CharField(max_length=64)))
users = serializers.ListField(
child=serializers.ListField(child=serializers.CharField(max_length=64))
)
class UserAdminSerializer(serializers.ModelSerializer):
@@ -34,10 +63,10 @@ class UserAdminSerializer(serializers.ModelSerializer):
"real_name",
"create_time",
"last_login",
"two_factor_auth",
"open_api",
"is_disabled",
"raw_password",
"class_name",
]
def get_real_name(self, obj):
@@ -61,9 +90,9 @@ class UserSerializer(serializers.ModelSerializer):
"problem_permission",
"create_time",
"last_login",
"two_factor_auth",
"open_api",
"is_disabled",
"class_name",
]
@@ -87,13 +116,19 @@ class EditUserSerializer(serializers.Serializer):
id = serializers.IntegerField()
username = serializers.CharField(max_length=32)
real_name = serializers.CharField(max_length=32, allow_blank=True, allow_null=True)
password = serializers.CharField(min_length=6, allow_blank=True, required=False, default=None)
password = serializers.CharField(
min_length=6, allow_blank=True, required=False, default=None
)
email = serializers.EmailField(max_length=64)
admin_type = serializers.ChoiceField(choices=AdminType.choices)
problem_permission = serializers.ChoiceField(choices=ProblemPermission.choices)
admin_type = serializers.ChoiceField(
choices=(AdminType.REGULAR_USER, AdminType.ADMIN, AdminType.SUPER_ADMIN)
)
problem_permission = serializers.ChoiceField(
choices=(ProblemPermission.NONE, ProblemPermission.OWN, ProblemPermission.ALL)
)
open_api = serializers.BooleanField()
two_factor_auth = serializers.BooleanField()
is_disabled = serializers.BooleanField()
class_name = serializers.CharField(required=False, allow_null=True, allow_blank=True)
class EditUserProfileSerializer(serializers.Serializer):
@@ -107,10 +142,33 @@ class EditUserProfileSerializer(serializers.Serializer):
language = serializers.CharField(max_length=32, allow_blank=True, required=False)
class ApplyResetPasswordSerializer(serializers.Serializer):
email = serializers.EmailField()
captcha = serializers.CharField()
class ResetPasswordSerializer(serializers.Serializer):
token = serializers.CharField()
password = serializers.CharField(min_length=6)
captcha = serializers.CharField()
class SSOSerializer(serializers.Serializer):
token = serializers.CharField()
class TwoFactorAuthCodeSerializer(serializers.Serializer):
code = serializers.IntegerField()
class ImageUploadForm(forms.Form):
image = forms.FileField()
class FileUploadForm(forms.Form):
file = forms.FileField()
class RankInfoSerializer(serializers.ModelSerializer):
user = UsernameSerializer()

22
account/tasks.py Normal file
View File

@@ -0,0 +1,22 @@
import logging
import dramatiq
from options.options import SysOptions
from utils.shortcuts import send_email, DRAMATIQ_WORKER_ARGS
logger = logging.getLogger(__name__)
@dramatiq.actor(**DRAMATIQ_WORKER_ARGS(max_retries=3))
def send_email_async(from_name, to_email, to_name, subject, content):
if not SysOptions.smtp_config:
return
try:
send_email(smtp_config=SysOptions.smtp_config,
from_name=from_name,
to_email=to_email,
to_name=to_name,
subject=subject,
content=content)
except Exception as e:
logger.exception(e)

View File

@@ -0,0 +1,31 @@
<div>
<table cellpadding="0" align="center"
style="overflow:hidden;background:#fff;margin:0 auto;text-align:left;position:relative;font-size:14px; font-family:'lucida Grande',Verdana;line-height:1.5;box-shadow:0 0 3px #ccc;border:1px solid #ccc;border-radius:5px;border-collapse:collapse;">
<tbody>
<tr>
<th valign="middle"
style="height:38px;color:#fff; font-size:14px;line-height:38px; font-weight:bold;text-align:left;padding:10px 24px 6px; border-bottom:1px solid #467ec3;background:#518bcb;border-radius:5px 5px 0 0;">
{{ website_name }}</th>
</tr>
<tr>
<td>
<div style="padding:20px 35px 40px;">
<h2 style="font-weight:bold;margin-bottom:5px;font-size:14px;">Hello, {{ username }}:</h2>
<p style="margin-top:20px">
Please click <a href="{{ link }}">{{ link }}</a> to reset your password in 20 minutes.
</p>
<p style="margin-top:20px">
To protect your account, please do not use simple passwords.
</p>
<p style="margin-top:20px">
If you still have any questions, please contract system administrator.
</p>
<p style="margin-left:2em;"></p>
<p style="text-indent:0;text-align:right;">{{ website_name }}</p>
</div>
</td>
</tr>
</tbody>
</table>
</div>

646
account/tests.py Normal file
View File

@@ -0,0 +1,646 @@
import time
from unittest import mock
from datetime import timedelta
from copy import deepcopy
from django.contrib import auth
from django.utils.timezone import now
from otpauth import OtpAuth
from utils.api.tests import APIClient, APITestCase
from utils.shortcuts import rand_str
from options.options import SysOptions
from .models import AdminType, ProblemPermission, User
from utils.constants import ContestRuleType
class PermissionDecoratorTest(APITestCase):
def setUp(self):
self.regular_user = User.objects.create(username="regular_user")
self.admin = User.objects.create(username="admin")
self.super_admin = User.objects.create(username="super_admin")
self.request = mock.MagicMock()
self.request.user.is_authenticated = mock.MagicMock()
def test_login_required(self):
self.request.user.is_authenticated.return_value = False
def test_admin_required(self):
pass
def test_super_admin_required(self):
pass
class DuplicateUserCheckAPITest(APITestCase):
def setUp(self):
user = self.create_user("test", "test123", login=False)
user.email = "test@test.com"
user.save()
self.url = self.reverse("check_username_or_email")
def test_duplicate_username(self):
resp = self.client.post(self.url, data={"username": "test"})
data = resp.data["data"]
self.assertEqual(data["username"], True)
resp = self.client.post(self.url, data={"username": "Test"})
self.assertEqual(resp.data["data"]["username"], True)
def test_ok_username(self):
resp = self.client.post(self.url, data={"username": "test1"})
data = resp.data["data"]
self.assertFalse(data["username"])
def test_duplicate_email(self):
resp = self.client.post(self.url, data={"email": "test@test.com"})
self.assertEqual(resp.data["data"]["email"], True)
resp = self.client.post(self.url, data={"email": "Test@Test.com"})
self.assertTrue(resp.data["data"]["email"])
def test_ok_email(self):
resp = self.client.post(self.url, data={"email": "aa@test.com"})
self.assertFalse(resp.data["data"]["email"])
class TFARequiredCheckAPITest(APITestCase):
def setUp(self):
self.url = self.reverse("tfa_required_check")
self.create_user("test", "test123", login=False)
def test_not_required_tfa(self):
resp = self.client.post(self.url, data={"username": "test"})
self.assertSuccess(resp)
self.assertEqual(resp.data["data"]["result"], False)
def test_required_tfa(self):
user = User.objects.first()
user.two_factor_auth = True
user.save()
resp = self.client.post(self.url, data={"username": "test"})
self.assertEqual(resp.data["data"]["result"], True)
class UserLoginAPITest(APITestCase):
def setUp(self):
self.username = self.password = "test"
self.user = self.create_user(username=self.username, password=self.password, login=False)
self.login_url = self.reverse("user_login_api")
def _set_tfa(self):
self.user.two_factor_auth = True
tfa_token = rand_str(32)
self.user.tfa_token = tfa_token
self.user.save()
return tfa_token
def test_login_with_correct_info(self):
response = self.client.post(self.login_url,
data={"username": self.username, "password": self.password})
self.assertDictEqual(response.data, {"error": None, "data": "Succeeded"})
user = auth.get_user(self.client)
self.assertTrue(user.is_authenticated)
def test_login_with_correct_info_upper_username(self):
resp = self.client.post(self.login_url, data={"username": self.username.upper(), "password": self.password})
self.assertDictEqual(resp.data, {"error": None, "data": "Succeeded"})
user = auth.get_user(self.client)
self.assertTrue(user.is_authenticated)
def test_login_with_wrong_info(self):
response = self.client.post(self.login_url,
data={"username": self.username, "password": "invalid_password"})
self.assertDictEqual(response.data, {"error": "error", "data": "Invalid username or password"})
user = auth.get_user(self.client)
self.assertFalse(user.is_authenticated)
def test_tfa_login(self):
token = self._set_tfa()
code = OtpAuth(token).totp()
if len(str(code)) < 6:
code = (6 - len(str(code))) * "0" + str(code)
response = self.client.post(self.login_url,
data={"username": self.username,
"password": self.password,
"tfa_code": code})
self.assertDictEqual(response.data, {"error": None, "data": "Succeeded"})
user = auth.get_user(self.client)
self.assertTrue(user.is_authenticated)
def test_tfa_login_wrong_code(self):
self._set_tfa()
response = self.client.post(self.login_url,
data={"username": self.username,
"password": self.password,
"tfa_code": "qqqqqq"})
self.assertDictEqual(response.data, {"error": "error", "data": "Invalid two factor verification code"})
user = auth.get_user(self.client)
self.assertFalse(user.is_authenticated)
def test_tfa_login_without_code(self):
self._set_tfa()
response = self.client.post(self.login_url,
data={"username": self.username,
"password": self.password})
self.assertDictEqual(response.data, {"error": "error", "data": "tfa_required"})
user = auth.get_user(self.client)
self.assertFalse(user.is_authenticated)
def test_user_disabled(self):
self.user.is_disabled = True
self.user.save()
resp = self.client.post(self.login_url, data={"username": self.username,
"password": self.password})
self.assertDictEqual(resp.data, {"error": "error", "data": "Your account has been disabled"})
class CaptchaTest(APITestCase):
def _set_captcha(self, session):
captcha = rand_str(4)
session["_django_captcha_key"] = captcha
session["_django_captcha_expires_time"] = int(time.time()) + 30
session.save()
return captcha
class UserRegisterAPITest(CaptchaTest):
def setUp(self):
self.client = APIClient()
self.register_url = self.reverse("user_register_api")
self.captcha = rand_str(4)
self.data = {"username": "test_user", "password": "testuserpassword",
"real_name": "real_name", "email": "test@qduoj.com",
"captcha": self._set_captcha(self.client.session)}
def test_website_config_limit(self):
SysOptions.allow_register = False
resp = self.client.post(self.register_url, data=self.data)
self.assertDictEqual(resp.data, {"error": "error", "data": "Register function has been disabled by admin"})
def test_invalid_captcha(self):
self.data["captcha"] = "****"
response = self.client.post(self.register_url, data=self.data)
self.assertDictEqual(response.data, {"error": "error", "data": "Invalid captcha"})
self.data.pop("captcha")
response = self.client.post(self.register_url, data=self.data)
self.assertTrue(response.data["error"] is not None)
def test_register_with_correct_info(self):
response = self.client.post(self.register_url, data=self.data)
self.assertDictEqual(response.data, {"error": None, "data": "Succeeded"})
def test_username_already_exists(self):
self.test_register_with_correct_info()
self.data["captcha"] = self._set_captcha(self.client.session)
self.data["email"] = "test1@qduoj.com"
response = self.client.post(self.register_url, data=self.data)
self.assertDictEqual(response.data, {"error": "error", "data": "Username already exists"})
def test_email_already_exists(self):
self.test_register_with_correct_info()
self.data["captcha"] = self._set_captcha(self.client.session)
self.data["username"] = "test_user1"
response = self.client.post(self.register_url, data=self.data)
self.assertDictEqual(response.data, {"error": "error", "data": "Email already exists"})
class SessionManagementAPITest(APITestCase):
def setUp(self):
self.create_user("test", "test123")
self.url = self.reverse("session_management_api")
# launch a request to provide session data
login_url = self.reverse("user_login_api")
self.client.post(login_url, data={"username": "test", "password": "test123"})
def test_get_sessions(self):
resp = self.client.get(self.url)
self.assertSuccess(resp)
data = resp.data["data"]
self.assertEqual(len(data), 1)
# def test_delete_session_key(self):
# resp = self.client.delete(self.url + "?session_key=" + self.session_key)
# self.assertSuccess(resp)
def test_delete_session_with_invalid_key(self):
resp = self.client.delete(self.url + "?session_key=aaaaaaaaaa")
self.assertDictEqual(resp.data, {"error": "error", "data": "Invalid session_key"})
class UserProfileAPITest(APITestCase):
def setUp(self):
self.url = self.reverse("user_profile_api")
def test_get_profile_without_login(self):
resp = self.client.get(self.url)
self.assertDictEqual(resp.data, {"error": None, "data": None})
def test_get_profile(self):
self.create_user("test", "test123")
resp = self.client.get(self.url)
self.assertSuccess(resp)
def test_update_profile(self):
self.create_user("test", "test123")
update_data = {"real_name": "zemal", "submission_number": 233, "language": "en-US"}
resp = self.client.put(self.url, data=update_data)
self.assertSuccess(resp)
data = resp.data["data"]
self.assertEqual(data["real_name"], "zemal")
self.assertEqual(data["submission_number"], 0)
self.assertEqual(data["language"], "en-US")
class TwoFactorAuthAPITest(APITestCase):
def setUp(self):
self.url = self.reverse("two_factor_auth_api")
self.create_user("test", "test123")
def _get_tfa_code(self):
user = User.objects.first()
code = OtpAuth(user.tfa_token).totp()
if len(str(code)) < 6:
code = (6 - len(str(code))) * "0" + str(code)
return code
def test_get_image(self):
resp = self.client.get(self.url)
self.assertSuccess(resp)
def test_open_tfa_with_invalid_code(self):
self.test_get_image()
resp = self.client.post(self.url, data={"code": "000000"})
self.assertDictEqual(resp.data, {"error": "error", "data": "Invalid code"})
def test_open_tfa_with_correct_code(self):
self.test_get_image()
code = self._get_tfa_code()
resp = self.client.post(self.url, data={"code": code})
self.assertSuccess(resp)
user = User.objects.first()
self.assertEqual(user.two_factor_auth, True)
def test_close_tfa_with_invalid_code(self):
self.test_open_tfa_with_correct_code()
resp = self.client.post(self.url, data={"code": "000000"})
self.assertDictEqual(resp.data, {"error": "error", "data": "Invalid code"})
def test_close_tfa_with_correct_code(self):
self.test_open_tfa_with_correct_code()
code = self._get_tfa_code()
resp = self.client.put(self.url, data={"code": code})
self.assertSuccess(resp)
user = User.objects.first()
self.assertEqual(user.two_factor_auth, False)
@mock.patch("account.views.oj.send_email_async.send")
class ApplyResetPasswordAPITest(CaptchaTest):
def setUp(self):
self.create_user("test", "test123", login=False)
user = User.objects.first()
user.email = "test@oj.com"
user.save()
self.url = self.reverse("apply_reset_password_api")
self.data = {"email": "test@oj.com", "captcha": self._set_captcha(self.client.session)}
def _refresh_captcha(self):
self.data["captcha"] = self._set_captcha(self.client.session)
def test_apply_reset_password(self, send_email_send):
resp = self.client.post(self.url, data=self.data)
self.assertSuccess(resp)
send_email_send.assert_called()
def test_apply_reset_password_twice_in_20_mins(self, send_email_send):
self.test_apply_reset_password()
send_email_send.reset_mock()
self._refresh_captcha()
resp = self.client.post(self.url, data=self.data)
self.assertDictEqual(resp.data, {"error": "error", "data": "You can only reset password once per 20 minutes"})
send_email_send.assert_not_called()
def test_apply_reset_password_again_after_20_mins(self, send_email_send):
self.test_apply_reset_password()
user = User.objects.first()
user.reset_password_token_expire_time = now() - timedelta(minutes=21)
user.save()
self._refresh_captcha()
self.test_apply_reset_password()
class ResetPasswordAPITest(CaptchaTest):
def setUp(self):
self.create_user("test", "test123", login=False)
self.url = self.reverse("reset_password_api")
user = User.objects.first()
user.reset_password_token = "online_judge?"
user.reset_password_token_expire_time = now() + timedelta(minutes=20)
user.save()
self.data = {"token": user.reset_password_token,
"captcha": self._set_captcha(self.client.session),
"password": "test456"}
def test_reset_password_with_correct_token(self):
resp = self.client.post(self.url, data=self.data)
self.assertSuccess(resp)
self.assertTrue(self.client.login(username="test", password="test456"))
def test_reset_password_with_invalid_token(self):
self.data["token"] = "aaaaaaaaaaa"
resp = self.client.post(self.url, data=self.data)
self.assertDictEqual(resp.data, {"error": "error", "data": "Token does not exist"})
def test_reset_password_with_expired_token(self):
user = User.objects.first()
user.reset_password_token_expire_time = now() - timedelta(seconds=30)
user.save()
resp = self.client.post(self.url, data=self.data)
self.assertDictEqual(resp.data, {"error": "error", "data": "Token has expired"})
class UserChangeEmailAPITest(APITestCase):
def setUp(self):
self.url = self.reverse("user_change_email_api")
self.user = self.create_user("test", "test123")
self.new_mail = "test@oj.com"
self.data = {"password": "test123", "new_email": self.new_mail}
def test_change_email_success(self):
resp = self.client.post(self.url, data=self.data)
self.assertSuccess(resp)
def test_wrong_password(self):
self.data["password"] = "aaaa"
resp = self.client.post(self.url, data=self.data)
self.assertDictEqual(resp.data, {"error": "error", "data": "Wrong password"})
def test_duplicate_email(self):
u = self.create_user("aa", "bb", login=False)
u.email = self.new_mail
u.save()
resp = self.client.post(self.url, data=self.data)
self.assertDictEqual(resp.data, {"error": "error", "data": "The email is owned by other account"})
class UserChangePasswordAPITest(APITestCase):
def setUp(self):
self.url = self.reverse("user_change_password_api")
# Create user at first
self.username = "test_user"
self.old_password = "testuserpassword"
self.new_password = "new_password"
self.user = self.create_user(username=self.username, password=self.old_password, login=False)
self.data = {"old_password": self.old_password, "new_password": self.new_password}
def _get_tfa_code(self):
user = User.objects.first()
code = OtpAuth(user.tfa_token).totp()
if len(str(code)) < 6:
code = (6 - len(str(code))) * "0" + str(code)
return code
def test_login_required(self):
response = self.client.post(self.url, data=self.data)
self.assertEqual(response.data, {"error": "permission-denied", "data": "Please login first"})
def test_valid_ola_password(self):
self.assertTrue(self.client.login(username=self.username, password=self.old_password))
response = self.client.post(self.url, data=self.data)
self.assertEqual(response.data, {"error": None, "data": "Succeeded"})
self.assertTrue(self.client.login(username=self.username, password=self.new_password))
def test_invalid_old_password(self):
self.assertTrue(self.client.login(username=self.username, password=self.old_password))
self.data["old_password"] = "invalid"
response = self.client.post(self.url, data=self.data)
self.assertEqual(response.data, {"error": "error", "data": "Invalid old password"})
def test_tfa_code_required(self):
self.user.two_factor_auth = True
self.user.tfa_token = "tfa_token"
self.user.save()
self.assertTrue(self.client.login(username=self.username, password=self.old_password))
self.data["tfa_code"] = rand_str(6)
resp = self.client.post(self.url, data=self.data)
self.assertEqual(resp.data, {"error": "error", "data": "Invalid two factor verification code"})
self.data["tfa_code"] = self._get_tfa_code()
resp = self.client.post(self.url, data=self.data)
self.assertSuccess(resp)
class UserRankAPITest(APITestCase):
def setUp(self):
self.url = self.reverse("user_rank_api")
self.create_user("test1", "test123", login=False)
self.create_user("test2", "test123", login=False)
test1 = User.objects.get(username="test1")
profile1 = test1.userprofile
profile1.submission_number = 10
profile1.accepted_number = 10
profile1.total_score = 240
profile1.save()
test2 = User.objects.get(username="test2")
profile2 = test2.userprofile
profile2.submission_number = 15
profile2.accepted_number = 10
profile2.total_score = 700
profile2.save()
def test_get_acm_rank(self):
resp = self.client.get(self.url, data={"rule": ContestRuleType.ACM})
self.assertSuccess(resp)
data = resp.data["data"]["results"]
self.assertEqual(data[0]["user"]["username"], "test1")
self.assertEqual(data[1]["user"]["username"], "test2")
def test_get_oi_rank(self):
resp = self.client.get(self.url, data={"rule": ContestRuleType.OI})
self.assertSuccess(resp)
data = resp.data["data"]["results"]
self.assertEqual(data[0]["user"]["username"], "test2")
self.assertEqual(data[1]["user"]["username"], "test1")
def test_admin_role_filted(self):
self.create_admin("admin", "admin123")
admin = User.objects.get(username="admin")
profile1 = admin.userprofile
profile1.submission_number = 20
profile1.accepted_number = 5
profile1.total_score = 300
profile1.save()
resp = self.client.get(self.url, data={"rule": ContestRuleType.ACM})
self.assertSuccess(resp)
self.assertEqual(len(resp.data["data"]), 2)
resp = self.client.get(self.url, data={"rule": ContestRuleType.OI})
self.assertSuccess(resp)
self.assertEqual(len(resp.data["data"]), 2)
class ProfileProblemDisplayIDRefreshAPITest(APITestCase):
def setUp(self):
pass
class AdminUserTest(APITestCase):
def setUp(self):
self.user = self.create_super_admin(login=True)
self.username = self.password = "test"
self.regular_user = self.create_user(username=self.username, password=self.password, login=False)
self.url = self.reverse("user_admin_api")
self.data = {"id": self.regular_user.id, "username": self.username, "real_name": "test_name",
"email": "test@qq.com", "admin_type": AdminType.REGULAR_USER,
"problem_permission": ProblemPermission.OWN, "open_api": True,
"two_factor_auth": False, "is_disabled": False}
def test_user_list(self):
response = self.client.get(self.url)
self.assertSuccess(response)
def test_edit_user_successfully(self):
response = self.client.put(self.url, data=self.data)
self.assertSuccess(response)
resp_data = response.data["data"]
self.assertEqual(resp_data["username"], self.username)
self.assertEqual(resp_data["email"], "test@qq.com")
self.assertEqual(resp_data["open_api"], True)
self.assertEqual(resp_data["two_factor_auth"], False)
self.assertEqual(resp_data["is_disabled"], False)
self.assertEqual(resp_data["problem_permission"], ProblemPermission.NONE)
self.assertTrue(self.regular_user.check_password("test"))
def test_edit_user_password(self):
data = self.data
new_password = "testpassword"
data["password"] = new_password
response = self.client.put(self.url, data=data)
self.assertSuccess(response)
user = User.objects.get(id=self.regular_user.id)
self.assertFalse(user.check_password(self.password))
self.assertTrue(user.check_password(new_password))
def test_edit_user_tfa(self):
data = self.data
self.assertIsNone(self.regular_user.tfa_token)
data["two_factor_auth"] = True
response = self.client.put(self.url, data=data)
self.assertSuccess(response)
resp_data = response.data["data"]
# if `tfa_token` is None, a new value will be generated
self.assertTrue(resp_data["two_factor_auth"])
token = User.objects.get(id=self.regular_user.id).tfa_token
self.assertIsNotNone(token)
response = self.client.put(self.url, data=data)
self.assertSuccess(response)
resp_data = response.data["data"]
# if `tfa_token` is not None, the value is not changed
self.assertTrue(resp_data["two_factor_auth"])
self.assertEqual(User.objects.get(id=self.regular_user.id).tfa_token, token)
def test_edit_user_openapi(self):
data = self.data
self.assertIsNone(self.regular_user.open_api_appkey)
data["open_api"] = True
response = self.client.put(self.url, data=data)
self.assertSuccess(response)
resp_data = response.data["data"]
# if `open_api_appkey` is None, a new value will be generated
self.assertTrue(resp_data["open_api"])
key = User.objects.get(id=self.regular_user.id).open_api_appkey
self.assertIsNotNone(key)
response = self.client.put(self.url, data=data)
self.assertSuccess(response)
resp_data = response.data["data"]
# if `openapi_app_key` is not None, the value is not changed
self.assertTrue(resp_data["open_api"])
self.assertEqual(User.objects.get(id=self.regular_user.id).open_api_appkey, key)
def test_import_users(self):
data = {"users": [["user1", "pass1", "eami1@e.com", "user1"],
["user2", "pass3", "eamil3@e.com", "user2"]]
}
resp = self.client.post(self.url, data)
self.assertSuccess(resp)
# successfully created 2 users
self.assertEqual(User.objects.all().count(), 4)
def test_import_duplicate_user(self):
data = {"users": [["user1", "pass1", "eami1@e.com", "user1"],
["user1", "pass1", "eami1@e.com", "user1"]]
}
resp = self.client.post(self.url, data)
self.assertFailed(resp, "DETAIL: Key (username)=(user1) already exists.")
# no user is created
self.assertEqual(User.objects.all().count(), 2)
def test_delete_users(self):
self.test_import_users()
user_ids = User.objects.filter(username__in=["user1", "user2"]).values_list("id", flat=True)
user_ids = ",".join([str(id) for id in user_ids])
resp = self.client.delete(self.url + "?id=" + user_ids)
self.assertSuccess(resp)
self.assertEqual(User.objects.all().count(), 2)
class GenerateUserAPITest(APITestCase):
def setUp(self):
self.create_super_admin()
self.url = self.reverse("generate_user_api")
self.data = {
"number_from": 100, "number_to": 105,
"prefix": "pre", "suffix": "suf",
"default_email": "test@test.com",
"password_length": 8
}
def test_error_case(self):
data = deepcopy(self.data)
data["prefix"] = "t" * 16
data["suffix"] = "s" * 14
resp = self.client.post(self.url, data=data)
self.assertEqual(resp.data["data"], "Username should not more than 32 characters")
data2 = deepcopy(self.data)
data2["number_from"] = 106
resp = self.client.post(self.url, data=data2)
self.assertEqual(resp.data["data"], "Start number must be lower than end number")
@mock.patch("account.views.admin.xlsxwriter.Workbook")
def test_generate_user_success(self, mock_workbook):
resp = self.client.post(self.url, data=self.data)
self.assertSuccess(resp)
mock_workbook.assert_called()
class OpenAPIAppkeyAPITest(APITestCase):
def setUp(self):
self.user = self.create_super_admin()
self.url = self.reverse("open_api_appkey_api")
def test_reset_appkey(self):
resp = self.client.post(self.url, data={})
self.assertFailed(resp)
self.user.open_api = True
self.user.save()
resp = self.client.post(self.url, data={})
self.assertSuccess(resp)
self.assertEqual(resp.data["data"]["appkey"], User.objects.get(username=self.user.username).open_api_appkey)

View File

@@ -1,8 +1,8 @@
from django.urls import path
from ..views.admin import ResetUserPasswordAPI, UserAdminAPI
from ..views.admin import UserAdminAPI, GenerateUserAPI
urlpatterns = [
path("user", UserAdminAPI.as_view()),
path("reset_password", ResetUserPasswordAPI.as_view()),
path("generate_user", GenerateUserAPI.as_view()),
]

View File

@@ -1,27 +1,54 @@
from django.urls import path
from ..views.oj import (
AvatarUploadAPI,
ApplyResetPasswordAPI,
ResetPasswordAPI,
UserChangePasswordAPI,
Metrics,
ProfileProblemDisplayIDRefreshAPI,
UserActivityRankAPI,
UserRegisterAPI,
UserChangeEmailAPI,
UserLoginAPI,
UserLogoutAPI,
UserProblemRankAPI,
UsernameOrEmailCheck,
AvatarUploadAPI,
TwoFactorAuthAPI,
UserProfileAPI,
UserRankAPI,
UserRegisterAPI,
UserActivityRankAPI,
CheckTFARequiredAPI,
SessionManagementAPI,
ProfileProblemDisplayIDRefreshAPI,
OpenAPIAppkeyAPI,
SSOAPI,
)
from utils.captcha.views import CaptchaAPIView
urlpatterns = [
path("login", UserLoginAPI.as_view()),
path("logout", UserLogoutAPI.as_view()),
path("register", UserRegisterAPI.as_view()),
path("change_password", UserChangePasswordAPI.as_view()),
path("change_email", UserChangeEmailAPI.as_view()),
path("apply_reset_password", ApplyResetPasswordAPI.as_view()),
path("reset_password", ResetPasswordAPI.as_view()),
path("captcha", CaptchaAPIView.as_view()),
path("check_username_or_email", UsernameOrEmailCheck.as_view()),
path("profile", UserProfileAPI.as_view(), name="user_profile_api"),
path("profile/fresh_display_id", ProfileProblemDisplayIDRefreshAPI.as_view()),
path("metrics", Metrics.as_view()),
path("upload_avatar", AvatarUploadAPI.as_view()),
path("tfa_required", CheckTFARequiredAPI.as_view()),
path(
"two_factor_auth",
TwoFactorAuthAPI.as_view(),
),
path("user_rank", UserRankAPI.as_view()),
path("user_activity_rank", UserActivityRankAPI.as_view()),
path("user_problem_rank", UserProblemRankAPI.as_view()),
path("sessions", SessionManagementAPI.as_view()),
path(
"open_api_appkey",
OpenAPIAppkeyAPI.as_view(),
),
path("sso", SSOAPI.as_view()),
]

View File

@@ -1,37 +1,24 @@
import os
import re
import xlsxwriter
from django.db import transaction, IntegrityError
from django.db.models import Q
from django.http import HttpResponse
from django.contrib.auth.hashers import make_password
from django.db import IntegrityError, transaction
from django.db.models import F, Q
from django.utils.crypto import get_random_string
from submission.models import Submission
from utils.api import APIError, APIView, validate_serializer
from utils.shortcuts import CLASS_NAME_MAX_DIGITS, CLASS_NAME_MIN_DIGITS, is_valid_class_name, rand_str
from utils.api import APIView, validate_serializer
from utils.shortcuts import rand_str
from ..decorators import super_admin_required
from ..models import AdminType, ProblemPermission, User, UserProfile
from ..serializers import (
EditUserSerializer,
ImportUserSerializer,
UserAdminSerializer,
GenerateUserSerializer,
)
# ks251XXX 或者 ks2510XX 返回 251 或者 2510。
# 不以 ks+数字 开头的(管理员、教师账号)返回 None。
# 位数不对就直接报错,不猜——猜错会把 class_name 存歪,
# 而剥前缀显示姓名、班级下拉、统计页都依赖它准确。
# 这里先用 \d+ 抓全再判位数,不能直接用 CLASS_NAME_RE 匹配:
# 那样 ks251001 会"匹配成功"并悄悄取前 4 位,正是要避免的猜测。
def get_class_name(username):
result = re.match(r"ks(\d+)", username)
if not result:
return None
class_name = result.group(1)
if not is_valid_class_name(class_name):
raise APIError(f"用户名 {username} 的班级号 {class_name}{len(class_name)} 位,必须是 {CLASS_NAME_MIN_DIGITS}~{CLASS_NAME_MAX_DIGITS} 位数字")
return class_name
from ..serializers import ImportUserSerializer
class UserAdminAPI(APIView):
@@ -53,14 +40,18 @@ class UserAdminAPI(APIView):
password=make_password(user_data[1]),
email=user_data[2],
raw_password=user_data[1],
class_name=get_class_name(user_data[0]),
)
)
try:
with transaction.atomic():
ret = User.objects.bulk_create(user_list)
UserProfile.objects.bulk_create([UserProfile(user=ret[i], real_name=data[i][3]) for i in range(len(ret))])
UserProfile.objects.bulk_create(
[
UserProfile(user=ret[i], real_name=data[i][3])
for i in range(len(ret))
]
)
return self.success()
except IntegrityError as e:
# Extract detail from exception message
@@ -79,22 +70,27 @@ class UserAdminAPI(APIView):
user = User.objects.get(id=data["id"])
except User.DoesNotExist:
return self.error("User does not exist")
if User.objects.filter(username=data["username"].lower()).exclude(id=user.id).exists():
if (
User.objects.filter(username=data["username"].lower())
.exclude(id=user.id)
.exists()
):
return self.error("Username already exists")
if User.objects.filter(email=data["email"].lower()).exclude(id=user.id).exists():
if (
User.objects.filter(email=data["email"].lower())
.exclude(id=user.id)
.exists()
):
return self.error("Email already exists")
pre_username = user.username
user.username = data["username"].lower()
user.class_name = get_class_name(data["username"])
user.email = data["email"].lower()
user.admin_type = data["admin_type"]
user.is_disabled = data["is_disabled"]
if data["admin_type"] == AdminType.STUDENT_ADMIN:
user.problem_permission = data["problem_permission"] or ProblemPermission.OWN
elif data["admin_type"] == AdminType.TEACHER_ADMIN:
user.problem_permission = data["problem_permission"] or ProblemPermission.OWN
if data["admin_type"] == AdminType.ADMIN:
user.problem_permission = data["problem_permission"]
elif data["admin_type"] == AdminType.SUPER_ADMIN:
user.problem_permission = ProblemPermission.ALL
else:
@@ -111,9 +107,20 @@ class UserAdminAPI(APIView):
user.open_api_appkey = None
user.open_api = data["open_api"]
if data["two_factor_auth"]:
# Avoid reset user tfa_token after saving changes
if not user.two_factor_auth:
user.tfa_token = rand_str()
else:
user.tfa_token = None
user.two_factor_auth = data["two_factor_auth"]
user.save()
if pre_username != user.username:
Submission.objects.filter(username=pre_username).update(username=user.username)
Submission.objects.filter(username=pre_username).update(
username=user.username
)
UserProfile.objects.filter(user=user).update(real_name=data["real_name"])
return self.success(UserAdminSerializer(user).data)
@@ -131,25 +138,20 @@ class UserAdminAPI(APIView):
return self.error("User does not exist")
return self.success(UserAdminSerializer(user).data)
# 获取排序参数
order_by = request.GET.get("order_by", "")
user = User.objects.all().order_by("-create_time")
# 根据排序参数设置排序规则
if order_by == "-last_login":
# 最近登录,将 None 值放在最后
user = User.objects.all().order_by(F("last_login").desc(nulls_last=True))
else:
# 默认按创建时间倒序
user = User.objects.all().order_by("-create_time")
is_admin = request.GET.get("admin", "0")
type = request.GET.get("type", "")
if type:
user = user.filter(admin_type=type)
if is_admin == "1":
user = user.exclude(admin_type=AdminType.REGULAR_USER)
keyword = request.GET.get("keyword", None)
if keyword:
user = user.filter(Q(username__icontains=keyword) | Q(userprofile__real_name__icontains=keyword) | Q(email__icontains=keyword))
user = user.filter(
Q(username__icontains=keyword)
| Q(userprofile__real_name__icontains=keyword)
| Q(email__icontains=keyword)
)
return self.success(self.paginate_data(request, user, UserAdminSerializer))
@super_admin_required
@@ -164,25 +166,76 @@ class UserAdminAPI(APIView):
return self.success()
class ResetUserPasswordAPI(APIView):
class GenerateUserAPI(APIView):
@super_admin_required
def get(self, request):
"""
download users excel
"""
file_id = request.GET.get("file_id")
if not file_id:
return self.error("Invalid Parameter, file_id is required")
if not re.match(r"^[a-zA-Z0-9]+$", file_id):
return self.error("Illegal file_id")
file_path = f"/tmp/{file_id}.xlsx"
if not os.path.isfile(file_path):
return self.error("File does not exist")
with open(file_path, "rb") as f:
raw_data = f.read()
os.remove(file_path)
response = HttpResponse(raw_data)
response["Content-Disposition"] = "attachment; filename=users.xlsx"
response["Content-Type"] = "application/xlsx"
return response
@validate_serializer(GenerateUserSerializer)
@super_admin_required
def post(self, request):
"""
重置用户密码为随机6位数字(不包括0)
Generate User
"""
data = request.data
user_id = data["id"]
number_max_length = max(
len(str(data["number_from"])), len(str(data["number_to"]))
)
if number_max_length + len(data["prefix"]) + len(data["suffix"]) > 32:
return self.error("Username should not more than 32 characters")
if data["number_from"] > data["number_to"]:
return self.error("Start number must be lower than end number")
file_id = rand_str(8)
filename = f"/tmp/{file_id}.xlsx"
workbook = xlsxwriter.Workbook(filename)
worksheet = workbook.add_worksheet()
worksheet.set_column("A:B", 20)
worksheet.write("A1", "Username")
worksheet.write("B1", "Password")
i = 1
user_list = []
for number in range(data["number_from"], data["number_to"] + 1):
raw_password = rand_str(data["password_length"])
user = User(
username=f"{data['prefix']}{number}{data['suffix']}",
password=make_password(raw_password),
)
user.raw_password = raw_password
user_list.append(user)
try:
user = User.objects.get(id=user_id)
except User.DoesNotExist:
return self.error("User does not exist")
# 生成6位随机数字密码(不包括0)
new_password = get_random_string(6, allowed_chars="123456789")
# 设置新密码
user.set_password(new_password)
user.save()
return self.success(new_password)
with transaction.atomic():
ret = User.objects.bulk_create(user_list)
UserProfile.objects.bulk_create(
[UserProfile(user=user) for user in ret]
)
for item in user_list:
worksheet.write_string(i, 0, item.username)
worksheet.write_string(i, 1, item.raw_password)
i += 1
workbook.close()
return self.success({"file_id": file_id})
except IntegrityError as e:
# Extract detail from exception message
# duplicate key value violates unique constraint "user_username_key"
# DETAIL: Key (username)=(root11) already exists.
return self.error(str(e).split("\n")[1])

View File

@@ -1,36 +1,54 @@
import asyncio
import os
from datetime import timedelta
from importlib import import_module
from django.conf import settings
from django.contrib import auth
from django.template.loader import render_to_string
from django.utils.decorators import method_decorator
from django.utils.timezone import now
from django.views.decorators.csrf import ensure_csrf_cookie, csrf_exempt
from django.db.models import Count, Q
from django.utils import timezone
from django.utils.decorators import method_decorator
from django.views.decorators.csrf import ensure_csrf_cookie
from options.options import SysOptions
import qrcode
from otpauth import TOTP
from problem.models import Problem
from submission.models import JudgeStatus, Submission
from utils.api import AsyncAPIView, validate_serializer
from utils.async_helpers import async_cache_get, async_cache_set
from utils.constants import CacheKey
from utils.shortcuts import datetime2str, rand_str
from submission.models import Submission, JudgeStatus
from utils.constants import ContestRuleType
from options.options import SysOptions
from utils.api import APIView, validate_serializer, CSRFExemptAPIView
from utils.captcha import Captcha
from utils.shortcuts import rand_str, img2base64, datetime2str
from ..decorators import login_required
from ..models import AdminType, User, UserProfile
from ..models import User, UserProfile, AdminType
from ..serializers import (
ApplyResetPasswordSerializer,
ResetPasswordSerializer,
UserChangePasswordSerializer,
UserLoginSerializer,
UserRegisterSerializer,
UsernameOrEmailCheckSerializer,
RankInfoSerializer,
UserChangeEmailSerializer,
SSOSerializer,
)
from ..serializers import (
TwoFactorAuthCodeSerializer,
UserProfileSerializer,
EditUserProfileSerializer,
ImageUploadForm,
RankInfoSerializer,
UserLoginSerializer,
UserProfileSerializer,
UserRegisterSerializer,
)
from ..tasks import send_email_async
class UserProfileAPI(AsyncAPIView):
class UserProfileAPI(APIView):
@method_decorator(ensure_csrf_cookie)
async def get(self, request, **kwargs):
def get(self, request, **kwargs):
"""
判断是否登录, 若登录返回用户信息
"""
user = request.user
if not user.is_authenticated:
return self.success()
@@ -38,51 +56,56 @@ class UserProfileAPI(AsyncAPIView):
username = request.GET.get("username")
try:
if username:
user = await User.objects.aget(username=username, is_disabled=False)
user = User.objects.get(username=username, is_disabled=False)
else:
user = request.user
# api返回的是自己的信息可以返real_name
show_real_name = True
except User.DoesNotExist:
return self.error("User does not exist")
profile = await UserProfile.objects.select_related("user").aget(user=user)
return self.success(UserProfileSerializer(profile, show_real_name=show_real_name).data)
return self.success(
UserProfileSerializer(user.userprofile, show_real_name=show_real_name).data
)
@login_required
@validate_serializer(EditUserProfileSerializer)
async def put(self, request):
@login_required
def put(self, request):
data = request.data
user_profile = await UserProfile.objects.select_related("user").aget(user=request.user)
user_profile = request.user.userprofile
for k, v in data.items():
setattr(user_profile, k, v)
await user_profile.asave()
return self.success(UserProfileSerializer(user_profile, show_real_name=True).data)
class Metrics(AsyncAPIView):
async def get(self, request):
userid = request.GET.get("userid")
qs = Submission.objects.filter(user_id=userid, contest_id__isnull=True)
count, latest, first = await asyncio.gather(
qs.acount(),
qs.order_by("-create_time").afirst(),
qs.order_by("create_time").afirst(),
)
if count == 0 or not latest or not first:
return self.error("暂无提交")
user_profile.save()
return self.success(
{
"now": datetime2str(timezone.now()),
"latest": datetime2str(latest.create_time),
"first": datetime2str(first.create_time),
}
UserProfileSerializer(user_profile, show_real_name=True).data
)
class AvatarUploadAPI(AsyncAPIView):
class Metrics(APIView):
def get(self, request):
userid = request.GET.get("userid")
submissions = Submission.objects.filter(user_id=userid, contest_id__isnull=True)
if submissions.count() == 0:
return self.error("暂无提交")
else:
latest_submission = submissions.first()
last_submission = submissions.last()
if last_submission and latest_submission:
return self.success(
{
"now": datetime2str(timezone.now()),
"latest": datetime2str(latest_submission.create_time),
"first": datetime2str(last_submission.create_time),
}
)
else:
return self.error("暂无提交")
class AvatarUploadAPI(APIView):
request_parsers = ()
@login_required
async def post(self, request):
def post(self, request):
form = ImageUploadForm(request.POST, request.FILES)
if form.is_valid():
avatar = form.cleaned_data["image"]
@@ -98,156 +121,412 @@ class AvatarUploadAPI(AsyncAPIView):
with open(os.path.join(settings.AVATAR_UPLOAD_DIR, name), "wb") as img:
for chunk in avatar:
img.write(chunk)
user_profile = await UserProfile.objects.aget(user=request.user)
user_profile = request.user.userprofile
user_profile.avatar = f"{settings.AVATAR_URI_PREFIX}/{name}"
await user_profile.asave()
user_profile.save()
return self.success("Succeeded")
class UserLoginAPI(AsyncAPIView):
@validate_serializer(UserLoginSerializer)
async def post(self, request):
class TwoFactorAuthAPI(APIView):
@login_required
def get(self, request):
"""
Get QR code
"""
user = request.user
if user.two_factor_auth:
return self.error("2FA is already turned on")
token = rand_str()
user.tfa_token = token
user.save()
label = f"{SysOptions.website_name_shortcut}:{user.username}"
image = qrcode.make(
TOTP(token).to_uri(
"totp", label, SysOptions.website_name.replace(" ", "")
)
)
return self.success(img2base64(image))
@login_required
@validate_serializer(TwoFactorAuthCodeSerializer)
def post(self, request):
"""
Open 2FA
"""
code = request.data["code"]
user = request.user
if TOTP(user.tfa_token).verify(code):
user.two_factor_auth = True
user.save()
return self.success("Succeeded")
else:
return self.error("Invalid code")
@login_required
@validate_serializer(TwoFactorAuthCodeSerializer)
def put(self, request):
code = request.data["code"]
user = request.user
if not user.two_factor_auth:
return self.error("2FA is already turned off")
if TOTP(user.tfa_token).verify(code):
user.two_factor_auth = False
user.save()
return self.success("Succeeded")
else:
return self.error("Invalid code")
class CheckTFARequiredAPI(APIView):
@validate_serializer(UsernameOrEmailCheckSerializer)
def post(self, request):
"""
Check TFA is required
"""
data = request.data
user = await auth.aauthenticate(username=data["username"], password=data["password"])
result = False
if data.get("username"):
try:
user = User.objects.get(username=data["username"])
result = user.two_factor_auth
except User.DoesNotExist:
pass
return self.success({"result": result})
class UserLoginAPI(APIView):
@validate_serializer(UserLoginSerializer)
def post(self, request):
"""
User login api
"""
data = request.data
user = auth.authenticate(username=data["username"], password=data["password"])
# None is returned if username or password is wrong
if user:
if user.is_disabled:
return self.error("Your account has been disabled")
prev_login = user.last_login
await auth.alogin(request, user)
request.session["prev_login"] = datetime2str(prev_login) if prev_login else ""
return self.success("Succeeded")
if not user.two_factor_auth:
auth.login(request, user)
return self.success("Succeeded")
# `tfa_code` not in post data
if user.two_factor_auth and "tfa_code" not in data:
return self.error("tfa_required")
if TOTP(user.tfa_token).verify(data["tfa_code"]):
auth.login(request, user)
return self.success("Succeeded")
else:
return self.error("Invalid two factor verification code")
else:
return self.error("Invalid username or password")
class UserLogoutAPI(AsyncAPIView):
async def get(self, request):
await auth.alogout(request)
class UserLogoutAPI(APIView):
def get(self, request):
auth.logout(request)
return self.success()
class UserRegisterAPI(AsyncAPIView):
class UsernameOrEmailCheck(APIView):
@validate_serializer(UsernameOrEmailCheckSerializer)
def post(self, request):
"""
check username or email is duplicate
"""
data = request.data
# True means already exist.
result = {"username": False, "email": False}
if data.get("username"):
result["username"] = User.objects.filter(
username=data["username"].lower()
).exists()
if data.get("email"):
result["email"] = User.objects.filter(email=data["email"].lower()).exists()
return self.success(result)
class UserRegisterAPI(APIView):
@validate_serializer(UserRegisterSerializer)
async def post(self, request):
if not await SysOptions.aget("allow_register"):
def post(self, request):
"""
User register api
"""
if not SysOptions.allow_register:
return self.error("Register function has been disabled by admin")
data = request.data
data["username"] = data["username"].lower()
data["email"] = data["email"].lower()
if await User.objects.filter(username=data["username"]).aexists():
captcha = Captcha(request)
if not captcha.check(data["captcha"]):
return self.error("Invalid captcha")
if User.objects.filter(username=data["username"]).exists():
return self.error("Username already exists")
if await User.objects.filter(email=data["email"]).aexists():
if User.objects.filter(email=data["email"]).exists():
return self.error("Email already exists")
user = await User.objects.acreate(username=data["username"], email=data["email"])
user = User.objects.create(username=data["username"], email=data["email"])
user.set_password(data["password"])
await user.asave()
await UserProfile.objects.acreate(user=user)
user.save()
UserProfile.objects.create(user=user)
return self.success("Succeeded")
class UserRankAPI(AsyncAPIView):
async def get(self, request):
class UserChangeEmailAPI(APIView):
@validate_serializer(UserChangeEmailSerializer)
@login_required
def post(self, request):
data = request.data
user = auth.authenticate(
username=request.user.username, password=data["password"]
)
if user:
if user.two_factor_auth:
if "tfa_code" not in data:
return self.error("tfa_required")
if not TOTP(user.tfa_token).verify(data["tfa_code"]):
return self.error("Invalid two factor verification code")
data["new_email"] = data["new_email"].lower()
if User.objects.filter(email=data["new_email"]).exists():
return self.error("The email is owned by other account")
user.email = data["new_email"]
user.save()
return self.success("Succeeded")
else:
return self.error("Wrong password")
class UserChangePasswordAPI(APIView):
@validate_serializer(UserChangePasswordSerializer)
@login_required
def post(self, request):
"""
User change password api
"""
data = request.data
username = request.user.username
user = auth.authenticate(username=username, password=data["old_password"])
if user:
if user.two_factor_auth:
if "tfa_code" not in data:
return self.error("tfa_required")
if not TOTP(user.tfa_token).verify(data["tfa_code"]):
return self.error("Invalid two factor verification code")
user.set_password(data["new_password"])
user.save()
return self.success("Succeeded")
else:
return self.error("Invalid old password")
class ApplyResetPasswordAPI(APIView):
@validate_serializer(ApplyResetPasswordSerializer)
def post(self, request):
if request.user.is_authenticated:
return self.error("You have already logged in, are you kidding me? ")
data = request.data
captcha = Captcha(request)
if not captcha.check(data["captcha"]):
return self.error("Invalid captcha")
try:
user = User.objects.get(email__iexact=data["email"])
except User.DoesNotExist:
return self.error("User does not exist")
if (
user.reset_password_token_expire_time
and 0
< int((user.reset_password_token_expire_time - now()).total_seconds())
< 20 * 60
):
return self.error("You can only reset password once per 20 minutes")
user.reset_password_token = rand_str()
user.reset_password_token_expire_time = now() + timedelta(minutes=20)
user.save()
render_data = {
"username": user.username,
"website_name": SysOptions.website_name,
"link": f"{SysOptions.website_base_url}/reset-password/{user.reset_password_token}",
}
email_html = render_to_string("reset_password_email.html", render_data)
send_email_async.send(
from_name=SysOptions.website_name_shortcut,
to_email=user.email,
to_name=user.username,
subject="Reset your password",
content=email_html,
)
return self.success("Succeeded")
class ResetPasswordAPI(APIView):
@validate_serializer(ResetPasswordSerializer)
def post(self, request):
data = request.data
captcha = Captcha(request)
if not captcha.check(data["captcha"]):
return self.error("Invalid captcha")
try:
user = User.objects.get(reset_password_token=data["token"])
except User.DoesNotExist:
return self.error("Token does not exist")
if user.reset_password_token_expire_time < now():
return self.error("Token has expired")
user.reset_password_token = None
user.two_factor_auth = False
user.set_password(data["password"])
user.save()
return self.success("Succeeded")
class SessionManagementAPI(APIView):
@login_required
def get(self, request):
engine = import_module(settings.SESSION_ENGINE)
session_store = engine.SessionStore
current_session = request.session.session_key
session_keys = request.user.session_keys
result = []
modified = False
for key in session_keys[:]:
session = session_store(key)
# session does not exist or is expiry
if not session._session:
session_keys.remove(key)
modified = True
continue
s = {}
if current_session == key:
s["current_session"] = True
s["ip"] = session["ip"]
s["user_agent"] = session["user_agent"]
s["last_activity"] = datetime2str(session["last_activity"])
s["session_key"] = key
result.append(s)
if modified:
request.user.save()
return self.success(result)
@login_required
def delete(self, request):
session_key = request.GET.get("session_key")
if not session_key:
return self.error("Parameter Error")
request.session.delete(session_key)
if session_key in request.user.session_keys:
request.user.session_keys.remove(session_key)
request.user.save()
return self.success("Succeeded")
else:
return self.error("Invalid session_key")
class UserRankAPI(APIView):
def get(self, request):
rule_type = request.GET.get("rule")
username = request.GET.get("username", "")
try:
n = int(request.GET.get("n", "0"))
except ValueError:
n = 0
profiles = (
UserProfile.objects.filter(
user__admin_type__in=[AdminType.REGULAR_USER, AdminType.STUDENT_ADMIN],
user__is_disabled=False,
user__username__icontains=username,
if rule_type not in ContestRuleType.choices():
rule_type = ContestRuleType.ACM
profiles = UserProfile.objects.filter(
user__admin_type=AdminType.REGULAR_USER,
user__is_disabled=False,
user__username__icontains=username,
).select_related("user")
if rule_type == ContestRuleType.ACM:
profiles = profiles.filter(accepted_number__gte=0).order_by(
"-accepted_number", "submission_number"
)
.select_related("user")
.filter(accepted_number__gte=0)
.order_by("-accepted_number", "submission_number")
)
else:
profiles = profiles.filter(total_score__gt=0).order_by("-total_score")
if n > 0:
profiles = profiles[:n]
return self.success(await self.async_paginate_data(request, profiles, RankInfoSerializer))
return self.success(self.paginate_data(request, profiles, RankInfoSerializer))
class UserActivityRankAPI(AsyncAPIView):
async def get(self, request):
class UserActivityRankAPI(APIView):
def get(self, request):
start = request.GET.get("start")
if not start:
return self.error("start time is required")
cache_key = f"{CacheKey.user_activity_rank}:{start}"
cached = await async_cache_get(cache_key)
if cached is not None:
return self.success(cached)
hidden_names = User.objects.filter(Q(admin_type=AdminType.SUPER_ADMIN) | Q(is_disabled=True)).values_list("username", flat=True)
hidden_names = User.objects.filter(
Q(admin_type=AdminType.SUPER_ADMIN)
| Q(admin_type=AdminType.ADMIN)
| Q(is_disabled=True)
).values_list("username", flat=True)
submissions = Submission.objects.filter(
contest_id__isnull=True,
create_time__gte=start,
result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED],
).exclude(username__in=hidden_names)
data = [row async for row in submissions.values("username").annotate(count=Count("problem_id", distinct=True)).order_by("-count")[:10]]
await async_cache_set(cache_key, data, 600)
return self.success(data)
class UserProblemRankAPI(AsyncAPIView):
async def get(self, request):
problem_id = request.GET.get("problem_id")
user = request.user
if not user.is_authenticated:
return self.error("User is not authenticated")
problem = await Problem.objects.aget(_id__iexact=problem_id, contest_id__isnull=True, visible=True)
submissions = Submission.objects.filter(problem=problem, result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED])
all_ac_count = await submissions.values("user_id").distinct().acount()
class_name = user.class_name or ""
class_ac_count = 0
if class_name:
users = User.objects.filter(class_name=user.class_name, is_disabled=False).values_list("id", flat=True)
user_ids = [user_id async for user_id in users]
submissions = submissions.filter(user_id__in=user_ids)
class_ac_count = await submissions.values("user_id").distinct().acount()
my_submissions = submissions.filter(user_id=user.id)
if not await my_submissions.aexists():
return self.success(
{
"class_name": class_name,
"rank": -1,
"class_ac_count": class_ac_count,
"all_ac_count": all_ac_count,
}
)
my_first_submission = await my_submissions.order_by("create_time").afirst()
rank = await submissions.filter(create_time__lte=my_first_submission.create_time).acount()
return self.success(
{
"class_name": class_name,
"rank": rank,
"class_ac_count": class_ac_count,
"all_ac_count": all_ac_count,
}
contest_id__isnull=True, create_time__gte=start, result=JudgeStatus.ACCEPTED
)
counts = (
submissions.values("username")
.annotate(count=Count("problem_id", distinct=True))
.order_by("-count")[: 10 + len(hidden_names)]
)
data = []
for count in counts:
if count["username"] not in hidden_names:
data.append(count)
return self.success(data[:10])
class ProfileProblemDisplayIDRefreshAPI(AsyncAPIView):
class ProfileProblemDisplayIDRefreshAPI(APIView):
@login_required
async def get(self, request):
profile = await UserProfile.objects.aget(user=request.user)
def get(self, request):
profile = request.user.userprofile
acm_problems = profile.acm_problems_status.get("problems", {})
ids = list(acm_problems.keys())
oi_problems = profile.oi_problems_status.get("problems", {})
ids = list(acm_problems.keys()) + list(oi_problems.keys())
if not ids:
return self.success()
display_ids = [did async for did in Problem.objects.filter(id__in=ids, visible=True).values_list("_id", flat=True)]
display_ids = Problem.objects.filter(id__in=ids, visible=True).values_list(
"_id", flat=True
)
id_map = dict(zip(ids, display_ids))
for k, v in acm_problems.items():
v["_id"] = id_map[k]
await profile.asave(update_fields=["acm_problems_status"])
for k, v in oi_problems.items():
v["_id"] = id_map[k]
profile.save(update_fields=["acm_problems_status", "oi_problems_status"])
return self.success()
class OpenAPIAppkeyAPI(APIView):
@login_required
def post(self, request):
user = request.user
if not user.open_api:
return self.error("OpenAPI function is truned off for you")
api_appkey = rand_str()
user.open_api_appkey = api_appkey
user.save()
return self.success({"appkey": api_appkey})
class SSOAPI(CSRFExemptAPIView):
@login_required
def get(self, request):
token = rand_str()
request.user.auth_token = token
request.user.save()
return self.success({"token": token})
@method_decorator(csrf_exempt)
@validate_serializer(SSOSerializer)
def post(self, request):
try:
user = User.objects.get(auth_token=request.data["token"])
except User.DoesNotExist:
return self.error("User does not exist")
return self.success(
{
"username": user.username,
"avatar": user.userprofile.avatar,
"admin_type": user.admin_type,
}
)

View File

@@ -1,10 +0,0 @@
from django.apps import AppConfig
class AchievementConfig(AppConfig):
default_auto_field = "django.db.models.BigAutoField"
name = "achievement"
def ready(self):
# 导入以触发 @metric 装饰器把指标注册进 METRIC_REGISTRY
import achievement.metrics # noqa: F401

View File

@@ -1,99 +0,0 @@
"""成就判定核心,被异步任务和管理命令共用。
判定刻意做成"一次查询取全部候选 + 内存比对",与 problemset 里逐条查询的
旧写法相反:一次判题只多 3~4 条 SQL。
"""
import logging
from django.db import transaction
from django.db.models import F
from achievement.metrics import META_METRICS, METRIC_REGISTRY, build_ctx
from achievement.models import Achievement, Operator, UserAchievement, UserStat
logger = logging.getLogger(__name__)
def evaluate(user, metrics, only_metrics=None):
"""返回该用户应解锁但尚未解锁的成就列表。"""
unlocked_ids = set(UserAchievement.objects.filter(user=user).values_list("achievement_id", flat=True))
qs = Achievement.objects.filter(visible=True).exclude(id__in=unlocked_ids)
if only_metrics is not None:
qs = qs.filter(metric__in=only_metrics)
hits = []
for achievement in qs:
value = metrics.get(achievement.metric)
# 指标从未产生有效值时 key 不存在,直接跳过:
# 否则极小值型指标(求 min、配 lte 用的那种)会对新用户恒成立
if value is None:
continue
if achievement.operator == Operator.GTE and value >= achievement.threshold:
hits.append(achievement)
elif achievement.operator == Operator.LTE and value <= achievement.threshold:
hits.append(achievement)
return hits
def unlock(user, achievements, backfilled=False, notified=False):
"""写入解锁记录并累加 unlock_count返回实际新建的记录。
刻意逐条 get_or_create 而不是 bulk_createunique_user_achievement 约束负责
并发竞态was_created 是"这一条确实是我新建的"的唯一可信判据。用
bulk_create(ignore_conflicts=True) 则无法区分新建与已存在,并发判题时会把
unlock_count 重复累加(获得率永久偏高),并对同一个奖杯重复推送通知。
循环次数是"本次新解锁的成就数",常态为 0因此常态零查询。
"""
created = []
for achievement in achievements:
record, was_created = UserAchievement.objects.get_or_create(
user=user,
achievement=achievement,
defaults={"backfilled": backfilled, "notified": notified},
)
if was_created:
Achievement.objects.filter(id=achievement.id).update(unlock_count=F("unlock_count") + 1)
created.append(record)
return created
def run_for_submission(user, submission):
"""判题后的完整流程:更新指标 → 第一轮判定 → 元指标第二轮判定。"""
ctx = build_ctx(user.id, submission)
if ctx["skip"]:
# 比赛提交不计入成就
return []
with transaction.atomic():
stat = UserStat.objects.select_for_update().get_or_create(user=user)[0]
for key, m in METRIC_REGISTRY.items():
if key in META_METRICS:
continue
m.on_submission(stat.metrics, submission, ctx)
stat.save(update_fields=["metrics", "update_time"])
first = unlock(user, evaluate(user, stat.metrics))
if not first:
return []
# 第二轮:只重算元指标、只判定依赖元指标的成就,不再有第三轮
meta_values = {key: METRIC_REGISTRY[key].recompute(user) for key in META_METRICS}
# 必须重新取锁并重新读一次 stat上面那个 stat 对象的 metrics 是解锁前的快照,
# 直接 save 会把整份字典写回,覆盖掉并发判题在这期间已提交的增量
# (同一用户两次提交并发判题时会让提交数/AC 数静默倒退,且无定期重算兜底)。
# 这里只合并元指标那几个 key。
with transaction.atomic():
stat = UserStat.objects.select_for_update().get(user=user)
for key, value in meta_values.items():
if value is None:
stat.metrics.pop(key, None)
else:
stat.metrics[key] = value
stat.save(update_fields=["metrics", "update_time"])
second = unlock(user, evaluate(user, stat.metrics, only_metrics=META_METRICS))
return first + second

View File

@@ -1,228 +0,0 @@
"""成就系统部署后自检。
只读,不修改任何数据,可以反复跑。
这套系统开发全程没有数据库可用,所有后端逻辑都只经过纸面审查。本命令覆盖
那些"错了也不报错、只是悄悄发错奖杯"的地方——每一项都对应一个已知的、
静态审查抓不到的失败模式。
python manage.py check_achievement_deploy
"""
from django.core.management.base import BaseCommand
from django.db.models import IntegerField
from django.db.models.fields.json import KeyTextTransform
from django.db.models.functions import Cast
from account.models import User
from achievement.metrics import META_METRICS, METRIC_REGISTRY
from achievement.models import Achievement, Operator, UserAchievement, UserStat
# on_submission 依赖的辅助键 -> 它应该与哪个顶层指标保持一致
STATE_PAIRS = {
"_active_dates": "active_days",
"_languages": "languages_used",
}
class Command(BaseCommand):
help = "成就系统部署后自检(只读)"
def handle(self, *args, **options):
self.failures = 0
self.warnings = 0
self._check_migration()
self._check_registry()
self._check_min_metric_absent()
self._check_jsonb_cast()
self._check_recompute_state()
self._check_unlock_count()
self._check_dead_achievements()
self.stdout.write("")
if self.failures:
self.stdout.write(self.style.ERROR(f"{self.failures} 项未通过,先别配成就。"))
elif self.warnings:
self.stdout.write(self.style.WARNING(f"全部通过,{self.warnings} 项提醒。"))
else:
self.stdout.write(self.style.SUCCESS("全部通过。"))
# ---------- 输出 ----------
def _ok(self, title, detail=""):
self.stdout.write(self.style.SUCCESS(f"[PASS] {title}") + (f" {detail}" if detail else ""))
def _fail(self, title, detail):
self.failures += 1
self.stdout.write(self.style.ERROR(f"[FAIL] {title}"))
self.stdout.write(f" {detail}")
def _skip(self, title, detail):
self.stdout.write(self.style.WARNING(f"[SKIP] {title}") + f" {detail}")
def _warn(self, title, detail):
self.warnings += 1
self.stdout.write(self.style.WARNING(f"[WARN] {title}"))
self.stdout.write(f" {detail}")
# ---------- 检查项 ----------
def _check_migration(self):
title = "迁移已应用"
try:
Achievement.objects.count()
UserStat.objects.count()
UserAchievement.objects.count()
except Exception as e:
self._fail(title, f"三张表至少有一张不存在:{e}\n 跑 python manage.py migrate achievement")
return
self._ok(title)
def _check_registry(self):
title = "指标注册表已加载"
count = len(METRIC_REGISTRY)
if count == 0:
self._fail(title, "METRIC_REGISTRY 是空的AppConfig.ready() 没有导入 metrics")
return
self._ok(title, f"{count} 个指标,元指标 {sorted(META_METRICS)}")
def _check_min_metric_absent(self):
"""lte 类成就用的指标,对没有 AC 记录的用户必须返回 None。
返回 0 的话,"最短 AC 代码 ≤ 50 字符"这类成就会白送给每一个从没做出过题的新生。
线上目前没有 lte 成就min_ac_code_chars 已删),本项因此会 SKIP
将来配了任何 lte 成就,它会自动开始检查对应的指标。
"""
title = "极小值指标对零 AC 用户返回 None"
metric_keys = sorted(set(Achievement.objects.filter(operator=Operator.LTE, visible=True).values_list("metric", flat=True)))
metrics = [(k, METRIC_REGISTRY[k]) for k in metric_keys if k in METRIC_REGISTRY]
if not metrics:
self._skip(title, "没有上架的 lte 类成就")
return
user = User.objects.filter(is_disabled=False, userprofile__accepted_number=0).first()
if user is None:
self._skip(title, "找不到 accepted_number=0 的用户,无法验证")
return
bad = [(key, value) for key, m in metrics if (value := m.recompute(user)) is not None]
if bad:
detail = "".join(f"{key} 得到 {value!r}" for key, value in bad)
self._fail(
title,
f"用户 {user.username}{detail},都应为 None。\n 这些 lte 成就正在白送给全部零 AC 用户。",
)
else:
self._ok(title, f"用户 {user.username},检查了 {len(metrics)} 个指标")
def _check_jsonb_cast(self):
"""JSONB 数字必须按整数比较,不能按字符串序。
没有显式 Cast 时 Postgres 会认为 "9" > "50",后台调低阈值触发的补发
会发给错误的人群,且不报任何错。
"""
title = "JSONB 阈值比较按整数而非字符串序"
key = "submission_count"
rows = list(UserStat.objects.filter(metrics__has_key=key).values_list("id", "metrics"))
if len(rows) < 2:
self._skip(title, f"{key} 的 UserStat 少于 2 条,无法验证")
return
values = [(rid, m.get(key)) for rid, m in rows if isinstance(m.get(key), int)]
if not values:
self._skip(title, f"{key} 没有整数值")
return
# 挑一个能区分整数序和字符串序的阈值:字符串比较下 "9" > "50"
threshold = 9
expected = {rid for rid, v in values if v >= threshold}
actual = set(
UserStat.objects.filter(metrics__has_key=key).annotate(v=Cast(KeyTextTransform(key, "metrics"), IntegerField())).filter(v__gte=threshold).values_list("id", flat=True)
)
if actual == expected:
self._ok(title, f"阈值 {threshold} 命中 {len(actual)} 人,与 Python 侧一致")
else:
missing = len(expected - actual)
extra = len(actual - expected)
self._fail(
title,
f"数据库筛出 {len(actual)} 人,正确答案是 {len(expected)} 人(少 {missing}、多 {extra})。\n rescan_achievement 会把成就补发给错误的人群。",
)
def _check_recompute_state(self):
"""重算必须一并重建 on_submission 依赖的辅助键。
辅助键缺失时,用户的下一次提交会把 active_days / languages_used
打回 1且要等到下一次重算才恢复。
"""
title = "重算重建了增量辅助键"
stat = UserStat.objects.filter(metrics__has_key="active_days").exclude(metrics__active_days=0).first()
if stat is None:
self._skip(title, "还没有任何用户的 active_days先跑 recompute_achievements --user <id>")
return
problems = []
for state_key, top_key in STATE_PAIRS.items():
top = stat.metrics.get(top_key)
state = stat.metrics.get(state_key)
if top is None:
continue
if state is None:
problems.append(f"{state_key} 缺失({top_key}={top}")
elif len(state) != top:
problems.append(f"len({state_key})={len(state)} != {top_key}={top}")
if problems:
self._fail(
title,
"".join(problems) + f"\n 用户 id={stat.user_id} 的下一次提交会让这些指标回退。",
)
else:
self._ok(title, f"用户 id={stat.user_id} 的辅助键与顶层值一致")
def _check_unlock_count(self):
"""unlock_count 是独立计数器,必须与实际解锁人数相等。
并发解锁重复累加、或删过 UserAchievement 但没重置计数器,
都会让获得率虚高。
"""
title = "unlock_count 与实际解锁人数一致"
if not Achievement.objects.exists():
self._skip(title, "还没有配置任何成就")
return
drifted = []
for a in Achievement.objects.all():
real = UserAchievement.objects.filter(achievement=a).count()
if a.unlock_count != real:
drifted.append(f"{a.name}」计数器 {a.unlock_count} vs 实际 {real}")
if drifted:
self._fail(
title,
"".join(drifted[:5]) + ("" if len(drifted) > 5 else "") + "\n 获得率显示会不准。修Achievement.objects.update(unlock_count=0) 后重跑重算。",
)
else:
self._ok(title, f"{Achievement.objects.count()} 条成就全部一致")
def _check_dead_achievements(self):
"""长期零解锁的成就多半是阈值配错了。"""
title = "没有疑似配错阈值的成就"
if not Achievement.objects.exists():
self._skip(title, "还没有配置任何成就")
return
dead = list(Achievement.objects.filter(visible=True, unlock_count=0, hidden=False).values_list("name", "metric", "operator", "threshold"))
if dead:
self._warn(
title,
f"{len(dead)} 条公开成就零解锁:"
+ "".join(f"{n}{m} {o} {t}" for n, m, o, t in dead[:5])
+ ("" if len(dead) > 5 else "")
+ "\n 刚上线属正常;配置一周后仍为 0 就要怀疑阈值。",
)
else:
self._ok(title)

View File

@@ -1,65 +0,0 @@
"""全量重算指标并补发成就。
用途:
1. 系统首次上线,给存量用户补发(用 --silent否则学生一登录被 30 个奖杯糊脸)
2. 新增指标或改了口径后重算(不带 --silent新解锁会正常弹出
这也是增量逻辑的安全网:怀疑 UserStat 漂移时跑一遍即可。
"""
from django.core.management.base import BaseCommand
from account.models import User
from achievement import checker
from achievement.metrics import META_METRICS, METRIC_REGISTRY
from achievement.models import UserStat
class Command(BaseCommand):
help = "重算全部用户的成就指标并补发成就"
def add_arguments(self, parser):
parser.add_argument("--user", type=int, default=None, help="只处理指定用户 id")
parser.add_argument("--silent", action="store_true", help="补发的成就标记为已通知,不弹窗")
def handle(self, *args, **options):
users = User.objects.filter(is_disabled=False)
if options["user"]:
users = users.filter(id=options["user"])
silent = options["silent"]
total_users = users.count()
total_unlocked = 0
for index, user in enumerate(users.iterator(), start=1):
stat, _ = UserStat.objects.get_or_create(user=user)
metrics = {}
for key, m in METRIC_REGISTRY.items():
if key in META_METRICS:
continue
value = m.recompute(user)
# None 表示该指标无有效值key 必须缺席而不是置 0
if value is not None:
metrics[key] = value
# 一并重建增量辅助键,否则重算后的第一次判题会把
# active_days / languages_used / max_ac_in_one_day 打回 1
metrics.update(m.recompute_state(user))
stat.metrics = metrics
stat.save(update_fields=["metrics", "update_time"])
records = checker.unlock(user, checker.evaluate(user, metrics), backfilled=True, notified=silent)
# 元指标第二轮
for key in META_METRICS:
value = METRIC_REGISTRY[key].recompute(user)
if value is not None:
metrics[key] = value
stat.metrics = metrics
stat.save(update_fields=["metrics", "update_time"])
records += checker.unlock(user, checker.evaluate(user, metrics, only_metrics=META_METRICS), backfilled=True, notified=silent)
total_unlocked += len(records)
if index % 50 == 0:
self.stdout.write(f"processed {index}/{total_users}")
self.stdout.write(self.style.SUCCESS(f"完成:{total_users} 个用户,新解锁 {total_unlocked}"))

View File

@@ -1,403 +0,0 @@
"""成就指标注册表。
这里定义"能测量什么",后台定义"多少算达成"。管理员在后台看到的指标下拉框
就是 METRIC_REGISTRY 的 key 列表。加一个新维度必须改本文件并部署,
之后在该维度上加任意多条成就都是纯配置。
约定指标从未产生过有效值时recompute 返回 None调用方删除该 key。
metrics 字典里 key 不存在 == 未达标,判定时直接跳过。
"""
import logging
from django.db.models import Count, Q
from django.utils import timezone
from submission.models import JudgeStatus, Submission, is_accepted
from utils.constants import Difficulty
logger = logging.getLogger(__name__)
METRIC_REGISTRY = {}
META_METRICS = set()
# AC 口径与全项目一致AST_CHECK_FAILED代码结构检查未通过但测试点全过也算通过。
# 见 submission/models.py 的 is_accepted()。ORM 过滤用这个常量,标量比较用 is_accepted()。
ACCEPTED_RESULTS = (JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED)
class BaseMetric:
key = ""
name = ""
help_text = ""
def on_submission(self, metrics, sub, ctx):
"""判题完成后增量更新 metrics原地修改"""
raise NotImplementedError
def recompute(self, user):
"""全量重算。返回 None 表示该用户此指标无有效值。"""
raise NotImplementedError
def recompute_state(self, user):
"""重算时一并重建增量所需的辅助键,返回 {key: value}。
默认无辅助状态。凡是 on_submission 依赖 `_` 前缀辅助键的指标都必须重写
本方法:否则全量重算把辅助键丢掉后,下一次判题会从零重新累积,指标当场
回退active_days 从 45 掉回 1且要等到下一次重算才恢复。
"""
return {}
def metric(key, name, help_text="", meta=False):
def deco(cls):
cls.key = key
cls.name = name
cls.help_text = help_text
METRIC_REGISTRY[key] = cls()
if meta:
META_METRICS.add(key)
return cls
return deco
def _practice_submissions(user_id):
"""成就只统计平时练习的提交,比赛提交不计入。"""
return Submission.objects.filter(user_id=user_id, contest_id__isnull=True)
def build_ctx(user_id, sub):
"""判题后预查一次,所有指标复用,避免每个指标各自查库。
比赛提交不参与成就统计,此时返回 skip=True调用方直接跳过整轮判定。
"""
if sub.contest_id is not None:
return {"skip": True}
prior = _practice_submissions(user_id).filter(problem_id=sub.problem_id).exclude(id=sub.id)
prior_stats = prior.aggregate(
total=Count("id"),
accepted=Count("id", filter=Q(result__in=ACCEPTED_RESULTS)),
)
local_now = timezone.localtime(sub.create_time)
sub_is_accepted = is_accepted(sub.result)
is_first_ac_of_problem = sub_is_accepted and prior_stats["accepted"] == 0
# 难度只有首次 AC 时才用得上,其余情况不查这一次库——
# 绝大多数提交都不是首次 AC放在外面等于给每次判题白加一条 SQL
difficulty = None
if is_first_ac_of_problem:
from problem.models import Problem
difficulty = Problem.objects.filter(id=sub.problem_id).values_list("difficulty", flat=True).first()
return {
"skip": False,
"is_accepted": sub_is_accepted,
# 该题此前的提交次数与 AC 次数
"prior_count": prior_stats["total"],
"prior_accepted": prior_stats["accepted"],
# 首次 AC 这道题(此前从未 AC 过)
"is_first_ac_of_problem": is_first_ac_of_problem,
# 一发入魂:此前无任何提交且本次 AC
"is_first_try_ac": sub_is_accepted and prior_stats["total"] == 0,
# 本题难度,仅首次 AC 时有值
"problem_difficulty": difficulty,
"local_date": local_now.date().isoformat(),
"local_hour": local_now.hour,
}
@metric("accepted_count", "AC 题目数", "去重后通过的题目数量(不含比赛)")
class AcceptedCount(BaseMetric):
def on_submission(self, metrics, sub, ctx):
if ctx["is_first_ac_of_problem"]:
metrics["accepted_count"] = metrics.get("accepted_count", 0) + 1
def recompute(self, user):
return _practice_submissions(user.id).filter(result__in=ACCEPTED_RESULTS).order_by().values("problem_id").distinct().count()
class _DifficultyAcCount(BaseMetric):
"""按难度去重统计 AC 题数。子类只需指定 difficulty。
增量靠 ctx["problem_difficulty"],它只在首次 AC 时才有值——与
is_first_ac_of_problem 是同一个条件,所以两者一起判即可。
"""
difficulty = ""
def on_submission(self, metrics, sub, ctx):
if ctx["is_first_ac_of_problem"] and ctx["problem_difficulty"] == self.difficulty:
metrics[self.key] = metrics.get(self.key, 0) + 1
def recompute(self, user):
return _practice_submissions(user.id).filter(result__in=ACCEPTED_RESULTS, problem__difficulty=self.difficulty).order_by().values("problem_id").distinct().count()
@metric("mid_ac_count", "中等题 AC 数", "去重后通过的中等难度题目数(不含比赛)")
class MidAcCount(_DifficultyAcCount):
difficulty = Difficulty.MID
@metric("hard_ac_count", "困难题 AC 数", "去重后通过的困难题目数(不含比赛)")
class HardAcCount(_DifficultyAcCount):
difficulty = Difficulty.HIGH
@metric("submission_count", "提交总数", "提交次数(不含比赛)")
class SubmissionCount(BaseMetric):
def on_submission(self, metrics, sub, ctx):
metrics["submission_count"] = metrics.get("submission_count", 0) + 1
def recompute(self, user):
return _practice_submissions(user.id).count()
@metric("active_days", "活跃天数", "有过提交的累计天数")
class ActiveDays(BaseMetric):
def on_submission(self, metrics, sub, ctx):
seen = metrics.get("_active_dates", [])
if ctx["local_date"] not in seen:
seen.append(ctx["local_date"])
metrics["_active_dates"] = seen
metrics["active_days"] = len(seen)
def recompute(self, user):
dates = {timezone.localtime(t).date().isoformat() for t in _practice_submissions(user.id).values_list("create_time", flat=True)}
return len(dates)
def recompute_state(self, user):
dates = sorted({timezone.localtime(t).date().isoformat() for t in _practice_submissions(user.id).values_list("create_time", flat=True)})
return {"_active_dates": dates}
@metric("max_ac_streak_days", "最长连续 AC 天数", "连续每天至少 AC 一题的最长天数")
class MaxAcStreakDays(BaseMetric):
def on_submission(self, metrics, sub, ctx):
if not ctx["is_accepted"]:
return
today = ctx["local_date"]
last = metrics.get("_last_ac_date")
if last == today:
return
current = metrics.get("_current_ac_streak", 0)
if last and (timezone.datetime.fromisoformat(today) - timezone.datetime.fromisoformat(last)).days == 1:
current += 1
else:
current = 1
metrics["_last_ac_date"] = today
metrics["_current_ac_streak"] = current
metrics["max_ac_streak_days"] = max(metrics.get("max_ac_streak_days", 0), current)
def recompute(self, user):
dates = sorted({timezone.localtime(t).date() for t in _practice_submissions(user.id).filter(result__in=ACCEPTED_RESULTS).values_list("create_time", flat=True)})
if not dates:
return None
best = current = 1
for prev, cur in zip(dates, dates[1:]):
current = current + 1 if (cur - prev).days == 1 else 1
best = max(best, current)
return best
def recompute_state(self, user):
dates = sorted({timezone.localtime(t).date() for t in _practice_submissions(user.id).filter(result__in=ACCEPTED_RESULTS).values_list("create_time", flat=True)})
if not dates:
return {}
current = 1
for prev, cur in zip(dates, dates[1:]):
current = current + 1 if (cur - prev).days == 1 else 1
return {"_last_ac_date": dates[-1].isoformat(), "_current_ac_streak": current}
@metric("languages_used", "使用语言数", "用过多少种编程语言")
class LanguagesUsed(BaseMetric):
def on_submission(self, metrics, sub, ctx):
seen = metrics.get("_languages", [])
if sub.language not in seen:
seen.append(sub.language)
metrics["_languages"] = seen
metrics["languages_used"] = len(seen)
def recompute(self, user):
return _practice_submissions(user.id).order_by().values("language").distinct().count()
def recompute_state(self, user):
# order_by() 不能省Submission.Meta 有默认排序 ("-create_time",)
# Django 会把排序字段并入 DISTINCT于是每条提交各成一行——
# 实测某用户返回 659 条而不是 5 种语言。
# recompute 侥幸正确只是因为 .count() 会清掉排序,不能依赖这一点。
return {"_languages": list(_practice_submissions(user.id).order_by().values_list("language", flat=True).distinct())}
@metric("contest_joined", "参赛场次", "参加过的比赛数量(本指标是比赛维度,不受比赛提交不计入的限制)")
class ContestJoined(BaseMetric):
def on_submission(self, metrics, sub, ctx):
# 比赛提交在 build_ctx 就被跳过,本指标只走 recompute
return
def recompute(self, user):
return Submission.objects.filter(user_id=user.id, contest_id__isnull=False).order_by().values("contest_id").distinct().count()
@metric("badge_count", "题单奖章数", "获得的题单奖章数量")
class BadgeCount(BaseMetric):
def on_submission(self, metrics, sub, ctx):
# 奖章由题单流程颁发,本指标只走 recompute见 Task 7
return
def recompute(self, user):
from problemset.models import UserBadge
return UserBadge.objects.filter(user=user).count()
@metric("problemset_completed", "完成题单数", "完成的题单数量")
class ProblemSetCompleted(BaseMetric):
def on_submission(self, metrics, sub, ctx):
return
def recompute(self, user):
from problemset.models import ProblemSetProgress
return ProblemSetProgress.objects.filter(user=user, complete_time__isnull=False).count()
@metric("first_try_ac_count", "一发入魂次数", "首次提交即通过的次数")
class FirstTryAcCount(BaseMetric):
def on_submission(self, metrics, sub, ctx):
if ctx["is_first_try_ac"]:
metrics["first_try_ac_count"] = metrics.get("first_try_ac_count", 0) + 1
def recompute(self, user):
count = 0
seen = set()
for s in _practice_submissions(user.id).order_by("create_time").values("problem_id", "result"):
if s["problem_id"] in seen:
continue
seen.add(s["problem_id"])
if is_accepted(s["result"]):
count += 1
return count
@metric("midnight_submissions", "凌晨提交次数", "0:005:00 之间的提交次数")
class MidnightSubmissions(BaseMetric):
def on_submission(self, metrics, sub, ctx):
if 0 <= ctx["local_hour"] < 5:
metrics["midnight_submissions"] = metrics.get("midnight_submissions", 0) + 1
def recompute(self, user):
return sum(1 for t in _practice_submissions(user.id).values_list("create_time", flat=True) if 0 <= timezone.localtime(t).hour < 5)
@metric("early_bird_submissions", "早起提交次数", "5:007:00 之间的提交次数")
class EarlyBirdSubmissions(BaseMetric):
"""与 midnight_submissions 对称的作息维度。5 点是两者的分界,不重叠。"""
def on_submission(self, metrics, sub, ctx):
if 5 <= ctx["local_hour"] < 7:
metrics["early_bird_submissions"] = metrics.get("early_bird_submissions", 0) + 1
def recompute(self, user):
return sum(1 for t in _practice_submissions(user.id).values_list("create_time", flat=True) if 5 <= timezone.localtime(t).hour < 7)
@metric("compile_error_count", "编译错误次数", "累计编译错误的次数")
class CompileErrorCount(BaseMetric):
def on_submission(self, metrics, sub, ctx):
if sub.result == JudgeStatus.COMPILE_ERROR:
metrics["compile_error_count"] = metrics.get("compile_error_count", 0) + 1
def recompute(self, user):
return _practice_submissions(user.id).filter(result=JudgeStatus.COMPILE_ERROR).count()
@metric("max_wa_before_ac", "屡败屡战", "单题失败最多多少次后终于通过")
class MaxWaBeforeAc(BaseMetric):
def on_submission(self, metrics, sub, ctx):
if ctx["is_first_ac_of_problem"]:
metrics["max_wa_before_ac"] = max(metrics.get("max_wa_before_ac", 0), ctx["prior_count"])
def recompute(self, user):
best = None
attempts = {}
for s in _practice_submissions(user.id).order_by("create_time").values("problem_id", "result"):
pid = s["problem_id"]
if pid in attempts and attempts[pid] is None:
continue
if is_accepted(s["result"]):
best = max(best or 0, attempts.get(pid, 0))
attempts[pid] = None
else:
attempts[pid] = attempts.get(pid, 0) + 1
return best
@metric("max_ac_in_one_day", "单日最多 AC", "一天之内最多通过多少题")
class MaxAcInOneDay(BaseMetric):
def on_submission(self, metrics, sub, ctx):
if not ctx["is_first_ac_of_problem"]:
return
counts = metrics.get("_ac_per_day", {})
counts[ctx["local_date"]] = counts.get(ctx["local_date"], 0) + 1
metrics["_ac_per_day"] = counts
metrics["max_ac_in_one_day"] = max(counts.values())
def recompute(self, user):
counts = {}
seen = set()
for s in _practice_submissions(user.id).filter(result__in=ACCEPTED_RESULTS).order_by("create_time").values("problem_id", "create_time"):
if s["problem_id"] in seen:
continue
seen.add(s["problem_id"])
day = timezone.localtime(s["create_time"]).date().isoformat()
counts[day] = counts.get(day, 0) + 1
return max(counts.values()) if counts else None
def recompute_state(self, user):
counts = {}
seen = set()
for s in _practice_submissions(user.id).filter(result__in=ACCEPTED_RESULTS).order_by("create_time").values("problem_id", "create_time"):
if s["problem_id"] in seen:
continue
seen.add(s["problem_id"])
day = timezone.localtime(s["create_time"]).date().isoformat()
counts[day] = counts.get(day, 0) + 1
return {"_ac_per_day": counts}
# 曾经这里有 min_ac_code_chars最短 AC 代码)。线上实测 1314 个用户的分布,
# 最小值 8、p5=10有道题 8 个字符就能通过,于是它测的是"谁做过那道水题"
# 而不是"谁写得简洁"配不出有意义的成就2026-08-05 删除。
# 要重新引入,得先按题目难度加权,或排除掉那类水题。
@metric("max_code_lines", "最长代码行数", "提交过的最长代码有多少行")
class MaxCodeLines(BaseMetric):
def on_submission(self, metrics, sub, ctx):
lines = len(sub.code.splitlines())
metrics["max_code_lines"] = max(metrics.get("max_code_lines", 0), lines)
def recompute(self, user):
counts = [len(c.splitlines()) for c in _practice_submissions(user.id).values_list("code", flat=True)]
return max(counts) if counts else None
@metric("achievement_unlocked_count", "已解锁成就数", "已解锁的成就数量(不含白金档)", meta=True)
class AchievementUnlockedCount(BaseMetric):
"""自引用指标:解锁成就会改变它。
因此判定流程限定为最多两轮(见 checker.py且口径排除白金档自身
避免「集齐 N 个成就」这类白金奖杯把自己算进分子。
"""
def on_submission(self, metrics, sub, ctx):
# 由 checker 在第一轮解锁后显式重算,不参与增量更新
return
def recompute(self, user):
from achievement.models import Rarity, UserAchievement
return UserAchievement.objects.filter(user=user).exclude(achievement__rarity=Rarity.PLATINUM).count()

View File

@@ -1,70 +0,0 @@
# Generated by Django 6.0.4 on 2026-08-04 06:35
import django.db.models.deletion
from django.conf import settings
from django.db import migrations, models
class Migration(migrations.Migration):
initial = True
dependencies = [
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
]
operations = [
migrations.CreateModel(
name='Achievement',
fields=[
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
('name', models.TextField(verbose_name='成就名称')),
('description', models.TextField(verbose_name='成就描述')),
('icon', models.TextField(verbose_name='图标')),
('rarity', models.TextField(choices=[('bronze', '青铜'), ('silver', '白银'), ('gold', '黄金'), ('platinum', '白金')], default='bronze', verbose_name='稀有度')),
('hidden', models.BooleanField(db_default=False, default=False, verbose_name='是否隐藏')),
('metric', models.TextField(verbose_name='指标名')),
('operator', models.TextField(choices=[('gte', '大于等于'), ('lte', '小于等于')], default='gte', verbose_name='比较符')),
('threshold', models.IntegerField(verbose_name='阈值')),
('visible', models.BooleanField(db_default=True, default=True, verbose_name='是否上架')),
('unlock_count', models.IntegerField(db_default=0, default=0, verbose_name='已解锁人数')),
('order', models.IntegerField(db_default=0, default=0, verbose_name='排序')),
('create_time', models.DateTimeField(auto_now_add=True)),
],
options={
'verbose_name': '成就',
'verbose_name_plural': '成就',
'db_table': 'achievement',
'ordering': ('order', 'id'),
},
),
migrations.CreateModel(
name='UserStat',
fields=[
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
('metrics', models.JSONField(db_default=models.Value({}, output_field=models.JSONField()), default=dict)),
('update_time', models.DateTimeField(auto_now=True)),
('user', models.OneToOneField(on_delete=django.db.models.deletion.CASCADE, related_name='achievement_stat', to=settings.AUTH_USER_MODEL)),
],
options={
'db_table': 'user_stat',
},
),
migrations.CreateModel(
name='UserAchievement',
fields=[
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
('unlock_time', models.DateTimeField(auto_now_add=True)),
('backfilled', models.BooleanField(db_default=False, default=False)),
('notified', models.BooleanField(db_default=False, default=False)),
('achievement', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='achievement.achievement')),
('user', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='achievements', to=settings.AUTH_USER_MODEL)),
],
options={
'db_table': 'user_achievement',
'ordering': ('-unlock_time',),
'indexes': [models.Index(fields=['user', '-unlock_time'], name='user_achv_time_idx'), models.Index(fields=['user', 'notified'], name='user_achv_notified_idx')],
'constraints': [models.UniqueConstraint(fields=('user', 'achievement'), name='unique_user_achievement')],
},
),
]

View File

@@ -1,123 +0,0 @@
"""把 Achievement.icon 从 emoji 字符换成 iconify 图标名。
机房里的 Chrome 91 之类的老浏览器缺 emoji 字体emoji 会渲染成方块。
iconify 出来的是 SVG跟浏览器字体无关。
映射到 notoNoto Color Emoji 的 SVG 版),视觉上和原 emoji 基本一致。
表里没有的 emoji 保持原值不动,前端会按纯文本兜底渲染,之后在管理后台手工补。
"""
from django.db import migrations
# 带 variation selector-16的 emoji 后台可能存成带或不带两种形态,
# 迁移时统一去掉再查表,所以这里的 key 一律不带
EMOJI_TO_ICONIFY = {
# 计划里首批成就用到的
"🌱": "noto:seedling",
"📗": "noto:green-book",
"📚": "noto:books",
"📅": "noto:calendar",
"🗓": "noto:spiral-calendar",
"🌐": "noto:globe-with-meridians",
"🎖": "noto:military-medal",
"🎯": "noto:direct-hit",
"🏹": "noto:bow-and-arrow",
"🦉": "noto:owl",
"🌙": "noto:crescent-moon",
"💥": "noto:collision",
"🔥": "noto:fire",
"": "noto:high-voltage",
"": "noto:scissors",
"📜": "noto:scroll",
"💎": "noto:gem-stone",
"🏆": "noto:trophy",
"": "noto:red-question-mark",
# 其余常用图标,后台可能已经用上了
"": "noto:white-question-mark",
"📖": "noto:open-book",
"🥇": "noto:1st-place-medal",
"🥈": "noto:2nd-place-medal",
"🥉": "noto:3rd-place-medal",
"🏅": "noto:sports-medal",
"": "noto:star",
"🌟": "noto:glowing-star",
"": "noto:sparkles",
"🎉": "noto:party-popper",
"🎊": "noto:confetti-ball",
"🚀": "noto:rocket",
"💡": "noto:light-bulb",
"🐛": "noto:bug",
"👑": "noto:crown",
"🥷": "noto:ninja",
"🧠": "noto:brain",
"💪": "noto:flexed-biceps",
"🎓": "noto:graduation-cap",
"📈": "noto:chart-increasing",
"🎁": "noto:wrapped-gift",
"🔮": "noto:crystal-ball",
"🦄": "noto:unicorn",
"🐢": "noto:turtle",
"": "noto:hot-beverage",
"🌈": "noto:rainbow",
"": "noto:alarm-clock",
"📝": "noto:memo",
"": "noto:check-mark-button",
"🔒": "noto:locked",
"🎮": "noto:video-game",
"🧩": "noto:puzzle-piece",
"🔧": "noto:wrench",
"🐍": "noto:snake",
"💻": "noto:laptop",
"": "noto:keyboard",
"": "noto:sun",
"😴": "noto:sleeping-face",
"🤖": "noto:robot",
"🔍": "noto:magnifying-glass-tilted-left",
"🎨": "noto:artist-palette",
"": "noto:crossed-swords",
"🛡": "noto:shield",
"🏰": "noto:castle",
"🗿": "noto:moai",
"🌊": "noto:water-wave",
"🍀": "noto:four-leaf-clover",
"🔁": "noto:repeat-button",
"🎪": "noto:circus-tent",
"🦅": "noto:eagle",
"🐉": "noto:dragon",
"💫": "noto:dizzy",
"🧭": "noto:compass",
"🪄": "noto:magic-wand",
"🔔": "noto:bell",
"📌": "noto:pushpin",
"🏁": "noto:chequered-flag",
"🎢": "noto:roller-coaster",
}
ICONIFY_TO_EMOJI = {v: k for k, v in EMOJI_TO_ICONIFY.items()}
def _convert(apps, schema_editor, table):
Achievement = apps.get_model("achievement", "Achievement")
updated = []
for a in Achievement.objects.all():
key = (a.icon or "").replace("", "").strip()
new = table.get(key)
if new and new != a.icon:
a.icon = new
updated.append(a)
if updated:
Achievement.objects.bulk_update(updated, ["icon"])
def emoji_to_iconify(apps, schema_editor):
_convert(apps, schema_editor, EMOJI_TO_ICONIFY)
def iconify_to_emoji(apps, schema_editor):
_convert(apps, schema_editor, ICONIFY_TO_EMOJI)
class Migration(migrations.Migration):
dependencies = [("achievement", "0001_initial")]
operations = [migrations.RunPython(emoji_to_iconify, iconify_to_emoji)]

View File

@@ -1,75 +0,0 @@
from django.db import models
from account.models import User
from utils.models import JSONField
class Rarity(models.TextChoices):
BRONZE = "bronze", "青铜"
SILVER = "silver", "白银"
GOLD = "gold", "黄金"
PLATINUM = "platinum", "白金"
class Operator(models.TextChoices):
GTE = "gte", "大于等于"
LTE = "lte", "小于等于"
class Achievement(models.Model):
name = models.TextField(verbose_name="成就名称")
description = models.TextField(verbose_name="成就描述")
# iconify 图标名(如 noto:owl不是 emoji 字符:
# 机房的老浏览器缺 emoji 字体会渲染成方块,前端统一渲染成 SVG
icon = models.TextField(verbose_name="图标")
rarity = models.TextField(default=Rarity.BRONZE, choices=Rarity.choices, verbose_name="稀有度")
hidden = models.BooleanField(default=False, db_default=False, verbose_name="是否隐藏")
metric = models.TextField(verbose_name="指标名")
operator = models.TextField(default=Operator.GTE, choices=Operator.choices, verbose_name="比较符")
threshold = models.IntegerField(verbose_name="阈值")
visible = models.BooleanField(default=True, db_default=True, verbose_name="是否上架")
unlock_count = models.IntegerField(default=0, db_default=0, verbose_name="已解锁人数")
order = models.IntegerField(default=0, db_default=0, verbose_name="排序")
create_time = models.DateTimeField(auto_now_add=True)
class Meta:
db_table = "achievement"
ordering = ("order", "id")
verbose_name = "成就"
verbose_name_plural = "成就"
class UserStat(models.Model):
"""成就系统唯一的指标源,不复用 UserProfile 的计数器,避免两处口径漂移。
metrics 为 {指标名: 数值};指标从未产生过有效值时 key 不存在(而非置 0
否则极小值型指标(求 min、配 lte 用的那种)会对新用户恒成立。
"""
user = models.OneToOneField(User, on_delete=models.CASCADE, related_name="achievement_stat")
metrics = JSONField(default=dict, db_default=models.Value({}, output_field=models.JSONField()))
update_time = models.DateTimeField(auto_now=True)
class Meta:
db_table = "user_stat"
class UserAchievement(models.Model):
user = models.ForeignKey(User, on_delete=models.CASCADE, related_name="achievements")
achievement = models.ForeignKey(Achievement, on_delete=models.CASCADE)
unlock_time = models.DateTimeField(auto_now_add=True)
# 上线补发的记录:前端不显示具体日期,只显示"已获得"
backfilled = models.BooleanField(default=False, db_default=False)
# 是否已向用户弹过奖杯pending 端点据此查询
notified = models.BooleanField(default=False, db_default=False)
class Meta:
db_table = "user_achievement"
ordering = ("-unlock_time",)
constraints = [
models.UniqueConstraint(fields=["user", "achievement"], name="unique_user_achievement"),
]
indexes = [
models.Index(fields=["user", "-unlock_time"], name="user_achv_time_idx"),
models.Index(fields=["user", "notified"], name="user_achv_notified_idx"),
]

View File

@@ -1,59 +0,0 @@
"""解锁通知。
通知走推拉结合UserAchievement.notified 是唯一的真相来源:
- 拉(主):前端在布局层拉 /api/achievements/pending覆盖全部场景
- 推增强WebSocket 只负责把"当场那一下"的延迟压到几百毫秒
必须推拉结合的原因:前端 WebSocket 不是常驻连接useSubmissionWebSocket
只在问题页且有提交监听时建连,纯推会丢消息。
"""
import logging
from utils.websocket import push_to_user
logger = logging.getLogger(__name__)
def notify_achievements(user_id, records):
"""records: list[UserAchievement],已带 select_related('achievement')。"""
if not records:
return
payload = [
{
"id": r.achievement_id,
"name": r.achievement.name,
"description": r.achievement.description,
"icon": r.achievement.icon,
"rarity": r.achievement.rarity,
"kind": "achievement",
}
for r in records
]
_push(user_id, payload)
def notify_badges(user_id, badges):
"""badges: list[ProblemSetBadge]。题单奖章复用同一个弹窗组件。"""
if not badges:
return
payload = [
{
"id": b.id,
"name": b.name,
"description": b.description,
"icon": b.icon,
"rarity": "bronze",
"kind": "badge",
}
for b in badges
]
_push(user_id, payload)
def _push(user_id, payload):
# 推送失败不影响已入库的解锁记录,前端下次拉 pending 时仍会补弹
try:
push_to_user(user_id, "achievement_unlocked", {"achievements": payload})
except Exception as e:
logger.error(f"Failed to push achievement notification: user_id={user_id}, error={e}")

View File

@@ -1,97 +0,0 @@
from django.core.cache import cache
from rest_framework import serializers
from account.models import User
from achievement.models import Achievement, UserAchievement
def get_active_user_count():
"""获得率的分母:未禁用用户总数,缓存 1 小时。
不用"有过提交的用户数"这类动态口径:分母波动会让同一个成就的获得率
忽高忽低,学生会当成 bug。
"""
count = cache.get("achievement_active_user_count")
if count is None:
count = User.objects.filter(is_disabled=False).count()
cache.set("achievement_active_user_count", count, 3600)
return count
class AchievementSerializer(serializers.ModelSerializer):
unlocked = serializers.BooleanField(read_only=True)
unlock_time = serializers.DateTimeField(read_only=True, allow_null=True)
backfilled = serializers.BooleanField(read_only=True)
progress = serializers.IntegerField(read_only=True, allow_null=True)
unlock_rate = serializers.SerializerMethodField()
name = serializers.SerializerMethodField()
description = serializers.SerializerMethodField()
icon = serializers.SerializerMethodField()
# 条件三件套也必须跟着遮掉:只遮名称和描述、却明文下发
# metric/operator/threshold学生打开 DevTools 就知道"凌晨提交 10 次"
# 隐藏成就的意义全部作废
metric = serializers.SerializerMethodField()
operator = serializers.SerializerMethodField()
threshold = serializers.SerializerMethodField()
class Meta:
model = Achievement
fields = (
"id",
"name",
"description",
"icon",
"rarity",
"hidden",
"metric",
"operator",
"threshold",
"unlocked",
"unlock_time",
"backfilled",
"progress",
"unlock_rate",
)
def _masked(self, obj):
return obj.hidden and not getattr(obj, "unlocked", False)
def get_name(self, obj):
return "???" if self._masked(obj) else obj.name
def get_description(self, obj):
return "达成条件保密" if self._masked(obj) else obj.description
def get_icon(self, obj):
# icon 存的是 iconify 图标名,不是 emoji 字符:机房里的老浏览器缺 emoji
# 字体会渲染成方块,前端统一按 iconify 渲染成 SVG
return "noto:red-question-mark" if self._masked(obj) else obj.icon
def get_metric(self, obj):
return None if self._masked(obj) else obj.metric
def get_operator(self, obj):
return None if self._masked(obj) else obj.operator
def get_threshold(self, obj):
return None if self._masked(obj) else obj.threshold
def get_unlock_rate(self, obj):
total = get_active_user_count()
if not total:
return 0.0
return round(obj.unlock_count / total * 100, 1)
class PendingAchievementSerializer(serializers.ModelSerializer):
"""待弹窗的解锁记录,无需打码(已解锁)。"""
id = serializers.IntegerField(source="achievement_id", read_only=True)
name = serializers.CharField(source="achievement.name", read_only=True)
description = serializers.CharField(source="achievement.description", read_only=True)
icon = serializers.CharField(source="achievement.icon", read_only=True)
rarity = serializers.CharField(source="achievement.rarity", read_only=True)
class Meta:
model = UserAchievement
fields = ("id", "name", "description", "icon", "rarity")

View File

@@ -1,65 +0,0 @@
import logging
import dramatiq
from account.models import User
from achievement import checker
from achievement.notify import notify_achievements
from submission.models import Submission
from utils.shortcuts import DRAMATIQ_WORKER_ARGS
logger = logging.getLogger(__name__)
@dramatiq.actor(**DRAMATIQ_WORKER_ARGS())
def check_achievements(user_id, submission_id):
"""判题完成后的成就判定。
所有异常在此吞掉:成就算错绝不能影响判题结果,这也是选异步的意义。
"""
try:
user = User.objects.get(id=user_id)
submission = Submission.objects.get(id=submission_id)
records = checker.run_for_submission(user, submission)
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:
# backfilled=True这是补发不是学生刚刚挣到的。
# 前端据此只显示"已获得"而不显示具体日期——否则一次补发会给几百人
# 盖上同一个时间戳,把"最近获得"板块彻底冲垮
records = checker.unlock(stat.user, [achievement], backfilled=True)
notify_achievements(stat.user_id, records)
except Exception as e:
logger.error(f"rescan_achievement failed for user {stat.user_id}: {e}")

View File

@@ -1,8 +0,0 @@
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"),
]

View File

@@ -1,9 +0,0 @@
from django.urls import path
from achievement.views.oj import AchievementListAPI, AchievementPendingAPI, AchievementSummaryAPI
urlpatterns = [
path("achievements", AchievementListAPI.as_view(), name="achievement_list_api"),
path("achievements/summary", AchievementSummaryAPI.as_view(), name="achievement_summary_api"),
path("achievements/pending", AchievementPendingAPI.as_view(), name="achievement_pending_api"),
]

View File

@@ -1,126 +0,0 @@
from account.decorators import super_admin_required
from achievement.metrics import METRIC_REGISTRY
from achievement.models import Achievement, Rarity
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)
before = (achievement.metric, achievement.operator, achievement.threshold, achievement.visible)
for field in ("name", "description", "icon", "rarity", "hidden", "metric", "operator", "threshold", "visible", "order"):
if field in data:
setattr(achievement, field, data[field])
achievement.save()
# 只要"谁能达成"这件事可能变了就补发,不去精细判断是否放宽。
# 补发是幂等的后台任务unlock 用 get_or_create多跑一次只花一次扫描
# 漏跑却是学生已达标却拿不到,两个方向代价不对称。
# 早先的 loosened 谓词只看 operator/threshold会漏掉两种情况
# 换了 metric换了维度、以及从下架改成上架草稿期已达标的人
after = (achievement.metric, achievement.operator, achievement.threshold, achievement.visible)
if achievement.visible and before != after:
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 "比较符不合法"
# rarity 不校验的话,一个乱填的值会让 AchievementSummaryAPI 的四档统计
# 对不上:它按 Rarity.choices 遍历,野值算进总数却不出现在任何一档里
if data["rarity"] not in Rarity.values:
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,
# 必须 isoformat():这里是手写的 dict 而不是 DRF 序列化器,
# 原始 datetime 交给 JSON 编码器会抛 TypeError整个管理接口 500。
# 表为空时列表接口看着正常(不进循环),一旦有数据就全挂。
"create_time": a.create_time.isoformat() if a.create_time else None,
}

View File

@@ -1,100 +0,0 @@
from django.db.models import Count
from account.decorators import login_required
from account.models import User
from achievement.models import Achievement, Rarity, UserAchievement, UserStat
from achievement.serializers import AchievementSerializer, PendingAchievementSerializer
from utils.api import APIView
def _resolve_user(request):
"""?name=<username> 指定他人,不传则为自己。与 /api/profile 的约定一致。"""
username = request.GET.get("name")
if username:
return User.objects.filter(username=username, is_disabled=False).first()
return request.user
def _decorate(achievements, unlocked_map, metrics):
"""给成就对象挂上该用户的解锁状态与进度,供序列化器读取。"""
for a in achievements:
record = unlocked_map.get(a.id)
a.unlocked = record is not None
a.unlock_time = record.unlock_time if record else None
a.backfilled = record.backfilled if record else False
value = metrics.get(a.metric)
# 隐藏且未解锁的成就不下发进度,否则能反推出条件
a.progress = None if (a.hidden and not a.unlocked) else (value or 0)
return achievements
class AchievementListAPI(APIView):
@login_required
def get(self, request):
user = _resolve_user(request)
if user is None:
return self.error("用户不存在")
achievements = list(Achievement.objects.filter(visible=True))
unlocked_map = {r.achievement_id: r for r in UserAchievement.objects.filter(user=user)}
stat = UserStat.objects.filter(user=user).first()
metrics = stat.metrics if stat else {}
_decorate(achievements, unlocked_map, metrics)
return self.success(
{
"username": user.username,
"achievements": AchievementSerializer(achievements, many=True).data,
}
)
class AchievementSummaryAPI(APIView):
@login_required
def get(self, request):
user = _resolve_user(request)
if user is None:
return self.error("用户不存在")
total_by_rarity = dict(Achievement.objects.filter(visible=True).values_list("rarity").annotate(c=Count("id")))
unlocked_by_rarity = dict(UserAchievement.objects.filter(user=user, achievement__visible=True).values_list("achievement__rarity").annotate(c=Count("id")))
total = sum(total_by_rarity.values())
unlocked = sum(unlocked_by_rarity.values())
recent = list(UserAchievement.objects.filter(user=user, achievement__visible=True).select_related("achievement").order_by("-unlock_time")[:10])
return self.success(
{
"username": user.username,
"total": total,
"unlocked": unlocked,
"percent": round(unlocked / total * 100, 1) if total else 0.0,
"rarity": [
{
"rarity": value,
"label": label,
"total": total_by_rarity.get(value, 0),
"unlocked": unlocked_by_rarity.get(value, 0),
}
for value, label in Rarity.choices
],
"recent": PendingAchievementSerializer(recent, many=True).data,
}
)
class AchievementPendingAPI(APIView):
@login_required
def get(self, request):
"""返回尚未弹过的解锁记录。前端在布局层路由切换时拉取。"""
records = UserAchievement.objects.filter(user=request.user, notified=False, achievement__visible=True).select_related("achievement").order_by("unlock_time")
return self.success(PendingAchievementSerializer(records, many=True).data)
@login_required
def post(self, request):
"""弹完后标记已读避免下次导航重复弹。body: {"ids": [成就 id]}"""
ids = request.data.get("ids") or []
if not isinstance(ids, list):
return self.error("参数错误")
UserAchievement.objects.filter(user=request.user, achievement_id__in=ids).update(notified=True)
return self.success("ok")

View File

View File

@@ -1,6 +0,0 @@
from django.apps import AppConfig
class AiConfig(AppConfig):
default_auto_field = 'django.db.models.BigAutoField'
name = 'ai'

View File

@@ -1,34 +0,0 @@
# Generated by Django 5.2.3 on 2025-09-24 12:59
import django.db.models.deletion
from django.conf import settings
from django.db import migrations, models
class Migration(migrations.Migration):
initial = True
dependencies = [
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
]
operations = [
migrations.CreateModel(
name='AIAnalysis',
fields=[
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
('provider', models.TextField(default='deepseek')),
('data', models.JSONField()),
('system_prompt', models.TextField()),
('user_prompt', models.TextField()),
('analysis', models.TextField()),
('create_time', models.DateTimeField(auto_now_add=True)),
('user', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to=settings.AUTH_USER_MODEL)),
],
options={
'db_table': 'ai_analysis',
'ordering': ['-create_time'],
},
),
]

View File

@@ -1,18 +0,0 @@
# Generated by Django 5.2.3 on 2025-09-24 13:02
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('ai', '0001_initial'),
]
operations = [
migrations.AddField(
model_name='aianalysis',
name='model',
field=models.TextField(default='deepseek-chat'),
),
]

View File

@@ -1,18 +0,0 @@
# Generated by Django 6.0 on 2026-04-27 12:31
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('ai', '0002_aianalysis_model'),
]
operations = [
migrations.AlterField(
model_name='aianalysis',
name='model',
field=models.TextField(default='deepseek-v4-flash'),
),
]

View File

@@ -1,18 +0,0 @@
# Generated by Django 6.0.4 on 2026-06-04 14:39
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('ai', '0003_alter_aianalysis_model'),
]
operations = [
migrations.AddField(
model_name='aianalysis',
name='is_pinned',
field=models.BooleanField(default=False),
),
]

View File

@@ -1,19 +0,0 @@
from django.db import models
from account.models import User
class AIAnalysis(models.Model):
user = models.ForeignKey(User, on_delete=models.CASCADE)
provider = models.TextField(default="deepseek")
model = models.TextField(default="deepseek-v4-flash")
data = models.JSONField()
system_prompt = models.TextField()
user_prompt = models.TextField()
analysis = models.TextField()
is_pinned = models.BooleanField(default=False)
create_time = models.DateTimeField(auto_now_add=True)
class Meta:
db_table = "ai_analysis"
ordering = ["-create_time"]

View File

@@ -1,27 +0,0 @@
from utils.api import serializers
from .models import AIAnalysis
class AIAnalysisListSerializer(serializers.ModelSerializer):
username = serializers.CharField(source="user.username")
analysis_excerpt = serializers.SerializerMethodField()
class Meta:
model = AIAnalysis
fields = ["id", "create_time", "username", "analysis_excerpt", "is_pinned"]
def get_analysis_excerpt(self, obj):
if not obj.analysis:
return ""
text = " ".join(obj.analysis.split())
return text[:120] if len(text) <= 120 else text[:120] + ""
class AIAnalysisDetailSerializer(serializers.ModelSerializer):
username = serializers.CharField(source="user.username")
class_name = serializers.CharField(source="user.class_name")
class Meta:
model = AIAnalysis
fields = ["id", "create_time", "username", "class_name", "analysis"]

View File

View File

@@ -1,7 +0,0 @@
from django.urls import path
from ..views.admin import AIAnalysisAdminAPI
urlpatterns = [
path("ai/reports", AIAnalysisAdminAPI.as_view()),
]

View File

@@ -1,25 +0,0 @@
from django.urls import path
from ..views.oj import (
AIAnalysisAPI,
AIDetailDataAPI,
AIDurationDataAPI,
AIHeatmapDataAPI,
AIHintAPI,
AILoginSummaryAPI,
AIPinnedReportAPI,
ClassPKAnalysisAPI,
SingleClassAnalysisAPI,
)
urlpatterns = [
path("ai/detail", AIDetailDataAPI.as_view()),
path("ai/duration", AIDurationDataAPI.as_view()),
path("ai/analysis", AIAnalysisAPI.as_view()),
path("ai/hint", AIHintAPI.as_view()),
path("ai/heatmap", AIHeatmapDataAPI.as_view()),
path("ai/login_summary", AILoginSummaryAPI.as_view()),
path("ai/pinned", AIPinnedReportAPI.as_view()),
path("ai/class_pk", ClassPKAnalysisAPI.as_view()),
path("ai/class_single", SingleClassAnalysisAPI.as_view()),
]

View File

View File

@@ -1,45 +0,0 @@
from account.decorators import teacher_admin_required
from utils.api import APIView
from ..models import AIAnalysis
from ..serializers import AIAnalysisDetailSerializer, AIAnalysisListSerializer
class AIAnalysisAdminAPI(APIView):
@teacher_admin_required
def get(self, request):
report_id = request.GET.get("id")
if report_id:
try:
report = AIAnalysis.objects.select_related("user").get(id=report_id)
except AIAnalysis.DoesNotExist:
return self.error("AIAnalysis not found")
return self.success(AIAnalysisDetailSerializer(report).data)
qs = AIAnalysis.objects.select_related("user").order_by("-create_time")
username = request.GET.get("username")
if username:
qs = qs.filter(user__username__icontains=username)
if request.GET.get("pinned_only") == "true":
pinned = qs.filter(is_pinned=True)
return self.success(AIAnalysisListSerializer(pinned, many=True).data)
return self.success(self.paginate_data(request, qs, AIAnalysisListSerializer))
@teacher_admin_required
def post(self, request):
report_id = request.data.get("id")
try:
report = AIAnalysis.objects.select_related("user").get(id=report_id)
except AIAnalysis.DoesNotExist:
return self.error("AIAnalysis not found")
if report.is_pinned:
report.is_pinned = False
else:
AIAnalysis.objects.filter(user=report.user, is_pinned=True).update(is_pinned=False)
report.is_pinned = True
report.save(update_fields=["is_pinned"])
return self.success({"is_pinned": report.is_pinned})

View File

@@ -1,907 +0,0 @@
import calendar
import hashlib
import json
from collections import defaultdict
from datetime import datetime, timedelta
from django.core.cache import cache
from django.db.models import Count, Min
from django.db.models.functions import TruncDate
from django.http import StreamingHttpResponse
from django.utils import timezone
from django.utils.dateparse import parse_datetime
from account.decorators import login_required, teacher_admin_required
from account.models import User
from ai.models import AIAnalysis
from ai.serializers import AIAnalysisDetailSerializer
from flowchart.models import FlowchartSubmission, FlowchartSubmissionStatus
from problem.models import Problem
from submission.models import JudgeStatus, Submission
from utils.api import APIView
from utils.openai import get_ai_client, get_async_ai_client
from utils.shortcuts import datetime2str
CACHE_TIMEOUT = 300
DIFFICULTY_MAP = {"Low": "简单", "Mid": "中等", "High": "困难"}
DEFAULT_CLASS_SIZE = 45
# 评级阈值配置:(百分位上限, 评级)
GRADE_THRESHOLDS = [
(10, "S"), # 前10%: S级 - 卓越
(35, "A"), # 前35%: A级 - 优秀
(75, "B"), # 前75%: B级 - 良好
(100, "C"), # 其余: C级 - 及格
]
# 小规模参与惩罚配置:(最小人数, 等级降级映射)
SMALL_SCALE_PENALTY = {
"threshold": 10,
"downgrade": {"S": "A", "A": "B"},
}
# 等级权重映射(用于加权平均计算)
GRADE_WEIGHTS = {"S": 4, "A": 3, "B": 2, "C": 1}
# 平均等级阈值:(最小权重, 等级)
AVERAGE_GRADE_THRESHOLDS = [(3.5, "S"), (2.5, "A"), (1.5, "B")]
def shift_months(dt, months):
"""按月平移。落到不存在的日期时收缩到当月最后一天1月31日 + 1个月 = 2月28/29日"""
month_index = dt.month - 1 + months
year = dt.year + month_index // 12
month = month_index % 12 + 1
day = min(dt.day, calendar.monthrange(year, month)[1])
return dt.replace(year=year, month=month, day=day)
def get_cache_key(prefix, *args):
return hashlib.md5(f"{prefix}:{'_'.join(map(str, args))}".encode()).hexdigest()
def get_difficulty(difficulty):
return DIFFICULTY_MAP.get(difficulty, "中等")
def get_grade(rank, submission_count, reference_count=None):
"""
计算题目完成评级
评级标准:
- S级前10%卓越水平10%的人)
- A级前35%优秀水平25%的人)
- B级前75%良好水平40%的人)
- C级75%之后及格水平25%的人)
特殊规则:
- 小规模惩罚用 reference_count全时段人数判断避免同期窗口窄导致惩罚误触发
- reference_count 未传时退化为 submission_count
"""
if not rank or rank <= 0 or submission_count <= 0:
return "C"
percentile = (rank - 1) / submission_count * 100
base_grade = "C"
for threshold, grade in GRADE_THRESHOLDS:
if percentile < threshold:
base_grade = grade
break
penalty_count = reference_count if reference_count is not None else submission_count
if penalty_count < SMALL_SCALE_PENALTY["threshold"]:
base_grade = SMALL_SCALE_PENALTY["downgrade"].get(base_grade, base_grade)
return base_grade
def calculate_average_grade(grades):
"""根据等级列表计算加权平均等级"""
scores = [GRADE_WEIGHTS[g] for g in grades if g in GRADE_WEIGHTS]
if not scores:
return ""
avg = sum(scores) / len(scores)
for threshold, grade in AVERAGE_GRADE_THRESHOLDS:
if avg >= threshold:
return grade
return "C"
def find_user_rank(ranking_list, user_id):
"""在排名列表中找到用户的排名1-based未找到返回 None"""
return next(
(idx + 1 for idx, rec in enumerate(ranking_list) if rec["user_id"] == user_id),
None,
)
def get_class_user_ids(user):
if not user.class_name:
return []
cache_key = get_cache_key("class_users", user.class_name)
user_ids = cache.get(cache_key)
if user_ids is None:
user_ids = list(User.objects.filter(class_name=user.class_name).values_list("id", flat=True))
cache.set(cache_key, user_ids, CACHE_TIMEOUT)
return user_ids
def get_user_first_ac_submissions(user_id, start, end, class_user_ids=None, use_class_scope=False, include_all_time=True):
# 用户自己的 AC 记录按时间范围过滤
user_first_ac = list(
Submission.objects.filter(
result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED],
user_id=user_id,
create_time__gte=start,
create_time__lte=end,
)
.values("problem_id")
.annotate(first_ac_time=Min("create_time"))
)
if not user_first_ac:
return [], {}, []
problem_ids = [item["problem_id"] for item in user_first_ac]
if not include_all_time:
return user_first_ac, {}, problem_ids
# 排名基于全局数据(不限时间),后注册的学生与所有人公平竞争
rank_qs = Submission.objects.filter(
result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED],
problem_id__in=problem_ids,
)
if use_class_scope and class_user_ids:
rank_qs = rank_qs.filter(user_id__in=class_user_ids)
ranked_first_ac = list(rank_qs.values("user_id", "problem_id").annotate(first_ac_time=Min("create_time")))
by_problem = defaultdict(list)
for item in ranked_first_ac:
by_problem[item["problem_id"]].append(item)
for submissions in by_problem.values():
submissions.sort(key=lambda x: (x["first_ac_time"], x["user_id"]))
return user_first_ac, by_problem, problem_ids
async def stream_ai_response(client, system_prompt, user_prompt, on_complete=None):
"""SSE 流式响应异步生成器on_complete(full_text) 在流结束时调用。
必须是异步生成器:在 ASGI 下,同步生成器会被 Django 的
StreamingHttpResponse 通过 sync_to_async(list) 一次性消费完才发送,
导致整段内容生成完毕后才返回,失去流式效果。
"""
try:
stream = await client.chat.completions.create(
model="deepseek-v4-flash",
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt},
],
stream=True,
extra_body={"thinking": {"type": "disabled"}},
)
except Exception as exc:
yield f"data: {json.dumps({'type': 'error', 'message': str(exc)})}\n\n"
yield "event: end\n\n"
return
yield "event: start\n\n"
chunks = []
try:
async for chunk in stream:
if not chunk.choices:
continue
choice = chunk.choices[0]
if choice.finish_reason:
if on_complete:
await on_complete("".join(chunks).strip())
yield f"data: {json.dumps({'type': 'done'})}\n\n"
break
content = choice.delta.content
if content:
chunks.append(content)
yield f"data: {json.dumps({'type': 'delta', 'content': content})}\n\n"
except Exception as exc:
yield f"data: {json.dumps({'type': 'error', 'message': str(exc)})}\n\n"
finally:
yield "event: end\n\n"
def make_sse_response(generator):
"""创建 SSE StreamingHttpResponse"""
response = StreamingHttpResponse(
streaming_content=generator,
content_type="text/event-stream",
)
response["Cache-Control"] = "no-cache"
# 关闭反向代理(如 nginx对流式响应的缓冲
response["X-Accel-Buffering"] = "no"
return response
class AIDetailDataAPI(APIView):
@login_required
def get(self, request):
start = request.GET.get("start")
end = request.GET.get("end")
username = request.GET.get("username")
if not start or not end:
return self.error("参数 start 和 end 不能为空")
if not parse_datetime(start):
return self.error("start 格式无效,请使用 ISO 8601 格式")
if not parse_datetime(end):
return self.error("end 格式无效,请使用 ISO 8601 格式")
user = request.user
if username and request.user.is_teacher_or_above():
try:
user = User.objects.get(username=username)
except User.DoesNotExist:
return self.error("User not found")
cache_key = get_cache_key("ai_detail", user.id, user.class_name or "", start, end)
cached_result = cache.get(cache_key)
if cached_result:
return self.success(cached_result)
class_user_ids = get_class_user_ids(user)
use_class_scope = bool(user.class_name) and len(class_user_ids) > 1
user_first_ac, by_problem, problem_ids = get_user_first_ac_submissions(user.id, start, end, class_user_ids, use_class_scope)
# 同期排名:只统计时间窗口内解题的人
by_problem_period = defaultdict(list)
if problem_ids:
period_qs = Submission.objects.filter(
result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED],
problem_id__in=problem_ids,
create_time__gte=start,
create_time__lte=end,
)
if use_class_scope and class_user_ids:
period_qs = period_qs.filter(user_id__in=class_user_ids)
for item in period_qs.values("user_id", "problem_id").annotate(first_ac_time=Min("create_time")):
by_problem_period[item["problem_id"]].append(item)
for lst in by_problem_period.values():
lst.sort(key=lambda x: (x["first_ac_time"], x["user_id"]))
result = {
"user": user.username,
"class_name": user.class_name,
"start": start,
"end": end,
"solved": [],
"flowcharts": [],
"grade": "",
"tags": {},
"difficulty": {},
"contest_count": 0,
}
if user_first_ac:
problems = {p.id: p for p in Problem.objects.filter(id__in=problem_ids).select_related("contest").prefetch_related("tags")}
solved, contest_ids = self._build_solved_records(user_first_ac, by_problem, by_problem_period, problems, user.id)
# 查找 flowchart submissions
flowcharts_query = FlowchartSubmission.objects.filter(
user_id=user,
status=FlowchartSubmissionStatus.COMPLETED,
)
# 添加时间范围过滤
if start:
flowcharts_query = flowcharts_query.filter(create_time__gte=start)
if end:
flowcharts_query = flowcharts_query.filter(create_time__lte=end)
flowcharts = flowcharts_query.select_related("problem").only(
"id",
"create_time",
"ai_score",
"ai_grade",
"problem___id",
"problem__title",
)
# 按problem分组
problem_groups = defaultdict(list)
for flowchart in flowcharts:
problem_id = flowchart.problem._id
problem_groups[problem_id].append(flowchart)
flowcharts_data = []
for problem_id, submissions in problem_groups.items():
if not submissions:
continue
# 获取第一个提交的基本信息
first_submission = submissions[0]
# 计算统计数据
scores = [s.ai_score for s in submissions if s.ai_score is not None]
times = [s.create_time for s in submissions]
# 找到最高分和对应的等级
best_score = max(scores) if scores else 0
best_submission = next((s for s in submissions if s.ai_score == best_score), submissions[0])
best_grade = best_submission.ai_grade or ""
# 计算平均分
avg_score = sum(scores) / len(scores) if scores else 0
# 最新提交时间
latest_time = max(times) if times else first_submission.create_time
merged_item = {
"problem__id": problem_id,
"problem_title": first_submission.problem.title,
"submission_count": len(submissions),
"best_score": best_score,
"best_grade": best_grade,
"latest_submission_time": latest_time.isoformat() if latest_time else None,
"avg_score": round(avg_score, 0),
}
flowcharts_data.append(merged_item)
# 按最新提交时间排序
flowcharts_data.sort(key=lambda x: x["latest_submission_time"] or "", reverse=True)
result.update(
{
"solved": solved,
"flowcharts": flowcharts_data,
"grade": calculate_average_grade([s["grade"] for s in solved]),
"tags": self._calculate_top_tags(problems.values()),
"difficulty": self._calculate_difficulty_distribution(problems.values()),
"contest_count": len(set(contest_ids)),
}
)
cache.set(cache_key, result, CACHE_TIMEOUT)
return self.success(result)
def _build_solved_records(self, user_first_ac, by_problem, by_problem_period, problems, user_id):
solved, contest_ids = [], []
for item in user_first_ac:
pid = item["problem_id"]
problem = problems.get(pid)
if not problem:
continue
ranking_list = by_problem.get(pid, [])
rank = find_user_rank(ranking_list, user_id)
period_ranking_list = by_problem_period.get(pid, [])
period_rank = find_user_rank(period_ranking_list, user_id)
if problem.contest_id:
contest_ids.append(problem.contest_id)
solved.append(
{
"problem": {
"display_id": problem._id,
"title": problem.title,
"contest_id": problem.contest_id,
"contest_title": getattr(problem.contest, "title", ""),
},
"ac_time": timezone.localtime(item["first_ac_time"]).isoformat(),
"rank": rank,
"ac_count": len(ranking_list),
"grade": get_grade(period_rank, len(period_ranking_list), reference_count=len(ranking_list)),
"period_rank": period_rank,
"period_ac_count": len(period_ranking_list),
"difficulty": get_difficulty(problem.difficulty),
}
)
return sorted(solved, key=lambda x: x["ac_time"]), contest_ids
def _calculate_top_tags(self, problems):
tags_counter = defaultdict(int)
for problem in problems:
for tag in problem.tags.all():
if tag.name:
tags_counter[tag.name] += 1
return dict(sorted(tags_counter.items(), key=lambda x: x[1], reverse=True)[:5])
def _calculate_difficulty_distribution(self, problems):
diff_counter = {"Low": 0, "Mid": 0, "High": 0}
for problem in problems:
diff_counter[problem.difficulty if problem.difficulty in diff_counter else "Mid"] += 1
return {get_difficulty(k): v for k, v in sorted(diff_counter.items(), key=lambda x: x[1], reverse=True)}
class AIDurationDataAPI(APIView):
@login_required
def get(self, request):
end_iso = request.GET.get("end")
duration = request.GET.get("duration")
username = request.GET.get("username")
user = request.user
if username and request.user.is_teacher_or_above():
try:
user = User.objects.get(username=username)
except User.DoesNotExist:
return self.error("User not found")
cache_key = get_cache_key("ai_duration", user.id, user.class_name or "", end_iso, duration)
cached_result = cache.get(cache_key)
if cached_result:
return self.success(cached_result)
class_user_ids = get_class_user_ids(user)
use_class_scope = bool(user.class_name) and len(class_user_ids) > 1
time_config = self._parse_duration(duration)
start = time_config["rewind"](datetime.fromisoformat(end_iso))
duration_data = []
for i in range(time_config["show_count"]):
start = time_config["advance"](start)
period_end = time_config["advance"](start)
submission_count = Submission.objects.filter(user_id=user.id, create_time__gte=start, create_time__lte=period_end).count()
period_data = {
"unit": time_config["show_unit"],
"index": time_config["show_count"] - 1 - i,
"start": start.isoformat(),
"end": period_end.isoformat(),
"problem_count": 0,
"submission_count": submission_count,
"grade": "",
}
if submission_count > 0:
user_first_ac, _, problem_ids = get_user_first_ac_submissions(
user.id,
start.isoformat(),
period_end.isoformat(),
class_user_ids,
use_class_scope,
include_all_time=False,
)
if user_first_ac:
period_data["problem_count"] = len(problem_ids)
by_problem_period = defaultdict(list)
period_qs = Submission.objects.filter(
result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED],
problem_id__in=problem_ids,
create_time__gte=start,
create_time__lte=period_end,
)
if use_class_scope and class_user_ids:
period_qs = period_qs.filter(user_id__in=class_user_ids)
for row in period_qs.values("user_id", "problem_id").annotate(first_ac_time=Min("create_time")):
by_problem_period[row["problem_id"]].append(row)
for lst in by_problem_period.values():
lst.sort(key=lambda x: (x["first_ac_time"], x["user_id"]))
grades = [
get_grade(
find_user_rank(by_problem_period.get(item["problem_id"], []), user.id),
len(by_problem_period.get(item["problem_id"], [])),
)
for item in user_first_ac
]
period_data["grade"] = calculate_average_grade(grades)
duration_data.append(period_data)
cache.set(cache_key, duration_data, CACHE_TIMEOUT)
return self.success(duration_data)
def _parse_duration(self, duration):
unit, count = duration.split(":")
count = int(count)
# rewind 把结束时间倒推到区间起点advance 前进一格。
# 按月的档位不能用 timedelta月长不固定用 shift_months 处理
configs = {
("months", 2): {
"show_count": 8,
"show_unit": "weeks",
"rewind": lambda dt: dt - timedelta(weeks=9),
"advance": lambda dt: dt + timedelta(weeks=1),
},
("months", 6): {
"show_count": 6,
"show_unit": "months",
"rewind": lambda dt: shift_months(dt, -7),
"advance": lambda dt: shift_months(dt, 1),
},
("years", 1): {
"show_count": 12,
"show_unit": "months",
"rewind": lambda dt: shift_months(dt, -13),
"advance": lambda dt: shift_months(dt, 1),
},
}
return configs.get(
(unit, count),
{
"show_count": 4,
"show_unit": "weeks",
"rewind": lambda dt: dt - timedelta(weeks=5),
"advance": lambda dt: dt + timedelta(weeks=1),
},
)
class AILoginSummaryAPI(APIView):
@login_required
def get(self, request):
user = request.user
end_time = timezone.now()
start_time = self._resolve_start_time(request, user, end_time)
problems_qs = Problem.objects.filter(
create_time__gte=start_time,
create_time__lte=end_time,
contest_id__isnull=True,
visible=True,
)
new_problem_count = problems_qs.count()
submissions_qs = Submission.objects.filter(user_id=user.id, create_time__gte=start_time, create_time__lte=end_time)
submission_count = submissions_qs.count()
accepted_count = submissions_qs.filter(result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED]).count()
solved_count = submissions_qs.filter(result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED]).values("problem_id").distinct().count()
flowchart_submission_count = FlowchartSubmission.objects.filter(user_id=user.id, create_time__gte=start_time, create_time__lte=end_time).count()
summary = {
"start": datetime2str(start_time),
"end": datetime2str(end_time),
"new_problem_count": new_problem_count,
"submission_count": submission_count,
"accepted_count": accepted_count,
"solved_count": solved_count,
"flowchart_submission_count": flowchart_submission_count,
}
analysis = ""
analysis_error = ""
if submission_count >= 3:
analysis, analysis_error = self._get_ai_analysis(summary)
data = {"summary": summary, "analysis": analysis}
if analysis_error:
data["analysis_error"] = analysis_error
return self.success(data)
def _resolve_start_time(self, request, user, end_time):
start_raw = request.session.get("prev_login") or request.GET.get("start")
start_time = parse_datetime(start_raw) if start_raw else None
if start_time and timezone.is_naive(start_time):
start_time = timezone.make_aware(start_time, timezone.get_current_timezone())
if not start_time:
if user.last_login and user.last_login < end_time:
start_time = user.last_login
elif user.create_time:
start_time = user.create_time
else:
start_time = end_time - timedelta(days=7)
if start_time >= end_time:
start_time = end_time - timedelta(days=1)
return start_time
def _get_ai_analysis(self, summary):
try:
client = get_ai_client()
except Exception as exc:
return "", str(exc)
system_prompt = "你是 OnlineJudge 的学习助教。请根据统计数据给出简短分析(1-2句),再给出一行结论,结论用“结论:”开头。"
user_prompt = (
f"时间范围:{summary['start']}{summary['end']}\n"
f"新题目数:{summary['new_problem_count']}\n"
f"提交次数:{summary['submission_count']}\n"
f"AC 次数:{summary['accepted_count']}\n"
f"AC 题目数:{summary['solved_count']}\n"
f"流程图提交数:{summary['flowchart_submission_count']}\n"
)
try:
completion = client.chat.completions.create(
model="deepseek-v4-flash",
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt},
],
extra_body={"thinking": {"type": "disabled"}},
)
except Exception as exc:
return "", str(exc)
if not completion.choices:
return "", ""
content = completion.choices[0].message.content or ""
return content.strip(), ""
class AIAnalysisAPI(APIView):
@login_required
def post(self, request):
details = request.data.get("details")
duration = request.data.get("duration")
client = get_async_ai_client()
system_prompt = (
"你是一个风趣的编程老师,学生使用判题狗平台进行编程练习。"
"请根据学生提供的详细数据和每周数据,给出用户的学习建议,最后写一句鼓励学生的话。"
"请使用 markdown 格式输出,不要在代码块中输出。"
)
user_prompt = f"这段时间内的详细数据: {details}\n(其中部分字段含义是 flowcharts:流程图的提交,solved:代码的提交)\n每周或每月的数据: {duration}"
user = request.user
async def on_complete(full_text):
await AIAnalysis.objects.acreate(
user=user,
provider="deepseek",
model="deepseek-v4-flash",
data={"details": details, "duration": duration},
system_prompt=system_prompt,
user_prompt="这段时间内的详细数据,每周或每月的数据。",
analysis=full_text,
)
return make_sse_response(stream_ai_response(client, system_prompt, user_prompt, on_complete))
class ClassPKAnalysisAPI(APIView):
@teacher_admin_required
def post(self, request):
comparisons = request.data.get("comparisons")
time_range_label = request.data.get("time_range_label", "全部时间")
if not comparisons or len(comparisons) < 2:
return self.error("至少需要2个班级的数据")
client = get_async_ai_client()
system_prompt = (
"你是一位经验丰富的编程教育数据分析专家,专注于职业院校计算机编程教学效果评估。"
"请根据在线评测系统OJ提供的班级对比数据给出深入、专业的分析报告。"
"请使用 markdown 格式输出,不要在代码块中输出。"
)
user_prompt = self._build_prompt(comparisons, time_range_label)
return make_sse_response(stream_ai_response(client, system_prompt, user_prompt))
def _build_prompt(self, comparisons, time_range_label):
def fmt_class(name):
return f"{name[:2]}计算机{name[2:]}"
lines = [
f"时间范围:{time_range_label}",
"",
"## 指标说明",
"- 总AC数班级累计通过的题目总数",
'- 平均AC / 中位数AC平均受尖子生影响中位数反映"典型学生"的真实水平',
"- 前10%/中间80%/后10%均值:梯队分层,判断是尖子班还是均衡班",
"- 优秀率 / 及格率:达到优秀线、及格线的学生比例",
"- 参与度:有过任意提交的学生比例(反映主动性)",
"- AC率提交通过率反映代码编写质量",
"- 综合分综合多项指标计算的总评分满分100",
"",
"## 班级数据",
]
for i, c in enumerate(comparisons):
class_display = fmt_class(c["class_name"])
lines.append(f"\n### 第{i + 1}名:{class_display}(综合分 {c['composite_score']:.1f}")
lines.append(f"- 人数:{c['user_count']}")
lines.append(f"- 总AC数{c['total_ac']},总提交数:{c['total_submission']}AC率{c['ac_rate']:.1f}%")
lines.append(f"- 平均AC{c['avg_ac']:.2f}中位数AC{c['median_ac']:.2f}")
lines.append(f"- Q1{c['q1_ac']:.2f}Q3{c['q3_ac']:.2f}IQR四分位距{c['iqr']:.2f},标准差:{c['std_dev']:.2f}")
lines.append(f"- 前10%均值:{c['top_10_avg']:.2f}中间80%均值:{c['middle_80_avg']:.2f}后10%均值:{c['bottom_10_avg']:.2f}")
lines.append(f"- 优秀率:{c['excellent_rate']:.1f}%,及格率:{c['pass_rate']:.1f}%,参与度:{c['active_rate']:.1f}%")
if c.get("recent_total_ac") is not None:
lines.append(
f"- 时间段内新增总AC {c['recent_total_ac']}"
f"平均AC {c.get('recent_avg_ac', 0):.2f}"
f"中位数AC {c.get('recent_median_ac', 0):.2f}"
f"前10%平均 {c.get('recent_top_10_avg', 0):.2f}"
f"活跃学生数 {c.get('recent_active_count', 0)}"
)
lines += [
"",
"## 请从以下7个维度分析输出中文报告",
"",
"**1. 总体排名**:直接给出排名,说明各班综合分差距大小。",
"",
"**2. 参与积极性**:对比参与度和总提交数,谁的班学生更积极主动?",
"",
'**3. 典型学生水平**重点用中位数AC数对比而非平均值分析谁班的"普通学生"更强。若均值明显高于中位数,说明均值被少数强者拉高,需指出。',
"",
'**4. 班级内部均衡性**结合标准差、IQR、前10%与后10%差距,判断哪个班是"均衡型",哪个班是"两极型"',
"",
"**5. 梯队深度对比**对比各班前10%均值尖子生天花板和后10%均值(薄弱学生水平),分析各班在培养尖子生和帮扶后进生上的差异。",
"",
'**6. 代码提交质量**对比AC率是否有班级存在"凑提交次数但不思考"的问题?',
"",
"**7. 综合结论与建议**用1句话明确说明胜负对落后班级给出2~3条具体可操作的改进建议点出领先班级1条值得借鉴的做法。",
"",
"分析对象是班级任课教师,语言专业但不过分学术。",
]
return "\n".join(lines)
class SingleClassAnalysisAPI(APIView):
@login_required
def post(self, request):
comparison = request.data.get("comparison")
if not comparison:
return self.error("缺少班级数据")
client = get_async_ai_client()
class_name = comparison.get("class_name", "")
class_display = f"{class_name[:2]}计算机{class_name[2:]}"
system_prompt = (
"你是一位经验丰富的编程教育数据分析专家,专注于职业院校计算机编程教学效果评估。"
"请根据在线评测系统OJ提供的班级数据给出深入、专业的分析报告。"
"请使用 markdown 格式输出,不要在代码块中输出。"
)
lines = [
f"## 班级:{class_display}",
"",
"## 指标说明",
"- 总AC数班级累计通过的题目总数",
'- 平均AC / 中位数AC平均受尖子生影响中位数反映"典型学生"的真实水平',
"- 前10%/中间80%/后10%均值:梯队分层",
"- 优秀率 / 及格率分别基于班级内部Q3/Q1阈值统计",
"- 参与度:有过任意提交的学生比例",
"- AC率提交通过率",
"- 综合分综合多项指标计算的总评分满分100",
"",
"## 班级数据",
f"- 人数:{comparison['user_count']}",
f"- 总AC数{comparison['total_ac']},总提交数:{comparison['total_submission']}AC率{comparison['ac_rate']:.1f}%",
f"- 平均AC{comparison['avg_ac']:.2f}中位数AC{comparison['median_ac']:.2f}",
f"- Q1{comparison['q1_ac']:.2f}Q3{comparison['q3_ac']:.2f}IQR{comparison['iqr']:.2f},标准差:{comparison['std_dev']:.2f}",
f"- 前10%均值:{comparison['top_10_avg']:.2f}中间80%均值:{comparison['middle_80_avg']:.2f}后10%均值:{comparison['bottom_10_avg']:.2f}",
f"- 优秀率:{comparison['excellent_rate']:.1f}%,及格率:{comparison['pass_rate']:.1f}%,参与度:{comparison['active_rate']:.1f}%",
f"- 综合分:{comparison['composite_score']:.1f}",
"",
"## 请从以下5个维度分析输出中文报告",
"",
"**1. 整体水平**基于总AC数、均值与中位数评价班级的整体编程水平。若均值明显高于中位数需指出被少数强者拉高的情况。",
"",
"**2. 参与积极性**结合参与度和AC率评价学生的学习主动性和代码质量。",
"",
'**3. 班级内部均衡性**结合标准差、IQR、前10%与后10%差距,判断是"均衡型"还是"两极型"班级。',
"",
"**4. 梯队分析**分析前10%尖子生、中间80%中坚力量、后10%(需帮扶学生)的水平差异,给出针对性建议。",
"",
"**5. 改进建议**结合优秀率、及格率等数据给出3条具体可操作的教学改进建议。最后用一句话鼓励这个班级。",
"",
"分析对象是班级任课教师,语言专业但不过分学术。",
]
user_prompt = "\n".join(lines)
return make_sse_response(stream_ai_response(client, system_prompt, user_prompt))
class AIPinnedReportAPI(APIView):
@login_required
def get(self, request):
try:
report = AIAnalysis.objects.get(user=request.user, is_pinned=True)
except AIAnalysis.DoesNotExist:
return self.success(None)
return self.success(AIAnalysisDetailSerializer(report).data)
class AIHintAPI(APIView):
@login_required
def post(self, request):
submission_id = request.data.get("submission_id")
if not submission_id:
return self.error("submission_id is required")
try:
submission = Submission.objects.get(id=submission_id, user_id=request.user.id)
except Submission.DoesNotExist:
return self.error("Submission not found")
problem = submission.problem
client = get_async_ai_client()
# 获取参考答案(同语言优先,否则取第一个)
answers = problem.answers or []
ref_answer = next(
(a["code"] for a in answers if a["language"] == submission.language),
answers[0]["code"] if answers else "",
)
system_prompt = (
"你是编程助教。你知道题目的参考答案,请按照以下规则给学生提示:\n\n"
"【核心规则】\n"
"- 【绝对禁止】直接给出答案或核心算法代码,也不能暗示完整解法。\n"
"- 提示要循序渐进:先指出问题所在的方向,再给出一个小的思考点,让学生自己推导。\n"
"- 对照参考答案分析学生代码,找出最关键的一个问题重点提示,不要一次列出所有问题。\n\n"
"【输入处理例外】\n"
"- 如果学生的代码在【读取输入】部分有错误(例如:输入格式解析错误、"
"未正确读取多组输入、split/scanf 使用有误等),"
"则【直接给出正确的输入读取代码片段】,并解释为什么这样写。"
"输入处理不属于算法核心,可以直接告诉学生。\n\n"
"【回复格式】\n"
"语气鼓励,使用 Markdown 格式回复简洁不超过6句话"
)
user_prompt = (
f"题目:{problem.title}\n"
f"题目描述:{problem.description[:500]}\n"
f"参考答案(仅供你分析,不可透露给学生):\n```\n{ref_answer[:2000]}\n```\n"
f"学生提交语言:{submission.language}\n"
f"判题结果:{submission.result}\n"
f"错误信息:{submission.statistic_info.get('err_info', '')}\n"
f"学生代码:\n```\n{submission.code[:2000]}\n```"
)
return make_sse_response(stream_ai_response(client, system_prompt, user_prompt))
class AIHeatmapDataAPI(APIView):
@login_required
def get(self, request):
username = request.GET.get("username")
user = request.user
if username and request.user.is_teacher_or_above():
try:
user = User.objects.get(username=username)
except User.DoesNotExist:
return self.error("User not found")
end = timezone.now()
today = end.date().isoformat()
cache_key = get_cache_key("ai_heatmap", user.id, user.class_name or "", today)
cached_result = cache.get(cache_key)
if cached_result:
return self.success(cached_result)
start = end - timedelta(days=365)
# 使用单次查询获取所有数据,按日期分组统计
submission_counts = (
Submission.objects.filter(user_id=user.id, create_time__gte=start, create_time__lte=end)
.annotate(date=TruncDate("create_time"))
.values("date")
.annotate(count=Count("id"))
.order_by("date")
)
# 将查询结果转换为字典,便于快速查找
submission_dict = {item["date"]: item["count"] for item in submission_counts}
# 生成365天的热力图数据
heatmap_data = []
current_date = start.date()
for i in range(365):
day_date = current_date + timedelta(days=i)
submission_count = submission_dict.get(day_date, 0)
heatmap_data.append(
{
"timestamp": int(datetime.combine(day_date, datetime.min.time()).timestamp() * 1000),
"value": submission_count,
}
)
cache.set(cache_key, heatmap_data, CACHE_TIMEOUT)
return self.success(heatmap_data)

View File

@@ -1,19 +0,0 @@
# Generated by Django 6.0 on 2026-04-23 20:07
from django.conf import settings
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('announcement', '0001_initial'),
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
]
operations = [
migrations.AddIndex(
model_name='announcement',
index=models.Index(fields=['visible', '-top', '-create_time'], name='announcement_list_idx'),
),
]

View File

@@ -17,10 +17,4 @@ class Announcement(models.Model):
class Meta:
db_table = "announcement"
ordering = (
"-top",
"-create_time",
)
indexes = [
models.Index(fields=["visible", "-top", "-create_time"], name="announcement_list_idx"),
]
ordering = ("-top", "-create_time",)

View File

@@ -25,7 +25,7 @@ class AnnouncementListSerializer(serializers.ModelSerializer):
class Meta:
model = Announcement
exclude = ["content"]
exclude = ['content']
class EditAnnouncementSerializer(serializers.Serializer):

48
announcement/tests.py Normal file
View File

@@ -0,0 +1,48 @@
from utils.api.tests import APITestCase
from .models import Announcement
class AnnouncementAdminTest(APITestCase):
def setUp(self):
self.user = self.create_super_admin()
self.url = self.reverse("announcement_admin_api")
def test_announcement_list(self):
response = self.client.get(self.url)
self.assertSuccess(response)
def create_announcement(self):
return self.client.post(self.url, data={"title": "test", "content": "test", "visible": True})
def test_create_announcement(self):
resp = self.create_announcement()
self.assertSuccess(resp)
return resp
def test_edit_announcement(self):
data = {"id": self.create_announcement().data["data"]["id"], "title": "ahaha", "content": "test content",
"visible": False}
resp = self.client.put(self.url, data=data)
self.assertSuccess(resp)
resp_data = resp.data["data"]
self.assertEqual(resp_data["title"], "ahaha")
self.assertEqual(resp_data["content"], "test content")
self.assertEqual(resp_data["visible"], False)
def test_delete_announcement(self):
id = self.test_create_announcement().data["data"]["id"]
resp = self.client.delete(self.url + "?id=" + str(id))
self.assertSuccess(resp)
self.assertFalse(Announcement.objects.filter(id=id).exists())
class AnnouncementAPITest(APITestCase):
def setUp(self):
self.user = self.create_super_admin()
Announcement.objects.create(title="title", content="content", visible=True, created_by=self.user)
self.url = self.reverse("announcement_api")
def test_get_announcement_list(self):
resp = self.client.get(self.url)
self.assertSuccess(resp)

View File

@@ -1,11 +1,12 @@
from account.decorators import super_admin_required
from utils.api import APIView, validate_serializer
from announcement.models import Announcement
from announcement.serializers import (
AnnouncementSerializer,
CreateAnnouncementSerializer,
EditAnnouncementSerializer,
)
from utils.api import APIView, validate_serializer
class AnnouncementAdminAPI(APIView):
@@ -52,7 +53,9 @@ class AnnouncementAdminAPI(APIView):
announcement = Announcement.objects.all().order_by("-create_time")
if request.GET.get("visible") == "true":
announcement = announcement.filter(visible=True)
return self.success(self.paginate_data(request, announcement, AnnouncementSerializer))
return self.success(
self.paginate_data(request, announcement, AnnouncementSerializer)
)
@super_admin_required
def delete(self, request):

View File

@@ -1,19 +1,20 @@
from utils.api import APIView
from announcement.models import Announcement
from announcement.serializers import AnnouncementListSerializer, AnnouncementSerializer
from utils.api import AsyncAPIView
from announcement.serializers import AnnouncementSerializer, AnnouncementListSerializer
class AnnouncementAPI(AsyncAPIView):
async def get(self, request):
class AnnouncementAPI(APIView):
def get(self, request):
id = request.GET.get("id")
if id:
try:
announcement = await Announcement.objects.select_related("created_by").filter(id=id, visible=True).afirst()
if announcement is None:
raise Announcement.DoesNotExist
return self.success(await self.async_serialize_data(AnnouncementSerializer, announcement))
announcement = Announcement.objects.get(id=id, visible=True)
return self.success(AnnouncementSerializer(announcement).data)
except Announcement.DoesNotExist:
return self.error("Announcement does not exist")
announcements = Announcement.objects.select_related("created_by").filter(visible=True)
return self.success(await self.async_paginate_data(request, announcements, AnnouncementListSerializer))
announcements = Announcement.objects.filter(visible=True)
return self.success(
self.paginate_data(request, announcements, AnnouncementListSerializer)
)

View File

@@ -1,35 +0,0 @@
from tree_sitter import Parser
from .engines import get_engine
from .mappings import get_language, get_mapping
def check_ast(code: str, language: str, rules: list[dict]) -> tuple[bool, list[dict]]:
if not rules:
return True, []
ts_language = get_language(language)
if ts_language is None:
return True, []
mapping = get_mapping(language)
try:
parser = Parser(ts_language)
tree = parser.parse(code.encode("utf-8"))
except Exception:
return True, []
results = []
all_passed = True
for rule in rules:
engine = get_engine(rule.get("engine", ""))
if engine is None:
continue
errors = engine.check(tree, rule, language, mapping)
passed = len(errors) == 0
if not passed:
all_passed = False
results.append({"description": engine.describe(rule, language, mapping), "passed": passed})
return all_passed, results

View File

@@ -1,23 +0,0 @@
from .function_call import CountFunctionCallEngine, MustCallFunctionEngine, MustNotCallFunctionEngine
from .method_call import MustCallMethodEngine, MustNotCallMethodEngine
from .nesting import MustHaveNestingEngine
from .node_count import CountNodeEngine
from .node_exists import MustExistNodeEngine, MustNotExistNodeEngine
from .operator import MustUseOperatorEngine
ENGINES = {
"must_exist_node": MustExistNodeEngine(),
"must_not_exist_node": MustNotExistNodeEngine(),
"count_node": CountNodeEngine(),
"must_call_function": MustCallFunctionEngine(),
"must_not_call_function": MustNotCallFunctionEngine(),
"count_function_call": CountFunctionCallEngine(),
"must_call_method": MustCallMethodEngine(),
"must_not_call_method": MustNotCallMethodEngine(),
"must_use_operator": MustUseOperatorEngine(),
"must_have_nesting": MustHaveNestingEngine(),
}
def get_engine(name: str):
return ENGINES.get(name)

View File

@@ -1,21 +0,0 @@
class BaseEngine:
@staticmethod
def collect_nodes(node, node_type):
results = []
if node.type == node_type:
results.append(node)
for child in node.children:
results.extend(BaseEngine.collect_nodes(child, node_type))
return results
@staticmethod
def has_node(node, node_type):
if node.type == node_type:
return True
return any(BaseEngine.has_node(child, node_type) for child in node.children)
def check(self, tree, rule, language, mapping) -> list[str]:
raise NotImplementedError
def describe(self, rule, language, mapping) -> str:
raise NotImplementedError

View File

@@ -1,78 +0,0 @@
from .base import BaseEngine
CALL_NODE_TYPES = {
"Python3": "call",
"C": "call_expression",
}
class _FunctionCallBase(BaseEngine):
def _find_function_calls(self, root, func_name, language):
call_type = CALL_NODE_TYPES.get(language, "call")
calls = self.collect_nodes(root, call_type)
matches = []
for call in calls:
func_node = call.child_by_field_name("function")
if func_node and func_node.type == "identifier" and func_node.text.decode() == func_name:
matches.append(call)
return matches
class MustCallFunctionEngine(_FunctionCallBase):
def _message(self, rule):
return rule.get("message") or f"必须调用 {rule['target']}()"
def check(self, tree, rule, language, mapping):
if not self._find_function_calls(tree.root_node, rule["target"], language):
return [self._message(rule)]
return []
def describe(self, rule, language, mapping):
return self._message(rule)
class MustNotCallFunctionEngine(_FunctionCallBase):
def _message(self, rule):
return rule.get("message") or f"不能调用 {rule['target']}()"
def check(self, tree, rule, language, mapping):
if self._find_function_calls(tree.root_node, rule["target"], language):
return [self._message(rule)]
return []
def describe(self, rule, language, mapping):
return self._message(rule)
class CountFunctionCallEngine(_FunctionCallBase):
def _message(self, rule, count):
target = rule["target"]
exact = rule.get("exact")
if exact is not None and count != exact:
return rule.get("message") or f"{target}() 需要调用 {exact} 次,当前 {count}"
min_count = rule.get("min")
max_count = rule.get("max")
if min_count is not None and count < min_count:
return rule.get("message") or f"{target}() 至少调用 {min_count} 次,当前 {count}"
if max_count is not None and count > max_count:
return rule.get("message") or f"{target}() 至多调用 {max_count} 次,当前 {count}"
return None
def check(self, tree, rule, language, mapping):
count = len(self._find_function_calls(tree.root_node, rule["target"], language))
msg = self._message(rule, count)
return [msg] if msg else []
def describe(self, rule, language, mapping):
target = rule["target"]
if rule.get("message"):
return rule["message"]
exact = rule.get("exact")
if exact is not None:
return f"{target}() 调用 {exact}"
parts = []
if rule.get("min") is not None:
parts.append(f"至少 {rule['min']}")
if rule.get("max") is not None:
parts.append(f"至多 {rule['max']}")
return f"{target}() " + "".join(parts)

View File

@@ -1,48 +0,0 @@
from .base import BaseEngine
CALL_NODE_TYPES = {
"Python3": "call",
"C": "call_expression",
}
class _MethodCallBase(BaseEngine):
def _find_method_calls(self, root, method_name, language):
if language == "C":
return []
call_type = CALL_NODE_TYPES.get(language, "call")
calls = self.collect_nodes(root, call_type)
matches = []
for call in calls:
func_node = call.child_by_field_name("function")
if func_node and func_node.type == "attribute":
attr_node = func_node.child_by_field_name("attribute")
if attr_node and attr_node.text.decode() == method_name:
matches.append(call)
return matches
class MustCallMethodEngine(_MethodCallBase):
def _message(self, rule):
return rule.get("message") or f"必须调用 .{rule['target']}()"
def check(self, tree, rule, language, mapping):
if not self._find_method_calls(tree.root_node, rule["target"], language):
return [self._message(rule)]
return []
def describe(self, rule, language, mapping):
return self._message(rule)
class MustNotCallMethodEngine(_MethodCallBase):
def _message(self, rule):
return rule.get("message") or f"不能调用 .{rule['target']}()"
def check(self, tree, rule, language, mapping):
if self._find_method_calls(tree.root_node, rule["target"], language):
return [self._message(rule)]
return []
def describe(self, rule, language, mapping):
return self._message(rule)

View File

@@ -1,34 +0,0 @@
from ast_checker.labels import label
from .base import BaseEngine
class MustHaveNestingEngine(BaseEngine):
def _has_inner_in_subtree(self, node, inner_type):
for child in node.children:
if self.has_node(child, inner_type):
return True
return False
def _message(self, rule):
if rule.get("message"):
return rule["message"]
outer = rule.get("outer", "")
inner = rule.get("inner", "")
outer_label = label(outer)
inner_label = label(inner)
if outer == inner:
return f"必须使用 {outer_label} 嵌套"
return f"必须在 {outer_label} 中嵌套使用 {inner_label}"
def check(self, tree, rule, language, mapping):
outer_type = mapping.get(rule["outer"], rule["outer"])
inner_type = mapping.get(rule["inner"], rule["inner"])
outer_nodes = self.collect_nodes(tree.root_node, outer_type)
for outer_node in outer_nodes:
if self._has_inner_in_subtree(outer_node, inner_type):
return []
return [self._message(rule)]
def describe(self, rule, language, mapping):
return self._message(rule)

View File

@@ -1,39 +0,0 @@
from ast_checker.labels import label
from .base import BaseEngine
class CountNodeEngine(BaseEngine):
def _message(self, rule, count):
name = rule.get("label") or label(rule["target"])
exact = rule.get("exact")
if exact is not None and count != exact:
return rule.get("message") or f"{name} 需要出现 {exact} 次,当前 {count}"
min_count = rule.get("min")
max_count = rule.get("max")
if min_count is not None and count < min_count:
return rule.get("message") or f"{name} 至少出现 {min_count} 次,当前 {count}"
if max_count is not None and count > max_count:
return rule.get("message") or f"{name} 至多出现 {max_count} 次,当前 {count}"
return None
def check(self, tree, rule, language, mapping):
target = rule["target"]
node_type = mapping.get(target, target)
count = len(self.collect_nodes(tree.root_node, node_type))
msg = self._message(rule, count)
return [msg] if msg else []
def describe(self, rule, language, mapping):
name = rule.get("label") or label(rule["target"])
if rule.get("message"):
return rule["message"]
exact = rule.get("exact")
if exact is not None:
return f"{name} 出现 {exact}"
parts = []
if rule.get("min") is not None:
parts.append(f"至少 {rule['min']}")
if rule.get("max") is not None:
parts.append(f"至多 {rule['max']}")
return f"{name} " + "".join(parts)

View File

@@ -1,31 +0,0 @@
from ast_checker.labels import label
from .base import BaseEngine
class MustExistNodeEngine(BaseEngine):
def _message(self, rule):
return rule.get("message") or f"必须使用 {rule.get('label') or label(rule['target'])}"
def check(self, tree, rule, language, mapping):
node_type = mapping.get(rule["target"], rule["target"])
if not self.has_node(tree.root_node, node_type):
return [self._message(rule)]
return []
def describe(self, rule, language, mapping):
return self._message(rule)
class MustNotExistNodeEngine(BaseEngine):
def _message(self, rule):
return rule.get("message") or f"不能使用 {rule.get('label') or label(rule['target'])}"
def check(self, tree, rule, language, mapping):
node_type = mapping.get(rule["target"], rule["target"])
if self.has_node(tree.root_node, node_type):
return [self._message(rule)]
return []
def describe(self, rule, language, mapping):
return self._message(rule)

View File

@@ -1,15 +0,0 @@
from .base import BaseEngine
class MustUseOperatorEngine(BaseEngine):
def _message(self, rule):
return rule.get("message") or f"必须使用 {rule['target']} 运算符"
def check(self, tree, rule, language, mapping):
mapped_op = mapping.get(rule["target"], rule["target"])
if not self.has_node(tree.root_node, mapped_op):
return [self._message(rule)]
return []
def describe(self, rule, language, mapping):
return self._message(rule)

View File

@@ -1,21 +0,0 @@
TARGET_LABELS: dict[str, str] = {
"for_loop": "for 循环",
"while_loop": "while 循环",
"if_statement": "if 条件",
"else_clause": "else 子句",
"function_definition": "函数定义",
"return": "return 语句",
"break": "break 语句",
"continue": "continue 语句",
"list_comprehension": "列表推导式",
"list_literal": "列表",
"dict_literal": "字典",
"set_literal": "集合",
"f_string": "f-string",
"try_except": "try-except",
"class_definition": "类定义",
}
def label(target: str) -> str:
return TARGET_LABELS.get(target, target)

View File

@@ -1,37 +0,0 @@
from tree_sitter import Language
from .c import C_MAPPING
from .python import PYTHON_MAPPING
_MAPPINGS = {
"Python3": PYTHON_MAPPING,
"C": C_MAPPING,
}
_LANGUAGES: dict[str, Language] = {}
def _init_languages():
try:
import tree_sitter_python as tspython
_LANGUAGES["Python3"] = Language(tspython.language())
except ImportError:
pass
try:
import tree_sitter_c as tsc
_LANGUAGES["C"] = Language(tsc.language())
except ImportError:
pass
_init_languages()
def get_mapping(language: str) -> dict:
return _MAPPINGS.get(language, {})
def get_language(language: str) -> Language | None:
return _LANGUAGES.get(language)

View File

@@ -1,39 +0,0 @@
C_MAPPING = {
"for_loop": "for_statement",
"while_loop": "while_statement",
"do_while": "do_statement",
"if_statement": "if_statement",
"else_clause": "else_clause",
"break": "break_statement",
"continue": "continue_statement",
"function_definition": "function_definition",
"return": "return_statement",
"switch_statement": "switch_statement",
"case_statement": "case_statement",
"assignment": "assignment_expression",
"struct": "struct_specifier",
"include": "preproc_include",
"+": "+",
"-": "-",
"*": "*",
"/": "/",
"%": "%",
"+=": "+=",
"-=": "-=",
"*=": "*=",
"/=": "/=",
"%=": "%=",
"==": "==",
"!=": "!=",
">": ">",
">=": ">=",
"<": "<",
"<=": "<=",
"and": "&&",
"or": "||",
"not": "!",
"&": "&",
"|": "|",
"++": "++",
"--": "--",
}

View File

@@ -1,45 +0,0 @@
PYTHON_MAPPING = {
"for_loop": "for_statement",
"while_loop": "while_statement",
"if_statement": "if_statement",
"else_clause": "else_clause",
"elif_clause": "elif_clause",
"break": "break_statement",
"continue": "continue_statement",
"function_definition": "function_definition",
"return": "return_statement",
"try_except": "try_statement",
"with_statement": "with_statement",
"list_comprehension": "list_comprehension",
"list_literal": "list",
"dict_literal": "dictionary",
"set_literal": "set",
"f_string": "format_string",
"import": "import_statement",
"import_from": "import_from_statement",
"assignment": "assignment",
"class_definition": "class_definition",
"+": "+",
"-": "-",
"*": "*",
"/": "/",
"//": "//",
"%": "%",
"**": "**",
"+=": "+=",
"-=": "-=",
"*=": "*=",
"/=": "/=",
"%=": "%=",
"==": "==",
"!=": "!=",
">": ">",
">=": ">=",
"<": "<",
"<=": "<=",
"and": "and",
"or": "or",
"not": "not",
"&": "&",
"|": "|",
}

View File

View File

@@ -1 +0,0 @@
# Register your models here.

View File

@@ -1,7 +0,0 @@
from django.apps import AppConfig
class ClassPkConfig(AppConfig):
default_auto_field = 'django.db.models.BigAutoField'
name = 'class_pk'
verbose_name = '班级PK'

View File

@@ -1,2 +0,0 @@
# 空文件

View File

@@ -1,2 +0,0 @@
# 如果需要存储班级PK历史记录可以在这里定义模型
# 目前暂时不需要,因为都是实时计算

View File

@@ -1,3 +0,0 @@
# 如果需要序列化器,可以在这里定义
# 目前使用APIView的paginate_data方法暂时不需要

View File

@@ -1,2 +0,0 @@
# 空文件

View File

@@ -1,9 +0,0 @@
from django.urls import path
from ..views.oj import ClassPKAPI, ClassRankAPI, UserClassRankAPI
urlpatterns = [
path("class_rank", ClassRankAPI.as_view()),
path("user_class_rank", UserClassRankAPI.as_view()),
path("class_pk", ClassPKAPI.as_view()),
]

View File

@@ -1,337 +0,0 @@
import math
import statistics
from datetime import datetime
from django.db.models import Avg, Sum
from django.utils import timezone
from account.decorators import login_required
from account.models import AdminType, User, UserProfile
from submission.models import JudgeStatus, Submission
from utils.api import APIView
class ClassRankAPI(APIView):
"""获取班级排名列表"""
def get(self, request):
# 获取年级参数
grade = int(request.GET.get("grade"))
# 获取所有有用户的班级
classes = (
User.objects.filter(
class_name__isnull=False,
is_disabled=False,
admin_type__in=[AdminType.REGULAR_USER, AdminType.STUDENT_ADMIN],
class_name__startswith=str(grade),
)
.values("class_name")
.distinct()
)
class_stats = []
for class_info in classes:
class_name = class_info["class_name"]
users = User.objects.filter(
class_name=class_name,
is_disabled=False,
admin_type__in=[AdminType.REGULAR_USER, AdminType.STUDENT_ADMIN],
)
user_ids = list(users.values_list("id", flat=True))
profiles = UserProfile.objects.filter(user_id__in=user_ids)
total_ac = profiles.aggregate(total=Sum("accepted_number"))["total"] or 0
total_submission = profiles.aggregate(total=Sum("submission_number"))["total"] or 0
avg_ac = profiles.aggregate(avg=Avg("accepted_number"))["avg"] or 0
user_count = users.count()
class_stats.append(
{
"class_name": class_name,
"user_count": user_count,
"total_ac": int(total_ac),
"total_submission": int(total_submission),
"avg_ac": round(avg_ac, 2),
"ac_rate": round(total_ac / total_submission * 100, 2) if total_submission > 0 else 0,
}
)
# 按总AC数排序
class_stats.sort(key=lambda x: (-x["total_ac"], x["total_submission"]))
# 添加排名
for i, stat in enumerate(class_stats):
stat["rank"] = i + 1
return self.success(class_stats)
class UserClassRankAPI(APIView):
"""获取用户在班级中的排名"""
@login_required
def get(self, request):
user = request.user
if not user.class_name:
return self.error("用户没有班级信息")
scope = request.GET.get("scope", "").lower()
show_all = scope == "all"
try:
limit = int(request.GET.get("limit", "10"))
except ValueError:
limit = 10
if limit <= 0 or limit > 250:
limit = 10
try:
offset = int(request.GET.get("offset", "0"))
except ValueError:
offset = 0
if offset < 0:
offset = 0
# 获取同班所有用户
class_users = User.objects.filter(
class_name=user.class_name,
is_disabled=False,
admin_type__in=[AdminType.REGULAR_USER, AdminType.STUDENT_ADMIN],
).select_related("userprofile")
user_ranks = []
for class_user in class_users:
profile = class_user.userprofile
user_ranks.append(
{
"user_id": class_user.id,
"username": class_user.username,
"accepted_number": profile.accepted_number,
"submission_number": profile.submission_number,
}
)
# 按AC数排序
user_ranks.sort(key=lambda x: (-x["accepted_number"], x["submission_number"]))
# 添加排名
my_rank = -1
for i, rank_info in enumerate(user_ranks):
rank_info["rank"] = i + 1
if rank_info["user_id"] == user.id:
my_rank = i + 1
trimmed_ranks = user_ranks
if not show_all and my_rank > 0 and len(user_ranks) > 10:
center_index = my_rank - 1
start = max(0, center_index - 5)
end = start + 10
if end > len(user_ranks):
end = len(user_ranks)
start = max(0, end - 10)
trimmed_ranks = user_ranks[start:end]
elif show_all:
trimmed_ranks = user_ranks[offset : offset + limit]
return self.success(
{
"class_name": user.class_name,
"my_rank": my_rank,
"total": len(user_ranks),
"ranks": trimmed_ranks,
}
)
class ClassPKAPI(APIView):
"""班级PK比较 - 多维度教育评价"""
def post(self, request):
class_names = request.data.get("class_name", [])
if not class_names or len(class_names) < 1:
return self.error("至少需要选择1个班级")
# 获取时间段参数
start_time = request.data.get("start_time")
end_time = request.data.get("end_time")
# 将时间字符串转换为datetime对象
# 处理空字符串、None 或 undefined 的情况
if start_time and isinstance(start_time, str) and start_time.strip():
try:
start_time = datetime.fromisoformat(start_time.replace("Z", "+00:00"))
if timezone.is_naive(start_time):
start_time = timezone.make_aware(start_time)
except (ValueError, AttributeError):
start_time = None
else:
start_time = None
if end_time and isinstance(end_time, str) and end_time.strip():
try:
end_time = datetime.fromisoformat(end_time.replace("Z", "+00:00"))
if timezone.is_naive(end_time):
end_time = timezone.make_aware(end_time)
except (ValueError, AttributeError):
end_time = None
else:
end_time = None
class_comparisons = []
# 预计算全局阈值所有参与PK班级的学生AC数合并
all_user_ids = list(
User.objects.filter(
class_name__in=class_names,
is_disabled=False,
admin_type__in=[AdminType.REGULAR_USER, AdminType.STUDENT_ADMIN],
).values_list("id", flat=True)
)
all_ac_list = sorted(
[p.accepted_number for p in UserProfile.objects.filter(user_id__in=all_user_ids)],
reverse=True,
)
if len(all_ac_list) > 1:
_quantiles = statistics.quantiles(all_ac_list, n=4)
global_q1 = _quantiles[0]
global_q3 = _quantiles[2]
else:
global_q1 = all_ac_list[0] if all_ac_list else 0
global_q3 = all_ac_list[0] if all_ac_list else 0
for class_name in class_names:
users = User.objects.filter(
class_name=class_name,
is_disabled=False,
admin_type__in=[AdminType.REGULAR_USER, AdminType.STUDENT_ADMIN],
)
user_ids = list(users.values_list("id", flat=True))
# 获取所有学生的AC数列表用于统计计算
profiles = UserProfile.objects.filter(user_id__in=user_ids)
ac_list = sorted([p.accepted_number for p in profiles], reverse=True)
submission_list = sorted([p.submission_number for p in profiles], reverse=True)
user_count = len(ac_list)
if user_count == 0:
continue
# 基础统计
total_ac = sum(ac_list)
total_submission = sum(submission_list)
avg_ac = statistics.mean(ac_list) if ac_list else 0
# 中位数和分位数
median_ac = statistics.median(ac_list) if ac_list else 0
q1_ac = statistics.quantiles(ac_list, n=4)[0] if len(ac_list) > 1 else 0
q3_ac = statistics.quantiles(ac_list, n=4)[2] if len(ac_list) > 1 else 0
iqr = q3_ac - q1_ac
# 标准差
std_dev = statistics.stdev(ac_list) if len(ac_list) > 1 else 0
# 前10%和后10%统计
top_10_count = max(1, math.ceil(user_count * 0.10))
bottom_10_count = max(1, math.ceil(user_count * 0.10))
top_10_avg = statistics.mean(ac_list[:top_10_count]) if top_10_count > 0 else 0
bottom_10_avg = statistics.mean(ac_list[-bottom_10_count:]) if bottom_10_count > 0 else 0
# 中间80%均值截尾均值去掉前10%和后10%
if top_10_count + bottom_10_count < user_count:
middle_list = ac_list[top_10_count:-bottom_10_count]
else:
middle_list = ac_list
middle_80_avg = statistics.mean(middle_list) if middle_list else avg_ac
# 优秀率AC数 >= 全局Q3即超过PK组所有学生的前25%
excellent_count = sum(1 for ac in ac_list if ac >= global_q3)
excellent_rate = (excellent_count / user_count * 100) if user_count > 0 else 0
# 及格率AC数 >= 全局Q1即超过PK组所有学生的后25%
pass_count = sum(1 for ac in ac_list if ac >= global_q1)
pass_rate = (pass_count / user_count * 100) if user_count > 0 else 0
# 参与度(有提交记录的学生比例)
active_count = sum(1 for sub in submission_list if sub > 0)
active_rate = (active_count / user_count * 100) if user_count > 0 else 0
# 时间段内的统计(如果提供了时间段)
recent_stats = {}
if start_time and end_time:
submissions = Submission.objects.filter(
user_id__in=user_ids,
create_time__gte=start_time,
create_time__lte=end_time,
)
recent_ac = submissions.filter(result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED]).values("user_id", "problem_id").distinct().count()
recent_submission = submissions.count()
# 时间段内的用户AC数列表
recent_user_ac = {}
for user_id in user_ids:
user_recent_ac = submissions.filter(user_id=user_id, result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED]).values("problem_id").distinct().count()
recent_user_ac[user_id] = user_recent_ac
recent_ac_list = sorted(recent_user_ac.values(), reverse=True)
if recent_ac_list:
recent_stats = {
"recent_total_ac": recent_ac,
"recent_total_submission": recent_submission,
"recent_avg_ac": statistics.mean(recent_ac_list),
"recent_median_ac": statistics.median(recent_ac_list),
"recent_top_10_avg": statistics.mean(recent_ac_list[: max(1, math.ceil(len(recent_ac_list) * 0.10))]) if recent_ac_list else 0,
"recent_active_count": sum(1 for ac in recent_ac_list if ac > 0),
}
class_comparisons.append(
{
"class_name": class_name,
"user_count": user_count,
# 基础统计
"total_ac": int(total_ac),
"total_submission": int(total_submission),
"avg_ac": round(avg_ac, 2),
# 中位数和分位数
"median_ac": round(median_ac, 2),
"q1_ac": round(q1_ac, 2),
"q3_ac": round(q3_ac, 2),
"iqr": round(iqr, 2),
# 标准差
"std_dev": round(std_dev, 2),
# 分层统计
"top_10_avg": round(top_10_avg, 2),
"middle_80_avg": round(middle_80_avg, 2),
"bottom_10_avg": round(bottom_10_avg, 2),
# 比率统计
"excellent_rate": round(excellent_rate, 2),
"pass_rate": round(pass_rate, 2),
"active_rate": round(active_rate, 2),
# 正确率
"ac_rate": round(total_ac / total_submission * 100, 2) if total_submission > 0 else 0,
# 时间段统计(如果有)
**recent_stats,
}
)
# 计算综合分(需要所有班级数据就绪后才能归一化)
max_median = max((c["median_ac"] for c in class_comparisons), default=1) or 1
max_middle = max((c["middle_80_avg"] for c in class_comparisons), default=1) or 1
for c in class_comparisons:
score = (
0.40 * (c["median_ac"] / max_median * 100)
+ 0.15 * (c["middle_80_avg"] / max_middle * 100)
+ 0.20 * c["active_rate"]
+ 0.15 * c["pass_rate"]
+ 0.10 * c["excellent_rate"]
)
c["composite_score"] = round(score, 1)
# 按综合分排序(主),中位数(次)
class_comparisons.sort(key=lambda x: (-x["composite_score"], -x["median_ac"]))
return self.success(
{
"comparisons": class_comparisons,
"has_time_range": bool(start_time and end_time),
}
)

View File

@@ -0,0 +1,38 @@
# Generated by Django 5.2.3 on 2025-06-14 08:51
import django.db.models.deletion
from django.conf import settings
from django.db import migrations, models
class Migration(migrations.Migration):
initial = True
dependencies = [
('problem', '0001_initial'),
('submission', '0001_initial'),
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
]
operations = [
migrations.CreateModel(
name='Comment',
fields=[
('id', models.AutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
('language', models.CharField(choices=[('Python', 'Python'), ('C', 'C'), ('C++', 'C++'), ('Java', 'Java')], default='Python', max_length=10, verbose_name='解决这道题使用的语言')),
('description_rating', models.PositiveSmallIntegerField(default=5, verbose_name='题目描述的分数')),
('difficulty_rating', models.PositiveSmallIntegerField(default=5, verbose_name='题目难度的分数')),
('comprehensive_rating', models.PositiveSmallIntegerField(default=5, verbose_name='综合的分数')),
('content', models.TextField(blank=True, null=True)),
('create_time', models.DateTimeField(auto_now_add=True)),
('problem', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='problem.problem')),
('submission', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='submission.submission')),
('user', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to=settings.AUTH_USER_MODEL)),
],
options={
'db_table': 'comment',
'ordering': ('-create_time',),
},
),
]

44
comment/models.py Normal file
View File

@@ -0,0 +1,44 @@
from django.db import models
from account.models import User
from problem.models import Problem
from submission.models import Submission
class Languages(models.TextChoices):
Python = "Python", "Python"
C = "C", "C"
Cpp = "C++", "C++"
Java = "Java", "Java"
class Comment(models.Model):
problem = models.ForeignKey(Problem, on_delete=models.CASCADE)
user = models.ForeignKey(User, on_delete=models.CASCADE)
submission = models.ForeignKey(Submission, on_delete=models.CASCADE)
language = models.CharField(
max_length=10,
default=Languages.Python,
choices=Languages.choices,
verbose_name="解决这道题使用的语言",
)
description_rating = models.PositiveSmallIntegerField(
default=5,
verbose_name="题目描述的分数",
)
difficulty_rating = models.PositiveSmallIntegerField(
default=5,
verbose_name="题目难度的分数",
)
comprehensive_rating = models.PositiveSmallIntegerField(
default=5,
verbose_name="综合的分数",
)
content = models.TextField(null=True, blank=True)
create_time = models.DateTimeField(auto_now_add=True)
class Meta:
db_table = "comment"
ordering = ("-create_time",)

31
comment/serializers.py Normal file
View File

@@ -0,0 +1,31 @@
from comment.models import Comment
from utils.api import UsernameSerializer, serializers
class CreateCommentSerializer(serializers.Serializer):
problem_id = serializers.IntegerField()
description_rating = serializers.IntegerField()
difficulty_rating = serializers.IntegerField()
comprehensive_rating = serializers.IntegerField()
content = serializers.CharField(required=False, allow_blank=True)
class CommentSerializer(serializers.ModelSerializer):
class Meta:
model = Comment
fields = [
"comprehensive_rating",
"description_rating",
"difficulty_rating",
"content",
"create_time",
]
class CommentListSerializer(serializers.ModelSerializer):
problem = serializers.SlugRelatedField(read_only=True, slug_field="_id")
user = UsernameSerializer()
class Meta:
model = Comment
fields = "__all__"

8
comment/urls/admin.py Normal file
View File

@@ -0,0 +1,8 @@
from django.urls import path
from ..views.admin import CommentAPI
urlpatterns = [
path("comment", CommentAPI.as_view()),
]

9
comment/urls/oj.py Normal file
View File

@@ -0,0 +1,9 @@
from django.urls import path
from ..views.oj import CommentAPI, CommentStatisticsAPI
urlpatterns = [
path("comment", CommentAPI.as_view()),
path("comment/statistics", CommentStatisticsAPI.as_view()),
]

Some files were not shown because too many files have changed in this diff Show More